mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(tests): stop DATABASE_URL env pollution from read-replica tests breaking DB e2e tests (#32653)
This commit is contained in:
parent
41e9cc491e
commit
1fa200123f
4 changed files with 89 additions and 52 deletions
65
tests/test_litellm/proxy/db/conftest.py
Normal file
65
tests/test_litellm/proxy/db/conftest.py
Normal 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")
|
||||
|
|
@ -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"):
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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] = {}
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue