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 re
from typing import Any, Final
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(
aws_access_key_id: str | None = None,
@ -151,7 +157,21 @@ def init_rds_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
if client is None:
@ -167,7 +187,12 @@ def generate_iam_auth_token(db_host, db_port, db_user, client: Any | None = None
else:
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="")
return cleaned_token

View file

@ -405,13 +405,14 @@ class DatabaseURLSettings(BaseSettings):
"""Load the settings from ``os.environ`` (read at call time)."""
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.
Raises ``RuntimeError`` when both toggles are on, since the password can only
come from one source.
"""
return build_database_token_auth(
read_replica=read_replica,
iam_token_db_auth=self.iam_token_db_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
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:
missing: Final = tuple(
env

View file

@ -191,7 +191,15 @@ class PrismaWrapper:
):
# 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.
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
# 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:
"""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
def label(self) -> str:
return "RDS IAM token"
@ -179,7 +186,9 @@ def mint_database_token(auth: DatabaseTokenAuth, endpoint: IAMEndpoint) -> str:
case RdsIamTokenAuth():
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():
return _quote(auth.token_provider())
case _:
@ -251,20 +260,23 @@ def build_azure_entra_token_provider() -> Callable[[], str]:
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."""
if iam_token_db_auth and azure_postgresql_auth:
raise RuntimeError(CONFLICTING_TOKEN_AUTH_MESSAGE)
if azure_postgresql_auth:
return AzureEntraTokenAuth(token_provider=build_azure_entra_token_provider())
if iam_token_db_auth:
return RdsIamTokenAuth()
return RdsIamTokenAuth.from_env(read_replica=read_replica)
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."""
return build_database_token_auth(
read_replica=read_replica,
iam_token_db_auth=token_auth_flag_enabled(
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
if read_replica_url:
try:
reader_token_auth: Final = resolve_database_token_auth(read_replica=True)
# If token auth is enabled, the reader refreshes its own token on
# the same cadence as the writer. We parse the static endpoint
# 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
# path in `PrismaWrapper.__getattr__`, which deadlocks the event
# loop and times out after 30s.
if token_auth is not None and reader_iam_endpoint is not None:
reader_token: Final = mint_database_token(token_auth, reader_iam_endpoint)
if reader_token_auth is not None and reader_iam_endpoint is not None:
reader_token: Final = mint_database_token(reader_token_auth, reader_iam_endpoint)
read_replica_url = add_missing_query_params(
reader_iam_endpoint.build_url(reader_token),
token_refresh_params_from_url(read_replica_url),
@ -4595,7 +4596,7 @@ class PrismaClient:
reader_prisma = Prisma(datasource=reader_datasource)
reader_wrapper: Final = PrismaWrapper(
original_prisma=reader_prisma,
token_auth=token_auth,
token_auth=reader_token_auth,
db_url_env_var="DATABASE_URL_READ_REPLICA",
iam_endpoint=reader_iam_endpoint,
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):
"""If the operator pinned DATABASE_URL_READ_REPLICA (e.g. a non-IAM
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] = {}
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
return "TOKEN"
@ -730,7 +730,7 @@ def test_writer_get_rds_iam_token_uses_database_host_env_vars(monkeypatch, unset
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["port"] = db_port
captured["user"] = db_user
@ -770,7 +770,7 @@ def test_reader_iam_refresh_uses_parsed_endpoint(monkeypatch):
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["port"] = db_port
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._writer.token_auth, AzureEntraTokenAuth)
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 json
from datetime import datetime, timezone
from typing import Final
from urllib.parse import parse_qs, unquote, urlsplit
from unittest.mock import patch
import pytest
@ -68,9 +70,40 @@ def test_rds_mint_delegates_to_the_sigv4_token_generator():
db_host="writer.aurora.local",
db_port="5432",
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():
"""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