litellm/tests/unit/llms/bedrock/test_base_aws_llm.py
yuneng-jiang 5e6dc89ba1
test: move tests/test_litellm/llms into tests/unit/llms (#43191)
* ci: run the unit_selection.sh shard files on every event instead of only fork pull requests

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* ci: rename fork-flag to unit-flag now that it applies on every event

* test: move tests/test_litellm root and small trees into tests/unit

Pure renames, no content changes. Follow-up commits in this PR fix
references, merge the three files that already existed in tests/unit,
keep live-provider tests in tests/test_litellm and wire CI.

* test: carry tests/test_litellm conftest isolation into tests/unit

Callback lists, routing fallbacks, cached HTTP clients, logger state, AWS,
proxy-URL and keychain env, and session-end client cleanup now reset for
unit tests too. The environment isolation owns its MonkeyPatch so a test's
own monkeypatch is undone before the model-cost teardown runs.

* test: merge, split and prune the moved root and small-tree tests

Merge batches/test_batch_utils.py and the chat_completions and messages
dispatch tests into the files that already existed in tests/unit. Keep
the live Gemini interactions tests, the async image-fetch format test and
the OpenAI embedding scorer test in tests/test_litellm since they need
real network or keys. Put test_router.py under tests/unit/test_router so
the existing package no longer shadows it. Delete eight tests the audit
found superseded by stronger ones kept in this move.

* ci: run the moved root and small-tree tests under their legacy flags

Add the misc and responses-caching-types flags to unit_selection.sh and
CircleCI, extend enterprise-routing and mcp-integration, and point the
legacy GHA shards, Makefile, redis-compat workflow, merge smoke manifest
and change classifier at the new paths.

* test: make the new tests/unit directories packages

tests/unit/test_package_layout.py requires every directory to carry an
__init__.py, and without one the moved and retained
test_litellm_responses_bridge.py modules collide on import.

* test: scope the unit socket block to tests/unit in shared sessions

The GHA shards collect the legacy test-path and the unit selection in one
pytest session. The unit conftest's loopback-only block leaked into legacy
modules that reach the network at import. The legacy conftest now lifts the
restriction at collect and setup time, and the unit conftest re-applies it
when collecting its own modules.

* test: move tests/test_litellm/llms into tests/unit/llms

Rename-only. Moves the provider tests and the fine-tuning fixtures they
load, mirroring the old paths. Follow-up commits merge, split and wire them.

* test: merge, split and prune the moved llms tests

Merges the Databricks chat transformation tests into the existing unit
file, keeps the tests that need real keys or the network in
tests/test_litellm, deletes the audited tests a stronger unit test
already covers, and points imports at tests.unit.llms.

* ci: run the moved llms tests under their legacy flags

The Vertex AI and All Other Providers shards keep their legacy test-path
for the retained files and add the llm-vertex-ai and llm-other-providers
unit selections. CircleCI gets matching unit jobs.

* test: make the tests/unit/llms directories packages

Adds __init__.py to the moved dirs and drops the legacy ones whose
directories no longer hold tests.

* test: drop script runners and path hacks the llms split left dangling

The __main__ runners in the split openai_like files and the Databricks e2e
runner called tests that now live in the other half of the split or were
deleted. The retained legacy halves also no longer need sys.path edits.

* test: give the shard-script tests their own GITHUB_OUTPUT

They only passed where the runner set it. The CircleCI unit job's env
allowlist drops it, so the script's redirect failed there.

* test: point the router and module-deletion checks at tests/unit

router_code_coverage and code_qa_check_tests only searched tests/test_litellm,
so the moved router tests no longer counted. The two silent-experiment tests
the audit deleted were the only direct callers of those methods; they are
replaced with tests that assert the forwarded shadow request and the
recursion guard.

* test: keep the Databricks manual e2e runner and fix the SageMaker Nova run path

The Databricks e2e file is a manual script whose main() calls the tests
that were pruned, so pruning them broke the documented run. It is back to
its main version. The SageMaker Nova docstring now points at the file's
real location in tests/local_testing.

* test: keep the job's UNIT_FLAG out of the shard-script tests

---------

Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
2026-09-25 12:43:23 -07:00

3703 lines
139 KiB
Python

import asyncio
import json
from concurrent.futures import ThreadPoolExecutor
import os
import threading
import time
import pytest
from fastapi.testclient import TestClient
from collections.abc import Callable
from datetime import datetime, timedelta, timezone
from typing import Any, Dict, Optional
from unittest.mock import MagicMock, patch
from botocore.awsrequest import AWSPreparedRequest, AWSRequest
from botocore.auth import SigV4Auth
from botocore.credentials import Credentials
from botocore.exceptions import ClientError, NoCredentialsError
import litellm
from litellm.llms.bedrock.base_aws_llm import (
AwsAuthError,
BaseAWSLLM,
Boto3CredentialsInfo,
run_aws_signing,
sign_request_off_loop_if_aws,
)
from tests.unit.llms.bedrock.event_loop_probe import EventLoopProbe
# Global variable for the base_aws_llm.py file path
BASE_AWS_LLM_PATH = os.path.join(
os.path.dirname(__file__), "../../../../litellm/llms/bedrock/base_aws_llm.py"
)
@pytest.fixture(autouse=True)
def flush_shared_bedrock_iam_cache():
"""Process-wide IAM cache must not leak static/env credential entries across tests."""
BaseAWSLLM._shared_iam_cache.flush_cache()
yield
@pytest.fixture(autouse=True)
def _clean_ssl_env(monkeypatch):
"""get_ssl_verify reads these, so the sts client's verify= would otherwise depend on
the ambient environment. The published images set SSL_CERT_FILE."""
for env_var in ("SSL_CERT_FILE", "SSL_VERIFY"):
monkeypatch.delenv(env_var, raising=False)
def test_base_aws_llm_instances_share_process_wide_iam_cache():
"""Regression LIT-2662: new instances must reuse iam_cache (Bedrock passthrough is per-request)."""
first = BaseAWSLLM()
second = BaseAWSLLM()
assert first.iam_cache is second.iam_cache
assert first.iam_cache is BaseAWSLLM._shared_iam_cache
def test_static_access_key_credentials_use_iam_cache_across_calls():
"""Static access-key path hits shared iam_cache; second identical call does not refetch."""
base = BaseAWSLLM()
fake_creds = MagicMock()
with patch.object(
base,
"_auth_with_access_key_and_secret_key",
return_value=(fake_creds, 3600),
) as mock_static_auth:
base.get_credentials(
aws_access_key_id="AKIAEXAMPLE",
aws_secret_access_key="secret",
aws_region_name="us-east-1",
)
base.get_credentials(
aws_access_key_id="AKIAEXAMPLE",
aws_secret_access_key="secret",
aws_region_name="us-east-1",
)
mock_static_auth.assert_called_once()
def _os_environ_without_aws_keys() -> Dict[str, str]:
"""Strip AWS_* so get_credentials hits the ambient-env branch when no explicit keys are passed."""
return {k: v for k, v in os.environ.items() if not k.startswith("AWS_")}
def test_ambient_env_credentials_use_iam_cache_across_instances():
"""Else-branch env path uses shared iam_cache; second call on another instance does not refetch."""
base_a = BaseAWSLLM()
base_b = BaseAWSLLM()
fake_creds = MagicMock()
with patch.dict(os.environ, _os_environ_without_aws_keys(), clear=True):
with patch.object(
BaseAWSLLM,
"_auth_with_env_vars",
return_value=(fake_creds, None),
) as mock_env:
base_a.get_credentials()
base_b.get_credentials()
mock_env.assert_called_once()
def test_static_access_key_path_boto3_session_constructed_once_when_cached():
"""With real _auth_with_access_key_and_secret_key, boto3.Session is only built once per cache key."""
base_a = BaseAWSLLM()
base_b = BaseAWSLLM()
real_creds = Credentials("AKIAEXAMPLE", "secret-key-val", None)
mock_session_instance = MagicMock()
mock_session_instance.get_credentials.return_value = real_creds
with patch("boto3.Session", return_value=mock_session_instance) as mock_session_cls:
base_a.get_credentials(
aws_access_key_id="AKIAEXAMPLE",
aws_secret_access_key="secret-key-val",
aws_region_name="us-east-1",
)
base_b.get_credentials(
aws_access_key_id="AKIAEXAMPLE",
aws_secret_access_key="secret-key-val",
aws_region_name="us-east-1",
)
mock_session_cls.assert_called_once()
def test_ambient_env_path_boto3_session_constructed_once_when_cached():
"""Else branch: boto3.Session() inside _auth_with_env_vars runs once for two cache hits."""
base_a = BaseAWSLLM()
base_b = BaseAWSLLM()
real_creds = Credentials("AKIAENV", "secret-env", None)
mock_session_instance = MagicMock()
mock_session_instance.get_credentials.return_value = real_creds
with patch.dict(os.environ, _os_environ_without_aws_keys(), clear=True):
with patch(
"boto3.Session", return_value=mock_session_instance
) as mock_session_cls:
base_a.get_credentials()
base_b.get_credentials()
mock_session_cls.assert_called_once()
def test_explicit_session_token_tuple_not_cached_in_iam_cache():
"""Temporary key+secret+session paths must not use process-wide iam_cache between calls."""
base = BaseAWSLLM()
with patch.object(
base,
"_auth_with_aws_session_token",
return_value=(
Credentials("ak", "sk", "token"),
None,
),
) as mock_sess:
base.get_credentials(
aws_access_key_id="AKIA",
aws_secret_access_key="sec",
aws_session_token="tok",
)
base.get_credentials(
aws_access_key_id="AKIA",
aws_secret_access_key="sec",
aws_session_token="tok",
)
assert mock_sess.call_count == 2
def test_aws_profile_path_not_cached_in_iam_cache():
base = BaseAWSLLM()
with patch.object(
base,
"_auth_with_aws_profile",
return_value=(Credentials("prof-ak", "prof-sk", None), None),
) as mock_profile:
base.get_credentials(aws_profile_name="my-profile")
base.get_credentials(aws_profile_name="my-profile")
assert mock_profile.call_count == 2
def test_get_credentials_does_not_expand_request_env_reference():
"""
A parameter of the form os.environ/<VAR> reaching get_credentials is left as-is
rather than expanded against the process environment, so the downstream auth
helper only ever receives the literal value.
"""
env = _os_environ_without_aws_keys()
env["SERVER_ONLY_VALUE"] = "config-managed-value"
base = BaseAWSLLM()
with patch.dict(os.environ, env, clear=True), patch.object(
base,
"_auth_with_aws_profile",
return_value=(Credentials("ak", "sk", None), None),
) as mock_profile:
base.get_credentials(aws_profile_name="os.environ/SERVER_ONLY_VALUE")
assert mock_profile.call_args.args[0] == "os.environ/SERVER_ONLY_VALUE"
assert "config-managed-value" not in str(mock_profile.call_args)
def test_get_credentials_falls_back_to_ambient_aws_profile_name_env():
"""
The fixed AWS_* ambient fallback keeps working: an unset aws_profile_name
resolves from the AWS_PROFILE_NAME environment variable.
"""
env = _os_environ_without_aws_keys()
env["AWS_PROFILE_NAME"] = "ambient-profile"
base = BaseAWSLLM()
with patch.dict(os.environ, env, clear=True), patch.object(
base,
"_auth_with_aws_profile",
return_value=(Credentials("ak", "sk", None), None),
) as mock_profile:
base.get_credentials(aws_profile_name=None)
assert mock_profile.call_args.args[0] == "ambient-profile"
def test_get_credentials_ambient_fallback_resolves_aws_external_id():
"""
Each unset param falls back to its own AWS_* env var. Regression for an index
misalignment between the value list and the env-name list, which left
AWS_EXTERNAL_ID unresolved.
"""
env = _os_environ_without_aws_keys()
env["AWS_EXTERNAL_ID"] = "ext-from-env"
base = BaseAWSLLM()
with patch.dict(os.environ, env, clear=True), patch.object(
base,
"_auth_with_aws_role",
return_value=(Credentials("ak", "sk", "tok"), None),
) as mock_role:
base.get_credentials(
aws_role_name="arn:aws:iam::123456789012:role/x",
aws_session_name="s",
)
assert mock_role.call_args.kwargs["aws_external_id"] == "ext-from-env"
def _capturing_sts_client(captured: Dict[str, Any]) -> MagicMock:
sts = MagicMock()
def _assume(**params):
captured["WebIdentityToken"] = params.get("WebIdentityToken")
return {
"Credentials": {
"AccessKeyId": "AKIA",
"SecretAccessKey": "sk",
"SessionToken": "tok",
},
"PackedPolicySize": 10,
}
sts.assume_role_with_web_identity.side_effect = _assume
return sts
@pytest.mark.parametrize(
"token_ref",
["os.environ/SERVER_ONLY_VALUE", "SERVER_ONLY_VALUE"],
ids=["os_environ_prefix", "bare_env_name"],
)
def test_web_identity_token_env_reference_not_expanded(token_ref):
"""
A web-identity token that is an environment-variable reference (an os.environ/
prefix, or a bare name matching an env var) is rejected rather than expanded, so
the process-environment value is never used as the token.
"""
env = _os_environ_without_aws_keys()
env["SERVER_ONLY_VALUE"] = "server-only-value"
captured: Dict[str, Any] = {}
base = BaseAWSLLM()
with patch.dict(os.environ, env, clear=True), patch(
"boto3.client", side_effect=lambda *a, **k: _capturing_sts_client(captured)
), patch("boto3.Session", return_value=MagicMock()):
with pytest.raises(AwsAuthError):
base.get_credentials(
aws_web_identity_token=token_ref,
aws_role_name="arn:aws:iam::123456789012:role/x",
aws_session_name="s",
aws_sts_endpoint="https://custom-sts.example",
)
assert "server-only-value" not in str(captured)
def test_web_identity_token_oidc_reference_still_resolved():
"""
The env-reference guard does not over-reject: an oidc/ reference still flows to
get_secret (mocked to None here), surfacing the existing 401 rather than the 400
used for rejected env-var references.
"""
base = BaseAWSLLM()
env = _os_environ_without_aws_keys()
with patch.dict(os.environ, env, clear=True), patch(
"litellm.llms.bedrock.base_aws_llm.get_secret", return_value=None
):
with pytest.raises(AwsAuthError) as exc:
base.get_credentials(
aws_web_identity_token="oidc/circleci/",
aws_role_name="arn:aws:iam::123456789012:role/x",
aws_session_name="s",
)
assert exc.value.status_code == 401
def test_web_identity_credentials_cached_in_iam_cache():
"""
Web identity STS credentials are cached for their STS lifetime, so repeated
get_credentials calls (e.g. per-request guardrail auth) reuse the assumed-role
credentials instead of replaying the OIDC token to STS on every request.
"""
base = BaseAWSLLM()
with patch.object(
base,
"_auth_with_web_identity_token",
return_value=(Credentials("wi-ak", "wi-sk", "wi-tok"), 3540),
) as mock_wi:
kwargs = dict(
aws_web_identity_token="oidc/google/https://example.com/",
aws_role_name="arn:aws:iam::123456789012:role/WebIdentity",
aws_session_name="web-id-session",
)
first = base.get_credentials(**kwargs)
second = base.get_credentials(**kwargs)
assert mock_wi.call_count == 1
assert first is second
def test_web_identity_cache_is_keyed_on_credential_args():
"""
Two web identity configs that differ in any credential arg (here the role)
must not share cached credentials.
"""
base = BaseAWSLLM()
with patch.object(
base,
"_auth_with_web_identity_token",
side_effect=[
(Credentials("wi-ak-a", "wi-sk-a", "wi-tok-a"), 3540),
(Credentials("wi-ak-b", "wi-sk-b", "wi-tok-b"), 3540),
],
) as mock_wi:
first = base.get_credentials(
aws_web_identity_token="oidc/google/https://example.com/",
aws_role_name="arn:aws:iam::123456789012:role/RoleA",
aws_session_name="web-id-session",
)
second = base.get_credentials(
aws_web_identity_token="oidc/google/https://example.com/",
aws_role_name="arn:aws:iam::123456789012:role/RoleB",
aws_session_name="web-id-session",
)
assert mock_wi.call_count == 2
assert first.access_key != second.access_key
def test_boto3_init_tracer_wrapping():
"""
Test that all boto3 initializations are wrapped in tracer.trace or @tracer.wrap
Ensures observability of boto3 calls in litellm.
"""
# Get the source code of base_aws_llm.py
with open(BASE_AWS_LLM_PATH, "r") as f:
content = f.read()
# List all boto3 initialization patterns we want to check
boto3_init_patterns = ["boto3.client", "boto3.Session"]
lines = content.split("\n")
# Check each boto3 initialization is wrapped in tracer.trace
for line_number, line in enumerate(lines, 1):
if line.lstrip().startswith("#"):
continue
for pattern in boto3_init_patterns:
if pattern in line:
# Look back up to 5 lines for decorator or trace block
start_line = max(0, line_number - 5)
context_lines = lines[start_line:line_number]
has_trace = (
"tracer.trace" in line
or any("tracer.trace" in prev_line for prev_line in context_lines)
or any("@tracer.wrap" in prev_line for prev_line in context_lines)
)
if not has_trace:
print(f"\nContext for line {line_number}:")
for i, ctx_line in enumerate(context_lines, start=start_line + 1):
print(f"{i}: {ctx_line}")
assert (
has_trace
), f"boto3 initialization '{pattern}' on line {line_number} is not wrapped with tracer.trace or @tracer.wrap"
def test_auth_functions_tracer_wrapping():
"""
Test that all _auth functions in base_aws_llm.py are wrapped with @tracer.wrap
Ensures observability of AWS authentication calls in litellm.
"""
# Get the source code of base_aws_llm.py
with open(BASE_AWS_LLM_PATH, "r") as f:
content = f.read()
lines = content.split("\n")
# Check each line for _auth function definitions
for line_number, line in enumerate(lines, 1):
if line.strip().startswith("def _auth_"):
# Look back up to 2 lines for the @tracer.wrap decorator
start_line = max(0, line_number - 2)
context_lines = lines[start_line:line_number]
has_tracer_wrap = any(
"@tracer.wrap" in prev_line for prev_line in context_lines
)
if not has_tracer_wrap:
print(f"\nContext for line {line_number}:")
for i, ctx_line in enumerate(context_lines, start=start_line + 1):
print(f"{i}: {ctx_line}")
assert (
has_tracer_wrap
), f"Auth function on line {line_number} is not wrapped with @tracer.wrap: {line.strip()}"
def test_get_aws_region_name_boto3_fallback():
"""
Test the boto3 session fallback logic in _get_aws_region_name method.
This tests the specific code block that tries to get the region from boto3.Session()
when aws_region_name is None and not found in environment variables.
"""
base_aws_llm = BaseAWSLLM()
# Test case 1: boto3.Session() returns a configured region
with patch("litellm.llms.bedrock.base_aws_llm.get_secret") as mock_get_secret:
mock_get_secret.return_value = None # No region in env vars
with patch("boto3.Session") as mock_boto3_session:
mock_session = MagicMock()
mock_session.region_name = "us-east-1"
mock_boto3_session.return_value = mock_session
optional_params = {}
result = base_aws_llm._get_aws_region_name(optional_params)
assert result == "us-east-1"
mock_boto3_session.assert_called_once()
# Test case 2: boto3.Session() returns None for region (should default to us-west-2)
with patch("litellm.llms.bedrock.base_aws_llm.get_secret") as mock_get_secret:
mock_get_secret.return_value = None # No region in env vars
with patch("boto3.Session") as mock_boto3_session:
mock_session = MagicMock()
mock_session.region_name = None
mock_boto3_session.return_value = mock_session
optional_params = {}
result = base_aws_llm._get_aws_region_name(optional_params)
assert result == "us-west-2"
mock_boto3_session.assert_called_once()
# Test case 3: boto3 import/session creation raises exception (should default to us-west-2)
with patch("litellm.llms.bedrock.base_aws_llm.get_secret") as mock_get_secret:
mock_get_secret.return_value = None # No region in env vars
with patch("boto3.Session") as mock_boto3_session:
mock_boto3_session.side_effect = Exception("boto3 not available")
optional_params = {}
result = base_aws_llm._get_aws_region_name(optional_params)
assert result == "us-west-2"
mock_boto3_session.assert_called_once()
# Test case 4: aws_region_name is provided in optional_params (should not use boto3)
with patch("boto3.Session") as mock_boto3_session:
optional_params = {"aws_region_name": "eu-west-1"}
result = base_aws_llm._get_aws_region_name(optional_params)
assert result == "eu-west-1"
mock_boto3_session.assert_not_called()
# Test case 5: aws_region_name found in environment variables (should not use boto3)
with patch("litellm.llms.bedrock.base_aws_llm.get_secret") as mock_get_secret:
def side_effect(key, default=None):
if key == "AWS_REGION_NAME":
return "ap-southeast-1"
return default
mock_get_secret.side_effect = side_effect
with patch("boto3.Session") as mock_boto3_session:
optional_params = {}
result = base_aws_llm._get_aws_region_name(optional_params)
assert result == "ap-southeast-1"
mock_boto3_session.assert_not_called()
@pytest.mark.parametrize(
"bad_region",
[
"us-east-1@example.com/",
"us-east-1@example.com",
"us-east-1/path",
"us-east-1.example.com",
"us-east-1:8080",
"us-east-1#fragment",
"us-east-1?query=1",
"us-east-1\\path",
"US-EAST-1", # uppercase not allowed
"us east 1", # spaces not allowed
"", # empty string not allowed
"us-east-1\n", # trailing newline must not slip past $
],
)
def test_get_aws_region_name_rejects_malformed_region(bad_region):
"""
Region names are interpolated into endpoint URL templates, so any value
containing characters that would alter URL parsing must be rejected.
"""
base_aws_llm = BaseAWSLLM()
with pytest.raises(ValueError, match="Invalid AWS region format"):
base_aws_llm._get_aws_region_name(
optional_params={"aws_region_name": bad_region}
)
@pytest.mark.parametrize(
"valid_region",
[
"us-east-1",
"eu-west-2",
"ap-southeast-1",
"us-gov-west-1",
"cn-north-1",
"me-south-1",
],
)
def test_get_aws_region_name_accepts_valid_regions(valid_region):
"""Real AWS region formats must continue to work after the format guard."""
base_aws_llm = BaseAWSLLM()
result = base_aws_llm._get_aws_region_name(
optional_params={"aws_region_name": valid_region}
)
assert result == valid_region
def test_get_aws_region_name_rejects_malformed_region_from_env():
"""
A malformed AWS_REGION / AWS_REGION_NAME env value must also be rejected
before it can flow into a URL template.
"""
base_aws_llm = BaseAWSLLM()
with patch("litellm.llms.bedrock.base_aws_llm.get_secret") as mock_get_secret:
def side_effect(key, default=None):
if key == "AWS_REGION_NAME":
return "us-east-1@example.com/"
return default
mock_get_secret.side_effect = side_effect
with pytest.raises(ValueError, match="Invalid AWS region format"):
base_aws_llm._get_aws_region_name(optional_params={})
def test_get_aws_region_name_for_non_llm_api_calls_rejects_malformed_param():
"""
The non-LLM helper (used by Guardrails, Vector Stores, etc.) must validate
a region passed in directly so it can't flow into a URL template.
"""
base_aws_llm = BaseAWSLLM()
with pytest.raises(ValueError, match="Invalid AWS region format"):
base_aws_llm.get_aws_region_name_for_non_llm_api_calls(
aws_region_name="us-east-1@example.com/"
)
def test_get_aws_region_name_for_non_llm_api_calls_rejects_malformed_env():
"""
A malformed AWS_REGION / AWS_REGION_NAME env value must be rejected on the
non-LLM path too — Guardrails and Vector Stores read the same env vars.
"""
base_aws_llm = BaseAWSLLM()
with patch("litellm.llms.bedrock.base_aws_llm.get_secret") as mock_get_secret:
def side_effect(key, default=None):
if key == "AWS_REGION_NAME":
return "us-east-1@example.com/"
return default
mock_get_secret.side_effect = side_effect
with pytest.raises(ValueError, match="Invalid AWS region format"):
base_aws_llm.get_aws_region_name_for_non_llm_api_calls()
def test_get_aws_region_name_for_non_llm_api_calls_accepts_valid_region():
"""The non-LLM helper still returns valid regions unchanged."""
base_aws_llm = BaseAWSLLM()
assert (
base_aws_llm.get_aws_region_name_for_non_llm_api_calls(
aws_region_name="us-east-1"
)
== "us-east-1"
)
def test_get_aws_region_from_model_arn_rejects_malformed_region():
"""
If the region segment of a model ARN does not match the expected format,
the helper must return None so the caller falls back to env / default.
"""
base_aws_llm = BaseAWSLLM()
bad_arn = (
"arn:aws:bedrock:us-east-1@example.com:123456789012"
":foundation-model/anthropic.claude-3-sonnet"
)
assert base_aws_llm._get_aws_region_from_model_arn(bad_arn) is None
good_arn = (
"arn:aws:bedrock:us-east-1:123456789012"
":foundation-model/anthropic.claude-3-sonnet"
)
assert base_aws_llm._get_aws_region_from_model_arn(good_arn) == "us-east-1"
def test_sign_request_with_env_var_bearer_token():
# Create instance of actual class
llm = BaseAWSLLM()
# Test data
service_name = "bedrock"
headers = {"Custom-Header": "test"}
optional_params = {}
request_data = {"prompt": "test"}
api_base = "https://api.example.com"
# Mock environment variable
with patch.dict(os.environ, {"AWS_BEARER_TOKEN_BEDROCK": "test_token"}):
# Execute
result_headers, result_body = llm._sign_request(
service_name=service_name,
headers=headers,
optional_params=optional_params,
request_data=request_data,
api_base=api_base,
)
# Assert
assert result_headers["Authorization"] == "Bearer test_token"
assert result_headers["Content-Type"] == "application/json"
assert result_headers["Custom-Header"] == "test"
assert result_body == json.dumps(request_data).encode()
def test_sign_request_with_sigv4():
llm = BaseAWSLLM()
# Mock AWS credentials and SigV4 auth
mock_credentials = Credentials("test_key", "test_secret", "test_token")
mock_sigv4 = MagicMock()
mock_request = MagicMock()
mock_request.headers = {
"Authorization": "AWS4-HMAC-SHA256 Credential=test",
"Content-Type": "application/json",
}
mock_request.body = b'{"prompt": "test"}'
# Test data
service_name = "bedrock"
headers = {"Custom-Header": "test"}
optional_params = {
"aws_access_key_id": "test_key",
"aws_secret_access_key": "test_secret",
"aws_region_name": "us-west-2",
}
request_data = {"prompt": "test"}
api_base = "https://api.example.com"
# Mock the necessary components
with (
patch("botocore.auth.SigV4Auth", return_value=mock_sigv4),
patch("botocore.awsrequest.AWSRequest", return_value=mock_request),
patch.object(llm, "get_credentials", return_value=mock_credentials),
patch.object(llm, "_get_aws_region_name", return_value="us-west-2"),
):
result_headers, result_body = llm._sign_request(
service_name=service_name,
headers=headers,
optional_params=optional_params,
request_data=request_data,
api_base=api_base,
)
# Assert
assert "Authorization" in result_headers
assert result_headers["Authorization"] != "Bearer test_token"
assert result_headers["Content-Type"] == "application/json"
assert result_body == mock_request.body
def test_sign_request_with_api_key_bearer_token():
"""
Test that _sign_request uses the api_key parameter as a bearer token when provided
"""
llm = BaseAWSLLM()
# Test data
service_name = "bedrock"
headers = {"Custom-Header": "test"}
optional_params = {}
request_data = {"prompt": "test"}
api_base = "https://api.example.com"
api_key = "test_api_key"
# Execute with api_key parameter
result_headers, result_body = llm._sign_request(
service_name=service_name,
headers=headers,
optional_params=optional_params,
request_data=request_data,
api_base=api_base,
api_key=api_key,
)
# Assert
assert result_headers["Authorization"] == f"Bearer {api_key}"
assert result_headers["Content-Type"] == "application/json"
assert result_headers["Custom-Header"] == "test"
assert result_body == json.dumps(request_data).encode()
def test_get_request_headers_with_env_var_bearer_token():
# Setup
llm = BaseAWSLLM()
credentials = Credentials("test_key", "test_secret", "test_token")
headers = {"Content-Type": "application/json"}
headers_dict = headers.copy()
# Create mock request
mock_prepared_request = MagicMock(spec=AWSPreparedRequest)
mock_request = MagicMock(spec=AWSRequest)
mock_request.headers = headers_dict
mock_request.prepare.return_value = mock_prepared_request
def mock_aws_request_init(method, url, data, headers):
mock_request.headers.update(headers)
return mock_request
# Test with bearer token
with (
patch.dict(os.environ, {"AWS_BEARER_TOKEN_BEDROCK": "test_token"}),
patch("botocore.awsrequest.AWSRequest", side_effect=mock_aws_request_init),
):
result = llm.get_request_headers(
credentials=credentials,
aws_region_name="us-west-2",
extra_headers=None,
endpoint_url="https://api.example.com",
data='{"prompt": "test"}',
headers=headers_dict,
)
# Assert
assert mock_request.headers["Authorization"] == "Bearer test_token"
assert result == mock_prepared_request
def test_get_request_headers_with_sigv4():
# Setup
llm = BaseAWSLLM()
credentials = Credentials("test_key", "test_secret", "test_token")
headers = {"Content-Type": "application/json"}
# Create mock request and SigV4 instance
mock_request = MagicMock(spec=AWSRequest)
mock_request.headers = headers.copy()
mock_request.prepare.return_value = MagicMock(spec=AWSPreparedRequest)
mock_sigv4 = MagicMock()
# Test without bearer token (should use SigV4)
with (
patch.dict(os.environ, {}, clear=True),
patch("botocore.auth.SigV4Auth", return_value=mock_sigv4) as mock_sigv4_class,
patch("botocore.awsrequest.AWSRequest", return_value=mock_request),
):
result = llm.get_request_headers(
credentials=credentials,
aws_region_name="us-west-2",
extra_headers=None,
endpoint_url="https://api.example.com",
data='{"prompt": "test"}',
headers=headers,
)
# Verify SigV4 authentication and result
mock_sigv4_class.assert_called_once_with(credentials, "bedrock", "us-west-2")
mock_sigv4.add_auth.assert_called_once_with(mock_request)
assert result == mock_request.prepare.return_value
def test_get_request_headers_without_credentials_or_bearer_token_raises_no_credentials():
"""Bearer-token auth needs no SigV4 principal, so `credentials` may be None.
Reaching the SigV4 branch with neither must fail the way botocore always
has instead of signing with a missing principal."""
llm = BaseAWSLLM()
with patch.dict(os.environ, {}, clear=True), pytest.raises(NoCredentialsError):
llm.get_request_headers(
credentials=None,
aws_region_name="us-west-2",
extra_headers=None,
endpoint_url="https://api.example.com",
data='{"prompt": "test"}',
headers={"Content-Type": "application/json"},
)
def test_sigv4_matches_rust_golden_vector():
request = AWSRequest(
method="POST",
url="https://bedrock-runtime.us-east-1.amazonaws.com/model/amazon.titan-text-express-v1/invoke",
data=b'{"input":"hello"}',
headers={"Content-Type": "application/json"},
)
credentials = Credentials(
"AKIDEXAMPLE",
"wJalrXUtnFEMI/K7MDENG+bPxRfiCYEXAMPLEKEY",
"session-token",
)
with patch("botocore.auth.get_current_datetime", return_value=datetime(2024, 1, 2, 3, 4, 5)):
SigV4Auth(credentials, "bedrock", "us-east-1").add_auth(request)
assert request.headers["X-Amz-Date"] == "20240102T030405Z"
assert request.headers["X-Amz-Security-Token"] == "session-token"
assert (
request.headers["Authorization"]
== "AWS4-HMAC-SHA256 Credential=AKIDEXAMPLE/20240102/us-east-1/bedrock/aws4_request, "
"SignedHeaders=content-type;host;x-amz-date;x-amz-security-token, "
"Signature=55c027ef47527d3ad63f1735f9d099efdbc99f296ff914bd94e727e24ec0e464"
)
def test_get_request_headers_with_api_key_bearer_token():
"""
Test that get_request_headers uses the api_key parameter as a bearer token when provided
"""
# Setup
llm = BaseAWSLLM()
credentials = Credentials("test_key", "test_secret", "test_token")
headers = {"Content-Type": "application/json"}
headers_dict = headers.copy()
api_key = "test_api_key"
# Create mock request
mock_prepared_request = MagicMock(spec=AWSPreparedRequest)
mock_request = MagicMock(spec=AWSRequest)
mock_request.headers = headers_dict
mock_request.prepare.return_value = mock_prepared_request
def mock_aws_request_init(method, url, data, headers):
mock_request.headers.update(headers)
return mock_request
# Test with api_key parameter
with (
patch.dict(os.environ, {}, clear=True),
patch("botocore.awsrequest.AWSRequest", side_effect=mock_aws_request_init),
):
result = llm.get_request_headers(
credentials=credentials,
aws_region_name="us-west-2",
extra_headers=None,
endpoint_url="https://api.example.com",
data='{"prompt": "test"}',
headers=headers_dict,
api_key=api_key,
)
# Assert
assert mock_request.headers["Authorization"] == f"Bearer {api_key}"
assert result == mock_prepared_request
def test_role_assumption_without_session_name():
"""
Test for issue 12583: Role assumption should work when only aws_role_name is provided
without aws_session_name. The system should auto-generate a session name.
"""
base_aws_llm = BaseAWSLLM()
# Mock the boto3 STS client
mock_sts_client = MagicMock()
# Mock the STS response with proper expiration handling
mock_expiry = MagicMock()
mock_expiry.tzinfo = timezone.utc
# Create a timedelta object that returns 3600 when total_seconds() is called
time_diff = MagicMock()
time_diff.total_seconds.return_value = 3600
mock_expiry.__sub__ = MagicMock(return_value=time_diff)
mock_sts_response = {
"Credentials": {
"AccessKeyId": "assumed-access-key",
"SecretAccessKey": "assumed-secret-key",
"SessionToken": "assumed-session-token",
"Expiration": mock_expiry,
}
}
mock_sts_client.assume_role.return_value = mock_sts_response
# Test case 1: aws_role_name provided without aws_session_name
with patch("boto3.client", return_value=mock_sts_client):
credentials = base_aws_llm.get_credentials(
aws_role_name="arn:aws:iam::2222222222222:role/LitellmEvalBedrockRole"
)
# Verify assume_role was called
mock_sts_client.assume_role.assert_called_once()
# Check the call arguments
call_args = mock_sts_client.assume_role.call_args
assert (
call_args[1]["RoleArn"]
== "arn:aws:iam::2222222222222:role/LitellmEvalBedrockRole"
)
# Session name should be auto-generated with format "litellm-session-{timestamp}"
assert call_args[1]["RoleSessionName"].startswith("litellm-session-")
# Verify credentials are returned correctly
assert isinstance(credentials, Credentials)
assert credentials.access_key == "assumed-access-key"
assert credentials.secret_key == "assumed-secret-key"
assert credentials.token == "assumed-session-token"
# Test case 2: Both aws_role_name and aws_session_name provided (existing behavior)
mock_sts_client.reset_mock()
with patch("boto3.client", return_value=mock_sts_client):
credentials = base_aws_llm.get_credentials(
aws_role_name="arn:aws:iam::2222222222222:role/LitellmEvalBedrockRole",
aws_session_name="my-custom-session",
)
# Verify assume_role was called with custom session name
mock_sts_client.assume_role.assert_called_once()
call_args = mock_sts_client.assume_role.call_args
assert call_args[1]["RoleSessionName"] == "my-custom-session"
# Test case 3: identical AssumeRole args reuse the cached STS session instead of re-assuming.
BaseAWSLLM._shared_iam_cache.flush_cache()
mock_sts_client.reset_mock()
with patch("boto3.client", return_value=mock_sts_client):
credentials1 = base_aws_llm.get_credentials(
aws_role_name="arn:aws:iam::2222222222222:role/LitellmEvalBedrockRole"
)
credentials2 = base_aws_llm.get_credentials(
aws_role_name="arn:aws:iam::2222222222222:role/LitellmEvalBedrockRole"
)
assert mock_sts_client.assume_role.call_count == 1
assert credentials1.access_key == credentials2.access_key
def _assume_role_sts_mock(access_key: str = "assumed-access-key") -> MagicMock:
"""STS client mock whose caller identity is a different role, forcing the AssumeRole path."""
mock_sts_client = MagicMock()
mock_sts_client.assume_role.return_value = {
"Credentials": {
"AccessKeyId": access_key,
"SecretAccessKey": "assumed-secret-key",
"SessionToken": "assumed-session-token",
"Expiration": datetime.now(timezone.utc) + timedelta(hours=1),
}
}
mock_sts_client.get_caller_identity.return_value = {
"Arn": "arn:aws:sts::111111111111:assumed-role/SomeOtherRole/session-name",
}
return mock_sts_client
def _env_without_irsa() -> dict:
return {
k: v
for k, v in os.environ.items()
if k not in ("AWS_ROLE_ARN", "AWS_WEB_IDENTITY_TOKEN_FILE")
}
def test_assume_role_path_uses_process_iam_cache():
"""Repeat calls with identical AssumeRole args make no further STS calls of either kind."""
base_aws_llm = BaseAWSLLM()
mock_sts_client = _assume_role_sts_mock()
role_arn = "arn:aws:iam::2222222222222:role/LitellmEvalBedrockRole"
with patch.dict(os.environ, _env_without_irsa(), clear=True):
with patch("boto3.client", return_value=mock_sts_client):
first = base_aws_llm.get_credentials(
aws_role_name=role_arn, aws_session_name="attributed-user-1"
)
mock_sts_client.get_caller_identity.reset_mock()
second = base_aws_llm.get_credentials(
aws_role_name=role_arn, aws_session_name="attributed-user-1"
)
mock_sts_client.get_caller_identity.assert_not_called()
assert mock_sts_client.assume_role.call_count == 1
assert first.access_key == second.access_key
assert first.token == second.token
def test_assume_role_cache_is_scoped_per_session_name():
"""Each attributed identity gets its own STS session; no user is served another user's creds."""
base_aws_llm = BaseAWSLLM()
mock_sts_client = _assume_role_sts_mock()
mock_sts_client.assume_role.side_effect = [
{
"Credentials": {
"AccessKeyId": f"assumed-access-key-{session_name}",
"SecretAccessKey": "assumed-secret-key",
"SessionToken": f"assumed-session-token-{session_name}",
"Expiration": datetime.now(timezone.utc) + timedelta(hours=1),
}
}
for session_name in ("attributed-user-1", "attributed-user-2")
]
role_arn = "arn:aws:iam::2222222222222:role/LitellmEvalBedrockRole"
with patch.dict(os.environ, _env_without_irsa(), clear=True):
with patch("boto3.client", return_value=mock_sts_client):
user_1 = base_aws_llm.get_credentials(
aws_role_name=role_arn, aws_session_name="attributed-user-1"
)
user_2 = base_aws_llm.get_credentials(
aws_role_name=role_arn, aws_session_name="attributed-user-2"
)
user_1_again = base_aws_llm.get_credentials(
aws_role_name=role_arn, aws_session_name="attributed-user-1"
)
assert mock_sts_client.assume_role.call_count == 2
session_names = [
call.kwargs["RoleSessionName"] for call in mock_sts_client.assume_role.call_args_list
]
assert session_names == ["attributed-user-1", "attributed-user-2"]
assert user_1.access_key == "assumed-access-key-attributed-user-1"
assert user_2.access_key == "assumed-access-key-attributed-user-2"
assert user_1_again.access_key == user_1.access_key
def test_assume_role_cache_entry_expires_with_the_sts_session():
"""The cache entry is written with the STS lifetime, not an indefinite or default TTL."""
base_aws_llm = BaseAWSLLM()
mock_sts_client = _assume_role_sts_mock()
role_arn = "arn:aws:iam::2222222222222:role/LitellmEvalBedrockRole"
with patch.dict(os.environ, _env_without_irsa(), clear=True):
with patch("boto3.client", return_value=mock_sts_client):
base_aws_llm.get_credentials(
aws_role_name=role_arn, aws_session_name="attributed-user-1"
)
in_memory_cache = BaseAWSLLM._shared_iam_cache.in_memory_cache
assert len(in_memory_cache.ttl_dict) == 1
remaining_ttl = list(in_memory_cache.ttl_dict.values())[0] - time.time()
# 1h STS session minus the 60s safety margin, minus test execution time
assert 3400 < remaining_ttl <= 3540
def test_assume_role_credentials_expired_in_cache_trigger_a_fresh_sts_call():
"""Once the STS session lapses the next call re-assumes rather than serving dead credentials."""
base_aws_llm = BaseAWSLLM()
mock_sts_client = _assume_role_sts_mock()
role_arn = "arn:aws:iam::2222222222222:role/LitellmEvalBedrockRole"
with patch.dict(os.environ, _env_without_irsa(), clear=True):
with patch("boto3.client", return_value=mock_sts_client):
base_aws_llm.get_credentials(
aws_role_name=role_arn, aws_session_name="attributed-user-1"
)
in_memory_cache = BaseAWSLLM._shared_iam_cache.in_memory_cache
expired_key = list(in_memory_cache.ttl_dict.keys())[0]
in_memory_cache.ttl_dict[expired_key] = time.time() - 1
base_aws_llm.get_credentials(
aws_role_name=role_arn, aws_session_name="attributed-user-1"
)
assert mock_sts_client.assume_role.call_count == 2
def test_concurrent_cold_cache_calls_for_one_identity_make_a_single_sts_call():
"""Concurrent misses on one identity single-flight instead of each issuing its own AssumeRole."""
base_aws_llm = BaseAWSLLM()
role_arn = "arn:aws:iam::2222222222222:role/LitellmEvalBedrockRole"
calling_threads = set()
threads_lock = threading.Lock()
def slow_assume_role(**kwargs):
with threads_lock:
calling_threads.add(threading.get_ident())
# Hold the fetch open so every other worker is inside get_credentials while this one runs;
# without that overlap the test would pass against unsynchronised code.
time.sleep(0.2)
return {
"Credentials": {
"AccessKeyId": "assumed-access-key",
"SecretAccessKey": "assumed-secret-key",
"SessionToken": "assumed-session-token",
"Expiration": datetime.now(timezone.utc) + timedelta(hours=1),
}
}
mock_sts_client = _assume_role_sts_mock()
mock_sts_client.assume_role.side_effect = slow_assume_role
barrier = threading.Barrier(8)
resolved = []
def resolve():
barrier.wait()
resolved.append(
base_aws_llm.get_credentials(aws_role_name=role_arn, aws_session_name="attribution-1")
)
with patch.dict(os.environ, _env_without_irsa(), clear=True):
with patch("boto3.client", return_value=mock_sts_client):
workers = [threading.Thread(target=resolve) for _ in range(8)]
for worker in workers:
worker.start()
for worker in workers:
worker.join(timeout=30)
# The harness has to have produced real concurrency, or the assertion below proves nothing
assert len(resolved) == 8
assert len({id(credentials) for credentials in resolved}) == 1
assert mock_sts_client.assume_role.call_count == 1
assert len(calling_threads) == 1
def test_ambient_identity_matching_target_role_is_cached_without_repeating_the_probe():
"""The skip-AssumeRole branch caches too, so GetCallerIdentity stops running per request."""
base_aws_llm = BaseAWSLLM()
role_arn = "arn:aws:iam::2222222222222:role/LitellmEvalBedrockRole"
mock_sts_client = MagicMock()
mock_sts_client.get_caller_identity.return_value = {
"Arn": "arn:aws:sts::2222222222222:assumed-role/LitellmEvalBedrockRole/session-name",
}
mock_ambient_credentials = MagicMock()
mock_ambient_credentials.access_key = "ambient-access-key"
with patch.dict(os.environ, _env_without_irsa(), clear=True):
with patch("boto3.client", return_value=mock_sts_client):
with patch.object(
base_aws_llm,
"_auth_with_env_vars",
return_value=(mock_ambient_credentials, None),
) as mock_env_auth:
first = base_aws_llm.get_credentials(aws_role_name=role_arn)
second = base_aws_llm.get_credentials(aws_role_name=role_arn)
assert mock_sts_client.get_caller_identity.call_count == 1
assert mock_sts_client.assume_role.call_count == 0
mock_env_auth.assert_called_once()
assert first.access_key == second.access_key == "ambient-access-key"
def test_cache_keys_are_different_for_different_roles():
"""
Test that cache keys are different for different AWS roles.
This ensures that credentials for different roles don't get mixed up.
"""
base_aws_llm = BaseAWSLLM()
# Create arguments for two different roles
args1 = {
"aws_access_key_id": None,
"aws_secret_access_key": None,
"aws_role_name": "arn:aws:iam::1111111111111:role/LitellmRole",
"aws_session_name": "test-session-1",
}
args2 = {
"aws_access_key_id": None,
"aws_secret_access_key": None,
"aws_role_name": "arn:aws:iam::2222222222222:role/LitellmEvalBedrockRole",
"aws_session_name": "test-session-2",
}
# Generate cache keys
cache_key1 = base_aws_llm.get_cache_key(args1)
cache_key2 = base_aws_llm.get_cache_key(args2)
# Cache keys should be different because the role names are different
assert cache_key1 != cache_key2
def test_different_roles_without_session_names_should_not_share_cache():
"""
Test that different roles with auto-generated session names don't share cache.
This was the original issue where cache keys were the same for different roles.
"""
base_aws_llm = BaseAWSLLM()
# Create arguments for two different roles without session names
args1 = {
"aws_access_key_id": None,
"aws_secret_access_key": None,
"aws_role_name": "arn:aws:iam::1111111111111:role/LitellmRole",
"aws_session_name": None,
}
args2 = {
"aws_access_key_id": None,
"aws_secret_access_key": None,
"aws_role_name": "arn:aws:iam::2222222222222:role/LitellmEvalBedrockRole",
"aws_session_name": None,
}
# Generate cache keys
cache_key1 = base_aws_llm.get_cache_key(args1)
cache_key2 = base_aws_llm.get_cache_key(args2)
# Cache keys should be different because the role names are different
assert cache_key1 != cache_key2
@pytest.mark.parametrize(
"role_kwargs,expected_client_kwargs",
[
({}, {"verify": True}),
(
{"aws_region_name": "us-east-1"},
{"verify": True, "region_name": "us-east-1"},
),
(
{"aws_sts_endpoint": "https://sts.eu-west-1.amazonaws.com"},
{
"endpoint_url": "https://sts.eu-west-1.amazonaws.com",
"region_name": "eu-west-1",
"verify": True,
},
),
],
ids=["no_region_or_endpoint", "configured_region_is_sts_fallback", "explicit_sts_endpoint"],
)
def test_eks_irsa_ambient_credentials_used(role_kwargs, expected_client_kwargs):
"""
Test that in EKS/IRSA environments, ambient credentials are used when no explicit keys provided.
This allows web identity tokens to work automatically.
"""
# Isolate from ambient AWS_REGION/AWS_DEFAULT_REGION so no_region_or_endpoint is deterministic
env_without_aws_region = {
k: v
for k, v in os.environ.items()
if k not in ("AWS_REGION", "AWS_DEFAULT_REGION")
}
base_aws_llm = BaseAWSLLM()
mock_expiry = MagicMock()
mock_expiry.tzinfo = timezone.utc
time_diff = MagicMock()
time_diff.total_seconds.return_value = 3600
mock_expiry.__sub__ = MagicMock(return_value=time_diff)
mock_sts_response = {
"Credentials": {
"AccessKeyId": "assumed-access-key",
"SecretAccessKey": "assumed-secret-key",
"SessionToken": "assumed-session-token",
"Expiration": mock_expiry,
}
}
mock_sts_client = MagicMock()
mock_sts_client.assume_role.return_value = mock_sts_response
with patch.dict(os.environ, env_without_aws_region, clear=True):
with patch("boto3.client", return_value=mock_sts_client) as mock_boto3_client:
credentials, ttl = base_aws_llm._auth_with_aws_role(
aws_access_key_id=None,
aws_secret_access_key=None,
aws_session_token=None,
aws_role_name="arn:aws:iam::2222222222222:role/LitellmEvalBedrockRole",
aws_session_name="test-session",
**role_kwargs,
)
mock_boto3_client.assert_called_once_with("sts", **expected_client_kwargs)
mock_sts_client.assume_role.assert_called_once_with(
RoleArn="arn:aws:iam::2222222222222:role/LitellmEvalBedrockRole",
RoleSessionName="test-session",
)
assert credentials.access_key == "assumed-access-key"
assert ttl is not None
@pytest.mark.parametrize(
"endpoint,expected_region",
[
("https://sts.eu-west-1.amazonaws.com", "eu-west-1"),
("https://sts.us-east-1.amazonaws.com", "us-east-1"),
("https://sts-fips.us-east-1.amazonaws.com", "us-east-1"),
("https://sts-fips.us-gov-west-1.amazonaws.com", "us-gov-west-1"),
("https://sts.us-gov-west-1.amazonaws.com", "us-gov-west-1"),
("https://sts.cn-north-1.amazonaws.com.cn", "cn-north-1"),
(
"https://vpce-abc123.sts.eu-west-1.vpce.amazonaws.com",
"eu-west-1",
),
("https://sts.amazonaws.com", None),
("https://invalid.example.com", None),
],
)
def test_parse_sts_region_from_endpoint(endpoint, expected_region):
assert BaseAWSLLM._parse_sts_region_from_endpoint(endpoint) == expected_region
@pytest.mark.parametrize(
"env,aws_sts_endpoint,expected_region",
[
({}, None, None),
({"AWS_REGION": "us-east-1"}, None, "us-east-1"),
({"AWS_DEFAULT_REGION": "ap-southeast-1"}, None, "ap-southeast-1"),
({}, "https://sts.eu-west-1.amazonaws.com", "eu-west-1"),
(
{"AWS_REGION": "us-east-1"},
"https://sts.eu-west-1.amazonaws.com",
"eu-west-1",
),
({}, "https://sts.amazonaws.com", None),
(
{},
"https://vpce-abc.sts.eu-central-1.vpce.amazonaws.com",
"eu-central-1",
),
],
ids=[
"no_env_no_endpoint",
"env_region",
"env_default_region",
"parsed_from_endpoint",
"parsed_endpoint_over_env",
"global_endpoint",
"vpce_endpoint",
],
)
def test_resolve_sts_region(env, aws_sts_endpoint, expected_region):
with patch.dict(os.environ, env, clear=True):
assert (
BaseAWSLLM._resolve_sts_region(aws_sts_endpoint=aws_sts_endpoint)
== expected_region
)
@pytest.mark.parametrize(
"env,aws_sts_endpoint,ssl_verify,expected",
[
({}, None, None, {"verify": True}),
(
{"AWS_REGION": "us-east-1"},
None,
None,
{"verify": True, "region_name": "us-east-1"},
),
(
{},
"https://sts.eu-west-1.amazonaws.com",
None,
{
"verify": True,
"endpoint_url": "https://sts.eu-west-1.amazonaws.com",
"region_name": "eu-west-1",
},
),
(
{"AWS_REGION": "us-east-1"},
"https://sts.eu-west-1.amazonaws.com",
None,
{
"verify": True,
"endpoint_url": "https://sts.eu-west-1.amazonaws.com",
"region_name": "eu-west-1",
},
),
(
{},
"https://sts.amazonaws.com",
None,
{"verify": True, "endpoint_url": "https://sts.amazonaws.com"},
),
(
{},
"https://vpce-abc.sts.eu-central-1.vpce.amazonaws.com",
None,
{
"verify": True,
"endpoint_url": "https://vpce-abc.sts.eu-central-1.vpce.amazonaws.com",
"region_name": "eu-central-1",
},
),
({}, None, False, {"verify": False}),
(
{"AWS_DEFAULT_REGION": "ap-southeast-1"},
None,
None,
{"verify": True, "region_name": "ap-southeast-1"},
),
],
ids=[
"default_verify_only",
"env_region",
"endpoint_with_parsed_region",
"endpoint_parsed_over_env",
"global_endpoint_no_region",
"vpce_endpoint",
"ssl_verify_false",
"env_default_region",
],
)
def test_build_sts_client_kwargs(env, aws_sts_endpoint, ssl_verify, expected):
base_aws_llm = BaseAWSLLM()
with patch.dict(os.environ, env, clear=True):
assert (
base_aws_llm._build_sts_client_kwargs(
aws_sts_endpoint=aws_sts_endpoint,
ssl_verify=ssl_verify,
)
== expected
)
@pytest.mark.parametrize(
"env,aws_sts_endpoint,aws_region_name,expected_region",
[
({}, None, "cn-north-1", "cn-north-1"),
({"AWS_REGION": "eu-west-1"}, None, "cn-north-1", "eu-west-1"),
({"AWS_DEFAULT_REGION": "ap-southeast-1"}, None, "cn-north-1", "ap-southeast-1"),
({}, "https://sts.cn-north-1.amazonaws.com.cn", "us-east-1", "cn-north-1"),
({}, None, None, None),
],
ids=[
"configured_region_fallback",
"env_region_beats_configured",
"env_default_region_beats_configured",
"cn_endpoint_beats_configured",
"nothing_configured",
],
)
def test_resolve_sts_region_configured_region_fallback(
env: dict[str, str],
aws_sts_endpoint: str | None,
aws_region_name: str | None,
expected_region: str | None,
) -> None:
with patch.dict(os.environ, env, clear=True):
assert (
BaseAWSLLM._resolve_sts_region(
aws_sts_endpoint=aws_sts_endpoint,
aws_region_name=aws_region_name,
)
== expected_region
)
def test_build_sts_client_kwargs_configured_region_fallback() -> None:
base_aws_llm = BaseAWSLLM()
with patch.dict(os.environ, {}, clear=True):
assert base_aws_llm._build_sts_client_kwargs(aws_region_name="cn-north-1") == {
"verify": True,
"region_name": "cn-north-1",
}
with patch.dict(os.environ, {"AWS_REGION": "eu-west-1"}, clear=True):
assert base_aws_llm._build_sts_client_kwargs(aws_region_name="cn-north-1") == {
"verify": True,
"region_name": "eu-west-1",
}
def test_assume_role_sts_client_uses_configured_cn_region() -> None:
"""arn:aws-cn roles must resolve against a cn STS endpoint, not the commercial default."""
base_aws_llm = BaseAWSLLM()
mock_expiry = MagicMock()
mock_expiry.tzinfo = timezone.utc
time_diff = MagicMock()
time_diff.total_seconds.return_value = 3600
mock_expiry.__sub__ = MagicMock(return_value=time_diff)
mock_sts_client = MagicMock()
mock_sts_client.assume_role.return_value = {
"Credentials": {
"AccessKeyId": "assumed-access-key",
"SecretAccessKey": "assumed-secret-key",
"SessionToken": "assumed-session-token",
"Expiration": mock_expiry,
}
}
with patch.dict(os.environ, {}, clear=True):
with patch("boto3.client", return_value=mock_sts_client) as mock_boto3_client:
credentials, ttl = base_aws_llm._auth_with_aws_role(
aws_access_key_id=None,
aws_secret_access_key=None,
aws_session_token=None,
aws_role_name="arn:aws-cn:iam::2222222222222:role/LitellmBedrockRole",
aws_session_name="test-session",
aws_region_name="cn-north-1",
)
mock_boto3_client.assert_called_with(
"sts",
region_name="cn-north-1",
verify=True,
)
assert credentials.access_key == "assumed-access-key"
assert credentials.secret_key == "assumed-secret-key"
assert credentials.token == "assumed-session-token"
assert ttl is not None
@pytest.mark.parametrize(
"model,expected_region",
[
(
"arn:aws-cn:bedrock:cn-north-1:123456789012:application-inference-profile/p",
"cn-north-1",
),
(
"arn:aws-us-gov:bedrock:us-gov-west-1:123456789012:foundation-model/m",
"us-gov-west-1",
),
(
"bedrock/arn:aws-cn:bedrock:cn-northwest-1:123456789012:inference-profile/p",
"cn-northwest-1",
),
("anthropic.claude-3", None),
],
)
def test_get_aws_region_from_model_arn_partition_arns(model: str, expected_region: str | None) -> None:
assert BaseAWSLLM()._get_aws_region_from_model_arn(model) == expected_region
@pytest.mark.parametrize(
"endpoint_type,region,expected",
[
("runtime", "cn-north-1", "https://bedrock-runtime.cn-north-1.amazonaws.com.cn"),
("agent", "cn-north-1", "https://bedrock-agent-runtime.cn-north-1.amazonaws.com.cn"),
("agentcore", "cn-north-1", "https://bedrock-agentcore.cn-north-1.amazonaws.com.cn"),
("runtime", "us-east-1", "https://bedrock-runtime.us-east-1.amazonaws.com"),
("agent", "us-east-1", "https://bedrock-agent-runtime.us-east-1.amazonaws.com"),
("agentcore", "us-east-1", "https://bedrock-agentcore.us-east-1.amazonaws.com"),
("runtime", "us-gov-west-1", "https://bedrock-runtime.us-gov-west-1.amazonaws.com"),
],
)
def test_select_default_endpoint_url_partitions(endpoint_type: str, region: str, expected: str) -> None:
assert (
BaseAWSLLM()._select_default_endpoint_url(
endpoint_type=endpoint_type, aws_region_name=region
)
== expected
)
def test_irsa_cross_account_sts_client_uses_resolved_region():
"""IRSA cross-account path must use _build_sts_client_kwargs (env region, not Bedrock)."""
base_aws_llm = BaseAWSLLM()
import tempfile
with tempfile.NamedTemporaryFile(mode="w", delete=False) as f:
f.write("test-web-identity-token")
token_file = f.name
try:
with patch.dict(
os.environ,
{
"AWS_WEB_IDENTITY_TOKEN_FILE": token_file,
"AWS_ROLE_ARN": "arn:aws:iam::111111111111:role/eks-service-account-role",
"AWS_REGION": "eu-west-1",
},
clear=True,
):
mock_sts_client = MagicMock()
mock_sts_client.assume_role_with_web_identity.return_value = {
"Credentials": {
"AccessKeyId": "temp-key",
"SecretAccessKey": "temp-secret",
"SessionToken": "temp-token",
"Expiration": datetime.now(timezone.utc) + timedelta(hours=1),
}
}
mock_sts_client.assume_role.return_value = {
"Credentials": {
"AccessKeyId": "assumed-key",
"SecretAccessKey": "assumed-secret",
"SessionToken": "assumed-token",
"Expiration": datetime.now(timezone.utc) + timedelta(hours=1),
}
}
with patch(
"boto3.client", return_value=mock_sts_client
) as mock_boto3_client:
base_aws_llm._auth_with_aws_role(
aws_access_key_id=None,
aws_secret_access_key=None,
aws_session_token=None,
aws_role_name="arn:aws:iam::222222222222:role/target-role",
aws_session_name="test-session",
aws_region_name="eu-central-1",
)
for call in mock_boto3_client.call_args_list:
assert call.args == ("sts",)
assert call.kwargs["region_name"] == "eu-west-1"
assert call.kwargs["verify"] is True
finally:
os.unlink(token_file)
def test_web_identity_token_sts_client_uses_build_sts_client_kwargs():
base_aws_llm = BaseAWSLLM()
mock_sts_client = MagicMock()
mock_sts_client.assume_role_with_web_identity.return_value = {
"Credentials": {
"AccessKeyId": "key",
"SecretAccessKey": "secret",
"SessionToken": "token",
"Expiration": datetime.now(timezone.utc) + timedelta(hours=1),
},
"PackedPolicySize": 0,
}
with patch.dict(os.environ, {"AWS_REGION": "eu-west-1"}, clear=True):
with patch("boto3.client", return_value=mock_sts_client) as mock_boto3_client:
with patch(
"litellm.llms.bedrock.base_aws_llm.get_secret",
return_value="oidc-token",
):
base_aws_llm._auth_with_web_identity_token(
aws_web_identity_token="my-token",
aws_role_name="arn:aws:iam::111111111111:role/target",
aws_session_name="test-session",
aws_region_name="eu-central-1",
aws_sts_endpoint="https://sts.eu-west-1.amazonaws.com",
)
mock_boto3_client.assert_called_once_with(
"sts",
verify=True,
endpoint_url="https://sts.eu-west-1.amazonaws.com",
region_name="eu-west-1",
)
def test_sts_uses_workload_region_not_bedrock_region():
"""Air-gapped: Bedrock in eu-central-1, STS VPC endpoint in eu-west-1 via AWS_REGION."""
base_aws_llm = BaseAWSLLM()
mock_expiry = MagicMock()
mock_expiry.tzinfo = timezone.utc
time_diff = MagicMock()
time_diff.total_seconds.return_value = 3600
mock_expiry.__sub__ = MagicMock(return_value=time_diff)
mock_sts_client = MagicMock()
mock_sts_client.assume_role.return_value = {
"Credentials": {
"AccessKeyId": "assumed-access-key",
"SecretAccessKey": "assumed-secret-key",
"SessionToken": "assumed-session-token",
"Expiration": mock_expiry,
}
}
with patch.dict(os.environ, {"AWS_REGION": "eu-west-1"}, clear=True):
with patch("boto3.client", return_value=mock_sts_client) as mock_boto3_client:
base_aws_llm._auth_with_aws_role(
aws_access_key_id=None,
aws_secret_access_key=None,
aws_session_token=None,
aws_role_name="arn:aws:iam::2222222222222:role/LitellmEvalBedrockRole",
aws_session_name="test-session",
aws_region_name="eu-central-1",
)
mock_boto3_client.assert_called_with(
"sts",
region_name="eu-west-1",
verify=True,
)
def test_sts_endpoint_region_matches_bedrock_region_param():
"""aws_sts_endpoint signing region must not follow aws_region_name when they differ."""
base_aws_llm = BaseAWSLLM()
mock_expiry = MagicMock()
mock_expiry.tzinfo = timezone.utc
time_diff = MagicMock()
time_diff.total_seconds.return_value = 3600
mock_expiry.__sub__ = MagicMock(return_value=time_diff)
mock_sts_client = MagicMock()
mock_sts_client.assume_role.return_value = {
"Credentials": {
"AccessKeyId": "assumed-access-key",
"SecretAccessKey": "assumed-secret-key",
"SessionToken": "assumed-session-token",
"Expiration": mock_expiry,
}
}
env_without_irsa = {
k: v
for k, v in os.environ.items()
if k
not in (
"AWS_ROLE_ARN",
"AWS_WEB_IDENTITY_TOKEN_FILE",
"AWS_REGION",
"AWS_DEFAULT_REGION",
)
}
with patch.dict(env_without_irsa, clear=True):
with patch("boto3.client", return_value=mock_sts_client) as mock_boto3_client:
base_aws_llm._auth_with_aws_role(
aws_access_key_id=None,
aws_secret_access_key=None,
aws_session_token=None,
aws_role_name="arn:aws:iam::2222222222222:role/LitellmEvalBedrockRole",
aws_session_name="test-session",
aws_region_name="eu-central-1",
aws_sts_endpoint="https://sts.eu-west-1.amazonaws.com",
)
mock_boto3_client.assert_called_with(
"sts",
endpoint_url="https://sts.eu-west-1.amazonaws.com",
region_name="eu-west-1",
verify=True,
)
@pytest.mark.parametrize(
"role_kwargs,expected_client_kwargs",
[
(
{},
{
"aws_access_key_id": "explicit-access-key",
"aws_secret_access_key": "explicit-secret-key",
"aws_session_token": "assumed-session-token",
"verify": True,
},
),
(
{"aws_region_name": "us-east-1"},
{
"aws_access_key_id": "explicit-access-key",
"aws_secret_access_key": "explicit-secret-key",
"aws_session_token": "assumed-session-token",
"verify": True,
"region_name": "us-east-1",
},
),
(
{"aws_sts_endpoint": "https://sts.eu-west-1.amazonaws.com"},
{
"endpoint_url": "https://sts.eu-west-1.amazonaws.com",
"region_name": "eu-west-1",
"aws_access_key_id": "explicit-access-key",
"aws_secret_access_key": "explicit-secret-key",
"aws_session_token": "assumed-session-token",
"verify": True,
},
),
],
ids=["no_region_or_endpoint", "configured_region_is_sts_fallback", "explicit_sts_endpoint"],
)
def test_explicit_credentials_used_when_provided(role_kwargs, expected_client_kwargs):
"""
Test that explicit credentials are used when provided (non-EKS/IRSA scenario).
"""
# Isolate from ambient AWS_REGION/AWS_DEFAULT_REGION so no_region_or_endpoint is deterministic
env_without_aws_region = {
k: v
for k, v in os.environ.items()
if k not in ("AWS_REGION", "AWS_DEFAULT_REGION")
}
base_aws_llm = BaseAWSLLM()
mock_expiry = MagicMock()
mock_expiry.tzinfo = timezone.utc
# Create a timedelta object that returns 3600 when total_seconds() is called
time_diff = MagicMock()
time_diff.total_seconds.return_value = 3600
mock_expiry.__sub__ = MagicMock(return_value=time_diff)
mock_sts_response = {
"Credentials": {
"AccessKeyId": "assumed-access-key",
"SecretAccessKey": "assumed-secret-key",
"SessionToken": "assumed-session-token",
"Expiration": mock_expiry,
}
}
mock_sts_client = MagicMock()
mock_sts_client.assume_role.return_value = mock_sts_response
with patch.dict(os.environ, env_without_aws_region, clear=True):
with patch("boto3.client", return_value=mock_sts_client) as mock_boto3_client:
credentials, ttl = base_aws_llm._auth_with_aws_role(
aws_access_key_id="explicit-access-key",
aws_secret_access_key="explicit-secret-key",
aws_session_token="assumed-session-token",
aws_role_name="arn:aws:iam::2222222222222:role/LitellmEvalBedrockRole",
aws_session_name="test-session",
**role_kwargs,
)
mock_boto3_client.assert_called_once_with("sts", **expected_client_kwargs)
mock_sts_client.assume_role.assert_called_once_with(
RoleArn="arn:aws:iam::2222222222222:role/LitellmEvalBedrockRole",
RoleSessionName="test-session",
)
assert credentials.access_key == "assumed-access-key"
assert credentials.secret_key == "assumed-secret-key"
assert credentials.token == "assumed-session-token"
assert ttl is not None
def test_partial_credentials_still_use_ambient():
"""
Test that if only one credential is provided, we still use ambient credentials.
This handles edge cases where configuration might be incomplete.
"""
env_without_aws_region = {
k: v
for k, v in os.environ.items()
if k not in ("AWS_REGION", "AWS_DEFAULT_REGION")
}
base_aws_llm = BaseAWSLLM()
# Mock the boto3 STS client
mock_sts_client = MagicMock()
# Mock the STS response
mock_expiry = MagicMock()
mock_expiry.tzinfo = timezone.utc
time_diff = MagicMock()
time_diff.total_seconds.return_value = 3600
mock_expiry.__sub__ = MagicMock(return_value=time_diff)
mock_sts_response = {
"Credentials": {
"AccessKeyId": "assumed-access-key",
"SecretAccessKey": "assumed-secret-key",
"SessionToken": "assumed-session-token",
"Expiration": mock_expiry,
}
}
mock_sts_client.assume_role.return_value = mock_sts_response
with patch.dict(os.environ, env_without_aws_region, clear=True):
with patch("boto3.client", return_value=mock_sts_client) as mock_boto3_client:
# Call with only access key (missing secret key)
credentials, ttl = base_aws_llm._auth_with_aws_role(
aws_access_key_id="AKIAEXAMPLE",
aws_secret_access_key=None,
aws_session_token=None,
aws_role_name="arn:aws:iam::2222222222222:role/LitellmEvalBedrockRole",
aws_session_name="test-session",
)
# Should still pass partial credentials to boto3.client
mock_boto3_client.assert_called_once_with(
"sts",
aws_access_key_id="AKIAEXAMPLE",
aws_secret_access_key=None,
aws_session_token=None,
verify=True,
)
# Should still call assume_role
mock_sts_client.assume_role.assert_called_once_with(
RoleArn="arn:aws:iam::2222222222222:role/LitellmEvalBedrockRole",
RoleSessionName="test-session",
)
def test_cross_account_role_assumption():
"""
Test assuming a role in a different AWS account (common in multi-account setups).
"""
env_without_aws_region = {
k: v
for k, v in os.environ.items()
if k not in ("AWS_REGION", "AWS_DEFAULT_REGION")
}
base_aws_llm = BaseAWSLLM()
# Mock the boto3 STS client
mock_sts_client = MagicMock()
# Mock the STS response for cross-account role
mock_expiry = MagicMock()
mock_expiry.tzinfo = timezone.utc
time_diff = MagicMock()
time_diff.total_seconds.return_value = 3600
mock_expiry.__sub__ = MagicMock(return_value=time_diff)
mock_sts_response = {
"Credentials": {
"AccessKeyId": "cross-account-access-key",
"SecretAccessKey": "cross-account-secret-key",
"SessionToken": "cross-account-session-token",
"Expiration": mock_expiry,
}
}
mock_sts_client.assume_role.return_value = mock_sts_response
with patch.dict(os.environ, env_without_aws_region, clear=True):
with patch("boto3.client", return_value=mock_sts_client) as mock_boto3_client:
# Assume role in different account (EKS/IRSA scenario)
credentials, ttl = base_aws_llm._auth_with_aws_role(
aws_access_key_id=None,
aws_secret_access_key=None,
aws_session_token=None,
aws_role_name="arn:aws:iam::999999999999:role/CrossAccountRole",
aws_session_name="cross-account-session",
)
# Should use ambient credentials
mock_boto3_client.assert_called_once_with("sts", verify=True)
# Should call assume_role with cross-account role
mock_sts_client.assume_role.assert_called_once_with(
RoleArn="arn:aws:iam::999999999999:role/CrossAccountRole",
RoleSessionName="cross-account-session",
)
# Verify cross-account credentials are returned
assert credentials.access_key == "cross-account-access-key"
assert credentials.secret_key == "cross-account-secret-key"
assert credentials.token == "cross-account-session-token"
assert ttl is not None
def test_role_assumption_with_custom_session_name():
"""
Test role assumption with a custom session name.
"""
base_aws_llm = BaseAWSLLM()
# Mock the boto3 STS client
mock_sts_client = MagicMock()
# Mock the STS response
mock_expiry = MagicMock()
mock_expiry.tzinfo = timezone.utc
time_diff = MagicMock()
time_diff.total_seconds.return_value = 3600
mock_expiry.__sub__ = MagicMock(return_value=time_diff)
mock_sts_response = {
"Credentials": {
"AccessKeyId": "custom-session-access-key",
"SecretAccessKey": "custom-session-secret-key",
"SessionToken": "custom-session-token",
"Expiration": mock_expiry,
}
}
mock_sts_client.assume_role.return_value = mock_sts_response
with patch("boto3.client", return_value=mock_sts_client):
# Use custom session name
credentials, ttl = base_aws_llm._auth_with_aws_role(
aws_access_key_id=None,
aws_secret_access_key=None,
aws_session_token=None,
aws_role_name="arn:aws:iam::1111111111111:role/LitellmRole",
aws_session_name="evals-bedrock-session",
)
# Should call assume_role with custom session name
mock_sts_client.assume_role.assert_called_once_with(
RoleArn="arn:aws:iam::1111111111111:role/LitellmRole",
RoleSessionName="evals-bedrock-session",
)
# Verify credentials are returned
assert credentials.access_key == "custom-session-access-key"
assert credentials.secret_key == "custom-session-secret-key"
assert credentials.token == "custom-session-token"
def test_role_assumption_ttl_calculation():
"""
Test that TTL is calculated correctly from STS response expiration.
"""
base_aws_llm = BaseAWSLLM()
# Mock the boto3 STS client
mock_sts_client = MagicMock()
# Create a real datetime for expiration (1 hour from now)
expiration_time = datetime.now(timezone.utc) + timedelta(hours=1)
mock_sts_response = {
"Credentials": {
"AccessKeyId": "ttl-test-access-key",
"SecretAccessKey": "ttl-test-secret-key",
"SessionToken": "ttl-test-session-token",
"Expiration": expiration_time,
}
}
mock_sts_client.assume_role.return_value = mock_sts_response
with patch("boto3.client", return_value=mock_sts_client):
credentials, ttl = base_aws_llm._auth_with_aws_role(
aws_access_key_id=None,
aws_secret_access_key=None,
aws_session_token=None,
aws_role_name="arn:aws:iam::1111111111111:role/LitellmRole",
aws_session_name="ttl-test-session",
)
# TTL should be approximately 3540 seconds (1 hour - 60 second buffer)
assert ttl is not None
assert 3500 <= ttl <= 3600 # Allow some variance for test execution time
def test_role_assumption_access_denied_falls_back_when_same_role():
"""
Test that when AssumeRole fails with AccessDenied AND the caller is confirmed
to already be running as the target role, we fall back to ambient credentials.
"""
base_aws_llm = BaseAWSLLM()
# Mock the boto3 STS client to raise AccessDenied
mock_sts_client = MagicMock()
mock_sts_client.assume_role.side_effect = Exception(
"An error occurred (AccessDenied) when calling the AssumeRole operation: "
"Roles may not be assumed by root accounts."
)
# Mock _auth_with_env_vars to return fallback credentials
mock_creds = MagicMock()
mock_creds.access_key = "fallback-access-key"
mock_creds.secret_key = "fallback-secret-key"
with patch("boto3.client", return_value=mock_sts_client):
with patch.object(
base_aws_llm, "_auth_with_env_vars", return_value=(mock_creds, None)
) as mock_env_auth:
# _is_already_running_as_role returns True => fallback allowed
with patch.object(
base_aws_llm, "_is_already_running_as_role", return_value=True
):
credentials, ttl = base_aws_llm._auth_with_aws_role(
aws_access_key_id=None,
aws_secret_access_key=None,
aws_session_token=None,
aws_role_name="arn:aws:iam::1111111111111:role/UnauthorizedRole",
aws_session_name="error-test-session",
)
# Should have fallen back to env vars
mock_env_auth.assert_called_once()
assert credentials.access_key == "fallback-access-key"
def test_role_assumption_access_denied_raises_when_different_role():
"""
Test that when AssumeRole fails with AccessDenied but the caller is NOT
the same role, the error is re-raised (genuine permission failure).
"""
base_aws_llm = BaseAWSLLM()
mock_sts_client = MagicMock()
mock_sts_client.assume_role.side_effect = Exception(
"An error occurred (AccessDenied) when calling the AssumeRole operation: "
"User is not authorized to perform sts:AssumeRole"
)
with patch("boto3.client", return_value=mock_sts_client):
# _is_already_running_as_role returns False => do NOT fallback
with patch.object(
base_aws_llm, "_is_already_running_as_role", return_value=False
):
with pytest.raises(Exception, match='An error occurred \\(AccessDenied\\) when calling the') as exc_info:
base_aws_llm._auth_with_aws_role(
aws_access_key_id=None,
aws_secret_access_key=None,
aws_session_token=None,
aws_role_name="arn:aws:iam::999999999999:role/CrossAccountRole",
aws_session_name="error-test-session",
)
assert "AccessDenied" in str(exc_info.value)
def test_role_assumption_non_access_denied_error_propagated():
"""
Test that non-AccessDenied errors from AssumeRole are still propagated.
"""
base_aws_llm = BaseAWSLLM()
# Mock the boto3 STS client to raise a non-AccessDenied exception
mock_sts_client = MagicMock()
mock_sts_client.assume_role.side_effect = Exception(
"An error occurred (MalformedPolicyDocument) when calling the AssumeRole operation"
)
with patch("boto3.client", return_value=mock_sts_client):
with pytest.raises(Exception, match='An error occurred \\(MalformedPolicyDocument\\) when calling') as exc_info:
base_aws_llm._auth_with_aws_role(
aws_access_key_id=None,
aws_secret_access_key=None,
aws_session_token=None,
aws_role_name="arn:aws:iam::1111111111111:role/BadPolicyRole",
aws_session_name="error-test-session",
)
assert "MalformedPolicyDocument" in str(exc_info.value)
def test_multiple_role_assumptions_in_sequence():
"""
Test that multiple role assumptions work correctly in sequence.
This simulates the scenario where different models use different roles.
"""
base_aws_llm = BaseAWSLLM()
# Mock the boto3 STS client
mock_sts_client = MagicMock()
# Mock different responses for different roles
mock_expiry = MagicMock()
mock_expiry.tzinfo = timezone.utc
time_diff = MagicMock()
time_diff.total_seconds.return_value = 3600
mock_expiry.__sub__ = MagicMock(return_value=time_diff)
# First role response
mock_sts_response1 = {
"Credentials": {
"AccessKeyId": "role1-access-key",
"SecretAccessKey": "role1-secret-key",
"SessionToken": "role1-session-token",
"Expiration": mock_expiry,
}
}
# Second role response
mock_sts_response2 = {
"Credentials": {
"AccessKeyId": "role2-access-key",
"SecretAccessKey": "role2-secret-key",
"SessionToken": "role2-session-token",
"Expiration": mock_expiry,
}
}
# Configure mock to return different responses
mock_sts_client.assume_role.side_effect = [mock_sts_response1, mock_sts_response2]
with patch("boto3.client", return_value=mock_sts_client):
# First role assumption
credentials1, ttl1 = base_aws_llm._auth_with_aws_role(
aws_access_key_id=None,
aws_secret_access_key=None,
aws_session_token=None,
aws_role_name="arn:aws:iam::1111111111111:role/LitellmRole",
aws_session_name="session-1",
)
# Second role assumption
credentials2, ttl2 = base_aws_llm._auth_with_aws_role(
aws_access_key_id=None,
aws_secret_access_key=None,
aws_session_token=None,
aws_role_name="arn:aws:iam::2222222222222:role/LitellmEvalBedrockRole",
aws_session_name="session-2",
)
# Verify both role assumptions were made
assert mock_sts_client.assume_role.call_count == 2
# Verify first role credentials
assert credentials1.access_key == "role1-access-key"
assert credentials1.secret_key == "role1-secret-key"
assert credentials1.token == "role1-session-token"
# Verify second role credentials
assert credentials2.access_key == "role2-access-key"
assert credentials2.secret_key == "role2-secret-key"
assert credentials2.token == "role2-session-token"
def test_auth_with_aws_role_irsa_environment():
"""Test that _auth_with_aws_role detects and uses IRSA environment variables"""
base_llm = BaseAWSLLM()
# Create a temporary file to simulate the web identity token
import tempfile
with tempfile.NamedTemporaryFile(mode="w", delete=False) as f:
f.write("test-web-identity-token")
token_file = f.name
try:
# Set IRSA environment variables
with patch.dict(
os.environ,
{
"AWS_WEB_IDENTITY_TOKEN_FILE": token_file,
"AWS_ROLE_ARN": "arn:aws:iam::111111111111:role/eks-service-account-role",
"AWS_REGION": "us-east-1",
},
):
# Mock the boto3 STS client
mock_sts_client = MagicMock()
mock_assume_web_identity_response = {
"Credentials": {
"AccessKeyId": "irsa-temp-access-key",
"SecretAccessKey": "irsa-temp-secret-key",
"SessionToken": "irsa-temp-session-token",
"Expiration": datetime.now() + timedelta(hours=1),
}
}
mock_assume_role_response = {
"Credentials": {
"AccessKeyId": "irsa-access-key",
"SecretAccessKey": "irsa-secret-key",
"SessionToken": "irsa-session-token",
"Expiration": datetime.now() + timedelta(hours=1),
}
}
mock_sts_client.assume_role_with_web_identity.return_value = (
mock_assume_web_identity_response
)
mock_sts_client.assume_role.return_value = mock_assume_role_response
with patch(
"boto3.client", return_value=mock_sts_client
) as mock_boto3_client:
# Call _auth_with_aws_role without explicit credentials
creds, ttl = base_llm._auth_with_aws_role(
aws_access_key_id=None,
aws_secret_access_key=None,
aws_session_token=None,
aws_role_name="arn:aws:iam::222222222222:role/target-role",
aws_session_name="test-session",
)
# Verify boto3.client was called multiple times
# First for manual IRSA, then with IRSA credentials
assert mock_boto3_client.call_count >= 2
# Verify assume_role_with_web_identity was called
mock_sts_client.assume_role_with_web_identity.assert_called_once_with(
RoleArn="arn:aws:iam::111111111111:role/eks-service-account-role",
RoleSessionName="test-session",
WebIdentityToken="test-web-identity-token",
)
# Verify assume_role was called with correct parameters
mock_sts_client.assume_role.assert_called_once_with(
RoleArn="arn:aws:iam::222222222222:role/target-role",
RoleSessionName="test-session",
)
# Verify the returned credentials
assert creds.access_key == "irsa-access-key"
assert creds.secret_key == "irsa-secret-key"
assert creds.token == "irsa-session-token"
# 1h STS session minus the safety margin, so a cached entry retires before AWS
# expires the credential
assert ttl is not None
assert 3400 < ttl <= 3540
finally:
# Clean up the temporary file
os.unlink(token_file)
def test_auth_with_aws_role_same_role_irsa():
"""Test that when IRSA role matches the requested role, we skip assumption"""
base_llm = BaseAWSLLM()
# Set IRSA environment variables
with patch.dict(
os.environ,
{
"AWS_ROLE_ARN": "arn:aws:iam::111111111111:role/LitellmRole",
"AWS_WEB_IDENTITY_TOKEN_FILE": "/var/run/secrets/eks.amazonaws.com/serviceaccount/token",
},
):
# Mock the _auth_with_env_vars method
mock_creds = MagicMock()
mock_creds.access_key = "irsa-access-key"
mock_creds.secret_key = "irsa-secret-key"
mock_creds.token = "irsa-session-token"
with patch.object(
base_llm, "_auth_with_env_vars", return_value=(mock_creds, None)
) as mock_env_auth:
# Call get_credentials instead of _auth_with_aws_role directly
# This tests the full flow
creds = base_llm.get_credentials(
aws_access_key_id=None,
aws_secret_access_key=None,
aws_role_name="arn:aws:iam::111111111111:role/LitellmRole", # Same as AWS_ROLE_ARN
aws_session_name="test-session",
aws_region_name="us-east-1",
)
# Verify it used the env vars auth (no role assumption)
mock_env_auth.assert_called_once()
# Verify the returned credentials
assert creds.access_key == "irsa-access-key"
def test_assume_role_with_external_id():
"""Test that assume_role STS call includes ExternalId parameter when provided"""
base_aws_llm = BaseAWSLLM()
# Mock the boto3 STS client
mock_sts_client = MagicMock()
mock_expiry = datetime.now(timezone.utc) + timedelta(hours=1)
mock_sts_response = {
"Credentials": {
"AccessKeyId": "test-access-key",
"SecretAccessKey": "test-secret-key",
"SessionToken": "test-session-token",
"Expiration": mock_expiry,
}
}
mock_sts_client.assume_role.return_value = mock_sts_response
with patch("boto3.client", return_value=mock_sts_client):
# Call _auth_with_aws_role with external ID
credentials, ttl = base_aws_llm._auth_with_aws_role(
aws_access_key_id=None,
aws_secret_access_key=None,
aws_session_token=None,
aws_role_name="arn:aws:iam::123456789012:role/ExampleRole",
aws_session_name="test-session",
aws_external_id="UniqueExternalID123",
)
# Verify assume_role was called with ExternalId
mock_sts_client.assume_role.assert_called_once_with(
RoleArn="arn:aws:iam::123456789012:role/ExampleRole",
RoleSessionName="test-session",
ExternalId="UniqueExternalID123",
)
def test_assume_role_without_external_id():
"""Test that assume_role STS call excludes ExternalId parameter when not provided"""
base_aws_llm = BaseAWSLLM()
# Mock the boto3 STS client
mock_sts_client = MagicMock()
mock_expiry = datetime.now(timezone.utc) + timedelta(hours=1)
mock_sts_response = {
"Credentials": {
"AccessKeyId": "test-access-key",
"SecretAccessKey": "test-secret-key",
"SessionToken": "test-session-token",
"Expiration": mock_expiry,
}
}
mock_sts_client.assume_role.return_value = mock_sts_response
with patch("boto3.client", return_value=mock_sts_client):
# Call _auth_with_aws_role without external ID
credentials, ttl = base_aws_llm._auth_with_aws_role(
aws_access_key_id=None,
aws_secret_access_key=None,
aws_session_token=None,
aws_role_name="arn:aws:iam::123456789012:role/ExampleRole",
aws_session_name="test-session",
)
# Verify assume_role was called without ExternalId
mock_sts_client.assume_role.assert_called_once_with(
RoleArn="arn:aws:iam::123456789012:role/ExampleRole",
RoleSessionName="test-session",
)
_SESSION_TAGS = ({"Key": "team", "Value": "genai"}, {"Key": "env", "Value": "prod"})
_SORTED_SESSION_TAGS = ({"Key": "env", "Value": "prod"}, {"Key": "team", "Value": "genai"})
_TAGGED_ROLE_ARN = "arn:aws:iam::123456789012:role/TaggedRole"
class _TagAwareSTSClient:
"""STS stand-in for a trust policy that only admits sessions carrying exactly the expected tags."""
def __init__(self, expected_tags: tuple = (), access_key: str = "ASIATAGGEDSESSION") -> None:
self.expected_tags = expected_tags
self.access_key = access_key
self.assume_role_calls: list = []
def get_caller_identity(self):
return {"Arn": "arn:aws:iam::111111111111:user/litellm-proxy-pod"}
def assume_role_with_web_identity(self, **params):
return {
"Credentials": {
"AccessKeyId": "ASIAIRSATEMP",
"SecretAccessKey": "irsa-temp-secret-key",
"SessionToken": "irsa-temp-session-token",
"Expiration": datetime.now(timezone.utc) + timedelta(hours=1),
}
}
def assume_role(self, **params):
self.assume_role_calls.append(params)
if tuple(params.get("Tags", ())) != self.expected_tags:
raise ClientError(
{"Error": {"Code": "AccessDenied", "Message": "is not authorized to perform: sts:TagSession"}},
"AssumeRole",
)
return {
"Credentials": {
"AccessKeyId": self.access_key,
"SecretAccessKey": "assumed-secret-key",
"SessionToken": "assumed-session-token",
"Expiration": datetime.now(timezone.utc) + timedelta(hours=1),
}
}
def _irsa_env(tmp_path, irsa_role_arn: str) -> dict:
token_file = tmp_path / "web-identity-token"
token_file.write_text("test-web-identity-token")
return {
"AWS_WEB_IDENTITY_TOKEN_FILE": str(token_file),
"AWS_ROLE_ARN": irsa_role_arn,
"AWS_REGION": "us-east-1",
}
def test_assume_role_sends_session_tags():
"""The STS session carries the configured tags, so a trust policy gated on sts:TagSession admits it."""
sts = _TagAwareSTSClient(expected_tags=_SESSION_TAGS)
with patch("boto3.client", return_value=sts):
credentials, _ttl = BaseAWSLLM()._auth_with_aws_role(
aws_access_key_id=None,
aws_secret_access_key=None,
aws_session_token=None,
aws_role_name=_TAGGED_ROLE_ARN,
aws_session_name="test-session",
aws_session_tags=list(_SESSION_TAGS),
)
assert credentials.access_key == "ASIATAGGEDSESSION"
assert sts.assume_role_calls == [
{"RoleArn": _TAGGED_ROLE_ARN, "RoleSessionName": "test-session", "Tags": _SESSION_TAGS}
]
def test_assume_role_sends_session_tags_alongside_external_id():
sts = _TagAwareSTSClient(expected_tags=_SESSION_TAGS)
with patch("boto3.client", return_value=sts):
credentials, _ttl = BaseAWSLLM()._auth_with_aws_role(
aws_access_key_id=None,
aws_secret_access_key=None,
aws_session_token=None,
aws_role_name=_TAGGED_ROLE_ARN,
aws_session_name="test-session",
aws_external_id="UniqueExternalID123",
aws_session_tags=_SESSION_TAGS,
)
assert credentials.access_key == "ASIATAGGEDSESSION"
assert sts.assume_role_calls == [
{
"RoleArn": _TAGGED_ROLE_ARN,
"RoleSessionName": "test-session",
"ExternalId": "UniqueExternalID123",
"Tags": _SESSION_TAGS,
}
]
@pytest.mark.parametrize("aws_session_tags", [None, [], ()], ids=["none", "empty-list", "empty-tuple"])
def test_assume_role_omits_the_tags_key_without_session_tags(aws_session_tags):
"""Nothing configured means the AssumeRole request looks exactly as it did before tags existed."""
sts = _TagAwareSTSClient(expected_tags=())
with patch("boto3.client", return_value=sts):
credentials, _ttl = BaseAWSLLM()._auth_with_aws_role(
aws_access_key_id=None,
aws_secret_access_key=None,
aws_session_token=None,
aws_role_name=_TAGGED_ROLE_ARN,
aws_session_name="test-session",
aws_session_tags=aws_session_tags,
)
assert credentials.access_key == "ASIATAGGEDSESSION"
assert sts.assume_role_calls == [{"RoleArn": _TAGGED_ROLE_ARN, "RoleSessionName": "test-session"}]
def test_irsa_cross_account_assume_role_sends_session_tags(tmp_path):
irsa_role_arn = "arn:aws:iam::111111111111:role/eks-service-account-role"
sts = _TagAwareSTSClient(expected_tags=_SESSION_TAGS)
with patch.dict(os.environ, _irsa_env(tmp_path, irsa_role_arn)), patch("boto3.client", return_value=sts):
credentials, _ttl = BaseAWSLLM()._auth_with_aws_role(
aws_access_key_id=None,
aws_secret_access_key=None,
aws_session_token=None,
aws_role_name=_TAGGED_ROLE_ARN,
aws_session_name="test-session",
aws_session_tags=_SESSION_TAGS,
)
assert credentials.access_key == "ASIATAGGEDSESSION"
assert sts.assume_role_calls == [
{"RoleArn": _TAGGED_ROLE_ARN, "RoleSessionName": "test-session", "Tags": _SESSION_TAGS}
]
def test_irsa_same_account_assume_role_sends_session_tags(tmp_path):
sts = _TagAwareSTSClient(expected_tags=_SESSION_TAGS)
with patch.dict(os.environ, _irsa_env(tmp_path, _TAGGED_ROLE_ARN)), patch("boto3.client", return_value=sts):
credentials, _ttl = BaseAWSLLM()._auth_with_aws_role(
aws_access_key_id=None,
aws_secret_access_key=None,
aws_session_token=None,
aws_role_name=_TAGGED_ROLE_ARN,
aws_session_name="test-session",
aws_session_tags=_SESSION_TAGS,
)
assert credentials.access_key == "ASIATAGGEDSESSION"
assert sts.assume_role_calls == [
{"RoleArn": _TAGGED_ROLE_ARN, "RoleSessionName": "test-session", "Tags": _SESSION_TAGS}
]
def test_get_credentials_canonicalizes_session_tag_order_for_the_cache():
"""Two deployments listing the same tags in a different order share one STS session."""
base_aws_llm = BaseAWSLLM()
sts = _TagAwareSTSClient(expected_tags=_SORTED_SESSION_TAGS)
with patch.dict(os.environ, _os_environ_without_aws_keys(), clear=True), patch("boto3.client", return_value=sts):
first = base_aws_llm.get_credentials(
aws_role_name=_TAGGED_ROLE_ARN,
aws_session_name="team-session",
aws_session_tags=list(_SESSION_TAGS),
)
second = base_aws_llm.get_credentials(
aws_role_name=_TAGGED_ROLE_ARN,
aws_session_name="team-session",
aws_session_tags=list(reversed(_SESSION_TAGS)),
)
assert first.access_key == second.access_key == "ASIATAGGEDSESSION"
assert sts.assume_role_calls == [
{"RoleArn": _TAGGED_ROLE_ARN, "RoleSessionName": "team-session", "Tags": _SORTED_SESSION_TAGS}
]
def test_get_credentials_scopes_the_cache_per_session_tag_set():
"""Different tag sets are different principals to AWS, so each gets its own STS session."""
base_aws_llm = BaseAWSLLM()
mock_sts_client = _assume_role_sts_mock()
mock_sts_client.assume_role.side_effect = [
{
"Credentials": {
"AccessKeyId": f"assumed-access-key-{team}",
"SecretAccessKey": "assumed-secret-key",
"SessionToken": f"assumed-session-token-{team}",
"Expiration": datetime.now(timezone.utc) + timedelta(hours=1),
}
}
for team in ("genai", "platform")
]
with (
patch.dict(os.environ, _os_environ_without_aws_keys(), clear=True),
patch("boto3.client", return_value=mock_sts_client),
):
genai = base_aws_llm.get_credentials(
aws_role_name=_TAGGED_ROLE_ARN,
aws_session_name="team-session",
aws_session_tags=[{"Key": "team", "Value": "genai"}],
)
platform = base_aws_llm.get_credentials(
aws_role_name=_TAGGED_ROLE_ARN,
aws_session_name="team-session",
aws_session_tags=[{"Key": "team", "Value": "platform"}],
)
assert genai.access_key == "assumed-access-key-genai"
assert platform.access_key == "assumed-access-key-platform"
assert [call.kwargs["Tags"] for call in mock_sts_client.assume_role.call_args_list] == [
({"Key": "team", "Value": "genai"},),
({"Key": "team", "Value": "platform"},),
]
@pytest.mark.parametrize(
"aws_session_tags",
[
"team=genai",
{"team": "genai"},
[["team", "genai"]],
[{"key": "team", "value": "genai"}],
[{"Key": "team"}],
[{"Key": 1, "Value": "genai"}],
],
ids=["string", "flat-dict", "pair-list", "lowercase-keys", "missing-value", "non-string-key"],
)
def test_get_credentials_rejects_malformed_session_tags(aws_session_tags):
with pytest.raises(ValueError, match="Invalid 'aws_session_tags' value"):
BaseAWSLLM().get_credentials(
aws_role_name=_TAGGED_ROLE_ARN,
aws_session_name="team-session",
aws_session_tags=aws_session_tags,
)
def test_get_boto_credentials_from_optional_params_consumes_session_tags():
"""Tags feed the STS call and must not linger in optional_params to be serialized into the body."""
sts = _TagAwareSTSClient(expected_tags=_SORTED_SESSION_TAGS)
optional_params = {
"aws_region_name": "us-east-1",
"aws_role_name": _TAGGED_ROLE_ARN,
"aws_session_name": "team-session",
"aws_session_tags": list(_SESSION_TAGS),
}
with patch.dict(os.environ, _os_environ_without_aws_keys(), clear=True), patch("boto3.client", return_value=sts):
target = BaseAWSLLM()._get_boto_credentials_from_optional_params(optional_params)
assert target.credentials.access_key == "ASIATAGGEDSESSION"
assert "aws_session_tags" not in optional_params
def test_sign_request_signs_with_the_tagged_sts_session():
sts = _TagAwareSTSClient(expected_tags=_SORTED_SESSION_TAGS)
optional_params = {
"aws_region_name": "us-east-1",
"aws_role_name": _TAGGED_ROLE_ARN,
"aws_session_name": "team-session",
"aws_session_tags": list(_SESSION_TAGS),
}
with patch.dict(os.environ, _os_environ_without_aws_keys(), clear=True), patch("boto3.client", return_value=sts):
headers, _body = BaseAWSLLM()._sign_request(
service_name="bedrock",
headers={},
optional_params=optional_params,
request_data={"prompt": "hi"},
api_base="https://bedrock-runtime.us-east-1.amazonaws.com/model/anthropic.claude-opus-5/invoke",
)
assert "Credential=ASIATAGGEDSESSION/" in headers["Authorization"]
def test_converse_handler_external_id_extraction():
"""Test that BedrockConverseLLM properly extracts and passes aws_external_id parameter"""
from litellm.llms.bedrock.chat.converse_handler import BedrockConverseLLM
converse_llm = BedrockConverseLLM()
# Mock get_credentials to capture parameters
def mock_get_credentials(**kwargs):
mock_get_credentials.called_kwargs = kwargs
mock_credentials = MagicMock()
mock_credentials.access_key = "test-access-key"
mock_credentials.secret_key = "test-secret-key"
mock_credentials.token = "test-session-token"
return mock_credentials
with patch.object(
converse_llm, "get_credentials", side_effect=mock_get_credentials
):
with patch.object(
converse_llm, "_get_aws_region_name", return_value="us-west-2"
):
with patch.object(
converse_llm,
"get_runtime_endpoint",
return_value=("https://test", "https://test"),
):
with patch("litellm.AmazonConverseConfig") as mock_config:
mock_config.return_value._transform_request.return_value = {
"test": "data"
}
with patch.object(
converse_llm, "get_request_headers"
) as mock_headers:
mock_headers.return_value = MagicMock()
mock_headers.return_value.headers = {"Authorization": "test"}
with patch(
"litellm.llms.custom_httpx.http_handler._get_httpx_client"
) as mock_client:
mock_http_client = MagicMock()
mock_response = MagicMock()
mock_response.raise_for_status.return_value = None
mock_http_client.post.return_value = mock_response
mock_client.return_value = mock_http_client
# Mock the transform_response method
mock_config.return_value._transform_response.return_value = (
MagicMock()
)
# Call completion with aws_external_id in optional_params
optional_params = {
"aws_role_name": "arn:aws:iam::123456789012:role/ExampleRole",
"aws_session_name": "test-session",
"aws_external_id": "TestExternalID123",
}
try:
converse_llm.completion(
model="anthropic.claude-3-sonnet-20240229-v1:0",
messages=[{"role": "user", "content": "Hello"}],
api_base=None,
custom_prompt_dict={},
model_response=MagicMock(),
encoding="utf-8",
logging_obj=MagicMock(),
optional_params=optional_params,
acompletion=False,
timeout=None,
litellm_params={},
)
except Exception:
# We expect this to fail due to mocking, but that's OK
# We just want to verify the parameter extraction
pass
# Verify aws_external_id was extracted and passed to get_credentials
assert hasattr(mock_get_credentials, "called_kwargs")
assert (
"aws_external_id" in mock_get_credentials.called_kwargs
)
assert (
mock_get_credentials.called_kwargs["aws_external_id"]
== "TestExternalID123"
)
def test_is_already_running_as_role_irsa_same_role():
"""Test IRSA fast path: when AWS_ROLE_ARN matches target role."""
base_aws_llm = BaseAWSLLM()
with patch.dict(
os.environ,
{
"AWS_ROLE_ARN": "arn:aws:iam::123456789012:role/MyRole",
"AWS_WEB_IDENTITY_TOKEN_FILE": "/var/run/secrets/token",
},
):
assert (
base_aws_llm._is_already_running_as_role(
"arn:aws:iam::123456789012:role/MyRole"
)
is True
)
def test_is_already_running_as_role_irsa_different_role():
"""Test IRSA fast path: when AWS_ROLE_ARN does NOT match target role."""
base_aws_llm = BaseAWSLLM()
with patch.dict(
os.environ,
{
"AWS_ROLE_ARN": "arn:aws:iam::123456789012:role/MyRole",
"AWS_WEB_IDENTITY_TOKEN_FILE": "/var/run/secrets/token",
},
):
assert (
base_aws_llm._is_already_running_as_role(
"arn:aws:iam::999999999999:role/OtherRole"
)
is False
)
def test_is_already_running_as_role_ecs_task_role():
"""Test ECS/EC2 path: GetCallerIdentity shows assumed-role matching target."""
base_aws_llm = BaseAWSLLM()
mock_sts_client = MagicMock()
mock_sts_client.get_caller_identity.return_value = {
"Arn": "arn:aws:sts::123456789012:assumed-role/MyEcsTaskRole/ecs-task-id"
}
with patch.dict(os.environ, {}, clear=False):
# Ensure no IRSA env vars
env = {
k: v
for k, v in os.environ.items()
if k not in ("AWS_ROLE_ARN", "AWS_WEB_IDENTITY_TOKEN_FILE")
}
with patch.dict(os.environ, env, clear=True):
with patch("boto3.client", return_value=mock_sts_client):
assert (
base_aws_llm._is_already_running_as_role(
"arn:aws:iam::123456789012:role/MyEcsTaskRole"
)
is True
)
def test_is_already_running_as_role_ecs_different_role():
"""Test ECS/EC2 path: GetCallerIdentity shows a different role."""
base_aws_llm = BaseAWSLLM()
mock_sts_client = MagicMock()
mock_sts_client.get_caller_identity.return_value = {
"Arn": "arn:aws:sts::123456789012:assumed-role/MyEcsTaskRole/ecs-task-id"
}
with patch.dict(os.environ, {}, clear=False):
env = {
k: v
for k, v in os.environ.items()
if k not in ("AWS_ROLE_ARN", "AWS_WEB_IDENTITY_TOKEN_FILE")
}
with patch.dict(os.environ, env, clear=True):
with patch("boto3.client", return_value=mock_sts_client):
assert (
base_aws_llm._is_already_running_as_role(
"arn:aws:iam::999999999999:role/DifferentRole"
)
is False
)
def test_is_already_running_as_role_ecs_role_with_path():
"""Test ECS path with role that has a path prefix (e.g., /service-role/MyRole)."""
base_aws_llm = BaseAWSLLM()
mock_sts_client = MagicMock()
mock_sts_client.get_caller_identity.return_value = {
"Arn": "arn:aws:sts::123456789012:assumed-role/MyEcsTaskRole/ecs-task-id"
}
with patch.dict(os.environ, {}, clear=False):
env = {
k: v
for k, v in os.environ.items()
if k not in ("AWS_ROLE_ARN", "AWS_WEB_IDENTITY_TOKEN_FILE")
}
with patch.dict(os.environ, env, clear=True):
with patch("boto3.client", return_value=mock_sts_client):
# Role ARN with path
assert (
base_aws_llm._is_already_running_as_role(
"arn:aws:iam::123456789012:role/service-role/MyEcsTaskRole"
)
is True
)
def test_is_already_running_as_role_get_caller_identity_fails():
"""Test that when GetCallerIdentity fails, we return False (don't crash)."""
base_aws_llm = BaseAWSLLM()
mock_sts_client = MagicMock()
mock_sts_client.get_caller_identity.side_effect = Exception("No credentials found")
with patch.dict(os.environ, {}, clear=False):
env = {
k: v
for k, v in os.environ.items()
if k not in ("AWS_ROLE_ARN", "AWS_WEB_IDENTITY_TOKEN_FILE")
}
with patch.dict(os.environ, env, clear=True):
with patch("boto3.client", return_value=mock_sts_client):
assert (
base_aws_llm._is_already_running_as_role(
"arn:aws:iam::123456789012:role/SomeRole"
)
is False
)
def test_get_credentials_ecs_same_role_skips_assume_role():
"""
End-to-end test: when running on ECS with the same role as aws_role_name,
get_credentials should use ambient credentials and NOT call AssumeRole.
"""
base_aws_llm = BaseAWSLLM()
mock_creds = MagicMock()
mock_creds.access_key = "ecs-access-key"
mock_creds.secret_key = "ecs-secret-key"
mock_creds.token = "ecs-session-token"
with patch.object(
base_aws_llm,
"_is_already_running_as_role",
return_value=True,
) as mock_already_running:
with patch.object(
base_aws_llm,
"_auth_with_env_vars",
return_value=(mock_creds, None),
) as mock_env_auth:
with patch.object(
base_aws_llm,
"_auth_with_aws_role",
) as mock_role_auth:
credentials = base_aws_llm.get_credentials(
aws_role_name="arn:aws:iam::123456789012:role/MyEcsTaskRole",
aws_region_name="us-east-1",
)
base_aws_llm.get_credentials(
aws_role_name="arn:aws:iam::123456789012:role/MyEcsTaskRole",
aws_region_name="us-east-1",
)
# The identity probe and the env resolution are both behind the iam_cache, so the
# second call repeats neither.
mock_already_running.assert_called_once()
mock_env_auth.assert_called_once()
mock_role_auth.assert_not_called()
assert credentials.access_key == "ecs-access-key"
def test_get_credentials_role_reevaluates_identity_once_the_cached_entry_lapses():
"""
First request: already target role -> env path fills iam_cache.
Once that entry expires and the identity no longer matches, the next request must take the
AssumeRole path rather than serving the cached env resolution.
"""
base_aws_llm = BaseAWSLLM()
env_creds = MagicMock()
env_creds.access_key = "ambient-key"
env_creds.secret_key = "ambient-secret"
env_creds.token = "ambient-token"
assumed_creds = MagicMock()
assumed_creds.access_key = "assumed-key"
assumed_creds.secret_key = "assumed-secret"
assumed_creds.token = "assumed-token"
role_arn = "arn:aws:iam::123456789012:role/TargetRole"
with patch.object(
base_aws_llm,
"_is_already_running_as_role",
side_effect=[True, False],
) as mock_already:
with patch.object(
base_aws_llm,
"_auth_with_env_vars",
return_value=(env_creds, None),
) as mock_env_auth:
with patch.object(
base_aws_llm,
"_auth_with_aws_role",
return_value=(assumed_creds, 3600),
) as mock_role_auth:
first = base_aws_llm.get_credentials(
aws_role_name=role_arn,
aws_region_name="us-east-1",
)
in_memory_cache = BaseAWSLLM._shared_iam_cache.in_memory_cache
lapsed_key = list(in_memory_cache.ttl_dict.keys())[0]
in_memory_cache.ttl_dict[lapsed_key] = time.time() - 1
second = base_aws_llm.get_credentials(
aws_role_name=role_arn,
aws_region_name="us-east-1",
)
assert mock_already.call_count == 2
mock_env_auth.assert_called_once()
mock_role_auth.assert_called_once()
assert first.access_key == "ambient-key"
assert second.access_key == "assumed-key"
def test_parse_arn_account_and_role_name():
"""Test the ARN parser helper for various ARN formats."""
parse = BaseAWSLLM._parse_arn_account_and_role_name
# Standard IAM role ARN
assert parse("arn:aws:iam::123456789012:role/MyRole") == (
"aws",
"123456789012",
"MyRole",
)
# IAM role ARN with path
assert parse("arn:aws:iam::123456789012:role/service-role/MyRole") == (
"aws",
"123456789012",
"MyRole",
)
# Assumed-role ARN (from GetCallerIdentity)
assert parse("arn:aws:sts::123456789012:assumed-role/MyRole/session-id") == (
"aws",
"123456789012",
"MyRole",
)
# China partition
assert parse("arn:aws-cn:iam::123456789012:role/MyRole") == (
"aws-cn",
"123456789012",
"MyRole",
)
# GovCloud partition
assert parse("arn:aws-us-gov:iam::123456789012:role/MyRole") == (
"aws-us-gov",
"123456789012",
"MyRole",
)
# Invalid ARNs
assert parse("not-an-arn") is None
assert parse("arn:aws:iam::123456789012:user/MyUser") is None
assert parse("") is None
def test_is_already_running_as_role_cross_account_same_name():
"""
Test that same role NAME in different accounts does NOT match.
This is the cross-account false-match prevention.
"""
base_aws_llm = BaseAWSLLM()
mock_sts_client = MagicMock()
# Caller is in account 111111111111
mock_sts_client.get_caller_identity.return_value = {
"Arn": "arn:aws:sts::111111111111:assumed-role/MyRole/session-id"
}
with patch.dict(os.environ, {}, clear=False):
env = {
k: v
for k, v in os.environ.items()
if k not in ("AWS_ROLE_ARN", "AWS_WEB_IDENTITY_TOKEN_FILE")
}
with patch.dict(os.environ, env, clear=True):
with patch("boto3.client", return_value=mock_sts_client):
# Target is same role name but in account 222222222222
assert (
base_aws_llm._is_already_running_as_role(
"arn:aws:iam::222222222222:role/MyRole"
)
is False
)
def test_is_already_running_as_role_cross_partition():
"""
Test that same role name + account but different partition does NOT match.
"""
base_aws_llm = BaseAWSLLM()
mock_sts_client = MagicMock()
mock_sts_client.get_caller_identity.return_value = {
"Arn": "arn:aws:sts::123456789012:assumed-role/MyRole/session-id"
}
with patch.dict(os.environ, {}, clear=False):
env = {
k: v
for k, v in os.environ.items()
if k not in ("AWS_ROLE_ARN", "AWS_WEB_IDENTITY_TOKEN_FILE")
}
with patch.dict(os.environ, env, clear=True):
with patch("boto3.client", return_value=mock_sts_client):
# Same account and role but aws-cn partition
assert (
base_aws_llm._is_already_running_as_role(
"arn:aws-cn:iam::123456789012:role/MyRole"
)
is False
)
def test_is_already_running_as_role_invalid_target_arn():
"""
Test that an unparseable target ARN returns False immediately.
"""
base_aws_llm = BaseAWSLLM()
# Should return False without making any API calls
assert base_aws_llm._is_already_running_as_role("not-a-valid-arn") is False
def test_filter_headers_skips_none_values():
"""
Test that _filter_headers_for_aws_signature skips headers with None values.
Reproduces the issue where botocore's SigV4Auth crashes with
'NoneType' object has no attribute 'split' when a header value is None.
"""
llm = BaseAWSLLM()
headers = {
"Content-Type": "application/json",
"x-amz-security-token": None,
"x-amzn-bedrock-kb-session-id": None,
"host": None,
"x-amz-date": "20240101T000000Z",
"x-custom-header": None,
}
filtered = llm._filter_headers_for_aws_signature(headers)
assert filtered["Content-Type"] == "application/json"
assert filtered["x-amz-date"] == "20240101T000000Z"
assert "x-amz-security-token" not in filtered
assert "x-amzn-bedrock-kb-session-id" not in filtered
assert "host" not in filtered
# Non-AWS headers are excluded regardless
assert "x-custom-header" not in filtered
def test_sign_request_with_none_header_values():
"""
End-to-end test that _sign_request does not crash when headers contain
None values for x-amz-* keys.
This reproduces the Bedrock KB GovCloud issue where SigV4 signing failed
with 'NoneType' object has no attribute 'split'.
Also verifies that None-valued headers are NOT re-merged into the
returned headers dict (which would cause downstream HTTP client failures).
"""
llm = BaseAWSLLM()
mock_credentials = Credentials("test_key", "test_secret")
headers_with_nones = {
"Content-Type": "application/json",
"x-amzn-trace-id": None,
"x-forwarded-for": None,
}
with (
patch.object(llm, "get_credentials", return_value=mock_credentials),
patch.object(llm, "_get_aws_region_name", return_value="us-gov-west-1"),
):
result_headers, result_body = llm._sign_request(
service_name="bedrock",
headers=headers_with_nones,
optional_params={
"aws_access_key_id": "test_key",
"aws_secret_access_key": "test_secret",
"aws_region_name": "us-gov-west-1",
},
request_data={"retrievalQuery": {"text": "test query"}},
api_base="https://bedrock-agent-runtime.us-gov-west-1.amazonaws.com/knowledgebases/KB123/retrieve",
)
assert "Authorization" in result_headers
assert result_body is not None
# None-valued headers must NOT appear in the returned headers
for header_name, header_value in result_headers.items():
assert (
header_value is not None
), f"Header '{header_name}' has None value in returned headers"
def test_is_already_running_as_role_ssl_verify_passed():
"""
Test that ssl_verify parameter is correctly passed to the STS client.
"""
base_aws_llm = BaseAWSLLM()
mock_sts_client = MagicMock()
mock_sts_client.get_caller_identity.return_value = {
"Arn": "arn:aws:sts::123456789012:assumed-role/MyRole/session-id"
}
with patch.dict(os.environ, {}, clear=False):
env = {
k: v
for k, v in os.environ.items()
if k not in ("AWS_ROLE_ARN", "AWS_WEB_IDENTITY_TOKEN_FILE")
}
with patch.dict(os.environ, env, clear=True):
with patch(
"boto3.client", return_value=mock_sts_client
) as mock_boto3_client:
base_aws_llm._is_already_running_as_role(
"arn:aws:iam::123456789012:role/MyRole",
ssl_verify="/path/to/ca-bundle.crt",
)
mock_boto3_client.assert_called_once_with(
"sts", verify="/path/to/ca-bundle.crt"
)
# ---------------------------------------------------------------------------
# LIT-3274: get_bedrock_model_id must strip "bedrock/" prefix and URL-encode
# ARNs for the invoke path (invoke-with-response-stream). Without this fix
# the Bedrock API receives a malformed URL, returns a JSON error body, and
# botocore's EventStreamBuffer raises ChecksumMismatch instead of the real
# error. 0x223a7b22 == ':{\"' — the start of a JSON object.
# ---------------------------------------------------------------------------
class TestGetBedrockModelIdArnHandling:
"""Unit tests for get_bedrock_model_id with inference-profile ARNs."""
ARN = "arn:aws:bedrock:us-east-1:086734376398:inference-profile/global.anthropic.claude-sonnet-4-5-20250929-v1:0"
def _call(self, model: str, optional_params: dict | None = None) -> str:
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
provider = BaseAWSLLM.get_bedrock_invoke_provider(model)
return BaseAWSLLM.get_bedrock_model_id(
model=model,
provider=provider,
optional_params=optional_params or {},
)
def test_arn_with_bedrock_prefix_is_stripped_and_encoded(self):
"""bedrock/arn:... must not appear verbatim in the model_id."""
model_id = self._call(f"bedrock/{self.ARN}")
assert (
"bedrock/arn" not in model_id
), f"'bedrock/' prefix not stripped; got: {model_id}"
# Must be URL-encoded (colons → %3A)
assert "%3A" in model_id, f"ARN not URL-encoded; got: {model_id}"
assert "%2F" in model_id, f"ARN slashes not URL-encoded; got: {model_id}"
def test_arn_with_compound_bedrock_invoke_prefix_is_fully_stripped_and_encoded(
self,
):
"""bedrock/invoke/arn:... — compound prefix — must be fully stripped.
The old fix used ``break`` after the first matched prefix, so
``bedrock/invoke/arn:...`` would only strip ``bedrock/``, leaving
``invoke/arn:...``. The subsequent ``.replace('invoke/', '')`` call
then returned the bare unencoded ARN, reproducing the same
malformed-URL bug the fix aimed to prevent.
strip_bedrock_routing_prefix() has no break and handles this correctly.
"""
model_id = self._call(f"bedrock/invoke/{self.ARN}")
assert (
"invoke/" not in model_id
), f"'invoke/' prefix not stripped; got: {model_id}"
assert (
"bedrock/" not in model_id
), f"'bedrock/' prefix not stripped; got: {model_id}"
assert "%3A" in model_id, f"ARN not URL-encoded; got: {model_id}"
assert "%2F" in model_id, f"ARN slashes not URL-encoded; got: {model_id}"
def test_bare_arn_is_encoded(self):
"""Direct ARN without routing prefix must also be URL-encoded."""
model_id = self._call(self.ARN)
assert "%3A" in model_id, f"ARN not URL-encoded; got: {model_id}"
assert "%2F" in model_id, f"ARN slashes not URL-encoded; got: {model_id}"
def test_arn_url_matches_expected(self):
"""Full URL built from messages config must match expected encoded form."""
import urllib.parse
from litellm.llms.bedrock.messages.invoke_transformations.anthropic_claude3_transformation import (
AmazonAnthropicClaudeMessagesConfig,
)
config = AmazonAnthropicClaudeMessagesConfig()
url = config.get_complete_url(
api_base=None,
api_key=None,
model=f"bedrock/{self.ARN}",
optional_params={"aws_region_name": "us-east-1"},
litellm_params={},
stream=True,
)
encoded_arn = urllib.parse.quote(self.ARN, safe="")
expected = (
f"https://bedrock-runtime.us-east-1.amazonaws.com"
f"/model/{encoded_arn}/invoke-with-response-stream"
)
assert (
url == expected
), f"URL mismatch:\n got: {url}\n expected: {expected}"
def test_regular_model_id_unaffected(self):
"""Non-ARN model IDs must continue to work as before."""
model_id = self._call("anthropic.claude-3-sonnet-20240229-v1:0")
assert model_id == "anthropic.claude-3-sonnet-20240229-v1:0"
def test_invoke_prefixed_model_unaffected(self):
"""invoke/ prefix stripping still works after the fix."""
model_id = self._call("invoke/anthropic.claude-3-sonnet-20240229-v1:0")
assert model_id == "anthropic.claude-3-sonnet-20240229-v1:0"
def _recomputed_sigv4_signature(url: str, secret_key: str, authorization: str, headers: Dict[str, Any], body) -> str:
import hashlib
import hmac
from urllib.parse import urlparse
parsed = urlparse(url)
credential_scope = authorization.split("Credential=")[1].split(",")[0].split("/", 1)[1]
signed_header_names = authorization.split("SignedHeaders=")[1].split(",")[0].split(";")
header_lookup = {name.lower(): str(value) for name, value in headers.items()}
header_lookup["host"] = parsed.netloc
body_bytes = body if isinstance(body, bytes) else str(body).encode()
canonical_request = "\n".join(
[
"POST",
parsed.path or "/",
"",
"".join(f"{name}:{header_lookup[name]}\n" for name in signed_header_names),
";".join(signed_header_names),
hashlib.sha256(body_bytes).hexdigest(),
]
)
string_to_sign = "\n".join(
[
"AWS4-HMAC-SHA256",
header_lookup["x-amz-date"],
credential_scope,
hashlib.sha256(canonical_request.encode()).hexdigest(),
]
)
key = f"AWS4{secret_key}".encode()
for scope_part in credential_scope.split("/"):
key = hmac.new(key, scope_part.encode(), hashlib.sha256).digest()
return hmac.new(key, string_to_sign.encode(), hashlib.sha256).hexdigest()
class TestSignRequestResign:
"""Regression: retrying a Bedrock request with headers from a previous SigV4 sign
(e.g. the /v1/messages strip-thinking-and-retry path) must produce a fresh
Authorization / X-Amz-Date for the new body, not inherit the stale ones and 403."""
URL = "https://bedrock-runtime.us-east-1.amazonaws.com/model/test-model/invoke"
ACCESS_KEY = "AKIAIOSFODNN7EXAMPLE"
SECRET_KEY = "wJalrXUtnFEMI/K7MDENG/bPxRfiCYEXAMPLEKEY"
@pytest.fixture(autouse=True)
def _clean_aws_env(self, monkeypatch):
for env_var in ("AWS_BEARER_TOKEN_BEDROCK", "AWS_SESSION_TOKEN", "AWS_PROFILE"):
monkeypatch.delenv(env_var, raising=False)
def _optional_params(self) -> Dict[str, Any]:
return {
"aws_access_key_id": self.ACCESS_KEY,
"aws_secret_access_key": self.SECRET_KEY,
"aws_region_name": "us-east-1",
}
def _sign(self, headers: Dict[str, Any], request_data: Dict[str, Any]):
return BaseAWSLLM()._sign_request(
service_name="bedrock",
headers=headers,
optional_params=self._optional_params(),
request_data=request_data,
api_base=self.URL,
)
def test_resign_with_previously_signed_headers_replaces_stale_sigv4_headers(self):
original_body = {
"messages": [
{
"role": "assistant",
"content": [{"type": "thinking", "thinking": "x", "signature": ""}],
}
]
}
first_headers, _ = self._sign(headers={"Content-Type": "application/json"}, request_data=original_body)
assert first_headers["Authorization"].startswith("AWS4-HMAC-SHA256")
stale_headers = {**first_headers, "X-Amz-Date": "20200101T000000Z"}
stripped_body = {"messages": [{"role": "user", "content": "hi"}]}
second_headers, second_signed_body = self._sign(headers=stale_headers, request_data=stripped_body)
assert second_headers["X-Amz-Date"] != "20200101T000000Z"
assert second_headers["Authorization"] != stale_headers["Authorization"]
assert second_headers["Authorization"].split("Signature=")[1] == _recomputed_sigv4_signature(
url=self.URL,
secret_key=self.SECRET_KEY,
authorization=second_headers["Authorization"],
headers=second_headers,
body=second_signed_body,
)
def test_forwarded_headers_still_added_back_after_signing(self):
signed_headers, _ = self._sign(
headers={"Content-Type": "application/json", "anthropic-version": "bedrock-2023-05-31"},
request_data={"messages": []},
)
assert signed_headers["anthropic-version"] == "bedrock-2023-05-31"
assert signed_headers["Content-Type"] == "application/json"
def test_caller_supplied_bearer_authorization_survives_signing(self):
signed_headers, _ = self._sign(
headers={"Content-Type": "application/json", "Authorization": "Bearer caller-token"},
request_data={"messages": []},
)
assert signed_headers["Authorization"] == "Bearer caller-token"
class TestGetRequestHeadersResign:
"""Regression: get_request_headers (invoke/converse/embed/image paths) must not let
stale SigV4 values present in the input headers clobber the freshly computed signature."""
URL = "https://bedrock-runtime.us-east-1.amazonaws.com/model/test-model/converse"
ACCESS_KEY = "AKIAIOSFODNN7EXAMPLE"
SECRET_KEY = "wJalrXUtnFEMI/K7MDENG/bPxRfiCYEXAMPLEKEY"
SESSION_TOKEN = "fresh-session-token"
@pytest.fixture(autouse=True)
def _clean_aws_env(self, monkeypatch):
for env_var in ("AWS_BEARER_TOKEN_BEDROCK", "AWS_SESSION_TOKEN", "AWS_PROFILE"):
monkeypatch.delenv(env_var, raising=False)
def _prepare(self, headers: Dict[str, Any], data: str, extra_headers: Optional[Dict[str, str]] = None):
return BaseAWSLLM().get_request_headers(
credentials=Credentials(self.ACCESS_KEY, self.SECRET_KEY, self.SESSION_TOKEN),
aws_region_name="us-east-1",
extra_headers=extra_headers,
endpoint_url=self.URL,
data=data,
headers=headers,
)
def test_stale_sigv4_headers_in_input_replaced_by_fresh_signature(self):
first_prepped = self._prepare(
headers={"Content-Type": "application/json"},
data=json.dumps({"messages": [{"role": "user", "content": "original"}]}),
)
stale_authorization = first_prepped.headers["Authorization"]
assert stale_authorization.startswith("AWS4-HMAC-SHA256")
stale_headers = {
"Content-Type": "application/json",
"Authorization": stale_authorization,
"X-Amz-Date": "20200101T000000Z",
"X-Amz-Security-Token": "stale-session-token",
}
retry_data = json.dumps({"messages": [{"role": "user", "content": "retry"}]})
second_prepped = self._prepare(headers=stale_headers, data=retry_data)
assert second_prepped.headers["X-Amz-Date"] != "20200101T000000Z"
assert second_prepped.headers["X-Amz-Security-Token"] == self.SESSION_TOKEN
assert second_prepped.headers["Authorization"] != stale_authorization
assert second_prepped.headers["Authorization"].split("Signature=")[1] == _recomputed_sigv4_signature(
url=self.URL,
secret_key=self.SECRET_KEY,
authorization=second_prepped.headers["Authorization"],
headers=dict(second_prepped.headers),
body=retry_data,
)
def test_forwarded_headers_still_added_back_after_signing(self):
prepped = self._prepare(
headers={
"Content-Type": "application/json",
"anthropic-version": "bedrock-2023-05-31",
"user-agent": "litellm-test-client",
},
data=json.dumps({"messages": []}),
)
assert prepped.headers["anthropic-version"] == "bedrock-2023-05-31"
assert prepped.headers["user-agent"] == "litellm-test-client"
assert prepped.headers["Content-Type"] == "application/json"
def test_extra_headers_bearer_authorization_still_overrides_signature(self):
prepped = self._prepare(
headers={"Content-Type": "application/json"},
data=json.dumps({"messages": []}),
extra_headers={"Authorization": "Bearer foo"},
)
assert prepped.headers["Authorization"] == "Bearer foo"
@pytest.mark.asyncio
async def test_sign_request_off_loop_if_aws_keeps_the_loop_serving_while_credentials_refresh():
"""Regression for issue #40165: an AWS provider's signing (and the botocore credential refresh
inside it) must run off the event loop, so other requests keep being served meanwhile."""
probe = EventLoopProbe()
def sign(headers: dict[str, str]) -> dict[str, str]:
request = AWSRequest(
method="POST", url="https://bedrock-runtime.us-west-2.amazonaws.com/", data="{}", headers=headers
)
SigV4Auth(probe.credentials(), "bedrock", "us-west-2").add_auth(request)
return dict(request.headers)
release = asyncio.create_task(probe.release_refresh_from_the_loop())
signed = await sign_request_off_loop_if_aws(BaseAWSLLM(), sign, headers={"Content-Type": "application/json"})
await release
assert "Authorization" in signed
assert probe.served_during_refresh is True
def test_run_aws_signing_leaves_the_default_executor_free_for_other_providers():
"""A signing parked on botocore's refresh lock must not hold a default-executor thread, since every
other provider's async entry point hops through that same executor. The scenario runs on its own loop
so the one-thread default executor it pins never leaks into the session loop."""
async def scenario() -> tuple[str, str]:
loop = asyncio.get_running_loop()
loop.set_default_executor(ThreadPoolExecutor(max_workers=1))
signing_parked = asyncio.Event()
refresh_done = threading.Event()
def sign() -> str:
loop.call_soon_threadsafe(signing_parked.set)
refresh_done.wait()
return threading.current_thread().name
signing = asyncio.create_task(run_aws_signing(sign))
try:
await asyncio.wait_for(signing_parked.wait(), timeout=5)
other_provider = await asyncio.wait_for(loop.run_in_executor(None, threading.current_thread), timeout=5)
finally:
refresh_done.set()
return other_provider.name, await signing
other_provider, signing_thread = asyncio.run(scenario())
assert other_provider != signing_thread
assert signing_thread.startswith("aws-signing")
def _recording_boto3_client(recorded: dict[str, dict[str, object]]) -> Callable[..., MagicMock]:
"""boto3.client replacement that records the STS client kwargs and the assume-role params."""
def _client(service_name: str, **client_kwargs: object) -> MagicMock:
recorded["client_kwargs"] = client_kwargs
sts = MagicMock()
def _assume(**params: object) -> dict[str, object]:
recorded["assume_role"] = params
return {
"Credentials": {
"AccessKeyId": "ASIAASSUMED",
"SecretAccessKey": "assumed-secret",
"SessionToken": "assumed-token",
"Expiration": datetime.now(timezone.utc) + timedelta(minutes=30),
}
}
def _assume_web_identity(**params: object) -> dict[str, object]:
recorded["assume_role_with_web_identity"] = params
return {
"Credentials": {
"AccessKeyId": "ASIAWEBIDENTITY",
"SecretAccessKey": "assumed-secret",
"SessionToken": "assumed-token",
"Expiration": datetime.now(timezone.utc) + timedelta(minutes=30),
},
"PackedPolicySize": 10,
}
sts.assume_role.side_effect = _assume
sts.assume_role_with_web_identity.side_effect = _assume_web_identity
return sts
return _client
def test_resolve_credentials_forwards_static_keys_role_session_and_external_id():
"""Every field the role-assumption route reads must reach STS, so a dropped struct field fails here."""
from litellm.types.llms.bedrock import AwsAuthParams
auth_params = AwsAuthParams(
aws_access_key_id="AKIACALLER",
aws_secret_access_key="caller-secret",
aws_session_token="caller-token",
aws_role_name="arn:aws:iam::123456789012:role/litellm-target",
aws_session_name="litellm-session",
aws_external_id="litellm-external-id",
aws_sts_endpoint="https://custom-sts.example",
aws_session_tags=[{"Key": "team", "Value": "genai"}, {"Key": "cost-center", "Value": "42"}],
)
recorded: dict[str, dict[str, object]] = {}
with (
patch.dict(os.environ, _os_environ_without_aws_keys(), clear=True),
patch("boto3.client", side_effect=_recording_boto3_client(recorded)),
):
credentials = BaseAWSLLM().resolve_credentials(auth_params, "us-east-1")
assert recorded["client_kwargs"]["aws_access_key_id"] == "AKIACALLER"
assert recorded["client_kwargs"]["aws_secret_access_key"] == "caller-secret"
assert recorded["client_kwargs"]["aws_session_token"] == "caller-token"
assert recorded["client_kwargs"]["endpoint_url"] == "https://custom-sts.example"
assert recorded["assume_role"]["RoleArn"] == "arn:aws:iam::123456789012:role/litellm-target"
assert recorded["assume_role"]["RoleSessionName"] == "litellm-session"
assert recorded["assume_role"]["ExternalId"] == "litellm-external-id"
assert recorded["assume_role"]["Tags"] == (
{"Key": "cost-center", "Value": "42"},
{"Key": "team", "Value": "genai"},
)
assert credentials.access_key == "ASIAASSUMED"
@pytest.mark.parametrize(
"malformed_tags",
[
"team=genai",
{"team": "genai"},
[{"key": "team", "value": "genai"}],
[{"Key": "team"}],
],
)
def test_resolve_credentials_rejects_malformed_session_tags(malformed_tags):
"""A struct built from raw config must surface the friendly session-tag error before STS is called."""
from litellm.types.llms.bedrock import AwsAuthParams
auth_params = AwsAuthParams(
aws_role_name="arn:aws:iam::123456789012:role/litellm-target",
aws_session_name="litellm-session",
aws_session_tags=malformed_tags,
)
recorded: dict[str, dict[str, object]] = {}
with (
patch.dict(os.environ, _os_environ_without_aws_keys(), clear=True),
patch("boto3.client", side_effect=_recording_boto3_client(recorded)),
):
with pytest.raises(ValueError, match="Invalid 'aws_session_tags' value"):
BaseAWSLLM().resolve_credentials(auth_params, "us-east-1")
assert "assume_role" not in recorded
def test_resolve_credentials_forwards_web_identity_token():
"""A struct carrying a web-identity token must take the web-identity route, not plain role assumption."""
from litellm.types.llms.bedrock import AwsAuthParams
auth_params = AwsAuthParams(
aws_web_identity_token="unresolvable-oidc-token",
aws_role_name="arn:aws:iam::123456789012:role/litellm-wif",
aws_session_name="litellm-wif-session",
)
recorded: dict[str, dict[str, object]] = {}
with (
patch.dict(os.environ, _os_environ_without_aws_keys(), clear=True),
patch("boto3.client", side_effect=_recording_boto3_client(recorded)),
):
with pytest.raises(AwsAuthError) as exc:
BaseAWSLLM().resolve_credentials(auth_params, "us-east-1")
assert exc.value.status_code == 401
assert "assume_role" not in recorded
def test_resolve_credentials_forwards_profile_name():
"""The profile route must receive the struct's profile name rather than the ambient session."""
from litellm.types.llms.bedrock import AwsAuthParams
auth_params = AwsAuthParams(aws_profile_name="litellm-qa-profile")
session_instance = MagicMock()
session_instance.get_credentials.return_value = Credentials(
access_key="AKIAPROFILE", secret_key="profile-secret", token=None
)
with (
patch.dict(os.environ, _os_environ_without_aws_keys(), clear=True),
patch("boto3.Session", return_value=session_instance) as mock_session_cls,
):
credentials = BaseAWSLLM().resolve_credentials(auth_params, "us-east-1")
assert mock_session_cls.call_args.kwargs["profile_name"] == "litellm-qa-profile"
assert credentials.access_key == "AKIAPROFILE"