mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
feat(proxy): add AWS_RDS_IAM_IGNORE_WEB_IDENTITY_TOKEN to skip injected web identity token for RDS IAM auth
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
1a4a9c5ab3
commit
a3220c25e6
2 changed files with 207 additions and 1 deletions
|
|
@ -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
|
||||
|
|
|
|||
188
tests/test_litellm/proxy/auth/test_rds_iam_token.py
Normal file
188
tests/test_litellm/proxy/auth/test_rds_iam_token.py
Normal 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="")
|
||||
Loading…
Add table
Reference in a new issue