mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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:
parent
fde0c5e53f
commit
90e1647566
3 changed files with 63 additions and 3 deletions
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue