This commit is contained in:
devin-ai-integration[bot] 2026-10-03 16:30:23 -04:00 • committed by GitHub
commit 15352a5365
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 207 additions and 1 deletions

View file

@ -3,6 +3,16 @@ from typing import Any, Final
import httpx
AWS_RDS_IAM_IGNORE_WEB_IDENTITY_TOKEN_ENV_VAR: Final = "AWS_RDS_IAM_IGNORE_WEB_IDENTITY_TOKEN"
_IGNORE_WEB_IDENTITY_TOKEN_TRUTHY_VALUES: Final[frozenset[str]] = frozenset({"1", "true"})
def _ignore_web_identity_token() -> bool:
return (
os.getenv(AWS_RDS_IAM_IGNORE_WEB_IDENTITY_TOKEN_ENV_VAR, "").strip().lower()
in _IGNORE_WEB_IDENTITY_TOKEN_TRUTHY_VALUES
)
def init_rds_client(
aws_access_key_id: str | None = None,
@ -12,6 +22,7 @@ def init_rds_client(
aws_profile_name: str | None = None,
aws_role_name: str | None = None,
aws_web_identity_token: str | None = None,
aws_session_token: str | None = None,
timeout: float | httpx.Timeout | None = None,
):
from litellm.secret_managers.main import get_secret
@ -29,6 +40,7 @@ def init_rds_client(
aws_profile_name,
aws_role_name,
aws_web_identity_token,
aws_session_token,
]
# Iterate over parameters and update if needed
@ -44,6 +56,7 @@ def init_rds_client(
aws_profile_name,
aws_role_name,
aws_web_identity_token,
aws_session_token,
) = params_to_check
### SET REGION NAME
@ -106,6 +119,7 @@ def init_rds_client(
"sts",
aws_access_key_id=aws_access_key_id,
aws_secret_access_key=aws_secret_access_key,
aws_session_token=aws_session_token,
)
sts_response = sts_client.assume_role(RoleArn=aws_role_name, RoleSessionName=aws_session_name)
@ -126,6 +140,7 @@ def init_rds_client(
service_name="rds",
aws_access_key_id=aws_access_key_id,
aws_secret_access_key=aws_secret_access_key,
aws_session_token=aws_session_token,
region_name=region_name,
config=config,
)
@ -159,10 +174,13 @@ def generate_iam_auth_token(db_host, db_port, db_user, client: Any | None = None
aws_region_name=os.getenv("AWS_REGION_NAME"),
aws_access_key_id=os.getenv("AWS_ACCESS_KEY_ID"),
aws_secret_access_key=os.getenv("AWS_SECRET_ACCESS_KEY"),
aws_session_token=os.getenv("AWS_SESSION_TOKEN"),
aws_session_name=os.getenv("AWS_SESSION_NAME"),
aws_profile_name=os.getenv("AWS_PROFILE_NAME"),
aws_role_name=os.getenv("AWS_ROLE_NAME", os.getenv("AWS_ROLE_ARN")),
aws_web_identity_token=os.getenv("AWS_WEB_IDENTITY_TOKEN", os.getenv("AWS_WEB_IDENTITY_TOKEN_FILE")),
aws_web_identity_token=None
if _ignore_web_identity_token()
else os.getenv("AWS_WEB_IDENTITY_TOKEN", os.getenv("AWS_WEB_IDENTITY_TOKEN_FILE")),
)
else:
boto_client = client

View file

@ -0,0 +1,188 @@
import sys
import types
import urllib.parse
from dataclasses import dataclass, field
from pathlib import Path
from typing import Final
import pytest
from litellm.proxy.auth.rds_iam_token import (
generate_iam_auth_token, # pyright: ignore[reportUnknownVariableType] # source signature params are untyped upstream
)
FAKE_DB_AUTH_TOKEN: Final = "db:5432/?Action=connect&X-Amz-Credential=a/b"
AWS_ENV_VARS: Final = (
"AWS_ACCESS_KEY_ID",
"AWS_SECRET_ACCESS_KEY",
"AWS_SESSION_TOKEN",
"AWS_SESSION_NAME",
"AWS_REGION_NAME",
"AWS_REGION",
"AWS_PROFILE_NAME",
"AWS_ROLE_NAME",
"AWS_ROLE_ARN",
"AWS_WEB_IDENTITY_TOKEN",
"AWS_WEB_IDENTITY_TOKEN_FILE",
"AWS_RDS_IAM_IGNORE_WEB_IDENTITY_TOKEN",
)
@dataclass(slots=True)
class _CallRecord:
client_calls: list[tuple[str | None, dict[str, object]]] = field(default_factory=list)
sts_calls: list[tuple[str, dict[str, object]]] = field(default_factory=list)
class _FakeStsClient:
def __init__(self, record: _CallRecord) -> None:
self._record = record
def assume_role(self, **kwargs: object) -> dict[str, object]:
self._record.sts_calls.append(("assume_role", kwargs))
return {
"Credentials": {
"AccessKeyId": "ASIA-ASSUMED",
"SecretAccessKey": "assumed-secret",
"SessionToken": "assumed-session",
}
}
def assume_role_with_web_identity(self, **kwargs: object) -> dict[str, object]:
self._record.sts_calls.append(("assume_role_with_web_identity", kwargs))
return {
"Credentials": {
"AccessKeyId": "ASIA-WEB-IDENTITY",
"SecretAccessKey": "web-identity-secret",
"SessionToken": "web-identity-session",
}
}
class _FakeRdsClient:
def generate_db_auth_token(self, DBHostname: str, Port: int, DBUsername: str) -> str:
return FAKE_DB_AUTH_TOKEN
class _FakeSessionModule(types.ModuleType):
def Config(self, *args: object, **kwargs: object) -> object:
return object()
class _FakeBoto3Module(types.ModuleType):
def __init__(self, record: _CallRecord) -> None:
super().__init__("boto3")
self._record = record
self._sts_client = _FakeStsClient(record)
self.session = _FakeSessionModule("boto3.session")
def client(self, *args: str, **kwargs: object) -> _FakeStsClient | _FakeRdsClient:
service_name = kwargs.get("service_name", args[0] if args else None)
assert isinstance(service_name, str) or service_name is None
self._record.client_calls.append((service_name, kwargs))
if service_name == "sts":
return self._sts_client
return _FakeRdsClient()
def Session(self, *args: object, **kwargs: object) -> _FakeRdsClient:
return _FakeRdsClient()
@pytest.fixture
def aws_env(monkeypatch: pytest.MonkeyPatch) -> pytest.MonkeyPatch:
for var in AWS_ENV_VARS:
monkeypatch.delenv(var, raising=False)
monkeypatch.setenv("AWS_REGION_NAME", "us-east-1")
return monkeypatch
@pytest.fixture
def boto3_stub(monkeypatch: pytest.MonkeyPatch) -> _CallRecord:
record = _CallRecord()
monkeypatch.setitem(sys.modules, "boto3", _FakeBoto3Module(record))
return record
def _set_irsa_pod_env(monkeypatch: pytest.MonkeyPatch, tmp_path: Path) -> str:
token_file = tmp_path / "web_identity_token"
token_file.write_text("oidc-token-contents")
monkeypatch.setenv("AWS_WEB_IDENTITY_TOKEN_FILE", str(token_file))
monkeypatch.setenv("AWS_ROLE_ARN", "arn:pod")
monkeypatch.setenv("AWS_ROLE_NAME", "arn:rds")
monkeypatch.setenv("AWS_SESSION_NAME", "sess")
monkeypatch.setenv("AWS_ACCESS_KEY_ID", "AKIA-EXPLICIT")
monkeypatch.setenv("AWS_SECRET_ACCESS_KEY", "explicit-secret")
monkeypatch.setenv("AWS_SESSION_TOKEN", "explicit-session")
return "oidc-token-contents"
@pytest.mark.parametrize("flag_value", ["true", "1", "TRUE"])
def test_ignore_web_identity_token_uses_assume_role_with_session_token(
aws_env: pytest.MonkeyPatch, boto3_stub: _CallRecord, tmp_path: Path, flag_value: str
) -> None:
_set_irsa_pod_env(aws_env, tmp_path)
aws_env.setenv("AWS_RDS_IAM_IGNORE_WEB_IDENTITY_TOKEN", flag_value)
token = generate_iam_auth_token(db_host="db", db_port=5432, db_user="litellm")
sts_methods = [name for name, _ in boto3_stub.sts_calls]
assert sts_methods == ["assume_role"]
assert boto3_stub.sts_calls[0][1]["RoleArn"] == "arn:rds"
assert boto3_stub.sts_calls[0][1]["RoleSessionName"] == "sess"
sts_client_calls = [kw for svc, kw in boto3_stub.client_calls if svc == "sts"]
assert len(sts_client_calls) == 1
assert sts_client_calls[0]["aws_access_key_id"] == "AKIA-EXPLICIT"
assert sts_client_calls[0]["aws_secret_access_key"] == "explicit-secret"
assert sts_client_calls[0]["aws_session_token"] == "explicit-session"
rds_client_calls = [kw for svc, kw in boto3_stub.client_calls if svc == "rds"]
assert len(rds_client_calls) == 1
assert rds_client_calls[0]["aws_access_key_id"] == "ASIA-ASSUMED"
assert rds_client_calls[0]["aws_secret_access_key"] == "assumed-secret"
assert rds_client_calls[0]["aws_session_token"] == "assumed-session"
assert token == urllib.parse.quote(FAKE_DB_AUTH_TOKEN, safe="")
def test_ignore_web_identity_token_explicit_keys_forward_session_token(
aws_env: pytest.MonkeyPatch, boto3_stub: _CallRecord
) -> None:
aws_env.setenv("AWS_RDS_IAM_IGNORE_WEB_IDENTITY_TOKEN", "true")
aws_env.setenv("AWS_ACCESS_KEY_ID", "AKIA-EXPLICIT")
aws_env.setenv("AWS_SECRET_ACCESS_KEY", "explicit-secret")
aws_env.setenv("AWS_SESSION_TOKEN", "explicit-session")
token = generate_iam_auth_token(db_host="db", db_port=5432, db_user="litellm")
sts_client_calls = [kw for svc, kw in boto3_stub.client_calls if svc == "sts"]
assert sts_client_calls == []
assert boto3_stub.sts_calls == []
rds_client_calls = [kw for svc, kw in boto3_stub.client_calls if svc == "rds"]
assert len(rds_client_calls) == 1
assert rds_client_calls[0]["aws_access_key_id"] == "AKIA-EXPLICIT"
assert rds_client_calls[0]["aws_secret_access_key"] == "explicit-secret"
assert rds_client_calls[0]["aws_session_token"] == "explicit-session"
assert token == urllib.parse.quote(FAKE_DB_AUTH_TOKEN, safe="")
@pytest.mark.parametrize("flag_value", [None, "false"])
def test_web_identity_token_used_when_flag_not_truthy(
aws_env: pytest.MonkeyPatch, boto3_stub: _CallRecord, tmp_path: Path, flag_value: str | None
) -> None:
oidc_contents = _set_irsa_pod_env(aws_env, tmp_path)
if flag_value is not None:
aws_env.setenv("AWS_RDS_IAM_IGNORE_WEB_IDENTITY_TOKEN", flag_value)
token = generate_iam_auth_token(db_host="db", db_port=5432, db_user="litellm")
sts_methods = [name for name, _ in boto3_stub.sts_calls]
assert sts_methods == ["assume_role_with_web_identity"]
assert boto3_stub.sts_calls[0][1]["RoleArn"] == "arn:rds"
assert boto3_stub.sts_calls[0][1]["RoleSessionName"] == "sess"
assert boto3_stub.sts_calls[0][1]["WebIdentityToken"] == oidc_contents
assert token == urllib.parse.quote(FAKE_DB_AUTH_TOKEN, safe="")