fix: handle DB_USER, DB_PASSWORD, DB_HOST problem I faced, since this… (#10842)

* 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
This commit is contained in:
Caffeine Coder 2025-05-20 00:29:12 +08:00 • committed by GitHub
parent fde0c5e53f
commit 90e1647566
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 63 additions and 3 deletions

View file

@ -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",

View file

@ -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):

View file

@ -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