fix(bedrock): grant rerank, retrieve, agent, and agentcore actions in the web identity session policy

This commit is contained in:
mateo-berri 2026-09-14 19:03:55 -07:00
parent d2859e18d7
commit 8d5d5d5824
2 changed files with 124 additions and 54 deletions

View file

@ -11,6 +11,7 @@ from concurrent.futures import ThreadPoolExecutor
from datetime import datetime
from functools import partial
from threading import Lock
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, ClassVar, Final, Literal, ParamSpec, TypeVar, cast, get_args, overload
import httpx
@ -96,6 +97,76 @@ def _assume_role_params(
)
_SecureTransportBool = TypedDict("_SecureTransportBool", {"aws:SecureTransport": ReadOnly[Literal["true"]]})
class _SecureTransportCondition(TypedDict):
Bool: ReadOnly[_SecureTransportBool]
class _SessionPolicyStatement(TypedDict):
Sid: ReadOnly[str]
Effect: ReadOnly[Literal["Allow"]]
Action: ReadOnly[tuple[str, ...]]
Resource: ReadOnly[Literal["*"]]
Condition: ReadOnly[_SecureTransportCondition]
class WebIdentitySessionPolicy(TypedDict):
Version: ReadOnly[Literal["2012-10-17"]]
Statement: ReadOnly[tuple[_SessionPolicyStatement, ...]]
_WEB_IDENTITY_SESSION_POLICY_ACTIONS: Final[Mapping[str, tuple[str, ...]]] = MappingProxyType(
{
"BedrockLiteLLM": (
"bedrock:InvokeModel",
"bedrock:InvokeModelWithResponseStream",
"bedrock:CountTokens",
"bedrock:Rerank",
"bedrock:Retrieve",
"bedrock:ListKnowledgeBases",
"bedrock:InvokeAgent",
"bedrock:ApplyGuardrail",
"bedrock:GetGuardrail",
"bedrock:ListGuardrails",
),
"BedrockAgentCoreLiteLLM": (
"bedrock-agentcore:InvokeAgentRuntime",
"bedrock-agentcore:InvokeGateway",
),
"ClaudePlatformLiteLLM": (
"aws-external-anthropic:CreateInference",
"aws-external-anthropic:CreateBatchInference",
"aws-external-anthropic:CancelBatchInference",
"aws-external-anthropic:DeleteBatchInference",
"aws-external-anthropic:CountTokens",
"aws-external-anthropic:Get*",
"aws-external-anthropic:List*",
),
"BedrockMantleLiteLLM": ("bedrock-mantle:CreateInference",),
}
)
_SECURE_TRANSPORT_ONLY: Final = _SecureTransportCondition(Bool=_SecureTransportBool({"aws:SecureTransport": "true"}))
def build_web_identity_session_policy() -> WebIdentitySessionPolicy:
return WebIdentitySessionPolicy(
Version="2012-10-17",
Statement=tuple(
_SessionPolicyStatement(
Sid=sid,
Effect="Allow",
Action=actions,
Resource="*",
Condition=_SECURE_TRANSPORT_ONLY,
)
for sid, actions in _WEB_IDENTITY_SESSION_POLICY_ACTIONS.items()
),
)
class BedrockRequestTarget(BaseModel):
aws_region_name: str
aws_bedrock_runtime_endpoint: str | None
@ -940,60 +1011,12 @@ class BaseAWSLLM(SignsRequestsWithAWS):
# auth only (static creds + IRSA take other code paths).
# https://docs.aws.amazon.com/STS/latest/APIReference/API_AssumeRoleWithWebIdentity.html
# https://boto3.amazonaws.com/v1/documentation/api/latest/reference/services/sts/client/assume_role_with_web_identity.html
bedrock_session_policy: Final = {
"Version": "2012-10-17",
"Statement": [
{
"Sid": "BedrockLiteLLM",
"Effect": "Allow",
"Action": [
"bedrock:InvokeModel",
"bedrock:InvokeModelWithResponseStream",
"bedrock:CountTokens",
"bedrock:ApplyGuardrail",
"bedrock:GetGuardrail",
"bedrock:ListGuardrails",
],
"Resource": "*",
"Condition": {"Bool": {"aws:SecureTransport": "true"}},
},
# Claude Platform on AWS (added by #27678 for the
# ``bedrock/claude_platform/<model>`` route) lives under
# a separate IAM action namespace; without these entries
# the OIDC path 403s on every claude_platform request
# even with a fully permissive identity policy (#30200).
{
"Sid": "ClaudePlatformLiteLLM",
"Effect": "Allow",
"Action": [
"aws-external-anthropic:CreateInference",
"aws-external-anthropic:CreateBatchInference",
"aws-external-anthropic:CancelBatchInference",
"aws-external-anthropic:DeleteBatchInference",
"aws-external-anthropic:CountTokens",
"aws-external-anthropic:Get*",
"aws-external-anthropic:List*",
],
"Resource": "*",
"Condition": {"Bool": {"aws:SecureTransport": "true"}},
},
{
"Sid": "BedrockMantleLiteLLM",
"Effect": "Allow",
"Action": [
"bedrock-mantle:CreateInference",
],
"Resource": "*",
"Condition": {"Bool": {"aws:SecureTransport": "true"}},
},
],
}
assume_role_params: Final = {
"RoleArn": aws_role_name,
"RoleSessionName": aws_session_name,
"WebIdentityToken": oidc_token,
"DurationSeconds": 3600,
"Policy": json.dumps(bedrock_session_policy, separators=(",", ":")),
"Policy": json.dumps(build_web_identity_session_policy(), separators=(",", ":")),
}
# Add ExternalId parameter if provided

