From e4f872d7f5b68602f794b19f8d0523606bfe92ce Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Sat, 4 Apr 2026 18:05:21 -0700 Subject: [PATCH] fix(proxy): restore transport reconnect retry in get_generic_data() Transient Prisma transport errors (e.g. httpx.ReadError) in get_generic_data() now attempt a DB reconnect and retry once before surfacing as db_exceptions alerts, matching the behavior of other DB read methods. Fixes regression introduced in 1.83.x. Closes #25143 Co-Authored-By: Claude Opus 4.6 --- litellm/proxy/utils.py | 16 +++ tests/test_litellm/proxy/test_proxy_utils.py | 111 +++++++++++++++++++ 2 files changed, 127 insertions(+) diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 845919e9120..358785575c6 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -2601,6 +2601,7 @@ class PrismaClient: key: str, value: Any, table_name: Literal["users", "keys", "config", "spend"], + _transport_reconnect_retry: bool = True, ): """ Generic implementation of get data @@ -2627,6 +2628,21 @@ class PrismaClient: except Exception as e: import traceback + if ( + _transport_reconnect_retry + and PrismaDBExceptionHandler.is_database_transport_error(e) + ): + did_reconnect = await self.attempt_db_reconnect( + reason="get_generic_data_transport_error", + ) + if did_reconnect: + return await self.get_generic_data( + key=key, + value=value, + table_name=table_name, + _transport_reconnect_retry=False, + ) + error_msg = f"LiteLLM Prisma Client Exception get_generic_data: {str(e)}" verbose_proxy_logger.error(error_msg) error_msg = error_msg + "\nException Type: {}".format(type(e)) diff --git a/tests/test_litellm/proxy/test_proxy_utils.py b/tests/test_litellm/proxy/test_proxy_utils.py index 4b50e9a4d31..cda207fd8a9 100644 --- a/tests/test_litellm/proxy/test_proxy_utils.py +++ b/tests/test_litellm/proxy/test_proxy_utils.py @@ -190,3 +190,114 @@ def test_get_projected_spend_over_limit_includes_current_spend(monkeypatch): projected_spend, projected_exceeded_date = result assert projected_spend == 290.0 assert projected_exceeded_date == real_datetime.date(2026, 4, 21) + + +@pytest.mark.asyncio +async def test_get_generic_data_retries_on_transport_error(): + """ + Test that get_generic_data retries once after a successful DB reconnect + when a transport error (e.g. httpx.ReadError) occurs. + """ + import httpx + from unittest.mock import AsyncMock, patch + + from litellm.proxy.utils import PrismaClient + + prisma_client = PrismaClient( + database_url="postgresql://user:pass@localhost:5432/db", + proxy_logging_obj=ProxyLogging(user_api_key_cache=DualCache()), + ) + + mock_db = MagicMock() + # First call raises a transport error, second call succeeds + fake_result = MagicMock() + fake_result.param_name = "general_settings" + mock_find_first = AsyncMock( + side_effect=[httpx.ReadError("connection reset"), fake_result] + ) + mock_db.litellm_config.find_first = mock_find_first + prisma_client.db = mock_db + + # Mock attempt_db_reconnect to succeed + prisma_client.attempt_db_reconnect = AsyncMock(return_value=True) + + result = await prisma_client.get_generic_data( + key="param_name", + value="general_settings", + table_name="config", + ) + + assert result == fake_result + assert mock_find_first.call_count == 2 + prisma_client.attempt_db_reconnect.assert_awaited_once_with( + reason="get_generic_data_transport_error", + ) + + +@pytest.mark.asyncio +async def test_get_generic_data_no_retry_on_non_transport_error(): + """ + Test that get_generic_data does NOT retry on non-transport errors + (e.g. a PrismaError for invalid query). + """ + from unittest.mock import AsyncMock, patch + + from litellm.proxy.utils import PrismaClient + + prisma_client = PrismaClient( + database_url="postgresql://user:pass@localhost:5432/db", + proxy_logging_obj=ProxyLogging(user_api_key_cache=DualCache()), + ) + + mock_db = MagicMock() + mock_find_first = AsyncMock(side_effect=ValueError("some non-transport error")) + mock_db.litellm_config.find_first = mock_find_first + prisma_client.db = mock_db + + prisma_client.attempt_db_reconnect = AsyncMock(return_value=True) + + with pytest.raises(ValueError, match="some non-transport error"): + await prisma_client.get_generic_data( + key="param_name", + value="general_settings", + table_name="config", + ) + + # Should NOT have attempted reconnect for a non-transport error + assert mock_find_first.call_count == 1 + prisma_client.attempt_db_reconnect.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_get_generic_data_raises_after_failed_reconnect(): + """ + Test that get_generic_data raises the original error when reconnect fails. + """ + import httpx + from unittest.mock import AsyncMock + + from litellm.proxy.utils import PrismaClient + + prisma_client = PrismaClient( + database_url="postgresql://user:pass@localhost:5432/db", + proxy_logging_obj=ProxyLogging(user_api_key_cache=DualCache()), + ) + + mock_db = MagicMock() + mock_find_first = AsyncMock(side_effect=httpx.ReadError("connection reset")) + mock_db.litellm_config.find_first = mock_find_first + prisma_client.db = mock_db + + # Mock attempt_db_reconnect to fail + prisma_client.attempt_db_reconnect = AsyncMock(return_value=False) + + with pytest.raises(httpx.ReadError): + await prisma_client.get_generic_data( + key="param_name", + value="general_settings", + table_name="config", + ) + + # Should have tried reconnect once but not retried the query + assert mock_find_first.call_count == 1 + prisma_client.attempt_db_reconnect.assert_awaited_once()