fix(proxy): sign RDS IAM tokens for the database's own region (#44782)

* fix(proxy): sign RDS IAM tokens for the database's own region

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

* feat(proxy): add per-connection RDS IAM signing region overrides for writer and reader

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

* test(proxy): integration cells for per-connection RDS IAM signing regions

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

* test(proxy): opt-in real RDS e2e cells for cross-region IAM signing

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

* test(proxy): address review on RDS IAM region cells

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

* test(e2e): drop recording-front integration tests, mark RDS e2e suite e2e and register its model via /model/new

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

* test(e2e): assert the spend read lands on the replica, not just a live connection

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

* fix(proxy): type generate_iam_auth_token params and retry the replica spend-read check

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

* fix(proxy): derive RDS region from custom RDS Proxy endpoint hostnames

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

* test(e2e): add Subject, step labels, frozen rows and Final to the RDS IAM e2e suite

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

* test(e2e): drop the opt-in RDS IAM e2e suite in favor of unit coverage

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

---------

Co-authored-by: yassin <yassin@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
Co-authored-by: mrinal <mrinal@berri.ai>
This commit is contained in:
devin-ai-integration[bot] 2026-10-09 15:41:07 -07:00 • committed by GitHub
parent fe3ace9a02
commit da59afee44
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
9 changed files with 271 additions and 15 deletions

View file