View file

@ -32,6 +32,8 @@ action.
import base64
import json
from datetime import datetime, timedelta, timezone
from types import MappingProxyType
from typing import Final
from unittest.mock import MagicMock, patch
import pytest
@ -49,9 +51,9 @@ _CLAUDE_PLATFORM_ACTIONS = {
}
def _captured_policy() -> dict:
"""Run _auth_with_web_identity_token under mocks + return the parsed
Policy dict that was actually sent to STS."""
def _captured_policy_document() -> str:
"""Run _auth_with_web_identity_token under mocks + return the Policy
JSON document that was actually sent to STS."""
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
base = BaseAWSLLM()
@ -84,8 +86,15 @@ def _captured_policy() -> dict:
mock_sts.assume_role_with_web_identity.assert_called_once()
kwargs = mock_sts.assume_role_with_web_identity.call_args.kwargs
policy_str = kwargs["Policy"]
return json.loads(policy_str)
return kwargs["Policy"]
def _captured_policy() -> dict:
return json.loads(_captured_policy_document())
def _granted_actions(policy: dict) -> frozenset[str]:
return frozenset(action for stmt in policy["Statement"] for action in stmt["Action"])
def _statement_by_sid(policy: dict, sid: str) -> dict:
@ -308,3 +317,41 @@ class TestPolicyTransportConditions:
"ClaudePlatformLiteLLM must require aws:SecureTransport=true "
"to keep parity with the bedrock statement"
)
_STS_SESSION_POLICY_PLAINTEXT_LIMIT: Final = 2048
_BEDROCK_ROUTE_ACTIONS: Final = MappingProxyType(
{
"model/{model_id}/invoke": "bedrock:InvokeModel",
"model/{model_id}/invoke-with-response-stream": "bedrock:InvokeModelWithResponseStream",
"model/{model_id}/converse": "bedrock:InvokeModel",
"model/{model_id}/converse-stream": "bedrock:InvokeModelWithResponseStream",
"model/{model_id}/count-tokens": "bedrock:CountTokens",
"guardrail/{guardrail_id}/version/{version}/apply": "bedrock:ApplyGuardrail",
"rerank": "bedrock:Rerank",
"knowledgebases/{knowledge_base_id}/retrieve": "bedrock:Retrieve",
"knowledgebases": "bedrock:ListKnowledgeBases",
"agents/{agent_id}/agentAliases/{alias_id}/sessions/{session_id}/text": "bedrock:InvokeAgent",
"runtimes/{agent_runtime_arn}/invocations": "bedrock-agentcore:InvokeAgentRuntime",
"mcp": "bedrock-agentcore:InvokeGateway",
}
)
class TestSessionPolicyGrantsEveryBedrockRoute:
"""LIT-7348: ``/rerank`` authorizes against ``bedrock:Rerank``, which the
ceiling never granted, so rerank 403d on web identity auth while static
credentials and IRSA worked. Each route the bedrock package signs with the
web identity session maps to the IAM action it authorizes against, and the
ceiling must grant every one of them."""
@pytest.mark.parametrize(("route", "action"), sorted(_BEDROCK_ROUTE_ACTIONS.items()))
def test_route_action_is_granted_by_the_ceiling(self, route: str, action: str):
assert action in _granted_actions(_captured_policy()), (
f"/{route} authorizes against {action}, which the session policy does not grant, "
"so it 403s on web identity auth"
)
def test_policy_document_fits_the_sts_plaintext_limit(self):
assert len(_captured_policy_document()) <= _STS_SESSION_POLICY_PLAINTEXT_LIMIT