import base64 from tempfile import NamedTemporaryFile import pytest from cryptography.hazmat.backends import default_backend from cryptography.hazmat.primitives import serialization from cryptography.hazmat.primitives.asymmetric import rsa from .....src.connectors.snowflake_db import auth @pytest.fixture(scope="module") def mock_private_key(): """Return a mock private key for testing purposes.""" private_key = rsa.generate_private_key( public_exponent=65537, key_size=2048, backend=default_backend() ) # Get the private key as a PEM-encoded byte string _ = private_key.public_key().public_bytes( encoding=serialization.Encoding.PEM, format=serialization.PublicFormat.SubjectPublicKeyInfo, ) return private_key @pytest.fixture(scope="module") def mock_pem_private(mock_private_key): pem_private = mock_private_key.private_bytes( encoding=serialization.Encoding.PEM, format=serialization.PrivateFormat.TraditionalOpenSSL, encryption_algorithm=serialization.NoEncryption(), ) return pem_private class TestGetPrivateKeyDER: @staticmethod def _assert_is_valid_private_key(key_der: bytes) -> None: """Validates that the provided key is a valid DER format private RSA key.""" assert isinstance(key_der, bytes), "Key must be bytes." assert ( key_der[:2] == b"0\x82" ), "Key does not start with expected magic bytes for a private RSA key." def test_private_key_as_string(self, mock_pem_private): """ Test that `get_private_key_der` correctly handles a private key provided as a string. """ private_key_string = mock_pem_private.decode("utf-8") result = auth.get_private_key_der(private_key_string) self._assert_is_valid_private_key(result) def test_private_key_from_file(self, mock_pem_private): """Test that `get_private_key_der` correctly handles a private key file.""" with NamedTemporaryFile() as key_file: key_file.write(mock_pem_private) key_file.seek(0) # Reset file pointer to the start for reading result = auth.get_private_key_der(key_file.name) self._assert_is_valid_private_key(result) def test_private_key_as_base64(self, mock_pem_private): """Test that `get_private_key_der` correctly handles private keys encoded in Base64. """ private_key_base64 = base64.b64encode(mock_pem_private).decode("utf-8") result = auth.get_private_key_der(private_key_base64) self._assert_is_valid_private_key(result) @pytest.mark.parametrize("strip", [True, False]) @pytest.mark.parametrize("line_breaks", [True, False]) def test_get_private_key_der_handles_base64_encoded_der(self, strip, line_breaks): """Test that `get_private_key_der` correctly handles private keys provided as base64-encoded DER (without PEM headers or footers), which is common when private keys are programmatically extracted or transformed for use in configuration files or secrets management tools. """ base64_der_key = """ MIIEvQIBADANBgkqhkiG9w0BAQEFAASCBKcwggSjAgEAAoIBAQC7UKV/g1AeO7Ru W/YhAHfdWaR9cagmC7PN4BNH1jgOG3eSgZ3M9lKk4cQe5O86RCVfOhmmVQsPH4m4 z21pxAjFfg6puqiJhFf4iXXrg/2hqj0/A4n+QlpymGDCZq04CRBW4KjITw22bMrU LGBCoyO62hFVbO0eWdWq00hiEcyl3DN8LxoRJ+IEZlsGQAv32xUSxk16Rz7LBem0 KUqCma3xXfo1LivmQiScN9Y4JMHGnEXZPnbJ0uS1W4Z3z4cP8h0XEFlq0aGF/zvF RDc+mFo56C33axL6OvUeszDVSnsiz1mifoSmzkkDPF6eLECtaCt8YnQNWLht50QJ UveNIlTxAgMBAAECggEBAIbpJOncR/4efmvl7DYEwlA42lJHZBZS42PqebiJv8HZ Ypuwo1kWKJv8x2aA+RR8NHaLwCGequJvkp/7NhCxUqf7jNAOUokJOtsVFktUu36O wKH8CI2KeN6EhVw+24+6Z3xLPwVWesfjP4rlk0crfPZ9TXK7i2UKyuvIVR/dNgpW e78AVTSPWcIsO9Y+/+zcGudoSjScu2dIxJZ6eh4o9wuVIzEbMH1Esq1JEb0p7u8K NT79MK+/cHNl81QKhAvyrx1UE8ZTFGs7cVxjv+9DHYzM9PmIyZ/xVfa1B/cvTam2 fZLg72rqe7S6uCQ6VJvnj0BVzrq3gIYLzzfQU+XONskCgYEA830zmb5U16oKgZ3e /PNSwpWw3CQXaS5SYZKuFkf32FWLzNRgUHpVWXI8ui9HgwVVPueplsJvK0WHUF+O WkEuQLFI8xnMu1R0af1IHmB5JhSwbumqQlC+dJAUly5w87LJvaq+CGZ5OQCfgmtP Qqn77blaoMO1t2XdlKzq3IAHaVcCgYEAxPCLbdBPbp1GIi15s72wyyzJrbTHR7Yr S04OgMXY9pHaA/O/EGKOvVWLN5h4LBqNirpsFNDtFC2jbazANu+uh6vuSoCqVWUv gZcaJwR7ZhXd0268y6AFxX1S8xhNbLR/iItSkjTYiR8HkXUt4nXxy0h/REcpkK4+ nIsgjLHInvcCgYA6iM+91xa4XeV2sYgo0SoFI01ILtj97Sfa0xNMwfJgLHiWPjwF PNOF1EOrErCjW8XZgIGxf36QLX/RH5euNNs0rCn4XyxroGr4C+6pdtHJCNI9Z6kJ ljWi+fwpN/3paAK9uO1EQbZEsNSn2rpMMWciCBw0Z7gopbF2C3fOmGyvswKBgDxR D+MKMcnHEXvWQbfzGPqhNQOmooIsIQZnWbnG3rRl50fel14FUYJbeNAGOogHeeJL Rl75viK3953XkudAcUvMNKdM0N5mpy4hgTkB/mk9uTrQZ7JVyG67+3PIta3delHv mdJ9rPQSNNcv9GWvieagxZm70dcmBrcbRVTR/ofxAoGAG7/Zp8k5iLv863X4RCNU k2icAHV6lgrdGYuv3sT202jlMFb7fAg6ZP0IlYxXw4zgGvzxja1FE+APVD6ZSIG7 aLm7vcbdmlyJL/T6KglPAzDF2qqOpqlvgP5877SDv9uFcDsE+l5/nNHHLeEr+prO stIbkMgOFVz9LOx43J3pwyM= """ if strip: base64_der_key = base64_der_key.strip() if not line_breaks: base64_der_key = base64_der_key.replace("\n", "") result = auth.get_private_key_der(base64_der_key) self._assert_is_valid_private_key(result) class TestGetPrivateKeyContent: def test_get_private_key_content_from_pk_as_string(self, mock_pem_private): pem_private_as_str = mock_pem_private.decode("utf-8") result = auth._get_private_key_content(pem_private_as_str) assert isinstance(result, bytes) assert ( result == mock_pem_private.strip() ), "Result should be stripped of leading/trailing whitespace." def test_get_private_key_content_from_pk_as_file(self, mock_pem_private): with NamedTemporaryFile() as mock_private_key_file: mock_private_key_file.write(mock_pem_private) mock_private_key_file.seek(0) result = auth._get_private_key_content(mock_private_key_file.name) assert isinstance(result, bytes) assert result == mock_pem_private def test_get_private_key_content_from_pk_as_base64(self, mock_pem_private): pem_private_as_base64 = base64.b64encode(mock_pem_private).decode("utf-8") result = auth._get_private_key_content(pem_private_as_base64) assert isinstance(result, bytes) assert result == mock_pem_private