fix(tests): stop DATABASE_URL env pollution from read-replica tests breaking DB e2e tests (#32653)

This commit is contained in:
Mateo Wang 2026-07-09 14:37:49 -07:00 • committed by GitHub
parent 41e9cc491e
commit 1fa200123f
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 89 additions and 52 deletions

View file

@ -0,0 +1,65 @@
import os
from collections.abc import Generator
from typing import Optional
import pytest
DB_ENV_KEYS = (
"IAM_TOKEN_DB_AUTH",
"DATABASE_URL",
"DIRECT_URL",
"DATABASE_URL_READ_REPLICA",
"DATABASE_HOST",
"DATABASE_PORT",
"DATABASE_USER",
"DATABASE_USERNAME",
"DATABASE_NAME",
"DATABASE_SCHEMA",
"DATABASE_PASSWORD",
"DATABASE_HOST_READ_REPLICA",
"DATABASE_PORT_READ_REPLICA",
"DATABASE_USER_READ_REPLICA",
"DATABASE_USERNAME_READ_REPLICA",
"DATABASE_NAME_READ_REPLICA",
"DATABASE_SCHEMA_READ_REPLICA",
"DATABASE_PASSWORD_READ_REPLICA",
)
_db_env_snapshot_key = pytest.StashKey[dict[str, Optional[str]]]()
def _db_env_snapshot() -> dict[str, Optional[str]]:
return {key: os.environ.get(key) for key in DB_ENV_KEYS}
@pytest.hookimpl(wrapper=True)
def pytest_runtest_setup(item: pytest.Item) -> Generator[None, None, None]:
item.stash[_db_env_snapshot_key] = _db_env_snapshot()
return (yield)
@pytest.hookimpl(wrapper=True)
def pytest_runtest_teardown(item: pytest.Item, nextitem: Optional[pytest.Item]) -> Generator[None, None, None]:
result = yield
before = item.stash[_db_env_snapshot_key]
leaked = {key: value for key, value in _db_env_snapshot().items() if value != before[key]}
for key, original in before.items():
if original is None:
os.environ.pop(key, None)
else:
os.environ[key] = original
assert not leaked, (
f"{item.nodeid} leaked DB env vars past monkeypatch teardown: {leaked}. "
"Product code under test writes DATABASE_URL(_READ_REPLICA) into os.environ as a side effect; "
"monkeypatch only restores keys it has a record for, so a value written to a previously unset "
"key survives the test and poisons every later test in this pytest-xdist worker process "
"(DB-backed e2e tests arm themselves on DATABASE_URL and then fail to connect). "
"Use the unset_database_url fixture (or monkeypatch.setenv) so restoration is registered."
)
return result
@pytest.fixture
def unset_database_url(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setenv("DATABASE_URL", "about-to-be-unset")
monkeypatch.delenv("DATABASE_URL")

View file

@ -51,25 +51,21 @@ _MANAGED_DB_ENV_VARS = (
@pytest.fixture(autouse=True)
def _scrub_db_env():
def _scrub_db_env(monkeypatch):
"""Start each test from a clean slate and restore the original env afterward.
``apply_to_env`` writes ``DATABASE_URL`` straight into ``os.environ``, which
``monkeypatch`` cannot undo. Snapshotting and restoring here keeps a
synthesized URL (e.g. ``writer.example.com``) from leaking into later tests
that read ``DATABASE_URL`` to decide whether to hit a real database.
``apply_to_env`` writes ``DATABASE_URL`` straight into ``os.environ``.
Registering a setenv+delenv pair per var gives ``monkeypatch`` a restore
record even for previously unset keys, so a synthesized URL (e.g.
``writer.example.com``) cannot leak into later tests that read
``DATABASE_URL`` to decide whether to hit a real database. Restoring via
the same ``monkeypatch`` instance the tests use also keeps undo ordering
consistent (a hand-rolled snapshot/restore runs before ``monkeypatch``'s
own undo and gets clobbered by it).
"""
saved = {var: os.environ.get(var) for var in _MANAGED_DB_ENV_VARS}
for var in _MANAGED_DB_ENV_VARS:
os.environ.pop(var, None)
try:
yield
finally:
for var, value in saved.items():
if value is None:
os.environ.pop(var, None)
else:
os.environ[var] = value
monkeypatch.setenv(var, "scrubbed")
monkeypatch.delenv(var)
def _stub_iam_token(token: str = "FAKE_TOKEN"):

View file

@ -26,25 +26,13 @@ class TestPrismaWrapperTokenRefresh:
"""Tests for the PrismaWrapper RDS IAM token refresh implementation."""
@pytest.fixture
def setup_env(self):
def setup_env(self, monkeypatch, unset_database_url):
"""Setup environment variables for testing."""
os.environ["DATABASE_HOST"] = "test-host.rds.amazonaws.com"
os.environ["DATABASE_PORT"] = "5432"
os.environ["DATABASE_USER"] = "test_user"
os.environ["DATABASE_NAME"] = "test_db"
os.environ["IAM_TOKEN_DB_AUTH"] = "True"
yield
# Cleanup
for key in [
"DATABASE_HOST",
"DATABASE_PORT",
"DATABASE_USER",
"DATABASE_NAME",
"DATABASE_URL",
"IAM_TOKEN_DB_AUTH",
"DATABASE_SCHEMA",
]:
os.environ.pop(key, None)
monkeypatch.setenv("DATABASE_HOST", "test-host.rds.amazonaws.com")
monkeypatch.setenv("DATABASE_PORT", "5432")
monkeypatch.setenv("DATABASE_USER", "test_user")
monkeypatch.setenv("DATABASE_NAME", "test_db")
monkeypatch.setenv("IAM_TOKEN_DB_AUTH", "True")
def _generate_mock_token(self, expires_in_seconds: int = 900) -> str:
"""Generate a mock IAM token with expiration info."""
@ -172,22 +160,12 @@ class TestBackgroundRefreshLoop:
"""Tests for the background refresh loop timing."""
@pytest.fixture
def setup_env(self):
def setup_env(self, monkeypatch, unset_database_url):
"""Setup environment variables for testing."""
os.environ["DATABASE_HOST"] = "test-host.rds.amazonaws.com"
os.environ["DATABASE_PORT"] = "5432"
os.environ["DATABASE_USER"] = "test_user"
os.environ["DATABASE_NAME"] = "test_db"
yield
# Cleanup
for key in [
"DATABASE_HOST",
"DATABASE_PORT",
"DATABASE_USER",
"DATABASE_NAME",
"DATABASE_URL",
]:
os.environ.pop(key, None)
monkeypatch.setenv("DATABASE_HOST", "test-host.rds.amazonaws.com")
monkeypatch.setenv("DATABASE_PORT", "5432")
monkeypatch.setenv("DATABASE_USER", "test_user")
monkeypatch.setenv("DATABASE_NAME", "test_db")
@pytest.mark.asyncio
async def test_calculate_seconds_fallback_when_no_url(self, setup_env):

View file

@ -626,7 +626,7 @@ async def test_getattr_does_not_block_inside_running_loop_on_expired_token(monke
assert refresh_calls["count"] == 1
def test_writer_get_rds_iam_token_defaults_port_when_unset(monkeypatch):
def test_writer_get_rds_iam_token_defaults_port_when_unset(monkeypatch, unset_database_url):
"""When DATABASE_PORT is unset, the writer must default to the Postgres
standard port instead of passing `None` through. Passing None to
`generate_iam_auth_token` makes botocore embed the literal string
@ -639,7 +639,6 @@ def test_writer_get_rds_iam_token_defaults_port_when_unset(monkeypatch):
monkeypatch.setenv("DATABASE_USER", "litellm")
monkeypatch.setenv("DATABASE_NAME", "litellm")
monkeypatch.delenv("DATABASE_SCHEMA", raising=False)
monkeypatch.delenv("DATABASE_URL", raising=False)
captured: Dict[str, Any] = {}
@ -661,7 +660,7 @@ def test_writer_get_rds_iam_token_defaults_port_when_unset(monkeypatch):
assert ":5432/litellm" in (new_url or "")
def test_writer_get_rds_iam_token_uses_database_host_env_vars(monkeypatch):
def test_writer_get_rds_iam_token_uses_database_host_env_vars(monkeypatch, unset_database_url):
"""Writer's IAM path (no iam_endpoint configured) reads host/port/user/db
from the legacy DATABASE_HOST/PORT/USER/NAME env vars and writes the URL
back to DATABASE_URL — this is the pre-read-replica behavior the patch
@ -673,7 +672,6 @@ def test_writer_get_rds_iam_token_uses_database_host_env_vars(monkeypatch):
monkeypatch.setenv("DATABASE_USER", "litellm")
monkeypatch.setenv("DATABASE_NAME", "litellm")
monkeypatch.setenv("DATABASE_SCHEMA", "public")
monkeypatch.delenv("DATABASE_URL", raising=False)
captured: Dict[str, Any] = {}