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
1 change: 1 addition & 0 deletions providers/snowflake/docs/connections/snowflake.rst
Original file line number Diff line number Diff line change
Expand Up @@ -64,6 +64,7 @@ Extra (optional)
* ``insecure_mode``: Turn off OCSP certificate checks. For details, see: `How To: Turn Off OCSP Checking in Snowflake Client Drivers - Snowflake Community <https://community.snowflake.com/s/article/How-to-turn-off-OCSP-checking-in-Snowflake-client-drivers>`_.
* ``host``: Target Snowflake hostname to connect to (e.g., for local testing with LocalStack).
* ``port``: Target Snowflake port to connect to (e.g., for local testing with LocalStack).
* ``ocsp_fail_open``: Specify `ocsp_fail_open <https://docs.snowflake.com/en/developer-guide/python-connector/python-connector-connect#label-python-ocsp-choosing-fail-open-or-fail-close-mode>`_.

URI format example
^^^^^^^^^^^^^^^^^^
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -299,6 +299,12 @@ def _get_conn_params(self) -> dict[str, str | None]:
if snowflake_port:
conn_config["port"] = snowflake_port

# if a value for ocsp_fail_open is set, pass it along.
# Note the check is for `is not None` so that we can pass along `False` as a value.
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)

return conn_config

def get_uri(self) -> str:
Expand All @@ -320,6 +326,7 @@ def _conn_params_to_sqlalchemy_uri(self, conn_params: dict) -> str:
"client_request_mfa_token",
"client_store_temporary_credential",
"json_result_force_utf8_decoding",
"ocsp_fail_open",
]
}
)
Expand All @@ -345,6 +352,9 @@ def get_sqlalchemy_engine(self, engine_kwargs=None):
if "json_result_force_utf8_decoding" in conn_params:
engine_kwargs.setdefault("connect_args", {})
engine_kwargs["connect_args"]["json_result_force_utf8_decoding"] = True
if "ocsp_fail_open" in conn_params:
engine_kwargs.setdefault("connect_args", {})
engine_kwargs["connect_args"]["ocsp_fail_open"] = conn_params["ocsp_fail_open"]
for key in ["session_parameters", "private_key"]:
if conn_params.get(key):
engine_kwargs.setdefault("connect_args", {})
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -277,6 +277,60 @@ class TestPytestSnowflakeHook:
"json_result_force_utf8_decoding": True,
},
),
(
{
**BASE_CONNECTION_KWARGS,
"extra": {
**BASE_CONNECTION_KWARGS["extra"],
"ocsp_fail_open": True,
},
},
(
"snowflake://user:pw@airflow.af_region/db/public?"
"application=AIRFLOW&authenticator=snowflake&role=af_role&warehouse=af_wh"
),
{
"account": "airflow",
"application": "AIRFLOW",
"authenticator": "snowflake",
"database": "db",
"password": "pw",
"region": "af_region",
"role": "af_role",
"schema": "public",
"session_parameters": None,
"user": "user",
"warehouse": "af_wh",
"ocsp_fail_open": True,
},
),
(
{
**BASE_CONNECTION_KWARGS,
"extra": {
**BASE_CONNECTION_KWARGS["extra"],
"ocsp_fail_open": False,
},
},
(
"snowflake://user:pw@airflow.af_region/db/public?"
"application=AIRFLOW&authenticator=snowflake&role=af_role&warehouse=af_wh"
),
{
"account": "airflow",
"application": "AIRFLOW",
"authenticator": "snowflake",
"database": "db",
"password": "pw",
"region": "af_region",
"role": "af_role",
"schema": "public",
"session_parameters": None,
"user": "user",
"warehouse": "af_wh",
"ocsp_fail_open": False,
},
),
],
)
def test_hook_should_support_prepare_basic_conn_params_and_uri(
Expand Down Expand Up @@ -530,6 +584,23 @@ def test_get_sqlalchemy_engine_should_support_private_key_auth(self, non_encrypt
assert "private_key" in mock_create_engine.call_args.kwargs["connect_args"]
assert mock_create_engine.return_value == conn

def test_get_sqlalchemy_engine_should_support_ocsp_fail_open(self):
connection_kwargs = deepcopy(BASE_CONNECTION_KWARGS)
connection_kwargs["extra"]["ocsp_fail_open"] = "False"

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")
conn = hook.get_sqlalchemy_engine()
mock_create_engine.assert_called_once_with(
"snowflake://user:pw@airflow.af_region/db/public"
"?application=AIRFLOW&authenticator=snowflake&role=af_role&warehouse=af_wh",
connect_args={"ocsp_fail_open": False},
)
assert mock_create_engine.return_value == conn

def test_hook_parameters_should_take_precedence(self):
with mock.patch.dict(
"os.environ", AIRFLOW_CONN_TEST_CONN=Connection(**BASE_CONNECTION_KWARGS).get_uri()
Expand Down