mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
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:
parent
fe3ace9a02
commit
da59afee44
9 changed files with 271 additions and 15 deletions
|
|
@ -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
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
),
|
),
|
||||||
|
|
|
||||||
|
|
@ -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,
|
||||||
|
|
|
||||||
53
tests/unit/proxy/auth/test_rds_iam_token.py
Normal file
53
tests/unit/proxy/auth/test_rds_iam_token.py
Normal 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
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
|
|
@ -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"
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue