From de66f7ddb52ac97e4033ef4718815304bb2bcb85 Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Thu, 16 Jul 2026 15:55:03 +0000 Subject: [PATCH] fix(proxy): reuse Azure database credentials Co-Authored-By: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com> --- litellm/proxy/db/prisma_client.py | 20 +++++++++++++++++-- litellm/proxy/utils.py | 16 ++++++++++++++- .../proxy/db/test_rds_iam_token_expiry.py | 7 ++++++- .../proxy/db/test_routing_prisma_wrapper.py | 17 +++++++++++++--- 4 files changed, 53 insertions(+), 7 deletions(-) diff --git a/litellm/proxy/db/prisma_client.py b/litellm/proxy/db/prisma_client.py index 4d67f348c12..ad78ec30093 100644 --- a/litellm/proxy/db/prisma_client.py +++ b/litellm/proxy/db/prisma_client.py @@ -20,6 +20,10 @@ from typing import Any, Callable, Union from pydantic import BaseModel, ValidationError from litellm._logging import verbose_proxy_logger +from litellm.secret_managers.get_azure_ad_token_provider import ( + AzureTokenCredential, + build_azure_identity_credential, +) from litellm.secret_managers.main import str_to_bool @@ -101,11 +105,12 @@ def parse_iam_endpoint_from_url(url: str) -> IAMEndpoint: def build_database_token_auth_url( endpoint: IAMEndpoint, database_token_auth: DatabaseTokenAuth, + azure_credential: AzureTokenCredential | None = None, ) -> str: if database_token_auth == DatabaseTokenAuth.AZURE_ENTRA: from litellm.proxy.auth.azure_postgres_token import generate_azure_postgres_auth_token - token = generate_azure_postgres_auth_token() + token = generate_azure_postgres_auth_token(credential=azure_credential) else: from litellm.proxy.auth.rds_iam_token import generate_iam_auth_token @@ -147,10 +152,16 @@ class PrismaWrapper: recreate_uses_datasource: bool = False, log_prefix: str = "", database_token_auth: DatabaseTokenAuth | None = None, + azure_credential: AzureTokenCredential | None = None, ): self._original_prisma = original_prisma self.database_token_auth = database_token_auth or (DatabaseTokenAuth.RDS_IAM if iam_token_db_auth else None) self.iam_token_db_auth = self.database_token_auth is not None + self._azure_credential = ( + azure_credential or build_azure_identity_credential() + if self.database_token_auth == DatabaseTokenAuth.AZURE_ENTRA + else None + ) # Per-connection knobs so the same wrapper can be used for the writer # (defaults: DATABASE_URL env, IAM endpoint from DATABASE_HOST/etc., @@ -289,6 +300,7 @@ class PrismaWrapper: return token_created + timedelta(seconds=int(expires_str)) except ValueError as exc: verbose_proxy_logger.debug("Failed to parse RDS IAM token expiration: %s", exc) + return None parts = token.split(".") if len(parts) < 2: @@ -375,7 +387,11 @@ class PrismaWrapper: if self.database_token_auth is None: return None - db_url = build_database_token_auth_url(endpoint, self.database_token_auth) + db_url = build_database_token_auth_url( + endpoint, + self.database_token_auth, + azure_credential=self._azure_credential, + ) os.environ[self._db_url_env_var] = db_url return db_url diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 52fb53b40fb..e10a8dcb045 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -127,12 +127,16 @@ from litellm.proxy.db.exception_handler import ( ) from litellm.proxy.db.log_db_metrics import log_db_metrics from litellm.proxy.db.prisma_client import ( + DatabaseTokenAuth, PrismaWrapper, build_database_token_auth_url, parse_iam_endpoint_from_url, resolve_database_token_auth, ) from litellm.proxy.db.routing_prisma_wrapper import RoutingPrismaWrapper +from litellm.secret_managers.get_azure_ad_token_provider import ( + build_azure_identity_credential, +) from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import ( UnifiedLLMGuardrails, ) @@ -2806,6 +2810,9 @@ class PrismaClient: iam_token_db_auth=iam_flag, azure_postgresql_auth=azure_postgresql_auth, ) + azure_credential = ( + build_azure_identity_credential() if database_token_auth == DatabaseTokenAuth.AZURE_ENTRA else None + ) token_auth_enabled = database_token_auth is not None # When read-replica routing is on, tag log lines with [writer]/[reader] # so the two wrappers' interleaved IAM refresh logs can be told apart. @@ -2818,6 +2825,7 @@ class PrismaClient: iam_token_db_auth=token_auth_enabled, log_prefix=writer_log_prefix, database_token_auth=database_token_auth, + azure_credential=azure_credential, ) else: writer_wrapper = PrismaWrapper( @@ -2825,6 +2833,7 @@ class PrismaClient: iam_token_db_auth=token_auth_enabled, log_prefix=writer_log_prefix, database_token_auth=database_token_auth, + azure_credential=azure_credential, ) # Optional read-replica routing. When DATABASE_URL_READ_REPLICA is set, @@ -2849,7 +2858,11 @@ class PrismaClient: # `PrismaWrapper.__getattr__`, which deadlocks the event loop # and times out after 30s. if database_token_auth is not None and reader_iam_endpoint is not None: - read_replica_url = build_database_token_auth_url(reader_iam_endpoint, database_token_auth) + read_replica_url = build_database_token_auth_url( + reader_iam_endpoint, + database_token_auth, + azure_credential=azure_credential, + ) os.environ["DATABASE_URL_READ_REPLICA"] = read_replica_url reader_kwargs: Dict[str, Any] = {"datasource": {"url": read_replica_url}} if http_client is not None: @@ -2864,6 +2877,7 @@ class PrismaClient: recreate_uses_datasource=True, log_prefix="[reader]", database_token_auth=database_token_auth, + azure_credential=azure_credential, ) self.db = RoutingPrismaWrapper(writer=writer_wrapper, reader=reader_wrapper) verbose_proxy_logger.info( diff --git a/tests/test_litellm/proxy/db/test_rds_iam_token_expiry.py b/tests/test_litellm/proxy/db/test_rds_iam_token_expiry.py index 1b81d8490d8..4c7ac5c6fb9 100644 --- a/tests/test_litellm/proxy/db/test_rds_iam_token_expiry.py +++ b/tests/test_litellm/proxy/db/test_rds_iam_token_expiry.py @@ -152,6 +152,8 @@ class TestTokenExpirationParsing: assert wrapper._parse_token_expiration(None) is None assert wrapper._parse_token_expiration("no-query-params") is None assert wrapper._parse_token_expiration("?missing=params") is None + payload = base64.urlsafe_b64encode(json.dumps({"exp": 1893456000}).encode()).decode().rstrip("=") + assert wrapper._parse_token_expiration(f"header.{payload}.signature?missing=params") is None def test_parse_azure_postgres_jwt_expiration(self): from litellm.proxy.db.prisma_client import PrismaWrapper @@ -166,6 +168,7 @@ class TestTokenExpirationParsing: def test_azure_postgres_refresh_builds_database_url(self, unset_database_url): from litellm.proxy.db.prisma_client import DatabaseTokenAuth, IAMEndpoint, PrismaWrapper + credential = MagicMock() wrapper = PrismaWrapper( original_prisma=MagicMock(), iam_token_db_auth=True, @@ -176,14 +179,16 @@ class TestTokenExpirationParsing: name="litellm db", ), database_token_auth=DatabaseTokenAuth.AZURE_ENTRA, + azure_credential=credential, ) with patch( "litellm.proxy.auth.azure_postgres_token.generate_azure_postgres_auth_token", return_value="AZURE_TOKEN", - ): + ) as generate_token: database_url = wrapper.get_rds_iam_token() + generate_token.assert_called_once_with(credential=credential) assert ( database_url == "postgresql://user%40example.com:AZURE_TOKEN@server.postgres.database.azure.com:5432/litellm%20db" diff --git a/tests/test_litellm/proxy/db/test_routing_prisma_wrapper.py b/tests/test_litellm/proxy/db/test_routing_prisma_wrapper.py index 8f0fbefb3dd..12772729f8b 100644 --- a/tests/test_litellm/proxy/db/test_routing_prisma_wrapper.py +++ b/tests/test_litellm/proxy/db/test_routing_prisma_wrapper.py @@ -871,9 +871,16 @@ def test_prisma_client_configures_azure_token_refresh_for_writer_and_reader(monk fake_prisma_module.Prisma = FakePrisma monkeypatch.setitem(sys.modules, "prisma", fake_prisma_module) - with patch( - "litellm.proxy.auth.azure_postgres_token.generate_azure_postgres_auth_token", - return_value="AZURE_TOKEN", + credential = MagicMock() + with ( + patch( + "litellm.proxy.utils.build_azure_identity_credential", + return_value=credential, + ) as build_credential, + patch( + "litellm.proxy.auth.azure_postgres_token.generate_azure_postgres_auth_token", + return_value="AZURE_TOKEN", + ) as generate_token, ): from litellm.proxy.utils import PrismaClient @@ -883,6 +890,10 @@ def test_prisma_client_configures_azure_token_refresh_for_writer_and_reader(monk ) assert isinstance(client.db, RoutingPrismaWrapper) + build_credential.assert_called_once_with() + generate_token.assert_called_once_with(credential=credential) + assert client.db.writer._azure_credential is credential + assert client.db.reader._azure_credential is credential assert client.db.writer.database_token_auth == DatabaseTokenAuth.AZURE_ENTRA assert client.db.reader.database_token_auth == DatabaseTokenAuth.AZURE_ENTRA assert (