mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
* 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>
3703 lines
139 KiB
Python
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"
|