fix: validate and test database connection URLs (#5178)

* test: add unit test for database url validation

* feat: add function to validate database urls

* refactor: use new database url validation function

* fix: ruff errors

* refactor: validate database urls using sqlalchemy

* test: add more cases for database url validation
This commit is contained in:
Ítalo Johnny 2024-12-17 14:29:53 -03:00 • committed by GitHub
commit 4be6b04d8c
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 116 additions and 60 deletions

View file

@ -16,6 +16,7 @@ from pydantic_settings import BaseSettings, EnvSettingsSource, PydanticBaseSetti
from typing_extensions import override from typing_extensions import override
from langflow.services.settings.constants import VARIABLES_TO_GET_FROM_ENVIRONMENT from langflow.services.settings.constants import VARIABLES_TO_GET_FROM_ENVIRONMENT
from langflow.utils.util_strings import is_valid_database_url
# BASE_COMPONENTS_PATH = str(Path(__file__).parent / "components") # BASE_COMPONENTS_PATH = str(Path(__file__).parent / "components")
BASE_COMPONENTS_PATH = str(Path(__file__).parent.parent.parent / "components") BASE_COMPONENTS_PATH = str(Path(__file__).parent.parent.parent / "components")
@ -240,7 +241,10 @@ class Settings(BaseSettings):
@field_validator("database_url", mode="before") @field_validator("database_url", mode="before")
@classmethod @classmethod
def set_database_url(cls, value, info): def set_database_url(cls, value, info):
if not value: if value and not is_valid_database_url(value):
msg = f"Invalid database_url provided: '{value}'"
raise ValueError(msg)
logger.debug("No database_url provided, trying LANGFLOW_DATABASE_URL env variable") logger.debug("No database_url provided, trying LANGFLOW_DATABASE_URL env variable")
if langflow_database_url := os.getenv("LANGFLOW_DATABASE_URL"): if langflow_database_url := os.getenv("LANGFLOW_DATABASE_URL"):
value = langflow_database_url value = langflow_database_url

View file

@ -1,3 +1,5 @@
from sqlalchemy.engine import make_url
from langflow.utils import constants from langflow.utils import constants
@ -28,3 +30,23 @@ def truncate_long_strings(data, max_length=None):
truncate_long_strings(item, max_length) truncate_long_strings(item, max_length)
return data return data
def is_valid_database_url(url: str) -> bool:
"""Validate database connection URLs compatible with SQLAlchemy.
Args:
url (str): Database connection URL to validate
Returns:
bool: True if URL is valid, False otherwise
"""
try:
parsed_url = make_url(url)
parsed_url.get_dialect()
parsed_url.get_driver_name()
except Exception: # noqa: BLE001
return False
return True

View file

@ -0,0 +1,30 @@
import pytest
from langflow.utils import util_strings
@pytest.mark.parametrize(
("value", "expected"),
[
("sqlite:///test.db", True),
("sqlite:////var/folders/test.db", True),
("sqlite:///:memory:", True),
("sqlite+aiosqlite:////var/folders/test.db", True),
("postgresql://user:pass@localhost/dbname", True),
("postgresql+psycopg2://scott:tiger@localhost:5432/mydatabase", True),
("postgresql+pg8000://dbuser:kx%40jj5%2Fg@pghost10/appdb", True),
("mysql://user:pass@localhost/dbname", True),
("mysql+mysqldb://scott:tiger@localhost/foo", True),
("mysql+pymysql://scott:tiger@localhost/foo", True),
("oracle://scott:tiger@127.0.0.1:1521/?service_name=freepdb1", True),
("oracle+cx_oracle://scott:tiger@tnsalias", True),
("oracle+oracledb://scott:tiger@127.0.0.1:1521/?service_name=freepdb1", True),
("", False),
(" invalid ", False),
("not_a_url", False),
(None, False),
("invalid://database", False),
("invalid://:@/test", False),
],
)
def test_is_valid_database_url(value, expected):
assert util_strings.is_valid_database_url(value) == expected