From 90e16475663bf7716dddfbd7c9706d7893143874 Mon Sep 17 00:00:00 2001 From: Caffeine Coder Date: Tue, 20 May 2025 00:29:12 +0800 Subject: [PATCH] =?UTF-8?q?fix:=20handle=20DB=5FUSER,=20DB=5FPASSWORD,=20D?= =?UTF-8?q?B=5FHOST=20problem=20I=20faced,=20since=20this=E2=80=A6=20(#108?= =?UTF-8?q?42)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * fix: handle DB_USER, DB_PASSWORD, DB_HOST problem I faced, since this could be with special character, so better be url encoded, similar problem also happen in DATABASE_URL * test: added the test cases * add: test case of url with sepcial character --- litellm/proxy/proxy_cli.py | 10 +++++-- tests/litellm/proxy/test_proxy_cli.py | 43 +++++++++++++++++++++++++++ tests/litellm/test_main.py | 13 +++++++- 3 files changed, 63 insertions(+), 3 deletions(-) diff --git a/litellm/proxy/proxy_cli.py b/litellm/proxy/proxy_cli.py index 9f3123c0b28..4c022991f11 100644 --- a/litellm/proxy/proxy_cli.py +++ b/litellm/proxy/proxy_cli.py @@ -11,7 +11,7 @@ from typing import TYPE_CHECKING, Any, Optional, Union import click import httpx from dotenv import load_dotenv - +import urllib.parse if TYPE_CHECKING: from fastapi import FastAPI else: @@ -672,8 +672,14 @@ def run_server( # noqa: PLR0915 and database_password and database_name ): + # Handle the problem of special character escaping in the database URL + database_username_enc = urllib.parse.quote_plus(database_username) + database_password_enc = urllib.parse.quote_plus(database_password) + database_name_enc = urllib.parse.quote_plus(database_name) + # Construct DATABASE_URL from the provided variables - database_url = f"postgresql://{database_username}:{database_password}@{database_host}/{database_name}" + database_url = f"postgresql://{database_username_enc}:{database_password_enc}@{database_host}/{database_name_enc}" + os.environ["DATABASE_URL"] = database_url db_connection_pool_limit = general_settings.get( "database_connection_pool_limit", diff --git a/tests/litellm/proxy/test_proxy_cli.py b/tests/litellm/proxy/test_proxy_cli.py index 55509a4f78a..17be8d75c23 100644 --- a/tests/litellm/proxy/test_proxy_cli.py +++ b/tests/litellm/proxy/test_proxy_cli.py @@ -156,6 +156,49 @@ class TestProxyInitializationHelpers: with patch("sys.platform", "linux"): assert ProxyInitializationHelpers._get_loop_type() == "uvloop" + + @patch.dict(os.environ, {}, clear=True) + def test_database_url_construction_with_special_characters(self): + # Setup environment variables with special characters that need escaping + test_env = { + "DATABASE_HOST": "localhost:5432", + "DATABASE_USERNAME": "user@with+special", + "DATABASE_PASSWORD": "pass&word!@#$%", + "DATABASE_NAME": "db_name/test" + } + + with patch.dict(os.environ, test_env): + # Call the relevant function - we'll need to extract the database URL construction logic + # This is simulating what happens in the run_server function when database_url is None + from litellm.proxy.proxy_cli import append_query_params + import urllib.parse + + database_host = os.environ["DATABASE_HOST"] + database_username = os.environ["DATABASE_USERNAME"] + database_password = os.environ["DATABASE_PASSWORD"] + database_name = os.environ["DATABASE_NAME"] + + # Test the URL encoding part + database_username_enc = urllib.parse.quote_plus(database_username) + database_password_enc = urllib.parse.quote_plus(database_password) + database_name_enc = urllib.parse.quote_plus(database_name) + + # Construct DATABASE_URL from the provided variables + database_url = f"postgresql://{database_username_enc}:{database_password_enc}@{database_host}/{database_name_enc}" + + # Assert the correct URL was constructed with properly escaped characters + expected_url = "postgresql://user%40with%2Bspecial:pass%26word%21%40%23%24%25@localhost:5432/db_name%2Ftest" + assert database_url == expected_url + + # Test appending query parameters + params = { + "connection_limit": 10, + "pool_timeout": 60 + } + modified_url = append_query_params(database_url, params) + assert "connection_limit=10" in modified_url + assert "pool_timeout=60" in modified_url + @patch("uvicorn.run") @patch("builtins.print") def test_skip_server_startup(self, mock_print, mock_uvicorn_run): diff --git a/tests/litellm/test_main.py b/tests/litellm/test_main.py index db8241ded29..6f13561b544 100644 --- a/tests/litellm/test_main.py +++ b/tests/litellm/test_main.py @@ -12,7 +12,7 @@ sys.path.insert( ) # Adds the parent directory to the system path from unittest.mock import MagicMock, patch - +import urllib.parse import litellm @@ -388,6 +388,17 @@ async def test_openai_env_base( assert response.choices[0].message.content == "Hello from mocked response!" + +def build_database_url(username, password, host, dbname): + username_enc = urllib.parse.quote_plus(username) + password_enc = urllib.parse.quote_plus(password) + dbname_enc = urllib.parse.quote_plus(dbname) + return f"postgresql://{username_enc}:{password_enc}@{host}/{dbname_enc}" + +def test_build_database_url(): + url = build_database_url("user@name", "p@ss:word", "localhost", "db/name") + assert url == "postgresql://user%40name:p%40ss%3Aword@localhost/db%2Fname" + def test_bedrock_llama(): litellm._turn_on_debug() from litellm.types.utils import CallTypes