diff --git a/airflow/providers/google/cloud/utils/credentials_provider.py b/airflow/providers/google/cloud/utils/credentials_provider.py index f147d8fff9d96..1190b23bdef6f 100644 --- a/airflow/providers/google/cloud/utils/credentials_provider.py +++ b/airflow/providers/google/cloud/utils/credentials_provider.py @@ -205,6 +205,7 @@ def __init__( disable_logging: bool = False, target_principal: str | None = None, delegates: Sequence[str] | None = None, + lifetime: int | None = None, ) -> None: super().__init__() key_options = [key_path, key_secret_name, keyfile_dict] @@ -223,6 +224,7 @@ def __init__( self.disable_logging = disable_logging self.target_principal = target_principal self.delegates = delegates + self.lifetime = lifetime def get_credentials_and_project(self) -> tuple[google.auth.credentials.Credentials, str]: """ @@ -261,6 +263,16 @@ def get_credentials_and_project(self) -> tuple[google.auth.credentials.Credentia project_id = _get_project_id_from_service_account_email(self.target_principal) + if self.lifetime: + # Create new impersonated credentials with an extended expiration time + credentials = impersonated_credentials.Credentials( + source_credentials=credentials.source_credentials, + target_principal=credentials.target_principal, + delegates=credentials.delegates, + target_scopes=credentials.target_scopes, + lifetime=self.lifetime, + ) + return credentials, project_id def _get_credentials_using_keyfile_dict(self): diff --git a/airflow/providers/google/common/hooks/base_google.py b/airflow/providers/google/common/hooks/base_google.py index e68cb074e6a57..69a59f03a9bce 100644 --- a/airflow/providers/google/common/hooks/base_google.py +++ b/airflow/providers/google/common/hooks/base_google.py @@ -44,6 +44,7 @@ from googleapiclient import discovery from googleapiclient.errors import HttpError from googleapiclient.http import MediaIoBaseDownload, build_http, set_user_agent +from wtforms.validators import Optional from airflow import version from airflow.exceptions import AirflowException, AirflowProviderDeprecationWarning @@ -214,6 +215,11 @@ def get_connection_form_widgets() -> dict[str, Any]: widget=BS3TextFieldWidget(), default=5, ), + "lifetime": IntegerField( + lazy_gettext("Lifetime of a token in seconds"), + validators=[Optional(), NumberRange(min=0)], + widget=BS3TextFieldWidget(), + ), } @staticmethod @@ -262,6 +268,8 @@ def get_credentials_and_project_id(self) -> tuple[google.auth.credentials.Creden target_principal, delegates = _get_target_principal_and_delegates(self.impersonation_chain) + lifetime = self._get_field("lifetime", None) + credentials, project_id = get_credentials_and_project_id( key_path=key_path, keyfile_dict=keyfile_dict_json, @@ -272,6 +280,7 @@ def get_credentials_and_project_id(self) -> tuple[google.auth.credentials.Creden delegate_to=self.delegate_to, target_principal=target_principal, delegates=delegates, + lifetime=lifetime, ) overridden_project_id = self._get_field("project") diff --git a/docs/apache-airflow-providers-google/connections/gcp.rst b/docs/apache-airflow-providers-google/connections/gcp.rst index b2f7243fe71e2..eb08063332629 100644 --- a/docs/apache-airflow-providers-google/connections/gcp.rst +++ b/docs/apache-airflow-providers-google/connections/gcp.rst @@ -125,6 +125,9 @@ Number of Retries represents the last request. If zero (default), we attempt the request only once. +Lifetime (optional) + Integer, number of seconds token should be valid for. + When specifying the connection in environment variable you should specify it using URI syntax, with the following requirements: diff --git a/tests/providers/google/cloud/utils/test_credentials_provider.py b/tests/providers/google/cloud/utils/test_credentials_provider.py index 21976c6f6df43..0801e0168d4d5 100644 --- a/tests/providers/google/cloud/utils/test_credentials_provider.py +++ b/tests/providers/google/cloud/utils/test_credentials_provider.py @@ -390,6 +390,29 @@ def test_disable_logging(self, mock_default, mock_info, mock_file, caplog): ) assert not caplog.record_tuples + @mock.patch( + "airflow.providers.google.cloud.utils.credentials_provider.impersonated_credentials.Credentials" + ) + @mock.patch("google.auth.default") + def test_get_credentials_and_project_id_with_default_auth_and_lifetime( + self, mock_auth_default, mock_impersonated_credentials + ): + mock_credentials = mock.MagicMock() + mock_auth_default.return_value = (mock_credentials, self.test_project_id) + + result = get_credentials_and_project_id( + lifetime=7200, + ) + mock_auth_default.assert_called_once_with(scopes=None) + mock_impersonated_credentials.assert_called_once_with( + source_credentials=mock_credentials.source_credentials, + target_principal=mock_credentials.target_principal, + delegates=mock_credentials.delegates, + target_scopes=mock_credentials.target_scopes, + lifetime=7200, + ) + assert mock_impersonated_credentials.return_value == result[0] + class TestGetScopes: def test_get_scopes_with_default(self): diff --git a/tests/providers/google/common/hooks/test_base_google.py b/tests/providers/google/common/hooks/test_base_google.py index 6ca5f19fc1a12..df4cd42c7cc4e 100644 --- a/tests/providers/google/common/hooks/test_base_google.py +++ b/tests/providers/google/common/hooks/test_base_google.py @@ -370,6 +370,7 @@ def test_get_credentials_and_project_id_with_default_auth(self, mock_get_creds_a delegate_to=None, target_principal=None, delegates=None, + lifetime=None, ) assert ("CREDENTIALS", "PROJECT_ID") == result @@ -407,6 +408,7 @@ def test_get_credentials_and_project_id_with_service_account_file(self, mock_get delegate_to=None, target_principal=None, delegates=None, + lifetime=None, ) assert (mock_credentials, "PROJECT_ID") == result @@ -437,6 +439,7 @@ def test_get_credentials_and_project_id_with_service_account_info(self, mock_get delegate_to=None, target_principal=None, delegates=None, + lifetime=None, ) assert (mock_credentials, "PROJECT_ID") == result @@ -457,6 +460,7 @@ def test_get_credentials_and_project_id_with_default_auth_and_delegate(self, moc delegate_to="USER", target_principal=None, delegates=None, + lifetime=None, ) assert (mock_credentials, "PROJECT_ID") == result @@ -493,6 +497,7 @@ def test_get_credentials_and_project_id_with_default_auth_and_overridden_project delegate_to=None, target_principal=None, delegates=None, + lifetime=None, ) assert ("CREDENTIALS", "SECOND_PROJECT_ID") == result @@ -695,6 +700,7 @@ def test_get_credentials_and_project_id_with_impersonation_chain( delegate_to=None, target_principal=target_principal, delegates=delegates, + lifetime=None, ) assert (mock_credentials, PROJECT_ID) == result