Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
39 changes: 39 additions & 0 deletions providers/mysql/src/airflow/providers/mysql/hooks/mysql.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,6 +22,7 @@
import json
import logging
from typing import TYPE_CHECKING, Any, Union
from urllib.parse import quote_plus, urlencode

from airflow.exceptions import AirflowOptionalProviderFeatureException
from airflow.providers.common.sql.hooks.sql import DbApiHook
Expand Down Expand Up @@ -363,3 +364,41 @@ def get_openlineage_database_dialect(self, _):
def get_openlineage_default_schema(self):
"""MySQL has no concept of schema."""
return None

def get_uri(self) -> str:
"""Get URI for MySQL connection."""
conn = self.connection or self.get_connection(self.get_conn_id())
conn_schema = self.schema or conn.schema or ""
client_name = conn.extra_dejson.get("client", "mysqlclient")

# Determine URI prefix based on client
if client_name == "mysql-connector-python":
uri_prefix = "mysql+mysqlconnector://"
else: # default: mysqlclient
uri_prefix = "mysql://"

auth_part = ""
if conn.login:
auth_part = quote_plus(conn.login)
if conn.password:
auth_part = f"{auth_part}:{quote_plus(conn.password)}"
auth_part = f"{auth_part}@"

host_part = conn.host or "localhost"
if conn.port:
host_part = f"{host_part}:{conn.port}"

schema_part = f"/{quote_plus(conn_schema)}" if conn_schema else ""

uri = f"{uri_prefix}{auth_part}{host_part}{schema_part}"

# Add extra connection parameters
extra = conn.extra_dejson.copy()
if "client" in extra:
extra.pop("client")

query_params = {k: str(v) for k, v in extra.items() if v}
if query_params:
uri = f"{uri}?{urlencode(query_params)}"

return uri
95 changes: 89 additions & 6 deletions providers/mysql/tests/unit/mysql/hooks/test_mysql.py
Original file line number Diff line number Diff line change
Expand Up @@ -86,12 +86,95 @@ def test_dummy_connection_setter(self, mock_connect):
assert kwargs["db"] == "schema"

@mock.patch("MySQLdb.connect")
def test_get_uri(self, mock_connect):
self.connection.extra = json.dumps({"charset": "utf-8"})
self.db_hook.get_conn()
assert mock_connect.call_count == 1
args, kwargs = mock_connect.call_args
assert self.db_hook.get_uri() == "mysql://login:password@host/schema?charset=utf-8"
@pytest.mark.parametrize(
"connection_params, expected_uri",
[
pytest.param(
{
"login": "login",
"password": "password",
"host": "host",
"schema": "schema",
"port": None,
"extra": json.dumps({"charset": "utf-8"}),
},
"mysql://login:password@host/schema?charset=utf-8",
id="basic_connection_with_charset",
),
pytest.param(
{
"login": "user@domain",
"password": "pass/word!",
"host": "host",
"schema": "schema",
"port": None,
"extra": json.dumps({"charset": "utf-8"}),
},
"mysql://user%40domain:pass%2Fword%21@host/schema?charset=utf-8",
id="special_chars_in_credentials",
),
pytest.param(
{
"login": "user@domain",
"password": "password",
"host": "host",
"schema": "schema",
"port": None,
"extra": json.dumps({"client": "mysql-connector-python"}),
},
"mysql+mysqlconnector://user%40domain:password@host/schema",
id="mysql_connector_python",
),
pytest.param(
{
"login": "user@domain",
"password": "password",
"host": "host",
"schema": "schema",
"port": 3307,
"extra": json.dumps({"client": "mysql-connector-python"}),
},
"mysql+mysqlconnector://user%40domain:password@host:3307/schema",
id="mysql_connector_with_port",
),
pytest.param(
{
"login": "user@domain",
"password": "password",
"host": "host",
"schema": "db/name",
"port": 3307,
"extra": json.dumps({"client": "mysql-connector-python"}),
},
"mysql+mysqlconnector://user%40domain:password@host:3307/db%2Fname",
id="special_chars_in_schema",
),
pytest.param(
{
"login": "user@domain",
"password": "password",
"host": "host",
"schema": "schema",
"port": 3307,
"extra": json.dumps(
{
"client": "mysql-connector-python",
"ssl_ca": "/path/to/ca",
"ssl_cert": "/path/to/cert with space",
}
),
},
"mysql+mysqlconnector://user%40domain:password@host:3307/schema?ssl_ca=%2Fpath%2Fto%2Fca&ssl_cert=%2Fpath%2Fto%2Fcert+with+space",
id="ssl_parameters",
),
],
)
def test_get_uri(self, mock_connect, connection_params, expected_uri):
"""Test get_uri method with various connection parameters."""
for key, value in connection_params.items():
setattr(self.connection, key, value)

assert self.db_hook.get_uri() == expected_uri

@mock.patch("MySQLdb.connect")
def test_get_conn_from_connection(self, mock_connect):
Expand Down