This is an automated email from the ASF dual-hosted git repository.
shahar1 pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/airflow.git
The following commit(s) were added to refs/heads/main by this push:
new bee48965b3a Validate Snowflake account and region before building the
SQL API URL (#72174)
bee48965b3a is described below
commit bee48965b3a9ae28c1ed14b745b5a417c1b99086
Author: Jarek Potiuk <[email protected]>
AuthorDate: Mon Sep 21 11:38:37 2026 +0200
Validate Snowflake account and region before building the SQL API URL
(#72174)
* Validate Snowflake account and region before building the SQL API URL
The account and region connection fields are interpolated into the SQL API
URL,
with the .snowflakecomputing.com suffix appended as text. Restrict both to
the
characters a Snowflake account or region identifier is made of, so neither
can
carry URL-significant punctuation into the request address.
Four sql_api tests mocked _get_conn_params without a return value and so
built
their URLs from a MagicMock; they now use the CONN_PARAMS constant the rest
of
the file already uses.
* Validate Snowflake account in the OAuth and Cortex Agent URLs too
Review pointed out the SQL API is not the only place a bare account
value reaches a URL: the OAuth token endpoint and the Cortex Agent base
URL build their host from it the same way, so the same unchecked value
could redirect those requests as well.
The line drawn is between fields that *name* the account (account,
region), which are validated, and fields that *are* the address
(token_endpoint, host), which stay the address the user configured.
Generated-by: Claude Opus 5
---
.../airflow/providers/snowflake/hooks/snowflake.py | 32 ++++++++-
.../snowflake/hooks/snowflake_cortex_agent.py | 5 +-
.../tests/unit/snowflake/hooks/test_snowflake.py | 77 ++++++++++++++++++++++
.../snowflake/hooks/test_snowflake_cortex_agent.py | 13 ++++
.../unit/snowflake/hooks/test_snowflake_sql_api.py | 4 ++
5 files changed, 126 insertions(+), 5 deletions(-)
diff --git
a/providers/snowflake/src/airflow/providers/snowflake/hooks/snowflake.py
b/providers/snowflake/src/airflow/providers/snowflake/hooks/snowflake.py
index c8e3fd1d105..237c13c6df1 100644
--- a/providers/snowflake/src/airflow/providers/snowflake/hooks/snowflake.py
+++ b/providers/snowflake/src/airflow/providers/snowflake/hooks/snowflake.py
@@ -19,6 +19,7 @@ from __future__ import annotations
import base64
import os
+import re
from collections.abc import Callable, Iterable, Mapping
from contextlib import closing, contextmanager
from datetime import datetime, timedelta
@@ -86,6 +87,25 @@ def _is_retryable_oauth_error(exception: BaseException) ->
bool:
return False
+_ACCOUNT_COMPONENT_PATTERN = re.compile(r"\A[A-Za-z0-9._-]+\Z")
+
+
+def _validate_account_component(value: Any, field_name: str) -> str:
+ """
+ Check that an account or region value is a bare Snowflake identifier.
+
+ These values are interpolated into the Snowflake REST URLs, so a value
carrying
+ characters that are significant in a URL would change which host the
request is
+ addressed to. Snowflake account and region identifiers are made up of
letters,
+ digits, dots, underscores and hyphens, so anything else is rejected rather
than sent.
+ """
+ if not isinstance(value, str) or not
_ACCOUNT_COMPONENT_PATTERN.fullmatch(value):
+ raise ValueError(
+ f"Invalid Snowflake {field_name} {value!r}: only letters, digits,
'.', '_' and '-' are allowed."
+ )
+ return value
+
+
class _SnowflakeOAuthManager:
"""Encapsulates OAuth token lifecycle management for Snowflake
authentication."""
@@ -160,7 +180,11 @@ class _SnowflakeOAuthManager:
):
return self._oauth_token
- url = token_endpoint or
f"https://{conn_config['account']}.snowflakecomputing.com/oauth/token-request"
+ if token_endpoint:
+ url = token_endpoint
+ else:
+ account = _validate_account_component(conn_config["account"],
"account")
+ url =
f"https://{account}.snowflakecomputing.com/oauth/token-request"
data = {
"grant_type": grant_type,
@@ -366,10 +390,12 @@ class SnowflakeHook(DbApiHook):
def account_identifier(self) -> str:
"""Get snowflake account identifier."""
conn_config = self._get_conn_params()
- account_identifier = f"https://{conn_config['account']}"
+ account = _validate_account_component(conn_config["account"],
"account")
+ account_identifier = f"https://{account}"
if conn_config["region"]:
- account_identifier += f".{conn_config['region']}"
+ region = _validate_account_component(conn_config["region"],
"region")
+ account_identifier += f".{region}"
return account_identifier
diff --git
a/providers/snowflake/src/airflow/providers/snowflake/hooks/snowflake_cortex_agent.py
b/providers/snowflake/src/airflow/providers/snowflake/hooks/snowflake_cortex_agent.py
index 59a4ac0bfa1..aa13b081ca1 100644
---
a/providers/snowflake/src/airflow/providers/snowflake/hooks/snowflake_cortex_agent.py
+++
b/providers/snowflake/src/airflow/providers/snowflake/hooks/snowflake_cortex_agent.py
@@ -22,7 +22,7 @@ from urllib.parse import quote
import requests
-from airflow.providers.snowflake.hooks.snowflake import SnowflakeHook
+from airflow.providers.snowflake.hooks.snowflake import SnowflakeHook,
_validate_account_component
JsonDict = dict[str, Any]
JsonList = list[JsonDict]
@@ -39,7 +39,8 @@ class SnowflakeCortexAgentHook(SnowflakeHook):
if host:
return f"https://{host}"
- return f"https://{conn_config['account']}.snowflakecomputing.com"
+ account = _validate_account_component(conn_config["account"],
"account")
+ return f"https://{account}.snowflakecomputing.com"
def _get_access_token(self) -> str:
conn_config = self._get_conn_params()
diff --git a/providers/snowflake/tests/unit/snowflake/hooks/test_snowflake.py
b/providers/snowflake/tests/unit/snowflake/hooks/test_snowflake.py
index 5986f034e60..a090166a86b 100644
--- a/providers/snowflake/tests/unit/snowflake/hooks/test_snowflake.py
+++ b/providers/snowflake/tests/unit/snowflake/hooks/test_snowflake.py
@@ -1915,3 +1915,80 @@ class TestPytestSnowflakeHook:
assert "token" not in conn_params
assert "token_file_path" not in conn_params
+
+
+class TestAccountIdentifierValidation:
+ """``account`` and ``region`` are interpolated into the Snowflake REST
URLs.
+
+ Both are therefore restricted to the characters a Snowflake account or
region
+ identifier is actually made of, so neither can carry URL-significant
punctuation
+ into the address the request is sent to.
+ """
+
+ @staticmethod
+ def _identifier(account: str, region: str = "") -> str:
+ with mock.patch.object(
+ SnowflakeHook, "_get_conn_params", return_value={"account":
account, "region": region}
+ ):
+ return
SnowflakeHook(snowflake_conn_id="test_conn").account_identifier
+
+ @pytest.mark.parametrize(
+ ("account", "region", "expected"),
+ [
+ ("airflow", "", "https://airflow"),
+ ("airflow", "us-east-1", "https://airflow.us-east-1"),
+ ("my_org-my_account", "", "https://my_org-my_account"),
+ ("acct.us-east-1.aws", "", "https://acct.us-east-1.aws"),
+ ],
+ )
+ def test_identifiers_are_passed_through(self, account, region, expected):
+ assert self._identifier(account, region) == expected
+
+ @pytest.mark.parametrize(
+ "account",
+ [
+ "acct.example.com/x",
+ "acct/../other",
+ "acct?x=1",
+ "acct#fragment",
+ "acct:8080",
+ "acct@elsewhere",
+ "acct\\x",
+ "acct x",
+ ],
+ )
+ def test_account_outside_identifier_charset_is_rejected(self, account):
+ with pytest.raises(ValueError, match="Invalid Snowflake account"):
+ self._identifier(account)
+
+ @pytest.mark.parametrize("region", ["us-east-1/x", "r?x", "r@h", "r:1"])
+ def test_region_outside_identifier_charset_is_rejected(self, region):
+ with pytest.raises(ValueError, match="Invalid Snowflake region"):
+ self._identifier("airflow", region)
+
+ def test_empty_account_is_rejected(self):
+ """An empty account produced a meaningless host rather than an
error."""
+ with pytest.raises(ValueError, match="Invalid Snowflake account"):
+ self._identifier("")
+
+ @mock.patch("requests.post")
+ def test_oauth_token_url_rejects_account_outside_charset(self,
requests_post):
+ conn_config = CONN_PARAMS_OAUTH | {"account": "acct.example.com/x"}
+ hook = SnowflakeHook(snowflake_conn_id="mock_conn_id")
+
+ with pytest.raises(ValueError, match="Invalid Snowflake account"):
+ hook.get_oauth_token(conn_config=conn_config)
+
+ requests_post.assert_not_called()
+
+ @mock.patch("airflow.providers.snowflake.hooks.snowflake.HTTPBasicAuth")
+ @mock.patch("requests.post")
+ def test_oauth_token_endpoint_bypasses_account_validation(self,
requests_post, mock_auth):
+ """An explicit ``token_endpoint`` replaces the account-derived URL
entirely."""
+ conn_config = CONN_PARAMS_OAUTH | {"account": "acct.example.com/x"}
+ requests_post.return_value.status_code = 200
+ hook = SnowflakeHook(snowflake_conn_id="mock_conn_id")
+
+ hook.get_oauth_token(conn_config=conn_config,
token_endpoint="https://example.com/oauth/token")
+
+ assert requests_post.call_args.args[0] ==
"https://example.com/oauth/token"
diff --git
a/providers/snowflake/tests/unit/snowflake/hooks/test_snowflake_cortex_agent.py
b/providers/snowflake/tests/unit/snowflake/hooks/test_snowflake_cortex_agent.py
index 8fb60fcee0a..cd454ddba57 100644
---
a/providers/snowflake/tests/unit/snowflake/hooks/test_snowflake_cortex_agent.py
+++
b/providers/snowflake/tests/unit/snowflake/hooks/test_snowflake_cortex_agent.py
@@ -537,3 +537,16 @@ class TestSnowflakeCortexAgentHook:
params={"ifExists": expected},
timeout=REQUEST_TIMEOUT,
)
+
+ @mock.patch(
+ f"{HOOK_PATH}._get_static_conn_params",
+ new_callable=mock.PropertyMock,
+ )
+ def test_base_url_rejects_account_outside_charset(self,
mock_static_conn_params):
+ """``account`` is interpolated into the base URL, so it may not carry
URL punctuation."""
+ mock_static_conn_params.return_value = {"account":
"acct.example.com/x"}
+
+ hook = SnowflakeCortexAgentHook(snowflake_conn_id="mock_conn_id")
+
+ with pytest.raises(ValueError, match="Invalid Snowflake account"):
+ hook._get_base_url()
diff --git
a/providers/snowflake/tests/unit/snowflake/hooks/test_snowflake_sql_api.py
b/providers/snowflake/tests/unit/snowflake/hooks/test_snowflake_sql_api.py
index 79c901558e8..6f1b7b03904 100644
--- a/providers/snowflake/tests/unit/snowflake/hooks/test_snowflake_sql_api.py
+++ b/providers/snowflake/tests/unit/snowflake/hooks/test_snowflake_sql_api.py
@@ -289,6 +289,7 @@ class TestSnowflakeSqlApiHook:
mock_requests,
):
"""Test execute_query method, run query by mocking post request method
and return the query ids"""
+ mock_conn_param.return_value = CONN_PARAMS
mock_requests.codes.ok = 200
mock_requests.request.side_effect = [
create_successful_response_mock(expected_response),
@@ -318,6 +319,7 @@ class TestSnowflakeSqlApiHook:
expected_query_ids,
mock_requests,
):
+ mock_conn_param.return_value = CONN_PARAMS
mock_requests.codes.ok = 200
mock_requests.request.side_effect = [
create_successful_response_mock(expected_response),
@@ -340,6 +342,7 @@ class TestSnowflakeSqlApiHook:
self, mock_get_header, mock_conn_param, mock_requests
):
"""Test execute_query method, run query by mocking post request method
and return the query ids"""
+ mock_conn_param.return_value = CONN_PARAMS
sql, statement_count, expected_response, expected_query_ids = (
SQL_MULTIPLE_STMTS,
4,
@@ -393,6 +396,7 @@ class TestSnowflakeSqlApiHook:
without statementHandle in the response
"""
# status_code, json payload without statementHandle
+ mock_conn_param.return_value = CONN_PARAMS
mock_make_api_call.return_value = (None, {"foo": "bar"})
hook = SnowflakeSqlApiHook("mock_conn_id")