Skip to content
Prev Previous commit
Next Next commit
fix: Remove leftover private key kwargs before passing to snowflake.c…
…onnector.connect

Signed-off-by: Jia Le <5955220+jials@users.noreply.github.com>
  • Loading branch information
jials authored and ntkathole committed May 1, 2026
commit 945d7a99a6b59170eb36f03c3be9ad62b1fdfa73
6 changes: 3 additions & 3 deletions sdk/python/feast/infra/utils/snowflake/snowflake_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -86,9 +86,9 @@ def __enter__(self):
# https://docs.snowflake.com/en/user-guide/key-pair-auth.html#configuring-key-pair-authentication
if "private_key" in kwargs or "private_key_content" in kwargs:
kwargs["private_key"] = parse_private_key_path(
kwargs.get("private_key_passphrase"),
kwargs.get("private_key"),
kwargs.get("private_key_content"),
kwargs.pop("private_key_passphrase", None),
kwargs.pop("private_key", None),
kwargs.pop("private_key_content", None),
)

try:
Expand Down
57 changes: 57 additions & 0 deletions sdk/python/tests/unit/infra/registry/test_snowflake_registry.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,7 @@

from feast.entity import Entity
from feast.infra.registry.snowflake import SnowflakeRegistry, SnowflakeRegistryConfig
from feast.infra.utils.snowflake.snowflake_utils import GetSnowflakeConnection


@pytest.fixture
Expand Down Expand Up @@ -271,3 +272,59 @@ def test_sync_deduplicates_project_ids(self, mock_execute, mock_get_conn):
assert mock_apply.call_count == 2
applied_names = {call[0][0].name for call in mock_apply.call_args_list}
assert applied_names == {"project_a", "project_b"}


class _DictableConfig:
"""A config object that supports dict() conversion and attribute access."""

def __init__(self, data):
self._data = data
for k, v in data.items():
setattr(self, k, v)

def __iter__(self):
return iter(self._data)

def keys(self):
return self._data.keys()

def __getitem__(self, key):
return self._data[key]


class TestGetSnowflakeConnection:
@patch("feast.infra.utils.snowflake.snowflake_utils.parse_private_key_path")
@patch("feast.infra.utils.snowflake.snowflake_utils.snowflake.connector")
@patch("feast.infra.utils.snowflake.snowflake_utils._cache", {})
def test_private_key_kwargs_not_leaked_to_connect(
self, mock_connector, mock_parse_key
):
"""private_key_passphrase and private_key_content must not be passed to connect()."""
mock_parse_key.return_value = b"parsed_key_bytes"
mock_conn = MagicMock()
mock_connector.connect.return_value = mock_conn

config = _DictableConfig(
{
"type": "snowflake.registry",
"account": "test_account",
"user": "test_user",
"password": None,
"role": "test_role",
"warehouse": "test_wh",
"database": "test_db",
"schema_": "test_schema",
"config_path": "",
"private_key": "/path/to/key.p8",
"private_key_passphrase": "my_secret", # pragma: allowlist secret
"private_key_content": None,
}
)

with GetSnowflakeConnection(config):
pass

connect_kwargs = mock_connector.connect.call_args[1]
assert "private_key_passphrase" not in connect_kwargs
assert "private_key_content" not in connect_kwargs
assert connect_kwargs["private_key"] == b"parsed_key_bytes"