@ -1,8 +1,14 @@
import os import os
import re
from typing import Any, Final from typing import Any, Final
import httpx import httpx
_RDS_HOSTNAME_REGION_PATTERN: Final = re.compile(
r"(?:[^.]+\.)+(?P<region>[a-z]{2}(?:-[a-z]+)+-\d+)\.rds\.amazonaws\.com(?:\.cn)?\.?",
re.IGNORECASE,
)
def init_rds_client( def init_rds_client(
aws_access_key_id: str | None = None, aws_access_key_id: str | None = None,
@ -151,7 +157,21 @@ def init_rds_client(
return client return client
def generate_iam_auth_token(db_host, db_port, db_user, client: Any | None = None) -> str: def rds_region_from_hostname(db_host: str) -> str | None:
match: Final = _RDS_HOSTNAME_REGION_PATTERN.fullmatch(db_host)
if match is None:
return None
return match.group("region").lower()
def generate_iam_auth_token(
db_host: str,
db_port: str,
db_user: str,
client: Any | None = None,
*,
region: str | None = None,
) -> str:
from urllib.parse import quote from urllib.parse import quote
if client is None: if client is None:
@ -167,7 +187,12 @@ def generate_iam_auth_token(db_host, db_port, db_user, client: Any | None = None
else: else:
boto_client = client boto_client = client
token: Final = boto_client.generate_db_auth_token(DBHostname=db_host, Port=db_port, DBUsername=db_user) token: Final = boto_client.generate_db_auth_token(
DBHostname=db_host,
Port=db_port,
DBUsername=db_user,
Region=region or rds_region_from_hostname(db_host),
)
cleaned_token: Final = quote(token, safe="") cleaned_token: Final = quote(token, safe="")
return cleaned_token return cleaned_token

View file

@ -405,13 +405,14 @@ class DatabaseURLSettings(BaseSettings):
"""Load the settings from ``os.environ`` (read at call time).""" """Load the settings from ``os.environ`` (read at call time)."""
return cls() return cls()
def token_auth(self) -> DatabaseTokenAuth | None: def token_auth(self, *, read_replica: bool = False) -> DatabaseTokenAuth | None:
"""The token strategy the toggles ask for, or ``None`` for password auth. """The token strategy the toggles ask for, or ``None`` for password auth.
Raises ``RuntimeError`` when both toggles are on, since the password can only Raises ``RuntimeError`` when both toggles are on, since the password can only
come from one source. come from one source.
""" """
return build_database_token_auth( return build_database_token_auth(
read_replica=read_replica,
iam_token_db_auth=self.iam_token_db_auth, iam_token_db_auth=self.iam_token_db_auth,
azure_postgresql_auth=self.azure_postgresql_auth, azure_postgresql_auth=self.azure_postgresql_auth,
) )
@ -517,7 +518,7 @@ class DatabaseURLSettings(BaseSettings):
schema: Final = self.database_schema_read_replica or self.database_schema schema: Final = self.database_schema_read_replica or self.database_schema
password: Final = self.database_password_read_replica or self.database_password password: Final = self.database_password_read_replica or self.database_password
auth: Final = self.token_auth() auth: Final = self.token_auth(read_replica=True)
if auth is not None: if auth is not None:
missing: Final = tuple( missing: Final = tuple(
env env

View file

@ -191,7 +191,15 @@ class PrismaWrapper:
): ):
# Set before `_original_prisma` so the `iam_token_db_auth` property below can # Set before `_original_prisma` so the `iam_token_db_auth` property below can
# never send `__getattr__` looking for a half-built strategy on the raw client. # never send `__getattr__` looking for a half-built strategy on the raw client.
self._token_auth = token_auth if token_auth is not None else (RdsIamTokenAuth() if iam_token_db_auth else None) self._token_auth = (
token_auth
if token_auth is not None
else (
RdsIamTokenAuth.from_env(read_replica=db_url_env_var == "DATABASE_URL_READ_REPLICA")
if iam_token_db_auth
else None
)
)
self._original_prisma = original_prisma self._original_prisma = original_prisma
# Per-connection knobs so the same wrapper can be used for the writer # Per-connection knobs so the same wrapper can be used for the writer

View file

@ -142,6 +142,13 @@ def parse_iam_endpoint_from_url(url: str) -> IAMEndpoint:
class RdsIamTokenAuth: class RdsIamTokenAuth:
"""AWS RDS IAM auth: a SigV4-presigned token minted from the ambient AWS credentials.""" """AWS RDS IAM auth: a SigV4-presigned token minted from the ambient AWS credentials."""
region: str | None = None
@classmethod
def from_env(cls, *, read_replica: bool = False) -> "RdsIamTokenAuth":
env_var: Final = "AWS_RDS_READ_REPLICA_REGION" if read_replica else "AWS_RDS_REGION"
return cls(region=os.getenv(env_var, "").strip() or None)
@property @property
def label(self) -> str: def label(self) -> str:
return "RDS IAM token" return "RDS IAM token"
@ -179,7 +186,9 @@ def mint_database_token(auth: DatabaseTokenAuth, endpoint: IAMEndpoint) -> str:
case RdsIamTokenAuth(): case RdsIamTokenAuth():
from litellm.proxy.auth.rds_iam_token import generate_iam_auth_token from litellm.proxy.auth.rds_iam_token import generate_iam_auth_token
return generate_iam_auth_token(db_host=endpoint.host, db_port=endpoint.port, db_user=endpoint.user) return generate_iam_auth_token(
db_host=endpoint.host, db_port=endpoint.port, db_user=endpoint.user, region=auth.region
)
case AzureEntraTokenAuth(): case AzureEntraTokenAuth():
return _quote(auth.token_provider()) return _quote(auth.token_provider())
case _: case _:
@ -251,20 +260,23 @@ def build_azure_entra_token_provider() -> Callable[[], str]:
return get_azure_ad_token_provider(azure_scope=AZURE_POSTGRESQL_SCOPE) return get_azure_ad_token_provider(azure_scope=AZURE_POSTGRESQL_SCOPE)
def build_database_token_auth(*, iam_token_db_auth: bool, azure_postgresql_auth: bool) -> DatabaseTokenAuth | None: def build_database_token_auth(
*, iam_token_db_auth: bool, azure_postgresql_auth: bool, read_replica: bool = False
) -> DatabaseTokenAuth | None:
"""Pick the token strategy the two toggles ask for, or None when neither is on.""" """Pick the token strategy the two toggles ask for, or None when neither is on."""
if iam_token_db_auth and azure_postgresql_auth: if iam_token_db_auth and azure_postgresql_auth:
raise RuntimeError(CONFLICTING_TOKEN_AUTH_MESSAGE) raise RuntimeError(CONFLICTING_TOKEN_AUTH_MESSAGE)
if azure_postgresql_auth: if azure_postgresql_auth:
return AzureEntraTokenAuth(token_provider=build_azure_entra_token_provider()) return AzureEntraTokenAuth(token_provider=build_azure_entra_token_provider())
if iam_token_db_auth: if iam_token_db_auth:
return RdsIamTokenAuth() return RdsIamTokenAuth.from_env(read_replica=read_replica)
return None return None
def resolve_database_token_auth() -> DatabaseTokenAuth | None: def resolve_database_token_auth(*, read_replica: bool = False) -> DatabaseTokenAuth | None:
"""Resolve the token strategy from the environment, raising when both toggles are set.""" """Resolve the token strategy from the environment, raising when both toggles are set."""
return build_database_token_auth( return build_database_token_auth(
read_replica=read_replica,
iam_token_db_auth=token_auth_flag_enabled( iam_token_db_auth=token_auth_flag_enabled(
os.getenv(IAM_TOKEN_DB_AUTH_ENV_VAR), env_var=IAM_TOKEN_DB_AUTH_ENV_VAR os.getenv(IAM_TOKEN_DB_AUTH_ENV_VAR), env_var=IAM_TOKEN_DB_AUTH_ENV_VAR
), ),

View file

@ -4567,6 +4567,7 @@ class PrismaClient:
self.db: PrismaWrapper | RoutingPrismaWrapper self.db: PrismaWrapper | RoutingPrismaWrapper
if read_replica_url: if read_replica_url:
try: try:
reader_token_auth: Final = resolve_database_token_auth(read_replica=True)
# If token auth is enabled, the reader refreshes its own token on # If token auth is enabled, the reader refreshes its own token on
# the same cadence as the writer. We parse the static endpoint # the same cadence as the writer. We parse the static endpoint
# pieces (host/port/user/db) once from the reader URL — only # pieces (host/port/user/db) once from the reader URL — only
@ -4581,8 +4582,8 @@ class PrismaClient:
# and the first query falls through to the synchronous fallback # and the first query falls through to the synchronous fallback
# path in `PrismaWrapper.__getattr__`, which deadlocks the event # path in `PrismaWrapper.__getattr__`, which deadlocks the event
# loop and times out after 30s. # loop and times out after 30s.
if token_auth is not None and reader_iam_endpoint is not None: if reader_token_auth is not None and reader_iam_endpoint is not None:
reader_token: Final = mint_database_token(token_auth, reader_iam_endpoint) reader_token: Final = mint_database_token(reader_token_auth, reader_iam_endpoint)
read_replica_url = add_missing_query_params( read_replica_url = add_missing_query_params(
reader_iam_endpoint.build_url(reader_token), reader_iam_endpoint.build_url(reader_token),
token_refresh_params_from_url(read_replica_url), token_refresh_params_from_url(read_replica_url),
@ -4595,7 +4596,7 @@ class PrismaClient:
reader_prisma = Prisma(datasource=reader_datasource) reader_prisma = Prisma(datasource=reader_datasource)
reader_wrapper: Final = PrismaWrapper( reader_wrapper: Final = PrismaWrapper(
original_prisma=reader_prisma, original_prisma=reader_prisma,
token_auth=token_auth, token_auth=reader_token_auth,
db_url_env_var="DATABASE_URL_READ_REPLICA", db_url_env_var="DATABASE_URL_READ_REPLICA",
iam_endpoint=reader_iam_endpoint, iam_endpoint=reader_iam_endpoint,
recreate_uses_datasource=True, recreate_uses_datasource=True,

View file

@ -0,0 +1,53 @@
from typing import Final
from urllib.parse import parse_qs, unquote, urlsplit
import boto3
import pytest
from litellm.proxy.auth.rds_iam_token import generate_iam_auth_token
@pytest.mark.parametrize(
("db_host", "expected_region"),
[
("w.abc123.us-east-1.rds.amazonaws.com", "us-east-1"),
("lit-r.abc123xyz.ap-northeast-1.rds.amazonaws.com", "ap-northeast-1"),
("c.cluster-abc123.eu-central-1.rds.amazonaws.com", "eu-central-1"),
("c.cluster-ro-abc.eu-west-2.rds.amazonaws.com", "eu-west-2"),
("p.proxy-abc123.us-west-2.rds.amazonaws.com", "us-west-2"),
("ep1.endpoint.proxy-ab0cd1efghij.us-east-2.rds.amazonaws.com", "us-east-2"),
("d.abc123.cn-north-1.rds.amazonaws.com.cn", "cn-north-1"),
("d.abc123.us-gov-west-1.rds.amazonaws.com", "us-gov-west-1"),
("W.ABC123.US-EAST-1.RDS.AMAZONAWS.COM", "us-east-1"),
("lit.abc.us-east-1.rds.amazonaws.com.", "us-east-1"),
("writer.aurora.local", "ap-northeast-1"),
("db.internal.example.com", "ap-northeast-1"),
("10.0.0.5", "ap-northeast-1"),
("localhost", "ap-northeast-1"),
("us-east-1.rds.amazonaws.com.evil.example", "ap-northeast-1"),
],
)
def test_generate_iam_auth_token_signs_for_the_database_hostname_region(
db_host: str, expected_region: str, monkeypatch: pytest.MonkeyPatch
) -> None:
monkeypatch.delenv("AWS_PROFILE", raising=False)
monkeypatch.delenv("AWS_DEFAULT_PROFILE", raising=False)
client: Final = boto3.client(
"rds",
region_name="ap-northeast-1",
aws_access_key_id="AKIDEXAMPLE",
aws_secret_access_key="x" * 40,
)
token: Final = unquote(
generate_iam_auth_token(
db_host=db_host,
db_port="5432",
db_user="litellm",
client=client,
)
)
credential: Final = parse_qs(urlsplit(token).query)["X-Amz-Credential"][0]
assert token.split("?", maxsplit=1)[0] == f"{db_host}:5432/"
assert credential.split("/")[2] == expected_region

View file

@ -194,6 +194,71 @@ def test_reader_url_assembled_when_host_set_and_url_unset(monkeypatch):
) )
@pytest.mark.parametrize(
("writer_override", "reader_override", "writer_region", "reader_region"),
[
(None, None, "us-east-1", "ap-northeast-1"),
("eu-west-1", "us-west-2", "eu-west-1", "us-west-2"),
("eu-west-1", None, "eu-west-1", "ap-northeast-1"),
(None, "us-west-2", "us-east-1", "us-west-2"),
(" ", "", "us-east-1", "ap-northeast-1"),
(" eu-west-1 ", " us-west-2 ", "eu-west-1", "us-west-2"),
],
)
def test_writer_and_reader_urls_sign_in_their_endpoint_regions(
monkeypatch: pytest.MonkeyPatch,
writer_override: str | None,
reader_override: str | None,
writer_region: str,
reader_region: str,
) -> None:
for key, value in (("AWS_RDS_REGION", writer_override), ("AWS_RDS_READ_REPLICA_REGION", reader_override)):
if value is None:
monkeypatch.delenv(key, raising=False)
else:
monkeypatch.setenv(key, value)
monkeypatch.setenv("IAM_TOKEN_DB_AUTH", "true")
monkeypatch.setenv("AWS_REGION", "ap-northeast-1")
monkeypatch.setenv("AWS_ACCESS_KEY_ID", "AKIDEXAMPLE")
monkeypatch.setenv("AWS_SECRET_ACCESS_KEY", "secret")
monkeypatch.delenv("AWS_PROFILE", raising=False)
monkeypatch.delenv("AWS_DEFAULT_PROFILE", raising=False)
monkeypatch.delenv("AWS_REGION_NAME", raising=False)
monkeypatch.delenv("AWS_PROFILE_NAME", raising=False)
monkeypatch.delenv("AWS_ROLE_NAME", raising=False)
monkeypatch.delenv("AWS_ROLE_ARN", raising=False)
monkeypatch.delenv("AWS_SESSION_NAME", raising=False)
monkeypatch.delenv("AWS_WEB_IDENTITY_TOKEN", raising=False)
monkeypatch.delenv("AWS_WEB_IDENTITY_TOKEN_FILE", raising=False)
monkeypatch.delenv("AWS_SESSION_TOKEN", raising=False)
monkeypatch.setenv("DATABASE_HOST", "writer.abc123.us-east-1.rds.amazonaws.com")
monkeypatch.setenv("DATABASE_HOST_READ_REPLICA", "reader.abc123.ap-northeast-1.rds.amazonaws.com")
monkeypatch.setenv("DATABASE_USER", "litellm_rds")
monkeypatch.setenv("DATABASE_NAME", "litellm_db")
monkeypatch.delenv("DATABASE_URL", raising=False)
monkeypatch.delenv("DATABASE_URL_READ_REPLICA", raising=False)
settings: Final = DatabaseURLSettings.from_env()
writer_url: Final = settings.build_writer_url()
reader_url: Final = settings.build_reader_url()
assert writer_url is not None
assert reader_url is not None
writer_password: Final = urllib.parse.urlsplit(writer_url).password
reader_password: Final = urllib.parse.urlsplit(reader_url).password
assert writer_password is not None
assert reader_password is not None
writer_token: Final = urllib.parse.unquote(writer_password)
reader_token: Final = urllib.parse.unquote(reader_password)
writer_query: Final = urllib.parse.urlsplit(writer_token).query
reader_query: Final = urllib.parse.urlsplit(reader_token).query
writer_credential: Final = urllib.parse.parse_qs(writer_query)["X-Amz-Credential"][0]
reader_credential: Final = urllib.parse.parse_qs(reader_query)["X-Amz-Credential"][0]
assert writer_credential.split("/")[2] == writer_region
assert reader_credential.split("/")[2] == reader_region
def test_reader_url_not_clobbered_when_already_set(monkeypatch): def test_reader_url_not_clobbered_when_already_set(monkeypatch):
"""If the operator pinned DATABASE_URL_READ_REPLICA (e.g. a non-IAM """If the operator pinned DATABASE_URL_READ_REPLICA (e.g. a non-IAM
reader), the model must leave it untouched even though reader), the model must leave it untouched even though

View file

@ -697,7 +697,7 @@ def test_writer_get_rds_iam_token_defaults_port_when_unset(monkeypatch, unset_da
captured: Dict[str, Any] = {} captured: Dict[str, Any] = {}
def fake_generate(db_host=None, db_port=None, db_user=None): def fake_generate(db_host=None, db_port=None, db_user=None, *, region=None):
captured["port"] = db_port captured["port"] = db_port
return "TOKEN" return "TOKEN"
@ -730,7 +730,7 @@ def test_writer_get_rds_iam_token_uses_database_host_env_vars(monkeypatch, unset
captured: Dict[str, Any] = {} captured: Dict[str, Any] = {}
def fake_generate(db_host=None, db_port=None, db_user=None): def fake_generate(db_host=None, db_port=None, db_user=None, *, region=None):
captured["host"] = db_host captured["host"] = db_host
captured["port"] = db_port captured["port"] = db_port
captured["user"] = db_user captured["user"] = db_user
@ -770,7 +770,7 @@ def test_reader_iam_refresh_uses_parsed_endpoint(monkeypatch):
captured: Dict[str, Any] = {} captured: Dict[str, Any] = {}
def fake_generate(db_host=None, db_port=None, db_user=None): def fake_generate(db_host=None, db_port=None, db_user=None, *, region=None):
captured["host"] = db_host captured["host"] = db_host
captured["port"] = db_port captured["port"] = db_port
captured["user"] = db_user captured["user"] = db_user
@ -1136,3 +1136,61 @@ def test_prisma_client_premints_an_entra_token_for_the_reader(monkeypatch):
assert isinstance(client.db._reader.token_auth, AzureEntraTokenAuth) assert isinstance(client.db._reader.token_auth, AzureEntraTokenAuth)
assert isinstance(client.db._writer.token_auth, AzureEntraTokenAuth) assert isinstance(client.db._writer.token_auth, AzureEntraTokenAuth)
assert isinstance(client.db._writer, PrismaWrapper) assert isinstance(client.db._writer, PrismaWrapper)
@pytest.mark.parametrize("opaque_hosts", [False, True])
@pytest.mark.parametrize(
("writer_override", "reader_override"),
[(None, None), ("eu-west-1", None), (None, "us-west-2"), (" eu-west-1 ", "us-west-2"), ("", " ")],
)
def test_rds_regions_survive_reader_url_startup_and_refresh(
monkeypatch: pytest.MonkeyPatch,
opaque_hosts: bool,
writer_override: str | None,
reader_override: str | None,
) -> None:
from urllib.parse import unquote
from litellm.proxy.db.db_url_settings import DatabaseURLSettings
from litellm.proxy.db.routing_prisma_wrapper import RoutingPrismaWrapper
from litellm.proxy.utils import PrismaClient
for key in tuple(os.environ):
if key.startswith(("AWS_", "DATABASE_", "AZURE_POSTGRESQL_", "LITELLM_PGBOUNCER")):
monkeypatch.delenv(key)
monkeypatch.setenv("AWS_ACCESS_KEY_ID", "AKIDEXAMPLE")
monkeypatch.setenv("AWS_SECRET_ACCESS_KEY", "test-secret")
monkeypatch.setenv("AWS_REGION", "ap-south-1")
monkeypatch.setenv("IAM_TOKEN_DB_AUTH", "true")
for key, value in (("AWS_RDS_REGION", writer_override), ("AWS_RDS_READ_REPLICA_REGION", reader_override)):
if value is not None:
monkeypatch.setenv(key, value)
writer_host: Final = "writer.internal" if opaque_hosts else "w.abc.us-east-1.rds.amazonaws.com"
reader_host: Final = "reader.internal" if opaque_hosts else "r.abc.ap-northeast-1.rds.amazonaws.com"
monkeypatch.setenv("DATABASE_HOST", writer_host)
monkeypatch.setenv("DATABASE_USER", "iam_user")
monkeypatch.setenv("DATABASE_NAME", "litellm")
monkeypatch.setenv("DATABASE_URL", "")
monkeypatch.setenv("DATABASE_URL_READ_REPLICA", f"postgresql://iam_user@{reader_host}:5432/litellm?sslmode=require")
monkeypatch.setitem(sys.modules, "prisma", MagicMock())
writer_url: Final = DatabaseURLSettings.from_env().build_writer_url()
assert writer_url is not None
monkeypatch.setenv("DATABASE_URL", writer_url)
client: Final = PrismaClient(database_url=writer_url, proxy_logging_obj=MagicMock())
assert isinstance(client.db, RoutingPrismaWrapper)
writer_region: Final = (writer_override or "").strip() or ("ap-south-1" if opaque_hosts else "us-east-1")
reader_region: Final = (reader_override or "").strip() or ("ap-south-1" if opaque_hosts else "ap-northeast-1")
for url, host, region in (
(writer_url, writer_host, writer_region),
(os.environ["DATABASE_URL_READ_REPLICA"], reader_host, reader_region),
(client.db.writer.get_rds_iam_token(), writer_host, writer_region),
(client.db.reader.get_rds_iam_token(), reader_host, reader_region),
):
assert url is not None
password: Final = urlsplit(url).password
assert password is not None
token: Final = unquote(password)
assert token.split("?", 1)[0] == f"{host}:5432/"
assert parse_qs(urlsplit(token).query)["X-Amz-Credential"][0].split("/")[2] == region
assert parse_qs(urlsplit(os.environ["DATABASE_URL_READ_REPLICA"]).query)["sslmode"] == ["require"]
assert os.environ["AWS_REGION"] == "ap-south-1"

View file

@ -11,6 +11,8 @@ a connection URL.
import base64 import base64
import json import json
from datetime import datetime, timezone from datetime import datetime, timezone
from typing import Final
from urllib.parse import parse_qs, unquote, urlsplit
from unittest.mock import patch from unittest.mock import patch
import pytest import pytest
@ -68,9 +70,40 @@ def test_rds_mint_delegates_to_the_sigv4_token_generator():
db_host="writer.aurora.local", db_host="writer.aurora.local",
db_port="5432", db_port="5432",
db_user="litellm_rds", db_user="litellm_rds",
region=None,
) )
def test_rds_mint_signs_each_endpoint_in_its_hostname_region(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setenv("AWS_REGION", "ap-northeast-1")
monkeypatch.setenv("AWS_ACCESS_KEY_ID", "AKIDEXAMPLE")
monkeypatch.setenv("AWS_SECRET_ACCESS_KEY", "secret")
monkeypatch.delenv("AWS_PROFILE", raising=False)
monkeypatch.delenv("AWS_DEFAULT_PROFILE", raising=False)
monkeypatch.delenv("AWS_REGION_NAME", raising=False)
monkeypatch.delenv("AWS_SESSION_NAME", raising=False)
monkeypatch.delenv("AWS_PROFILE_NAME", raising=False)
monkeypatch.delenv("AWS_ROLE_NAME", raising=False)
monkeypatch.delenv("AWS_ROLE_ARN", raising=False)
monkeypatch.delenv("AWS_WEB_IDENTITY_TOKEN", raising=False)
monkeypatch.delenv("AWS_WEB_IDENTITY_TOKEN_FILE", raising=False)
monkeypatch.delenv("AWS_SESSION_TOKEN", raising=False)
writer_token: Final = mint_database_token(
RdsIamTokenAuth(),
_endpoint(host="writer.abc123.us-east-1.rds.amazonaws.com", user="litellm_rds"),
)
reader_token: Final = mint_database_token(
RdsIamTokenAuth(),
_endpoint(host="reader.abc123.ap-northeast-1.rds.amazonaws.com", user="litellm_rds"),
)
writer_credential: Final = parse_qs(urlsplit(unquote(writer_token)).query)["X-Amz-Credential"][0]
reader_credential: Final = parse_qs(urlsplit(unquote(reader_token)).query)["X-Amz-Credential"][0]
assert writer_credential.split("/")[2] == "us-east-1"
assert reader_credential.split("/")[2] == "ap-northeast-1"
def test_entra_mint_calls_the_injected_provider_and_encodes_the_token(): def test_entra_mint_calls_the_injected_provider_and_encodes_the_token():
"""A real compact JWT is already URL-safe, but the provider is an Azure SDK call """A real compact JWT is already URL-safe, but the provider is an Azure SDK call
whose output we do not control, and an unencoded ``/`` or ``=`` in a password whose output we do not control, and an unencoded ``/`` or ``=`` in a password