This is an automated email from the ASF dual-hosted git repository.
shunping pushed a commit to branch master
in repository https://gitbox.apache.org/repos/asf/beam.git
The following commit(s) were added to refs/heads/master by this push:
new 52702a15957 [python] Add Secret management module in
apache_beam.utils.secret (#39636)
52702a15957 is described below
commit 52702a15957a0277f04a4336b344dd1769fc5e72
Author: Shunping Huang <[email protected]>
AuthorDate: Tue Aug 18 09:58:04 2026 -0400
[python] Add Secret management module in apache_beam.utils.secret (#39636)
* Refactor Secret classes and tests into apache_beam.utils.secret
There is no functional changes in this commit. We also re-export
the secret classes in apache_beam.transforms.util for backward
compatibility.
* Add secret caching, string accessors, and __getstate__
* Add Secret.from_json factory, RawSecret implementation, and expanded test
coverage
* Refactor parse_secret option and unify the logic.
* Rename get to get_str in Secret.
---
sdks/python/apache_beam/transforms/util.py | 245 +------------
sdks/python/apache_beam/transforms/util_test.py | 157 +-------
sdks/python/apache_beam/utils/secret.py | 464 ++++++++++++++++++++++++
sdks/python/apache_beam/utils/secret_test.py | 454 +++++++++++++++++++++++
4 files changed, 926 insertions(+), 394 deletions(-)
diff --git a/sdks/python/apache_beam/transforms/util.py
b/sdks/python/apache_beam/transforms/util.py
index 2ea9df9399c..60295f68a92 100644
--- a/sdks/python/apache_beam/transforms/util.py
+++ b/sdks/python/apache_beam/transforms/util.py
@@ -82,6 +82,9 @@ from apache_beam.typehints.sharded_key_type import
ShardedKeyType
from apache_beam.utils import shared
from apache_beam.utils import windowed_value
from apache_beam.utils.annotations import deprecated
+from apache_beam.utils.secret import Secret
+from apache_beam.utils.secret import GcpSecret
+from apache_beam.utils.secret import GcpHsmGeneratedSecret
from apache_beam.utils.sharded_key import ShardedKey
from apache_beam.utils.timestamp import Timestamp
@@ -94,6 +97,7 @@ __all__ = [
'BatchElements',
'CoGroupByKey',
'Distinct',
+ 'GcpHsmGeneratedSecret',
'GcpSecret',
'GroupByEncryptedKey',
'Keys',
@@ -327,247 +331,6 @@ def RemoveDuplicates(pcoll):
return pcoll | 'RemoveDuplicates' >> Distinct()
-class Secret():
- """A secret management class used for handling sensitive data.
-
- This class provides a generic interface for secret management.
Implementations
- of this class should handle fetching secrets from a secret management system.
- """
- def get_secret_bytes(self) -> bytes:
- """Returns the secret as a byte string."""
- raise NotImplementedError()
-
- @staticmethod
- def generate_secret_bytes() -> bytes:
- """Generates a new secret key."""
- return Fernet.generate_key()
-
- @staticmethod
- def parse_secret_option(secret) -> 'Secret':
- """Parses a secret string and returns the appropriate secret type.
-
- The secret string should be formatted like:
- 'type:<secret_type>;<secret_param>:<value>'
-
- For example, 'type:GcpSecret;version_name:my_secret/versions/latest'
- would return a GcpSecret initialized with 'my_secret/versions/latest'.
- """
- param_map = {}
- for param in secret.split(';'):
- parts = param.split(':')
- param_map[parts[0]] = parts[1]
-
- if 'type' not in param_map:
- raise ValueError('Secret string must contain a valid type parameter')
-
- secret_type = param_map['type'].lower()
- del param_map['type']
- secret_class = Secret
- secret_params = None
- if secret_type == 'gcpsecret':
- secret_class = GcpSecret # type: ignore[assignment]
- secret_params = ['version_name']
- elif secret_type == 'gcphsmgeneratedsecret':
- secret_class = GcpHsmGeneratedSecret # type: ignore[assignment]
- secret_params = [
- 'project_id', 'location_id', 'key_ring_id', 'key_id', 'job_name'
- ]
- else:
- raise ValueError(
- f'Invalid secret type {secret_type}, currently only '
- 'GcpSecret and GcpHsmGeneratedSecret are supported')
-
- for param_name in param_map.keys():
- if param_name not in secret_params:
- raise ValueError(
- f'Invalid secret parameter {param_name}, '
- f'{secret_type} only supports the following '
- f'parameters: {secret_params}')
- return secret_class(**param_map)
-
-
-class GcpSecret(Secret):
- """A secret manager implementation that retrieves secrets from Google Cloud
- Secret Manager.
- """
- def __init__(self, version_name: str):
- """Initializes a GcpSecret object.
-
- Args:
- version_name: The full version name of the secret in Google Cloud Secret
- Manager. For example:
- projects/<id>/secrets/<secret_name>/versions/1.
- For more info, see
-
https://cloud.google.com/python/docs/reference/secretmanager/latest/google.cloud.secretmanager_v1beta1.services.secret_manager_service.SecretManagerServiceClient#google_cloud_secretmanager_v1beta1_services_secret_manager_service_SecretManagerServiceClient_access_secret_version
- """
- self._version_name = version_name
-
- def get_secret_bytes(self) -> bytes:
- try:
- from google.cloud import secretmanager
- client = secretmanager.SecretManagerServiceClient()
- response = client.access_secret_version(
- request={"name": self._version_name})
- secret = response.payload.data
- return secret
- except Exception as e:
- raise RuntimeError(
- 'Failed to retrieve secret bytes for secret '
- f'{self._version_name} with exception {e}')
-
- def __eq__(self, secret):
- return self._version_name == getattr(secret, '_version_name', None)
-
-
-class GcpHsmGeneratedSecret(Secret):
- """A secret manager implementation that generates a secret using a GCP HSM
key
- and stores it in Google Cloud Secret Manager. If the secret already exists,
- it will be retrieved.
- """
- def __init__(
- self,
- project_id: str,
- location_id: str,
- key_ring_id: str,
- key_id: str,
- job_name: str):
- """Initializes a GcpHsmGeneratedSecret object.
-
- Args:
- project_id: The GCP project ID.
- location_id: The GCP location ID for the HSM key.
- key_ring_id: The ID of the KMS key ring.
- key_id: The ID of the KMS key.
- job_name: The name of the job, used to generate a unique secret name.
- """
- self._project_id = project_id
- self._location_id = location_id
- self._key_ring_id = key_ring_id
- self._key_id = key_id
- self._secret_version_name = f'HsmGeneratedSecret_{job_name}'
-
- def get_secret_bytes(self) -> bytes:
- """Retrieves the secret bytes.
-
- If the secret version already exists in Secret Manager, it is retrieved.
- Otherwise, a new secret and version are created. The new secret is
- generated using the HSM key.
-
- Returns:
- The secret as a byte string.
- """
- try:
- from google.api_core import exceptions as api_exceptions
- from google.cloud import secretmanager
- client = secretmanager.SecretManagerServiceClient()
-
- project_path = f"projects/{self._project_id}"
- secret_path = f"{project_path}/secrets/{self._secret_version_name}"
- # Since we may generate multiple versions when doing this on workers,
- # just always take the first version added to maintain consistency.
- secret_version_path = f"{secret_path}/versions/1"
-
- try:
- response = client.access_secret_version(
- request={"name": secret_version_path})
- return response.payload.data
- except api_exceptions.NotFound:
- # Don't bother logging yet, we'll only log if we actually add the
- # secret version below
- pass
-
- try:
- client.create_secret(
- request={
- "parent": project_path,
- "secret_id": self._secret_version_name,
- "secret": {
- "replication": {
- "automatic": {}
- }
- },
- })
- except api_exceptions.AlreadyExists:
- # Don't bother logging yet, we'll only log if we actually add the
- # secret version below
- pass
-
- new_key = self.generate_dek()
- try:
- # Try one more time in case it was created while we were generating the
- # DEK.
- response = client.access_secret_version(
- request={"name": secret_version_path})
- return response.payload.data
- except api_exceptions.NotFound:
- _LOGGER.info(
- "Secret version %s not found. "
- "Creating new secret and version.",
- secret_version_path)
- client.add_secret_version(
- request={
- "parent": secret_path, "payload": {
- "data": new_key
- }
- })
- response = client.access_secret_version(
- request={"name": secret_version_path})
- return response.payload.data
-
- except Exception as e:
- raise RuntimeError(
- f'Failed to retrieve or create secret bytes for secret '
- f'{self._secret_version_name} with exception {e}')
-
- def generate_dek(self, dek_size: int = 32) -> bytes:
- """Generates a new Data Encryption Key (DEK) using an HSM-backed key.
-
- This function follows a key derivation process that incorporates entropy
- from the HSM-backed key into the nonce used for key derivation.
-
- Args:
- dek_size: The size of the DEK to generate.
-
- Returns:
- A new DEK of the specified size, url-safe base64-encoded.
- """
- try:
- import base64
- import os
-
- from cryptography.hazmat.primitives import hashes
- from cryptography.hazmat.primitives.kdf.hkdf import HKDF
- from google.cloud import kms
-
- # 1. Generate a random nonce (nonce_one)
- nonce_one = os.urandom(dek_size)
-
- # 2. Use the HSM-backed key to encrypt nonce_one to create nonce_two
- kms_client = kms.KeyManagementServiceClient()
- key_path = kms_client.crypto_key_path(
- self._project_id, self._location_id, self._key_ring_id, self._key_id)
- response = kms_client.encrypt(
- request={
- 'name': key_path, 'plaintext': nonce_one
- })
- nonce_two = response.ciphertext
-
- # 3. Generate a Derivation Key (DK)
- dk = os.urandom(dek_size)
-
- # 4. Use a KDF to derive the DEK using DK and nonce_two
- hkdf = HKDF(
- algorithm=hashes.SHA256(),
- length=dek_size,
- salt=nonce_two,
- info=None,
- )
- dek = hkdf.derive(dk)
- return base64.urlsafe_b64encode(dek)
- except Exception as e:
- raise RuntimeError(f'Failed to generate DEK with exception {e}')
-
-
class _EncryptMessage(DoFn):
"""A DoFn that encrypts the key and value of each element."""
def __init__(
diff --git a/sdks/python/apache_beam/transforms/util_test.py
b/sdks/python/apache_beam/transforms/util_test.py
index 63ce42726c1..446fe68e594 100644
--- a/sdks/python/apache_beam/transforms/util_test.py
+++ b/sdks/python/apache_beam/transforms/util_test.py
@@ -72,9 +72,9 @@ from apache_beam.transforms import window
from apache_beam.transforms.core import FlatMapTuple
from apache_beam.transforms.trigger import AfterCount
from apache_beam.transforms.trigger import Repeatedly
-from apache_beam.transforms.util import GcpHsmGeneratedSecret
-from apache_beam.transforms.util import GcpSecret
-from apache_beam.transforms.util import Secret
+from apache_beam.utils.secret import GcpHsmGeneratedSecret
+from apache_beam.utils.secret import GcpSecret
+from apache_beam.utils.secret import Secret
from apache_beam.transforms.util import _BatchSizeEstimator
from apache_beam.transforms.util import _GlobalWindowsBatchingDoFn
from apache_beam.transforms.window import FixedWindows
@@ -287,37 +287,6 @@ class
MockNoOpDecrypt(beam.transforms.util._DecryptMessage):
return final_elements
-class SecretTest(unittest.TestCase):
- @parameterized.expand([
- param(
-
secret_string='type:GcpSecret;version_name:my_secret/versions/latest',
- secret=GcpSecret('my_secret/versions/latest')),
- param(
- secret_string='type:GcpSecret;version_name:foo',
- secret=GcpSecret('foo')),
- param(
-
secret_string='type:gcpsecreT;version_name:my_secret/versions/latest',
- secret=GcpSecret('my_secret/versions/latest')),
- ])
- def test_secret_manager_parses_correctly(self, secret_string, secret):
- self.assertEqual(secret, Secret.parse_secret_option(secret_string))
-
- @parameterized.expand([
- param(
- secret_string='version_name:foo',
- exception_str='must contain a valid type parameter'),
- param(
- secret_string='type:gcpsecreT',
- exception_str='missing 1 required positional argument'),
- param(
- secret_string='type:gcpsecreT;version_name:foo;extra:val',
- exception_str='Invalid secret parameter extra'),
- ])
- def test_secret_manager_throws_on_invalid(self, secret_string,
exception_str):
- with self.assertRaisesRegex(Exception, exception_str):
- Secret.parse_secret_option(secret_string)
-
-
class GroupByEncryptedKeyTest(unittest.TestCase):
@classmethod
def setUpClass(cls):
@@ -387,7 +356,7 @@ class GroupByEncryptedKeyTest(unittest.TestCase):
result, equal_to([('a', ([1, 2])), ('b', ([3])), ('c', ([4]))]))
@mock.patch('apache_beam.transforms.util._DecryptMessage', MockNoOpDecrypt)
- @mock.patch('apache_beam.transforms.util.GcpSecret', FakeSecret)
+ @mock.patch('apache_beam.utils.secret.GcpSecret', FakeSecret)
def test_gbk_actually_does_encryption(self):
options = PipelineOptions()
# Version of GcpSecret doesn't matter since it is replaced by FakeSecret
@@ -435,124 +404,6 @@ class GroupByEncryptedKeyTest(unittest.TestCase):
result, equal_to([('a', ([1, 2])), ('b', ([3])), ('c', ([4]))]))
[email protected](secretmanager is None, 'GCP dependencies are not installed')
-class GcpHsmGeneratedSecretTest(unittest.TestCase):
- def setUp(self):
- self.mock_secret_manager_client = mock.MagicMock()
- self.mock_kms_client = mock.MagicMock()
-
- # Patch the clients
- self.secretmanager_patcher = mock.patch(
- 'google.cloud.secretmanager.SecretManagerServiceClient',
- return_value=self.mock_secret_manager_client)
- self.kms_patcher = mock.patch(
- 'google.cloud.kms.KeyManagementServiceClient',
- return_value=self.mock_kms_client)
- self.os_urandom_patcher = mock.patch('os.urandom', return_value=b'0' * 32)
- self.hkdf_patcher = mock.patch(
- 'cryptography.hazmat.primitives.kdf.hkdf.HKDF.derive',
- return_value=b'derived_key')
-
- self.secretmanager_patcher.start()
- self.kms_patcher.start()
- self.os_urandom_patcher.start()
- self.hkdf_patcher.start()
-
- def tearDown(self):
- self.secretmanager_patcher.stop()
- self.kms_patcher.stop()
- self.os_urandom_patcher.stop()
- self.hkdf_patcher.stop()
-
- def test_happy_path_secret_creation(self):
- from google.api_core import exceptions as api_exceptions
-
- project_id = 'test-project'
- location_id = 'global'
- key_ring_id = 'test-key-ring'
- key_id = 'test-key'
- job_name = 'test-job'
-
- secret = GcpHsmGeneratedSecret(
- project_id, location_id, key_ring_id, key_id, job_name)
-
- # Mock responses for secret creation path
- self.mock_secret_manager_client.access_secret_version.side_effect = [
- api_exceptions.NotFound('not found'), # first check
- api_exceptions.NotFound('not found'), # second check
- mock.MagicMock(payload=mock.MagicMock(data=b'derived_key'))
- ]
- self.mock_kms_client.encrypt.return_value = mock.MagicMock(
- ciphertext=b'encrypted_nonce')
-
- secret_bytes = secret.get_secret_bytes()
- self.assertEqual(secret_bytes, b'derived_key')
-
- # Assertions on mocks
- secret_version_path = (
- f'projects/{project_id}/secrets/{secret._secret_version_name}'
- '/versions/1')
- self.mock_secret_manager_client.access_secret_version.assert_any_call(
- request={'name': secret_version_path})
- self.assertEqual(
- self.mock_secret_manager_client.access_secret_version.call_count, 3)
- self.mock_secret_manager_client.create_secret.assert_called_once()
- self.mock_kms_client.encrypt.assert_called_once()
- self.mock_secret_manager_client.add_secret_version.assert_called_once()
-
- def test_secret_already_exists(self):
- from google.api_core import exceptions as api_exceptions
-
- project_id = 'test-project'
- location_id = 'global'
- key_ring_id = 'test-key-ring'
- key_id = 'test-key'
- job_name = 'test-job'
-
- secret = GcpHsmGeneratedSecret(
- project_id, location_id, key_ring_id, key_id, job_name)
-
- # Mock responses for secret creation path
- self.mock_secret_manager_client.access_secret_version.side_effect = [
- api_exceptions.NotFound('not found'),
- api_exceptions.NotFound('not found'),
- mock.MagicMock(payload=mock.MagicMock(data=b'derived_key'))
- ]
- self.mock_secret_manager_client.create_secret.side_effect = (
- api_exceptions.AlreadyExists('exists'))
- self.mock_kms_client.encrypt.return_value = mock.MagicMock(
- ciphertext=b'encrypted_nonce')
-
- secret_bytes = secret.get_secret_bytes()
- self.assertEqual(secret_bytes, b'derived_key')
-
- # Assertions on mocks
- self.mock_secret_manager_client.create_secret.assert_called_once()
- self.mock_secret_manager_client.add_secret_version.assert_called_once()
-
- def test_secret_version_already_exists(self):
- project_id = 'test-project'
- location_id = 'global'
- key_ring_id = 'test-key-ring'
- key_id = 'test-key'
- job_name = 'test-job'
-
- secret = GcpHsmGeneratedSecret(
- project_id, location_id, key_ring_id, key_id, job_name)
-
- self.mock_secret_manager_client.access_secret_version.return_value = (
- mock.MagicMock(payload=mock.MagicMock(data=b'existing_dek')))
-
- secret_bytes = secret.get_secret_bytes()
- self.assertEqual(secret_bytes, b'existing_dek')
-
- # Assertions
- self.mock_secret_manager_client.access_secret_version.assert_called_once()
- self.mock_secret_manager_client.create_secret.assert_not_called()
- self.mock_secret_manager_client.add_secret_version.assert_not_called()
- self.mock_kms_client.encrypt.assert_not_called()
-
-
class FakeClock(object):
def __init__(self, now=time.time()):
self._now = now
diff --git a/sdks/python/apache_beam/utils/secret.py
b/sdks/python/apache_beam/utils/secret.py
new file mode 100644
index 00000000000..c9c13f1d60f
--- /dev/null
+++ b/sdks/python/apache_beam/utils/secret.py
@@ -0,0 +1,464 @@
+#
+# Licensed to the Apache Software Foundation (ASF) under one or more
+# contributor license agreements. See the NOTICE file distributed with
+# this work for additional information regarding copyright ownership.
+# The ASF licenses this file to You under the Apache License, Version 2.0
+# (the "License"); you may not use this file except in compliance with
+# the License. You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+#
+
+"""Interface and implementations for Secret providers in Apache Beam."""
+
+import abc
+import json
+import logging
+import os
+import warnings
+from typing import Any, Dict, Optional, Union
+
+_LOGGER = logging.getLogger(__name__)
+
+
+class Secret(abc.ABC):
+ """A secret management class used for handling sensitive data.
+
+ This class provides a generic interface for secret management.
Implementations
+ of this class should handle fetching secrets from a secret management system.
+ """
+ def __init__(self):
+ self._cached_secret_bytes: Optional[bytes] = None
+
+ def get_str(self, cacheSecret: bool = False) -> str:
+ """Retrieve secret value as string.
+
+ Args:
+ cacheSecret: If True, caches secret value in memory after first fetch.
+
+ Returns:
+ The retrieved secret value as string.
+ """
+ return self.get_bytes(cacheSecret=cacheSecret).decode("utf-8")
+
+ def get_bytes(self, cacheSecret: bool = False) -> bytes:
+ """Retrieve secret value as bytes.
+
+ Args:
+ cacheSecret: If True, caches secret value in memory after first fetch.
+
+ Returns:
+ The retrieved secret value as bytes.
+ """
+ if cacheSecret and getattr(self, '_cached_secret_bytes', None) is not None:
+ return self._cached_secret_bytes
+
+ secret_val_bytes = self.get_secret_bytes()
+
+ if cacheSecret:
+ self._cached_secret_bytes = secret_val_bytes
+
+ return secret_val_bytes
+
+ @abc.abstractmethod
+ def get_secret_bytes(self) -> bytes:
+ """Returns the secret as a byte string."""
+ raise NotImplementedError()
+
+ def __getstate__(self):
+ """Strip cached secrets before pickling for pipeline
submission/transmission."""
+ state = self.__dict__.copy()
+ state['_cached_secret_bytes'] = None
+ return state
+
+ @staticmethod
+ def generate_secret_bytes() -> bytes:
+ """Generates a new secret key using Fernet."""
+ from cryptography.fernet import Fernet
+ return Fernet.generate_key()
+
+ @classmethod
+ def parse_secret_option(cls, secret: str) -> 'Secret':
+ """Parses a secret string and returns the appropriate secret type.
+
+ The secret string should be formatted like:
+ 'type:<secret_type>;<secret_param>:<value>'
+
+ For example, 'type:GcpSecret;version_name:my_secret/versions/latest'
+ would return a GcpSecret initialized with 'my_secret/versions/latest'.
+ """
+ param_map = {}
+ for param in secret.split(';'):
+ parts = param.split(':')
+ if len(parts) == 2:
+ param_map[parts[0]] = parts[1]
+
+ if 'type' not in param_map:
+ raise ValueError('Secret string must contain a valid type parameter')
+
+ raw_type = param_map.pop('type')
+ secret_type = raw_type.lower()
+ secret_manager = _SECRET_TYPE_TO_SECRET_MANAGER.get(secret_type)
+ if not secret_manager:
+ raise ValueError(
+ f'Invalid secret type {secret_type}, currently only '
+ 'GcpSecret and GcpHsmGeneratedSecret are supported')
+
+ return cls.from_json(json.dumps(param_map), secret_manager)
+
+ @classmethod
+ def from_json(
+ cls, spec: str, secret_manager: Optional[str] = None) -> 'Secret':
+ """Return a Secret instance based on secret_manager provider and secret
specification.
+
+ Args:
+ spec: Secret string (raw secret or JSON specification string).
+ secret_manager: Secret manager string (e.g. 'GoogleCloudSecretManager').
+
+ Returns:
+ An instance of Secret.
+ """
+ if not isinstance(spec, str):
+ raise TypeError(
+ f"Secret 'spec' must be a string, got {type(spec).__name__}")
+
+ secret_manager_name = (
+ secret_manager.strip()
+ if secret_manager and secret_manager.strip() else None)
+
+ spec_dict = None
+ try:
+ spec_dict = json.loads(spec)
+ if not isinstance(spec_dict, dict):
+ spec_dict = None
+ except Exception:
+ try:
+ import ast
+ spec_dict = ast.literal_eval(spec)
+ if not isinstance(spec_dict, dict):
+ spec_dict = None
+ except Exception:
+ pass
+
+ if secret_manager_name:
+ secret_cls_entry = _SECRET_CLASSES.get(secret_manager_name.lower())
+ if secret_cls_entry:
+ if isinstance(secret_cls_entry, str):
+ secret_cls = globals().get(secret_cls_entry, secret_cls_entry)
+ else:
+ secret_cls = secret_cls_entry
+ if isinstance(spec_dict, dict) and hasattr(secret_cls, 'from_dict'):
+ return secret_cls.from_dict(spec_dict)
+ elif isinstance(spec_dict, dict):
+ return secret_cls(**spec_dict)
+ else:
+ return secret_cls(spec)
+ else:
+ raise ValueError(
+ f"Unsupported secret manager: '{secret_manager_name}'. Currently
supported options: 'GoogleCloudSecretManager',
'GoogleCloudHsmGeneratedSecretManager'."
+ )
+
+ # If secret_manager is not set or empty, check if spec is a JSON
specification dict
+ if spec_dict is not None:
+ msg = (
+ "The 'spec' parameter appears to be a JSON specification, but "
+ "'secret_manager' is not set. Defaulting to Raw.")
+ _LOGGER.warning(msg)
+ warnings.warn(msg, UserWarning)
+
+ return RawSecret(spec)
+
+
+class RawSecret(Secret):
+ """Secret implementation wrapping a raw secret string or bytes directly."""
+ def __init__(self, secret: Union[str, bytes]):
+ super().__init__()
+ if isinstance(secret, str):
+ self._secret = secret.encode("utf-8")
+ else:
+ self._secret = secret
+
+ def get_secret_bytes(self) -> bytes:
+ return self._secret
+
+ def __eq__(self, other: Any) -> bool:
+ if not isinstance(other, RawSecret):
+ return False
+ return self._secret == other._secret
+
+
+class GcpSecret(Secret):
+ """A secret manager implementation that retrieves secrets from Google Cloud
+ Secret Manager.
+ """
+ def __init__(self, version_name: str):
+ """Initializes a GcpSecret object.
+
+ Args:
+ version_name: The full version name of the secret in Google Cloud Secret
+ Manager. For example:
+ projects/<id>/secrets/<secret_name>/versions/1.
+ For more info, see
+
https://cloud.google.com/python/docs/reference/secretmanager/latest/google.cloud.secretmanager_v1beta1.services.secret_manager_service.SecretManagerServiceClient#google_cloud_secretmanager_v1beta1_services_secret_manager_service_SecretManagerServiceClient_access_secret_version
+ """
+ super().__init__()
+ self._version_name = version_name
+
+ @classmethod
+ def from_dict(cls, spec_dict: Dict[str, str]) -> 'GcpSecret':
+ """Initialize GcpSecret from a dictionary specification."""
+ allowed_keys = {'version_name', 'name', 'project', 'version'}
+ invalid_keys = set(spec_dict.keys()) - allowed_keys
+ if invalid_keys:
+ raise ValueError(
+ f"Invalid secret parameter {', '.join(sorted(invalid_keys))}")
+ version_name = cls._parse_version_name(spec_dict)
+ return cls(version_name)
+
+ @classmethod
+ def _parse_version_name(cls, spec_dict: Dict[str, str]) -> str:
+ if "version_name" in spec_dict:
+ return spec_dict["version_name"]
+
+ secret_id = spec_dict.get("name")
+ if not secret_id:
+ raise ValueError("Secret name must be specified in secret spec.")
+
+ # Resolve project ID from spec, environment variables, or Application
Default Credentials
+ project_id = (
+ spec_dict.get("project") or os.environ.get("GOOGLE_CLOUD_PROJECT") or
+ os.environ.get("GCP_PROJECT"))
+
+ if not project_id:
+ try:
+ import google.auth
+ _, project_id = google.auth.default()
+ except Exception:
+ pass
+
+ version_id = spec_dict.get("version", "latest")
+
+ if not project_id:
+ raise ValueError(
+ f"Could not resolve GCP project ID for secret '{secret_id}'. "
+ "Please specify 'project' in the secret spec, set
GOOGLE_CLOUD_PROJECT environment variable, "
+ "or configure Application Default Credentials.")
+
+ return f"projects/{project_id}/secrets/{secret_id}/versions/{version_id}"
+
+ def get_secret_bytes(self) -> bytes:
+ try:
+ from google.cloud import secretmanager
+ client = secretmanager.SecretManagerServiceClient()
+ response = client.access_secret_version(
+ request={"name": self._version_name})
+ secret = response.payload.data
+ return secret
+ except Exception as e:
+ raise RuntimeError(
+ 'Failed to retrieve secret bytes for secret '
+ f'{self._version_name} with exception {e}')
+
+ def __eq__(self, secret):
+ return self._version_name == getattr(secret, '_version_name', None)
+
+
+class GcpHsmGeneratedSecret(Secret):
+ """A secret manager implementation that generates a secret using a GCP HSM
key
+ and stores it in Google Cloud Secret Manager. If the secret already exists,
+ it will be retrieved.
+ """
+ def __init__(
+ self,
+ project_id: str,
+ location_id: str,
+ key_ring_id: str,
+ key_id: str,
+ job_name: str):
+ """Initializes a GcpHsmGeneratedSecret object.
+
+ Args:
+ project_id: The GCP project ID.
+ location_id: The GCP location ID for the HSM key.
+ key_ring_id: The ID of the KMS key ring.
+ key_id: The ID of the KMS key.
+ job_name: The name of the job, used to generate a unique secret name.
+ """
+ super().__init__()
+ self._project_id = project_id
+ self._location_id = location_id
+ self._key_ring_id = key_ring_id
+ self._key_id = key_id
+ self._job_name = job_name
+ self._secret_version_name = f'HsmGeneratedSecret_{job_name}'
+
+ def __eq__(self, other: Any) -> bool:
+ if not isinstance(other, GcpHsmGeneratedSecret):
+ return False
+ return (
+ self._project_id == other._project_id and
+ self._location_id == other._location_id and
+ self._key_ring_id == other._key_ring_id and
+ self._key_id == other._key_id and
+ getattr(self, '_job_name', None) == getattr(other, '_job_name', None))
+
+ @classmethod
+ def from_dict(cls, spec_dict: Dict[str, str]) -> 'GcpHsmGeneratedSecret':
+ """Initialize GcpHsmGeneratedSecret from a dictionary specification."""
+ allowed_keys = {
+ 'project_id', 'location_id', 'key_ring_id', 'key_id', 'job_name'
+ }
+ missing = allowed_keys - set(spec_dict.keys())
+ if missing:
+ raise ValueError(
+ f"Missing required parameter(s) for GcpHsmGeneratedSecret:
{sorted(list(missing))}"
+ )
+ invalid_keys = set(spec_dict.keys()) - allowed_keys
+ if invalid_keys:
+ raise ValueError(
+ f"Invalid secret parameter {', '.join(sorted(invalid_keys))}")
+ return cls(
+ project_id=spec_dict['project_id'],
+ location_id=spec_dict['location_id'],
+ key_ring_id=spec_dict['key_ring_id'],
+ key_id=spec_dict['key_id'],
+ job_name=spec_dict['job_name'],
+ )
+
+ def get_secret_bytes(self) -> bytes:
+ """Retrieves the secret bytes.
+
+ If the secret version already exists in Secret Manager, it is retrieved.
+ Otherwise, a new secret and version are created. The new secret is
+ generated using the HSM key.
+
+ Returns:
+ The secret as a byte string.
+ """
+ try:
+ from google.api_core import exceptions as api_exceptions
+ from google.cloud import secretmanager
+ client = secretmanager.SecretManagerServiceClient()
+
+ project_path = f"projects/{self._project_id}"
+ secret_path = f"{project_path}/secrets/{self._secret_version_name}"
+ # Since we may generate multiple versions when doing this on workers,
+ # just always take the first version added to maintain consistency.
+ secret_version_path = f"{secret_path}/versions/1"
+
+ try:
+ response = client.access_secret_version(
+ request={"name": secret_version_path})
+ return response.payload.data
+ except api_exceptions.NotFound:
+ # Don't bother logging yet, we'll only log if we actually add the
+ # secret version below
+ pass
+
+ try:
+ client.create_secret(
+ request={
+ "parent": project_path,
+ "secret_id": self._secret_version_name,
+ "secret": {
+ "replication": {
+ "automatic": {}
+ }
+ },
+ })
+ except api_exceptions.AlreadyExists:
+ # Don't bother logging yet, we'll only log if we actually add the
+ # secret version below
+ pass
+
+ new_key = self.generate_dek()
+ try:
+ # Try one more time in case it was created while we were generating the
+ # DEK.
+ response = client.access_secret_version(
+ request={"name": secret_version_path})
+ return response.payload.data
+ except api_exceptions.NotFound:
+ _LOGGER.info(
+ "Secret version %s not found. "
+ "Creating new secret and version.",
+ secret_version_path)
+ client.add_secret_version(
+ request={
+ "parent": secret_path, "payload": {
+ "data": new_key
+ }
+ })
+ response = client.access_secret_version(
+ request={"name": secret_version_path})
+ return response.payload.data
+
+ except Exception as e:
+ raise RuntimeError(
+ f'Failed to retrieve or create secret bytes for secret '
+ f'{self._secret_version_name} with exception {e}')
+
+ def generate_dek(self, dek_size: int = 32) -> bytes:
+ """Generates a new Data Encryption Key (DEK) using an HSM-backed key.
+
+ This function follows a key derivation process that incorporates entropy
+ from the HSM-backed key into the nonce used for key derivation.
+
+ Args:
+ dek_size: The size of the DEK to generate.
+
+ Returns:
+ A new DEK of the specified size, url-safe base64-encoded.
+ """
+ try:
+ import base64
+ import os
+
+ from cryptography.hazmat.primitives import hashes
+ from cryptography.hazmat.primitives.kdf.hkdf import HKDF
+ from google.cloud import kms
+
+ # 1. Generate a random nonce (nonce_one)
+ nonce_one = os.urandom(dek_size)
+
+ # 2. Use the HSM-backed key to encrypt nonce_one to create nonce_two
+ kms_client = kms.KeyManagementServiceClient()
+ key_path = kms_client.crypto_key_path(
+ self._project_id, self._location_id, self._key_ring_id, self._key_id)
+ response = kms_client.encrypt(
+ request={
+ 'name': key_path, 'plaintext': nonce_one
+ })
+ nonce_two = response.ciphertext
+
+ # 3. Generate a Derivation Key (DK)
+ dk = os.urandom(dek_size)
+
+ # 4. Use a KDF to derive the DEK using DK and nonce_two
+ hkdf = HKDF(
+ algorithm=hashes.SHA256(),
+ length=dek_size,
+ salt=nonce_two,
+ info=None,
+ )
+ dek = hkdf.derive(dk)
+ return base64.urlsafe_b64encode(dek)
+ except Exception as e:
+ raise RuntimeError(f'Failed to generate DEK with exception {e}')
+
+
+_SECRET_TYPE_TO_SECRET_MANAGER: Dict[str, str] = {
+ "gcpsecret": "GoogleCloudSecretManager",
+ "gcphsmgeneratedsecret": "GoogleCloudHsmGeneratedSecretManager",
+}
+
+_SECRET_CLASSES: Dict[str, Any] = {
+ "googlecloudsecretmanager": "GcpSecret",
+ "googlecloudhsmgeneratedsecretmanager": "GcpHsmGeneratedSecret",
+}
\ No newline at end of file
diff --git a/sdks/python/apache_beam/utils/secret_test.py
b/sdks/python/apache_beam/utils/secret_test.py
new file mode 100644
index 00000000000..179b7ca3a5f
--- /dev/null
+++ b/sdks/python/apache_beam/utils/secret_test.py
@@ -0,0 +1,454 @@
+#
+# Licensed to the Apache Software Foundation (ASF) under one or more
+# contributor license agreements. See the NOTICE file distributed with
+# this work for additional information regarding copyright ownership.
+# The ASF licenses this file to You under the Apache License, Version 2.0
+# (the "License"); you may not use this file except in compliance with
+# the License. You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+#
+
+import json
+import unittest
+from unittest import mock
+
+from parameterized import param
+from parameterized import parameterized
+
+from apache_beam.utils.annotations import BeamDeprecationWarning
+from apache_beam.utils.secret import GcpHsmGeneratedSecret
+from apache_beam.utils.secret import GcpSecret
+from apache_beam.utils.secret import RawSecret
+from apache_beam.utils.secret import Secret
+
+try:
+ from google.cloud import secretmanager
+except ImportError:
+ secretmanager = None # type: ignore[assignment]
+
+
+class SecretTest(unittest.TestCase):
+ @parameterized.expand([
+ param(
+
secret_string='type:GcpSecret;version_name:my_secret/versions/latest',
+ secret=GcpSecret('my_secret/versions/latest')),
+ param(
+ secret_string='type:GcpSecret;version_name:foo',
+ secret=GcpSecret('foo')),
+ param(
+
secret_string='type:gcpsecreT;version_name:my_secret/versions/latest',
+ secret=GcpSecret('my_secret/versions/latest')),
+ ])
+ def test_secret_manager_parses_correctly(self, secret_string, secret):
+ self.assertEqual(secret, Secret.parse_secret_option(secret_string))
+
+ @parameterized.expand([
+ param(
+ secret_string='version_name:foo',
+ exception_str='must contain a valid type parameter'),
+ param(
+ secret_string='type:gcpsecreT',
+ exception_str='Secret name must be specified in secret spec'),
+ param(
+ secret_string='type:gcpsecreT;version_name:foo;extra:val',
+ exception_str='Invalid secret parameter extra'),
+ ])
+ def test_secret_manager_throws_on_invalid(self, secret_string,
exception_str):
+ with self.assertRaisesRegex(Exception, exception_str):
+ Secret.parse_secret_option(secret_string)
+
+
[email protected](secretmanager is None, 'GCP dependencies are not installed')
+class GcpSecretTest(unittest.TestCase):
+ @mock.patch("google.cloud.secretmanager.SecretManagerServiceClient")
+ def test_gcp_secret_success(self, mock_client_cls):
+ mock_client = mock.MagicMock()
+ mock_client_cls.return_value = mock_client
+ mock_response = mock.MagicMock()
+ mock_response.payload.data = b"secret-payload-value"
+ mock_client.access_secret_version.return_value = mock_response
+
+ spec_dict = {"name": "my-secret", "version": "1", "project": "my-project"}
+ secret = GcpSecret.from_dict(spec_dict)
+
+ secret_val = secret.get_str(cacheSecret=True)
+ self.assertEqual(secret_val, "secret-payload-value")
+ secret_bytes = secret.get_bytes(cacheSecret=True)
+ self.assertEqual(secret_bytes, b"secret-payload-value")
+ mock_client.access_secret_version.assert_called_once_with(
+ request={"name": "projects/my-project/secrets/my-secret/versions/1"})
+
+ # Second call with cacheSecret=True should return cached value without
calling client again
+ mock_client.reset_mock()
+ secret_val_cached = secret.get_str(cacheSecret=True)
+ self.assertEqual(secret_val_cached, "secret-payload-value")
+ mock_client.access_secret_version.assert_not_called()
+
+ @mock.patch("google.cloud.secretmanager.SecretManagerServiceClient")
+ def test_gcp_secret_get_bytes_uncached(self, mock_client_cls):
+ mock_client = mock.MagicMock()
+ mock_client_cls.return_value = mock_client
+ mock_response = mock.MagicMock()
+ mock_response.payload.data = b"secret-payload-value"
+ mock_client.access_secret_version.return_value = mock_response
+
+ spec_dict = {"name": "my-secret", "project": "my-project"}
+ secret = GcpSecret.from_dict(spec_dict)
+
+ secret_bytes = secret.get_bytes()
+ self.assertEqual(secret_bytes, b"secret-payload-value")
+ self.assertIsNone(secret._cached_secret_bytes)
+
+ @mock.patch("google.cloud.secretmanager.SecretManagerServiceClient")
+ def test_gcp_secret_getstate_clears_cached_secret(self, mock_client_cls):
+ mock_client = mock.MagicMock()
+ mock_client_cls.return_value = mock_client
+ mock_response = mock.MagicMock()
+ mock_response.payload.data = b"secret-payload-value"
+ mock_client.access_secret_version.return_value = mock_response
+
+ spec_dict = {"name": "my-secret", "project": "my-project"}
+ secret = GcpSecret.from_dict(spec_dict)
+
+ # Cache the secret in memory
+ secret.get_str(cacheSecret=True)
+ self.assertEqual(secret._cached_secret_bytes, b"secret-payload-value")
+
+ # When pickled / getstate is called during pipeline submission
+ state = secret.__getstate__()
+ self.assertIsNone(state["_cached_secret_bytes"])
+
+ @mock.patch.dict("os.environ", {"GOOGLE_CLOUD_PROJECT": "env-project-123"})
+ @mock.patch("google.cloud.secretmanager.SecretManagerServiceClient")
+ def test_gcp_secret_env_project_fallback(self, mock_client_cls):
+ mock_client = mock.MagicMock()
+ mock_client_cls.return_value = mock_client
+ mock_response = mock.MagicMock()
+ mock_response.payload.data = b"env-secret-val"
+ mock_client.access_secret_version.return_value = mock_response
+
+ # Project omitted from spec
+ spec_dict = {"name": "env-secret", "version": "latest"}
+ secret = GcpSecret.from_dict(spec_dict)
+
+ secret_val = secret.get_str(cacheSecret=False)
+ self.assertEqual(secret_val, "env-secret-val")
+ self.assertEqual(secret.get_bytes(cacheSecret=False), b"env-secret-val")
+ mock_client.access_secret_version.assert_called_with(
+ request={
+ "name":
"projects/env-project-123/secrets/env-secret/versions/latest"
+ })
+
+ @mock.patch("google.cloud.secretmanager.SecretManagerServiceClient")
+ def test_gcp_secret_failure_raises_exception(self, mock_client_cls):
+ mock_client = mock.MagicMock()
+ mock_client_cls.return_value = mock_client
+ mock_client.access_secret_version.side_effect = RuntimeError(
+ "Permission denied or secret not found")
+
+ spec_dict = {"name": "non-existent-secret", "project": "my-project"}
+ secret = GcpSecret.from_dict(spec_dict)
+
+ with self.assertRaises(RuntimeError) as ctx:
+ secret.get_str(cacheSecret=False)
+ self.assertIn("Permission denied or secret not found", str(ctx.exception))
+
+ @mock.patch.dict("os.environ", {}, clear=True)
+ @mock.patch("google.auth.default", side_effect=Exception("No ADC"))
+ def test_ill_formed_missing_project_raises_value_error(
+ self, mock_auth_default):
+ spec_dict = {"name": "my-secret"}
+ with self.assertRaises(ValueError) as ctx:
+ GcpSecret.from_dict(spec_dict)
+ self.assertIn("Could not resolve GCP project ID", str(ctx.exception))
+
+ def test_ill_formed_missing_secret_name_raises_value_error(self):
+ spec_dict = {"project": "my-project"}
+ with self.assertRaises(ValueError) as ctx:
+ GcpSecret.from_dict(spec_dict)
+ self.assertIn("Secret name must be specified", str(ctx.exception))
+
+
[email protected](secretmanager is None, 'GCP dependencies are not installed')
+class GcpHsmGeneratedSecretTest(unittest.TestCase):
+ def setUp(self):
+ self.mock_secret_manager_client = mock.MagicMock()
+ self.mock_kms_client = mock.MagicMock()
+
+ # Patch the clients
+ self.secretmanager_patcher = mock.patch(
+ 'google.cloud.secretmanager.SecretManagerServiceClient',
+ return_value=self.mock_secret_manager_client)
+ self.kms_patcher = mock.patch(
+ 'google.cloud.kms.KeyManagementServiceClient',
+ return_value=self.mock_kms_client)
+ self.os_urandom_patcher = mock.patch('os.urandom', return_value=b'0' * 32)
+ self.hkdf_patcher = mock.patch(
+ 'cryptography.hazmat.primitives.kdf.hkdf.HKDF.derive',
+ return_value=b'derived_key')
+
+ self.secretmanager_patcher.start()
+ self.kms_patcher.start()
+ self.os_urandom_patcher.start()
+ self.hkdf_patcher.start()
+
+ def tearDown(self):
+ self.secretmanager_patcher.stop()
+ self.kms_patcher.stop()
+ self.os_urandom_patcher.stop()
+ self.hkdf_patcher.stop()
+
+ def test_happy_path_secret_creation(self):
+ from google.api_core import exceptions as api_exceptions
+
+ project_id = 'test-project'
+ location_id = 'global'
+ key_ring_id = 'test-key-ring'
+ key_id = 'test-key'
+ job_name = 'test-job'
+
+ secret = GcpHsmGeneratedSecret(
+ project_id, location_id, key_ring_id, key_id, job_name)
+
+ # Mock responses for secret creation path
+ self.mock_secret_manager_client.access_secret_version.side_effect = [
+ api_exceptions.NotFound('not found'), # first check
+ api_exceptions.NotFound('not found'), # second check
+ mock.MagicMock(payload=mock.MagicMock(data=b'derived_key'))
+ ]
+ self.mock_kms_client.encrypt.return_value = mock.MagicMock(
+ ciphertext=b'encrypted_nonce')
+
+ secret_bytes = secret.get_secret_bytes()
+ self.assertEqual(secret_bytes, b'derived_key')
+
+ # Assertions on mocks
+ secret_version_path = (
+ f'projects/{project_id}/secrets/{secret._secret_version_name}'
+ '/versions/1')
+ self.mock_secret_manager_client.access_secret_version.assert_any_call(
+ request={'name': secret_version_path})
+ self.assertEqual(
+ self.mock_secret_manager_client.access_secret_version.call_count, 3)
+ self.mock_secret_manager_client.create_secret.assert_called_once()
+ self.mock_kms_client.encrypt.assert_called_once()
+ self.mock_secret_manager_client.add_secret_version.assert_called_once()
+
+ def test_secret_already_exists(self):
+ from google.api_core import exceptions as api_exceptions
+
+ project_id = 'test-project'
+ location_id = 'global'
+ key_ring_id = 'test-key-ring'
+ key_id = 'test-key'
+ job_name = 'test-job'
+
+ secret = GcpHsmGeneratedSecret(
+ project_id, location_id, key_ring_id, key_id, job_name)
+
+ # Mock responses for secret creation path
+ self.mock_secret_manager_client.access_secret_version.side_effect = [
+ api_exceptions.NotFound('not found'),
+ api_exceptions.NotFound('not found'),
+ mock.MagicMock(payload=mock.MagicMock(data=b'derived_key'))
+ ]
+ self.mock_secret_manager_client.create_secret.side_effect = (
+ api_exceptions.AlreadyExists('exists'))
+ self.mock_kms_client.encrypt.return_value = mock.MagicMock(
+ ciphertext=b'encrypted_nonce')
+
+ secret_bytes = secret.get_secret_bytes()
+ self.assertEqual(secret_bytes, b'derived_key')
+
+ # Assertions on mocks
+ self.mock_secret_manager_client.create_secret.assert_called_once()
+ self.mock_secret_manager_client.add_secret_version.assert_called_once()
+
+ def test_secret_version_already_exists(self):
+ project_id = 'test-project'
+ location_id = 'global'
+ key_ring_id = 'test-key-ring'
+ key_id = 'test-key'
+ job_name = 'test-job'
+
+ secret = GcpHsmGeneratedSecret(
+ project_id, location_id, key_ring_id, key_id, job_name)
+
+ self.mock_secret_manager_client.access_secret_version.return_value = (
+ mock.MagicMock(payload=mock.MagicMock(data=b'existing_dek')))
+
+ secret_bytes = secret.get_secret_bytes()
+ self.assertEqual(secret_bytes, b'existing_dek')
+
+ # Assertions
+ self.mock_secret_manager_client.access_secret_version.assert_called_once()
+ self.mock_secret_manager_client.create_secret.assert_not_called()
+ self.mock_secret_manager_client.add_secret_version.assert_not_called()
+ self.mock_kms_client.encrypt.assert_not_called()
+
+ def test_from_dict_success(self):
+ spec_dict = {
+ "project_id": "test-proj",
+ "location_id": "global",
+ "key_ring_id": "ring",
+ "key_id": "key",
+ "job_name": "my-job"
+ }
+ secret = GcpHsmGeneratedSecret.from_dict(spec_dict)
+ self.assertEqual(secret._project_id, "test-proj")
+ self.assertEqual(secret._location_id, "global")
+ self.assertEqual(secret._key_ring_id, "ring")
+ self.assertEqual(secret._key_id, "key")
+ self.assertEqual(secret._job_name, "my-job")
+ self.assertEqual(secret._secret_version_name, "HsmGeneratedSecret_my-job")
+
+ def test_from_dict_missing_params_raises_value_error(self):
+ spec_dict = {"project_id": "test-proj", "location_id": "global"}
+ with self.assertRaises(ValueError) as ctx:
+ GcpHsmGeneratedSecret.from_dict(spec_dict)
+ self.assertIn("Missing required parameter(s)", str(ctx.exception))
+
+ @mock.patch("google.cloud.secretmanager.SecretManagerServiceClient")
+ def test_get_bytes_cached(self, mock_sm_client_cls):
+ mock_client = mock.MagicMock()
+ mock_sm_client_cls.return_value = mock_client
+ mock_response = mock.MagicMock()
+ mock_response.payload.data = b"hsm-derived-key"
+ mock_client.access_secret_version.return_value = mock_response
+
+ secret = GcpHsmGeneratedSecret("p", "l", "r", "k", "j")
+ secret_bytes = secret.get_bytes(cacheSecret=True)
+ self.assertEqual(secret_bytes, b"hsm-derived-key")
+
+ # Second call uses cache
+ mock_client.reset_mock()
+ self.assertEqual(secret.get_bytes(cacheSecret=True), b"hsm-derived-key")
+ mock_client.access_secret_version.assert_not_called()
+
+ @mock.patch("google.cloud.secretmanager.SecretManagerServiceClient")
+ def test_getstate_clears_cached_secret(self, mock_sm_client_cls):
+ mock_client = mock.MagicMock()
+ mock_sm_client_cls.return_value = mock_client
+ mock_response = mock.MagicMock()
+ mock_response.payload.data = b"hsm-derived-key"
+ mock_client.access_secret_version.return_value = mock_response
+
+ secret = GcpHsmGeneratedSecret("p", "l", "r", "k", "j")
+ secret.get_bytes(cacheSecret=True)
+ self.assertEqual(secret._cached_secret_bytes, b"hsm-derived-key")
+
+ state = secret.__getstate__()
+ self.assertIsNone(state["_cached_secret_bytes"])
+
+
+class RawSecretTest(unittest.TestCase):
+ def test_raw_secret_str(self):
+ secret = RawSecret("STATIC_SECRET_")
+ self.assertEqual(secret.get_str(cacheSecret=True), "STATIC_SECRET_")
+ self.assertEqual(secret.get_bytes(cacheSecret=True), b"STATIC_SECRET_")
+
+ def test_raw_secret_bytes(self):
+ secret = RawSecret(b"STATIC_BYTES_")
+ self.assertEqual(secret.get_str(cacheSecret=True), "STATIC_BYTES_")
+ self.assertEqual(secret.get_bytes(cacheSecret=True), b"STATIC_BYTES_")
+
+
+class SecretFactoryTest(unittest.TestCase):
+ def test_secret_factory(self):
+ spec = json.dumps({"name": "test-secret", "project": "proj"})
+
+ # When provider is set to 'GoogleCloudSecretManager'
+ secret_gcp = Secret.from_json(
+ spec=spec, secret_manager="GoogleCloudSecretManager")
+ self.assertIsInstance(secret_gcp, GcpSecret)
+
+ # When spec is a valid JSON string
+ single_quoted_spec = "{\"name\": \"test-secret\", \"project\": \"proj\"}"
+ secret_single_quoted = Secret.from_json(
+ spec=single_quoted_spec, secret_manager="GoogleCloudSecretManager")
+ self.assertIsInstance(secret_single_quoted, GcpSecret)
+ self.assertEqual(
+ secret_single_quoted._version_name,
+ "projects/proj/secrets/test-secret/versions/latest")
+
+ # When spec is a single-quoted JSON string, we still allow it for
convienence
+ # though it is not a valid JSON string.
+ single_quoted_spec = "{'name': 'test-secret', 'project': 'proj'}"
+ secret_single_quoted = Secret.from_json(
+ spec=single_quoted_spec, secret_manager="GoogleCloudSecretManager")
+ self.assertIsInstance(secret_single_quoted, GcpSecret)
+ self.assertEqual(
+ secret_single_quoted._version_name,
+ "projects/proj/secrets/test-secret/versions/latest")
+
+ # When provider is None or empty with plain string
+ secret_raw = Secret.from_json(spec="STATIC_SECRET_", secret_manager=None)
+ self.assertIsInstance(secret_raw, RawSecret)
+
+ # Unsupported provider raises ValueError
+ with self.assertRaises(ValueError):
+ Secret.from_json(spec="spec", secret_manager="unsupported_provider")
+
+ # Non-string spec raises TypeError
+ spec_dict = {"name": "test-secret"}
+ with self.assertRaises(TypeError):
+ Secret.from_json(
+ spec=spec_dict, # type: ignore[arg-type]
+ secret_manager="GoogleCloudSecretManager")
+
+ def test_secret_factory_hsm(self):
+ hsm_spec = json.dumps({
+ "project_id": "p",
+ "location_id": "l",
+ "key_ring_id": "r",
+ "key_id": "k",
+ "job_name": "j"
+ })
+ secret_hsm = Secret.from_json(
+ spec=hsm_spec, secret_manager="GoogleCloudHsmGeneratedSecretManager")
+ self.assertIsInstance(secret_hsm, GcpHsmGeneratedSecret)
+ self.assertEqual(secret_hsm._project_id, "p")
+
+ def test_json_secret_without_secret_manager_warning(self):
+ json_spec = json.dumps({"name": "my-secret", "project": "my-proj"})
+ with self.assertWarns(UserWarning):
+ secret = Secret.from_json(spec=json_spec, secret_manager=None)
+ self.assertIsInstance(secret, RawSecret)
+
+ def test_generate_secret_bytes(self):
+ key = Secret.generate_secret_bytes()
+ self.assertIsInstance(key, bytes)
+ self.assertTrue(len(key) > 0)
+
+ def test_equality(self):
+ raw1 = RawSecret("secret_value")
+ raw2 = RawSecret("secret_value")
+ raw3 = RawSecret("other_value")
+ self.assertEqual(raw1, raw2)
+ self.assertNotEqual(raw1, raw3)
+ self.assertNotEqual(raw1, "secret_value")
+
+ gcp1 = GcpSecret.from_dict({"name": "sec", "project": "proj"})
+ gcp2 = GcpSecret.from_dict({"name": "sec", "project": "proj"})
+ gcp3 = GcpSecret.from_dict({"name": "other", "project": "proj"})
+ self.assertEqual(gcp1, gcp2)
+ self.assertNotEqual(gcp1, gcp3)
+ self.assertNotEqual(gcp1, raw1)
+
+ hsm1 = GcpHsmGeneratedSecret("p", "l", "r", "k", "j")
+ hsm2 = GcpHsmGeneratedSecret("p", "l", "r", "k", "j")
+ hsm3 = GcpHsmGeneratedSecret("p", "l", "r", "k", "other")
+ self.assertEqual(hsm1, hsm2)
+ self.assertNotEqual(hsm1, hsm3)
+ self.assertNotEqual(hsm1, gcp1)
+
+
+if __name__ == "__main__":
+ unittest.main()