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
4 changes: 4 additions & 0 deletions providers/ibm/db2/src/airflow/providers/ibm/db2/hooks/db2.py
Original file line number Diff line number Diff line change
Expand Up @@ -99,6 +99,8 @@ def get_conn(self) -> Any:
# Add all extra parameters to connection string
# Parameter names are automatically converted to uppercase for Db2
for key, value in extra.items():
if value is None:
continue
# Convert boolean values to appropriate strings
if isinstance(value, bool):
converted_value = "true" if value else "false"
Expand Down Expand Up @@ -134,6 +136,8 @@ def get_uri(self) -> str:
if extra:
query_params = {}
for key, value in extra.items():
if value is None:
continue
# Convert boolean values to appropriate strings
if isinstance(value, bool):
query_params[key.upper()] = "true" if value else "false"
Expand Down
33 changes: 33 additions & 0 deletions providers/ibm/db2/tests/unit/ibm/db2/hooks/test_db2.py
Original file line number Diff line number Diff line change
Expand Up @@ -99,6 +99,39 @@ def test_get_conn_with_ssl(self, mock_get_connection, mock_connection_with_extra
assert "SECURITY=SSL" in call_args
assert "SSLSERVERCERTIFICATE=/path/to/cert.crt" in call_args

@pytest.mark.parametrize(
("extra", "expected_absent"),
[
('{"SSLServerCertificate": null}', ["SSLSERVERCERTIFICATE=None", "SSLSERVERCERTIFICATE="]),
('{"SECURITY": "SSL", "SSLServerCertificate": null}', ["SSLSERVERCERTIFICATE=None"]),
],
)
@patch("airflow.providers.ibm.db2.hooks.db2.Db2Hook.get_connection")
def test_skips_none_extra_values(self, mock_get_connection, extra, expected_absent):
conn = Connection(
conn_id="db2_default",
conn_type="db2",
host="localhost",
login="db2user",
password="db2pass",
schema="testdb",
port=50000,
extra=extra,
)
mock_get_connection.return_value = conn
mock_ibm_db_dbi = MagicMock()
mock_ibm_db_dbi.connect.return_value = MagicMock()

with patch.dict(sys.modules, {"ibm_db_dbi": mock_ibm_db_dbi}):
hook = Db2Hook(db2_conn_id="db2_default")
hook.get_conn()
uri = hook.get_uri()

conn_str = mock_ibm_db_dbi.connect.call_args[0][0]
for absent in expected_absent:
assert absent not in conn_str
assert absent not in uri

@patch("airflow.providers.ibm.db2.hooks.db2.Db2Hook.get_connection")
def test_get_uri(self, mock_get_connection, mock_connection):
"""Test get_uri method."""
Expand Down