From d8ab7037e6477aebdd386f489135df4f61c76422 Mon Sep 17 00:00:00 2001 From: Jesse Mansfield Date: Mon, 12 Jan 2026 10:37:15 +1100 Subject: [PATCH 1/4] reformat add proxy support commit --- .../providers/snowflake/hooks/snowflake.py | 33 ++++- .../unit/snowflake/hooks/test_snowflake.py | 132 ++++++++++++++++++ 2 files changed, 164 insertions(+), 1 deletion(-) diff --git a/providers/snowflake/src/airflow/providers/snowflake/hooks/snowflake.py b/providers/snowflake/src/airflow/providers/snowflake/hooks/snowflake.py index df1b584ff9dc7..fc1c9b191d00c 100644 --- a/providers/snowflake/src/airflow/providers/snowflake/hooks/snowflake.py +++ b/providers/snowflake/src/airflow/providers/snowflake/hooks/snowflake.py @@ -111,7 +111,7 @@ def get_connection_form_widgets(cls) -> dict[str, Any]: BS3TextFieldWidget, ) from flask_babel import lazy_gettext - from wtforms import BooleanField, PasswordField, StringField + from wtforms import BooleanField, IntegerField, PasswordField, StringField return { "account": StringField(lazy_gettext("Account"), widget=BS3TextFieldWidget()), @@ -126,6 +126,10 @@ def get_connection_form_widgets(cls) -> dict[str, Any]: "insecure_mode": BooleanField( label=lazy_gettext("Insecure mode"), description="Turns off OCSP certificate checks" ), + "proxy_host": StringField(lazy_gettext("Proxy Host"), widget=BS3TextFieldWidget()), + "proxy_port": IntegerField(lazy_gettext("Proxy Port")), + "proxy_user": StringField(lazy_gettext("Proxy User"), widget=BS3TextFieldWidget()), + "proxy_password": PasswordField(lazy_gettext("Proxy Password"), widget=BS3PasswordFieldWidget()), } @classmethod @@ -148,6 +152,10 @@ def get_ui_field_behaviour(cls) -> dict[str, Any]: "token_endpoint": "token endpoint", "refresh_token": "refresh token", "scope": "scope", + "proxy_host": "proxy.example.com", + "proxy_port": "8080", + "proxy_user": "proxy_username", + "proxy_password": "proxy_password", }, indent=1, ), @@ -162,6 +170,10 @@ def get_ui_field_behaviour(cls) -> dict[str, Any]: "private_key_file": "Path of snowflake private key (PEM Format)", "private_key_content": "Content to snowflake private key (PEM format)", "insecure_mode": "insecure mode", + "proxy_host": "Proxy server hostname", + "proxy_port": "Proxy server port", + "proxy_user": "Proxy username (optional)", + "proxy_password": "Proxy password (optional)", }, } @@ -426,6 +438,21 @@ def _get_static_conn_params(self) -> dict[str, str | None]: ocsp_fail_open = extra_dict.get("ocsp_fail_open") if ocsp_fail_open is not None: conn_config["ocsp_fail_open"] = _try_to_boolean(ocsp_fail_open) + + # Add proxy configuration if specified + proxy_host = self._get_field(extra_dict, "proxy_host") + proxy_port = self._get_field(extra_dict, "proxy_port") + proxy_user = self._get_field(extra_dict, "proxy_user") + proxy_password = self._get_field(extra_dict, "proxy_password") + + if proxy_host: + conn_config["proxy_host"] = proxy_host + if proxy_port: + conn_config["proxy_port"] = int(proxy_port) if isinstance(proxy_port, str) else proxy_port + if proxy_user: + conn_config["proxy_user"] = proxy_user + if proxy_password: + conn_config["proxy_password"] = proxy_password return conn_config @@ -520,6 +547,10 @@ def _conn_params_to_sqlalchemy_uri(self, conn_params: dict) -> str: "client_store_temporary_credential", "json_result_force_utf8_decoding", "ocsp_fail_open", + "proxy_host", + "proxy_port", + "proxy_user", + "proxy_password", ] } ) diff --git a/providers/snowflake/tests/unit/snowflake/hooks/test_snowflake.py b/providers/snowflake/tests/unit/snowflake/hooks/test_snowflake.py index b44703c546937..44ccdbd571656 100644 --- a/providers/snowflake/tests/unit/snowflake/hooks/test_snowflake.py +++ b/providers/snowflake/tests/unit/snowflake/hooks/test_snowflake.py @@ -1309,3 +1309,135 @@ def test_oauth_token_refresh_after_expiry(self, mock_timezone_utcnow, mock_reque # Ensure refresh actually happened assert mock_requests_post.call_count == 2 + + + def test_get_conn_params_with_proxy_host_only(self): + """Test proxy configuration with only host specified.""" + connection_kwargs = deepcopy(BASE_CONNECTION_KWARGS) + connection_kwargs["extra"]["proxy_host"] = "proxy.example.com" + + with mock.patch.dict("os.environ", AIRFLOW_CONN_TEST_CONN=Connection(**connection_kwargs).get_uri()): + hook = SnowflakeHook(snowflake_conn_id="test_conn") + conn_params = hook._get_conn_params + + assert conn_params["proxy_host"] == "proxy.example.com" + assert "proxy_port" not in conn_params + assert "proxy_user" not in conn_params + assert "proxy_password" not in conn_params + + def test_get_conn_params_with_proxy_host_and_port(self): + """Test proxy configuration with host and port.""" + connection_kwargs = deepcopy(BASE_CONNECTION_KWARGS) + connection_kwargs["extra"]["proxy_host"] = "proxy.example.com" + connection_kwargs["extra"]["proxy_port"] = "8080" + + with mock.patch.dict("os.environ", AIRFLOW_CONN_TEST_CONN=Connection(**connection_kwargs).get_uri()): + hook = SnowflakeHook(snowflake_conn_id="test_conn") + conn_params = hook._get_conn_params + + assert conn_params["proxy_host"] == "proxy.example.com" + assert conn_params["proxy_port"] == 8080 + assert "proxy_user" not in conn_params + assert "proxy_password" not in conn_params + + def test_get_conn_params_with_proxy_port_as_int(self): + """Test proxy configuration with port as integer.""" + connection_kwargs = deepcopy(BASE_CONNECTION_KWARGS) + connection_kwargs["extra"]["proxy_host"] = "proxy.example.com" + connection_kwargs["extra"]["proxy_port"] = 8080 # Integer instead of string + + with mock.patch.dict("os.environ", AIRFLOW_CONN_TEST_CONN=Connection(**connection_kwargs).get_uri()): + hook = SnowflakeHook(snowflake_conn_id="test_conn") + conn_params = hook._get_conn_params + + assert conn_params["proxy_host"] == "proxy.example.com" + assert conn_params["proxy_port"] == 8080 + assert isinstance(conn_params["proxy_port"], int) + + def test_get_conn_params_with_proxy_full_config(self): + """Test proxy configuration with all parameters.""" + connection_kwargs = deepcopy(BASE_CONNECTION_KWARGS) + connection_kwargs["extra"]["proxy_host"] = "proxy.example.com" + connection_kwargs["extra"]["proxy_port"] = "8080" + connection_kwargs["extra"]["proxy_user"] = "proxy_username" + connection_kwargs["extra"]["proxy_password"] = "proxy_password" + + with mock.patch.dict("os.environ", AIRFLOW_CONN_TEST_CONN=Connection(**connection_kwargs).get_uri()): + hook = SnowflakeHook(snowflake_conn_id="test_conn") + conn_params = hook._get_conn_params + + assert conn_params["proxy_host"] == "proxy.example.com" + assert conn_params["proxy_port"] == 8080 + assert conn_params["proxy_user"] == "proxy_username" + assert conn_params["proxy_password"] == "proxy_password" + + def test_get_conn_params_with_proxy_backcompat_prefix(self): + """Test proxy configuration with backcompat prefix.""" + connection_kwargs = deepcopy(BASE_CONNECTION_KWARGS) + connection_kwargs["extra"]["extra__snowflake__proxy_host"] = "proxy.example.com" + connection_kwargs["extra"]["extra__snowflake__proxy_port"] = "8080" + connection_kwargs["extra"]["extra__snowflake__proxy_user"] = "proxy_username" + connection_kwargs["extra"]["extra__snowflake__proxy_password"] = "proxy_password" + + with mock.patch.dict("os.environ", AIRFLOW_CONN_TEST_CONN=Connection(**connection_kwargs).get_uri()): + hook = SnowflakeHook(snowflake_conn_id="test_conn") + conn_params = hook._get_conn_params + + assert conn_params["proxy_host"] == "proxy.example.com" + assert conn_params["proxy_port"] == 8080 + assert conn_params["proxy_user"] == "proxy_username" + assert conn_params["proxy_password"] == "proxy_password" + + def test_get_conn_with_proxy_should_call_connect(self): + """Test that proxy parameters are passed to connector.connect().""" + connection_kwargs = deepcopy(BASE_CONNECTION_KWARGS) + connection_kwargs["extra"]["proxy_host"] = "proxy.example.com" + connection_kwargs["extra"]["proxy_port"] = "8080" + connection_kwargs["extra"]["proxy_user"] = "proxy_user" + connection_kwargs["extra"]["proxy_password"] = "proxy_pass" + + with ( + mock.patch.dict("os.environ", AIRFLOW_CONN_TEST_CONN=Connection(**connection_kwargs).get_uri()), + mock.patch("airflow.providers.snowflake.hooks.snowflake.connector") as mock_connector, + ): + hook = SnowflakeHook(snowflake_conn_id="test_conn") + hook.get_conn() + + call_args = mock_connector.connect.call_args[1] + assert call_args["proxy_host"] == "proxy.example.com" + assert call_args["proxy_port"] == 8080 + assert call_args["proxy_user"] == "proxy_user" + assert call_args["proxy_password"] == "proxy_pass" + + def test_sqlalchemy_uri_excludes_proxy_params(self): + """Test that proxy parameters are excluded from SQLAlchemy URI.""" + connection_kwargs = deepcopy(BASE_CONNECTION_KWARGS) + connection_kwargs["extra"]["proxy_host"] = "proxy.example.com" + connection_kwargs["extra"]["proxy_port"] = "8080" + + with mock.patch.dict("os.environ", AIRFLOW_CONN_TEST_CONN=Connection(**connection_kwargs).get_uri()): + hook = SnowflakeHook(snowflake_conn_id="test_conn") + uri = hook.get_uri() + + # Proxy parameters should NOT appear in the URI + assert "proxy_host" not in uri + assert "proxy_port" not in uri + assert "proxy.example.com" not in uri + assert "8080" not in uri + + def test_get_sqlalchemy_engine_with_proxy(self): + """Test get_sqlalchemy_engine does not include proxy params in URI but passes to connect_args if needed.""" + connection_kwargs = deepcopy(BASE_CONNECTION_KWARGS) + connection_kwargs["extra"]["proxy_host"] = "proxy.example.com" + connection_kwargs["extra"]["proxy_port"] = "8080" + + with ( + mock.patch.dict("os.environ", AIRFLOW_CONN_TEST_CONN=Connection(**connection_kwargs).get_uri()), + mock.patch("airflow.providers.snowflake.hooks.snowflake.create_engine") as mock_create_engine, + ): + hook = SnowflakeHook(snowflake_conn_id="test_conn") + hook.get_sqlalchemy_engine() + + # Check that the URI doesn't contain proxy params + called_uri = mock_create_engine.call_args[0][0] + assert "proxy_host" not in str(called_uri) From 6ee994e28e727e4c1f1d1713c8d09fcd9599b47c Mon Sep 17 00:00:00 2001 From: Jesse Mansfield Date: Tue, 13 Jan 2026 10:21:44 +1100 Subject: [PATCH 2/4] static errors fix, fixes get_conn_params() causing method is not subscriptable error --- .../airflow/providers/snowflake/hooks/snowflake.py | 2 +- .../tests/unit/snowflake/hooks/test_snowflake.py | 13 ++++++------- 2 files changed, 7 insertions(+), 8 deletions(-) diff --git a/providers/snowflake/src/airflow/providers/snowflake/hooks/snowflake.py b/providers/snowflake/src/airflow/providers/snowflake/hooks/snowflake.py index fc1c9b191d00c..6f58c21c7cc63 100644 --- a/providers/snowflake/src/airflow/providers/snowflake/hooks/snowflake.py +++ b/providers/snowflake/src/airflow/providers/snowflake/hooks/snowflake.py @@ -438,7 +438,7 @@ def _get_static_conn_params(self) -> dict[str, str | None]: ocsp_fail_open = extra_dict.get("ocsp_fail_open") if ocsp_fail_open is not None: conn_config["ocsp_fail_open"] = _try_to_boolean(ocsp_fail_open) - + # Add proxy configuration if specified proxy_host = self._get_field(extra_dict, "proxy_host") proxy_port = self._get_field(extra_dict, "proxy_port") diff --git a/providers/snowflake/tests/unit/snowflake/hooks/test_snowflake.py b/providers/snowflake/tests/unit/snowflake/hooks/test_snowflake.py index 44ccdbd571656..8f057a46565fd 100644 --- a/providers/snowflake/tests/unit/snowflake/hooks/test_snowflake.py +++ b/providers/snowflake/tests/unit/snowflake/hooks/test_snowflake.py @@ -1310,7 +1310,6 @@ def test_oauth_token_refresh_after_expiry(self, mock_timezone_utcnow, mock_reque # Ensure refresh actually happened assert mock_requests_post.call_count == 2 - def test_get_conn_params_with_proxy_host_only(self): """Test proxy configuration with only host specified.""" connection_kwargs = deepcopy(BASE_CONNECTION_KWARGS) @@ -1318,7 +1317,7 @@ def test_get_conn_params_with_proxy_host_only(self): with mock.patch.dict("os.environ", AIRFLOW_CONN_TEST_CONN=Connection(**connection_kwargs).get_uri()): hook = SnowflakeHook(snowflake_conn_id="test_conn") - conn_params = hook._get_conn_params + conn_params = hook._get_conn_params() assert conn_params["proxy_host"] == "proxy.example.com" assert "proxy_port" not in conn_params @@ -1333,7 +1332,7 @@ def test_get_conn_params_with_proxy_host_and_port(self): with mock.patch.dict("os.environ", AIRFLOW_CONN_TEST_CONN=Connection(**connection_kwargs).get_uri()): hook = SnowflakeHook(snowflake_conn_id="test_conn") - conn_params = hook._get_conn_params + conn_params = hook._get_conn_params() assert conn_params["proxy_host"] == "proxy.example.com" assert conn_params["proxy_port"] == 8080 @@ -1348,7 +1347,7 @@ def test_get_conn_params_with_proxy_port_as_int(self): with mock.patch.dict("os.environ", AIRFLOW_CONN_TEST_CONN=Connection(**connection_kwargs).get_uri()): hook = SnowflakeHook(snowflake_conn_id="test_conn") - conn_params = hook._get_conn_params + conn_params = hook._get_conn_params() assert conn_params["proxy_host"] == "proxy.example.com" assert conn_params["proxy_port"] == 8080 @@ -1364,7 +1363,7 @@ def test_get_conn_params_with_proxy_full_config(self): with mock.patch.dict("os.environ", AIRFLOW_CONN_TEST_CONN=Connection(**connection_kwargs).get_uri()): hook = SnowflakeHook(snowflake_conn_id="test_conn") - conn_params = hook._get_conn_params + conn_params = hook._get_conn_params() assert conn_params["proxy_host"] == "proxy.example.com" assert conn_params["proxy_port"] == 8080 @@ -1381,7 +1380,7 @@ def test_get_conn_params_with_proxy_backcompat_prefix(self): with mock.patch.dict("os.environ", AIRFLOW_CONN_TEST_CONN=Connection(**connection_kwargs).get_uri()): hook = SnowflakeHook(snowflake_conn_id="test_conn") - conn_params = hook._get_conn_params + conn_params = hook._get_conn_params() assert conn_params["proxy_host"] == "proxy.example.com" assert conn_params["proxy_port"] == 8080 @@ -1402,7 +1401,7 @@ def test_get_conn_with_proxy_should_call_connect(self): ): hook = SnowflakeHook(snowflake_conn_id="test_conn") hook.get_conn() - + call_args = mock_connector.connect.call_args[1] assert call_args["proxy_host"] == "proxy.example.com" assert call_args["proxy_port"] == 8080 From 2412feaf6d52b6a1ef5739a38220fea29901a554 Mon Sep 17 00:00:00 2001 From: Jesse Mansfield Date: Wed, 14 Jan 2026 08:33:55 +1100 Subject: [PATCH 3/4] static --- .../src/airflow/providers/snowflake/hooks/snowflake.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/providers/snowflake/src/airflow/providers/snowflake/hooks/snowflake.py b/providers/snowflake/src/airflow/providers/snowflake/hooks/snowflake.py index 8d935968be2a6..9672547249304 100644 --- a/providers/snowflake/src/airflow/providers/snowflake/hooks/snowflake.py +++ b/providers/snowflake/src/airflow/providers/snowflake/hooks/snowflake.py @@ -442,7 +442,7 @@ def _get_static_conn_params(self) -> dict[str, str | None]: ocsp_fail_open = extra_dict.get("ocsp_fail_open") if ocsp_fail_open is not None: conn_config["ocsp_fail_open"] = _try_to_boolean(ocsp_fail_open) - + # Add proxy configuration if specified proxy_host = self._get_field(extra_dict, "proxy_host") proxy_port = self._get_field(extra_dict, "proxy_port") From 492c7a32cdf81ca758f07ee83681d1b65d90c851 Mon Sep 17 00:00:00 2001 From: Jesse Mansfield Date: Wed, 21 Jan 2026 14:03:31 +1100 Subject: [PATCH 4/4] add proxy_password as default sensitive field --- .../src/airflow_shared/secrets_masker/secrets_masker.py | 1 + 1 file changed, 1 insertion(+) diff --git a/shared/secrets_masker/src/airflow_shared/secrets_masker/secrets_masker.py b/shared/secrets_masker/src/airflow_shared/secrets_masker/secrets_masker.py index 6e8d556eb6dda..c99ad568e1c69 100644 --- a/shared/secrets_masker/src/airflow_shared/secrets_masker/secrets_masker.py +++ b/shared/secrets_masker/src/airflow_shared/secrets_masker/secrets_masker.py @@ -59,6 +59,7 @@ def to_dict(self) -> dict[str, Any]: ... "password", "private_key", "proxy", + "proxy_password", "proxies", "secret", "token",