diff --git a/providers/mysql/src/airflow/providers/mysql/hooks/mysql.py b/providers/mysql/src/airflow/providers/mysql/hooks/mysql.py index 30f45818e09b3..95eaebd0bc667 100644 --- a/providers/mysql/src/airflow/providers/mysql/hooks/mysql.py +++ b/providers/mysql/src/airflow/providers/mysql/hooks/mysql.py @@ -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 @@ -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 diff --git a/providers/mysql/tests/unit/mysql/hooks/test_mysql.py b/providers/mysql/tests/unit/mysql/hooks/test_mysql.py index 18ffe66abbecb..ccab88b87a74b 100644 --- a/providers/mysql/tests/unit/mysql/hooks/test_mysql.py +++ b/providers/mysql/tests/unit/mysql/hooks/test_mysql.py @@ -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):