diff --git a/docs/guides/README.md b/docs/guides/README.md index 18025871096..ad5a4230773 100644 --- a/docs/guides/README.md +++ b/docs/guides/README.md @@ -47,6 +47,9 @@ This directory contains specific developer guides for the ADK Python implementat ### Tools * [to_mcp_server](tools/mcp_tool/agent_to_mcp/index.md) - Expose an ADK agent as an MCP server so any MCP host can drive it as a single tool (the MCP counterpart of to_a2a). +### Security +* [Credentials Encryption](credentials_encryption.md) - Securely encrypting sensitive session credentials using GCP Cloud KMS. + ### Workflows * [Workflow](workflow/workflow/index.md) - Graph-based orchestration of complex, multi-step agent interactions. * [Workflow Graphs](workflow/graph/index.md) - Understanding nodes, edges, and graph structures in workflows. diff --git a/docs/guides/credentials_encryption.md b/docs/guides/credentials_encryption.md new file mode 100644 index 00000000000..747bdf5ff9a --- /dev/null +++ b/docs/guides/credentials_encryption.md @@ -0,0 +1,60 @@ +# Session Credentials Encryption Guide + +To prevent sensitive OAuth 2 credentials (like access tokens, refresh tokens, and client secrets) from being stored in plaintext inside the session state database, ADK supports encrypting them using Google Cloud KMS with **Envelope Encryption**. + +## How It Works + +1. **Envelope Encryption for Google OAuth Credentials**: + * **Data Encryption Key (DEK)**: A local 256-bit symmetric key (Fernet) is generated locally to encrypt the sensitive fields (`access_token`, `refresh_token`, `client_secret`). + * **Key Encryption Key (KEK)**: The Google Cloud KMS key acts as the KEK and is used to encrypt (wrap) the local DEK. + * **Storage**: The session stores the locally encrypted credentials, the public reference of the KMS key (`kms_key_name`), and the encrypted DEK (`wrapped_dek`). +2. **Direct KMS Encryption for Generic Credentials (`SessionStateCredentialService`)**: + * All non-OAuth credentials (API keys, HTTP Basic Auth, Bearer tokens, Service Account private keys) saved to session state via `SessionStateCredentialService` are automatically encrypted using Cloud KMS on save (`save_credential`) and decrypted on load (`load_credential`). + * Encrypted values are stored in state with a `kms:` prefix. +3. **In-Memory Caching (Zero Latency)**: + * To prevent performing a slow GCP KMS network request on every field encryption or decryption, the resolved plaintext DEK and its corresponding `wrapped_dek` are cached in-memory. + * On deserialization, KMS is called **exactly once** per session load, and subsequent decryptions are processed locally in-memory (instantaneous). On serialization, we reuse the cached wrapped DEK (zero KMS calls). +4. **Re-Authentication Fallback (No-Crash)**: + * If Cloud KMS decryption fails (e.g. key destroyed, IAM permission revoked, or key version unavailable), `SessionStateCredentialService` logs a warning and returns `None`, gracefully triggering user re-authentication instead of throwing validation errors. +5. **Backward Compatibility**: If no KMS key is configured or the stored credentials do not contain encrypted values, ADK automatically falls back to loading/saving them in plaintext without raising errors. + +--- + +## Configuration + +Set the environment variable `GOOGLE_CREDENTIAL_KMS_KEY` to point to your GCP KMS CryptoKey (optionally pinning a specific version): + +```bash +export GOOGLE_CREDENTIAL_KMS_KEY="projects/{project_id}/locations/{location}/keyRings/{key_ring_name}/cryptoKeys/{key_name}/cryptoKeyVersions/{version_id}" +``` + +Alternatively, you can configure it programmatically on any `CredentialsConfig` (like `BigQueryCredentialsConfig`): + +```python +oauth_credentials_config = BigQueryCredentialsConfig( + client_id=client_id, + client_secret=client_secret, + scopes=scopes, + kms_key_name="projects/{project_id}/locations/{location}/keyRings/{key_ring_name}/cryptoKeys/{key_name}/cryptoKeyVersions/{version_id}" +) +``` + +--- + +## Required IAM Permissions + +The Service Account running the ADK Agent / Runner must be granted the appropriate permissions to call the Cloud KMS API. + +### KMS Permissions +* **Role**: `Cloud KMS CryptoKey Encrypter/Decrypter` (`roles/cloudkms.cryptoKeyEncrypterDecrypter`) +* **Scope**: Must be granted on the specified CryptoKey or KeyRing. + +Example `gcloud` command to grant access: + +```bash +gcloud kms keys add-iam-policy-binding {key_name} \ + --location={location} \ + --keyring={key_ring_name} \ + --member="serviceAccount:{agent_service_account}@{project_id}.iam.gserviceaccount.com" \ + --role="roles/cloudkms.cryptoKeyEncrypterDecrypter" +``` diff --git a/pyproject.toml b/pyproject.toml index 46edcbc6949..a5de9b8ac39 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -37,6 +37,7 @@ dependencies = [ "click>=8.1.8,<9", "fastapi>=0.133,<1", "google-auth[pyopenssl]>=2.47", + "google-cloud-kms>=3,<4", "google-genai>=2.12.1,<3", "graphviz>=0.20.2,<1", "httpx>=0.27,<1", diff --git a/src/google/adk/auth/_kms_encryptor.py b/src/google/adk/auth/_kms_encryptor.py new file mode 100644 index 00000000000..8dc71d8be83 --- /dev/null +++ b/src/google/adk/auth/_kms_encryptor.py @@ -0,0 +1,207 @@ +# Copyright 2026 Google LLC +# +# Licensed 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. + +from __future__ import annotations + +import base64 +import logging +from typing import Dict +from typing import Tuple + +from cryptography.fernet import Fernet + +logger = logging.getLogger("google_adk." + __name__) + +# Cache mapping KMS key name -> tuple of (plaintext_dek: bytes, wrapped_dek: str) +_KMS_KEY_DEK_CACHE: Dict[str, Tuple[bytes, str]] = {} + +# Cache mapping wrapped_dek (str) -> Fernet instance +_DEK_FERNET_CACHE: Dict[str, Fernet] = {} + +# KMS Client cache +_KMS_CLIENT_CACHE: Dict[str, any] = {} + + +def _get_kms_client(kms_key_name: str): + """Gets or creates a cached Google Cloud KMS client.""" + if kms_key_name not in _KMS_CLIENT_CACHE: + from google.cloud import kms + + _KMS_CLIENT_CACHE[kms_key_name] = kms.KeyManagementServiceClient() + return _KMS_CLIENT_CACHE[kms_key_name] + + +def _get_or_create_dek(kms_key_name: str) -> Tuple[bytes, str]: + """Gets the cached DEK for a KMS key, or generates and wraps a new one.""" + if kms_key_name not in _KMS_KEY_DEK_CACHE: + try: + # Generate a new 32-byte Fernet key + plaintext_dek = Fernet.generate_key() + + # Wrap (encrypt) the DEK using Cloud KMS + client = _get_kms_client(kms_key_name) + response = client.encrypt( + request={ + "name": kms_key_name, + "plaintext": plaintext_dek, + } + ) + wrapped_dek = base64.b64encode(response.ciphertext).decode("utf-8") + _KMS_KEY_DEK_CACHE[kms_key_name] = (plaintext_dek, wrapped_dek) + # Also populate the Fernet cache for this wrapped DEK + _DEK_FERNET_CACHE[wrapped_dek] = Fernet(plaintext_dek) + except Exception as e: + logger.error( + "Failed to generate and wrap DEK using KMS key %s: %s", + kms_key_name, + e, + ) + raise e + + return _KMS_KEY_DEK_CACHE[kms_key_name] + + +def _get_crypto_key_name(kms_key_name: str) -> str: + """Returns the CryptoKey resource name by stripping any version suffix if present.""" + if "/cryptoKeyVersions/" in kms_key_name: + return kms_key_name.split("/cryptoKeyVersions/")[0] + return kms_key_name + + +def _get_fernet_for_wrapped_dek(kms_key_name: str, wrapped_dek: str) -> Fernet: + """Gets the cached Fernet instance for a wrapped DEK, unwrapping it with KMS if needed.""" + if wrapped_dek not in _DEK_FERNET_CACHE: + try: + # Unwrap (decrypt) the DEK using Cloud KMS (decrypt requires CryptoKey name, not version) + client = _get_kms_client(kms_key_name) + ciphertext_bytes = base64.b64decode(wrapped_dek.encode("utf-8")) + crypto_key_name = _get_crypto_key_name(kms_key_name) + response = client.decrypt( + request={ + "name": crypto_key_name, + "ciphertext": ciphertext_bytes, + } + ) + plaintext_dek = response.plaintext + _DEK_FERNET_CACHE[wrapped_dek] = Fernet(plaintext_dek) + except Exception as e: + logger.error("Failed to unwrap DEK using KMS key %s: %s", kms_key_name, e) + raise e + + return _DEK_FERNET_CACHE[wrapped_dek] + + +def encrypt_credentials( + kms_key_name: str, + token: str | None, + refresh_token: str | None, + client_secret: str | None, +) -> Tuple[str | None, str | None, str | None, str | None]: + """Encrypts the sensitive credential fields using envelope encryption. + + Returns a tuple of (encrypted_token, encrypted_refresh_token, encrypted_client_secret, wrapped_dek). + """ + try: + plaintext_dek, wrapped_dek = _get_or_create_dek(kms_key_name) + fernet = _DEK_FERNET_CACHE[wrapped_dek] + + enc_token = ( + fernet.encrypt(token.encode("utf-8")).decode("utf-8") if token else None + ) + enc_refresh = ( + fernet.encrypt(refresh_token.encode("utf-8")).decode("utf-8") + if refresh_token + else None + ) + enc_secret = ( + fernet.encrypt(client_secret.encode("utf-8")).decode("utf-8") + if client_secret + else None + ) + + return enc_token, enc_refresh, enc_secret, wrapped_dek + except Exception as e: + logger.error("Failed to encrypt credentials: %s", e) + raise e + + +def decrypt_credentials( + kms_key_name: str, + encrypted_token: str | None, + encrypted_refresh_token: str | None, + encrypted_client_secret: str | None, + wrapped_dek: str | None, +) -> Tuple[str | None, str | None, str | None]: + """Decrypts the sensitive credential fields using the wrapped DEK.""" + if not wrapped_dek: + # Backward compatibility + return encrypted_token, encrypted_refresh_token, encrypted_client_secret + + try: + fernet = _get_fernet_for_wrapped_dek(kms_key_name, wrapped_dek) + + dec_token = ( + fernet.decrypt(encrypted_token.encode("utf-8")).decode("utf-8") + if encrypted_token + else None + ) + dec_refresh = ( + fernet.decrypt(encrypted_refresh_token.encode("utf-8")).decode("utf-8") + if encrypted_refresh_token + else None + ) + dec_secret = ( + fernet.decrypt(encrypted_client_secret.encode("utf-8")).decode("utf-8") + if encrypted_client_secret + else None + ) + + return dec_token, dec_refresh, dec_secret + except Exception as e: + logger.error("Failed to decrypt credentials: %s", e) + raise e + + +def encrypt_value(kms_key_name: str, plaintext: str) -> str: + """Fallback/Direct encryption helper.""" + try: + client = _get_kms_client(kms_key_name) + response = client.encrypt( + request={ + "name": kms_key_name, + "plaintext": plaintext.encode("utf-8"), + } + ) + return base64.b64encode(response.ciphertext).decode("utf-8") + except Exception as e: + logger.error("Failed to encrypt value: %s", e) + raise e + + +def decrypt_value(kms_key_name: str, ciphertext: str) -> str: + """Fallback/Direct decryption helper.""" + try: + client = _get_kms_client(kms_key_name) + ciphertext_bytes = base64.b64decode(ciphertext.encode("utf-8")) + crypto_key_name = _get_crypto_key_name(kms_key_name) + response = client.decrypt( + request={ + "name": crypto_key_name, + "ciphertext": ciphertext_bytes, + } + ) + return response.plaintext.decode("utf-8") + except Exception as e: + logger.error("Failed to decrypt value: %s", e) + raise e diff --git a/src/google/adk/auth/auth_credential.py b/src/google/adk/auth/auth_credential.py index 8c5de6cced2..cedfe8a0620 100644 --- a/src/google/adk/auth/auth_credential.py +++ b/src/google/adk/auth/auth_credential.py @@ -22,6 +22,7 @@ from typing import List from typing import Literal +import google.oauth2.credentials from pydantic import alias_generators from pydantic import BaseModel from pydantic import ConfigDict @@ -315,3 +316,133 @@ class AuthCredential(BaseModelWithConfig): http: HttpAuth | None = None service_account: ServiceAccount | None = None oauth2: OAuth2Auth | None = None + kms_key_name: str | None = None + + +class KmsEncryptedCredentials(google.oauth2.credentials.Credentials): + """Subclass of Google Credentials that supports encrypting sensitive fields using KMS.""" + + def __init__( + self, + token, + refresh_token=None, + id_token=None, + token_uri=None, + client_id=None, + client_secret=None, + scopes=None, + default_scopes=None, + quota_project_id=None, + expiry=None, + rapt_token=None, + kms_key_name: str | None = None, + ): + import inspect + + import google.oauth2.credentials + + sig = inspect.signature(google.oauth2.credentials.Credentials.__init__) + kwargs = { + "token": token, + "refresh_token": refresh_token, + "id_token": id_token, + "token_uri": token_uri, + "client_id": client_id, + "client_secret": client_secret, + "scopes": scopes, + "default_scopes": default_scopes, + "quota_project_id": quota_project_id, + "expiry": expiry, + } + if "rapt_token" in sig.parameters: + kwargs["rapt_token"] = rapt_token + super().__init__(**kwargs) + self.kms_key_name = kms_key_name + + def to_json(self, strip=None): + """Serialize credentials to JSON, encrypting sensitive fields if kms_key_name is present.""" + serialized_json = super().to_json(strip=strip) + import json + + from ._kms_encryptor import encrypt_credentials + + data = json.loads(serialized_json) + + if self.kms_key_name: + token = data.get("token") + refresh_token = data.get("refresh_token") + client_secret = data.get("client_secret") + + enc_token, enc_refresh, enc_secret, wrapped_dek = encrypt_credentials( + self.kms_key_name, token, refresh_token, client_secret + ) + + if enc_token: + data["token"] = "kms:" + enc_token + if enc_refresh: + data["refresh_token"] = "kms:" + enc_refresh + if enc_secret: + data["client_secret"] = "kms:" + enc_secret + + if wrapped_dek: + data["wrapped_dek"] = wrapped_dek + data["kms_key_name"] = self.kms_key_name + + return json.dumps(data) + + @classmethod + def from_authorized_user_info(cls, info, scopes=None): + """Deserialize credentials from user info, decrypting sensitive fields if encrypted.""" + kms_key_name = info.get("kms_key_name") + wrapped_dek = info.get("wrapped_dek") + info_copy = dict(info) + + if kms_key_name: + from ._kms_encryptor import decrypt_credentials + + token = info_copy.get("token") + refresh_token = info_copy.get("refresh_token") + client_secret = info_copy.get("client_secret") + + enc_token = ( + token[4:] + if isinstance(token, str) and token.startswith("kms:") + else token + ) + enc_refresh = ( + refresh_token[4:] + if isinstance(refresh_token, str) and refresh_token.startswith("kms:") + else refresh_token + ) + enc_secret = ( + client_secret[4:] + if isinstance(client_secret, str) and client_secret.startswith("kms:") + else client_secret + ) + + dec_token, dec_refresh, dec_secret = decrypt_credentials( + kms_key_name, enc_token, enc_refresh, enc_secret, wrapped_dek + ) + + info_copy["token"] = dec_token + info_copy["refresh_token"] = dec_refresh + info_copy["client_secret"] = dec_secret + + # Some versions of google-auth might return google.oauth2.credentials.Credentials + import google.oauth2.credentials + + creds = google.oauth2.credentials.Credentials.from_authorized_user_info( + info_copy, scopes=scopes + ) + + return cls( + token=creds.token, + refresh_token=creds.refresh_token, + id_token=creds.id_token, + token_uri=creds.token_uri, + client_id=creds.client_id, + client_secret=creds.client_secret, + scopes=creds.scopes, + expiry=creds.expiry, + kms_key_name=kms_key_name, + ) diff --git a/src/google/adk/auth/credential_service/session_state_credential_service.py b/src/google/adk/auth/credential_service/session_state_credential_service.py index 5559ec60058..2131920c5b4 100644 --- a/src/google/adk/auth/credential_service/session_state_credential_service.py +++ b/src/google/adk/auth/credential_service/session_state_credential_service.py @@ -14,16 +14,190 @@ from __future__ import annotations +import logging +import os +from typing import Any from typing import Optional from typing_extensions import override from ...agents.callback_context import CallbackContext from ...utils.feature_decorator import experimental +from .._kms_encryptor import decrypt_value +from .._kms_encryptor import encrypt_value from ..auth_credential import AuthCredential from ..auth_tool import AuthConfig from .base_credential_service import BaseCredentialService +logger = logging.getLogger("google_adk." + __name__) + + +def _encrypt_auth_credential( + kms_key: str, cred: AuthCredential +) -> dict[str, Any]: + data = cred.model_dump(by_alias=True) + data["kmsKeyName"] = kms_key + + try: + if cred.api_key and not cred.api_key.startswith("kms:"): + data["apiKey"] = "kms:" + encrypt_value(kms_key, cred.api_key) + + if cred.http and cred.http.credentials: + if ( + cred.http.credentials.password + and not cred.http.credentials.password.startswith("kms:") + ): + data["http"]["credentials"]["password"] = "kms:" + encrypt_value( + kms_key, cred.http.credentials.password + ) + if ( + cred.http.credentials.token + and not cred.http.credentials.token.startswith("kms:") + ): + data["http"]["credentials"]["token"] = "kms:" + encrypt_value( + kms_key, cred.http.credentials.token + ) + + if cred.oauth2: + if cred.oauth2.access_token and not cred.oauth2.access_token.startswith( + "kms:" + ): + data["oauth2"]["accessToken"] = "kms:" + encrypt_value( + kms_key, cred.oauth2.access_token + ) + if ( + cred.oauth2.refresh_token + and not cred.oauth2.refresh_token.startswith("kms:") + ): + data["oauth2"]["refreshToken"] = "kms:" + encrypt_value( + kms_key, cred.oauth2.refresh_token + ) + if ( + cred.oauth2.client_secret + and not cred.oauth2.client_secret.startswith("kms:") + ): + data["oauth2"]["clientSecret"] = "kms:" + encrypt_value( + kms_key, cred.oauth2.client_secret + ) + + if cred.service_account and cred.service_account.service_account_credential: + pk = cred.service_account.service_account_credential.private_key + if pk and not pk.startswith("kms:"): + sa_dict = data.get("serviceAccount") or data.get("service_account") + if sa_dict: + sac_key = ( + "serviceAccountCredential" + if "serviceAccountCredential" in sa_dict + else "service_account_credential" + ) + sac_dict = sa_dict.get(sac_key) + if sac_dict: + sac_dict["privateKey"] = "kms:" + encrypt_value(kms_key, pk) + except Exception as e: + logger.error("Failed to encrypt AuthCredential with KMS: %s", e) + raise e + + return data + + +def _decrypt_auth_credential( + val: Any, default_kms_key: str | None +) -> AuthCredential | None: + if isinstance(val, dict): + cred_dict = dict(val) + elif isinstance(val, AuthCredential): + cred_dict = val.model_dump(by_alias=True) + else: + return None + + kms_key = ( + cred_dict.get("kms_key_name") + or cred_dict.get("kmsKeyName") + or default_kms_key + ) + + try: + for key in ("api_key", "apiKey"): + if cred_dict.get(key) and str(cred_dict[key]).startswith("kms:"): + if not kms_key: + logger.warning( + "Encrypted api_key found but no KMS key provided. Falling back to" + " re-auth." + ) + return None + cred_dict[key] = decrypt_value(kms_key, str(cred_dict[key])[4:]) + + if "http" in cred_dict and isinstance(cred_dict["http"], dict): + http_creds = cred_dict["http"].get("credentials") + if http_creds and isinstance(http_creds, dict): + for pwd_key in ("password",): + if http_creds.get(pwd_key) and str(http_creds[pwd_key]).startswith( + "kms:" + ): + if not kms_key: + logger.warning( + "Encrypted password found but no KMS key provided. Falling" + " back to re-auth." + ) + return None + http_creds[pwd_key] = decrypt_value( + kms_key, str(http_creds[pwd_key])[4:] + ) + for tok_key in ("token",): + if http_creds.get(tok_key) and str(http_creds[tok_key]).startswith( + "kms:" + ): + if not kms_key: + logger.warning( + "Encrypted token found but no KMS key provided. Falling back" + " to re-auth." + ) + return None + http_creds[tok_key] = decrypt_value( + kms_key, str(http_creds[tok_key])[4:] + ) + + if "oauth2" in cred_dict and isinstance(cred_dict["oauth2"], dict): + oa = cred_dict["oauth2"] + for secret_key in ("client_secret", "clientSecret"): + if oa.get(secret_key) and str(oa[secret_key]).startswith("kms:"): + if not kms_key: + return None + oa[secret_key] = decrypt_value(kms_key, str(oa[secret_key])[4:]) + for token_key in ("access_token", "accessToken"): + if oa.get(token_key) and str(oa[token_key]).startswith("kms:"): + if not kms_key: + return None + oa[token_key] = decrypt_value(kms_key, str(oa[token_key])[4:]) + for refresh_key in ("refresh_token", "refreshToken"): + if oa.get(refresh_key) and str(oa[refresh_key]).startswith("kms:"): + if not kms_key: + return None + oa[refresh_key] = decrypt_value(kms_key, str(oa[refresh_key])[4:]) + + sa_dict = cred_dict.get("serviceAccount") or cred_dict.get( + "service_account" + ) + if sa_dict and isinstance(sa_dict, dict): + sac_dict = sa_dict.get("serviceAccountCredential") or sa_dict.get( + "service_account_credential" + ) + if sac_dict and isinstance(sac_dict, dict): + for pk_key in ("private_key", "privateKey"): + if sac_dict.get(pk_key) and str(sac_dict[pk_key]).startswith("kms:"): + if not kms_key: + return None + sac_dict[pk_key] = decrypt_value(kms_key, str(sac_dict[pk_key])[4:]) + + return AuthCredential.model_validate(cred_dict) + except Exception as e: + logger.warning( + "Failed to decrypt AuthCredential from session state: %s. Falling back" + " to re-authentication.", + e, + ) + return None + @experimental class SessionStateCredentialService(BaseCredentialService): @@ -54,7 +228,14 @@ async def load_credential( Optional[AuthCredential]: the credential saved in the store. """ - return callback_context.state.get(auth_config.credential_key) + val = callback_context.state.get(auth_config.credential_key) + if not val: + return None + + kms_key = getattr(auth_config, "kms_key_name", None) or os.environ.get( + "GOOGLE_CREDENTIAL_KMS_KEY" + ) + return _decrypt_auth_credential(val, kms_key) @override async def save_credential( @@ -77,7 +258,19 @@ async def save_credential( Returns: None """ + cred = auth_config.exchanged_auth_credential + if not cred: + return - callback_context.state[auth_config.credential_key] = ( - auth_config.exchanged_auth_credential + kms_key = ( + getattr(auth_config, "kms_key_name", None) + or (cred.kms_key_name if hasattr(cred, "kms_key_name") else None) + or os.environ.get("GOOGLE_CREDENTIAL_KMS_KEY") ) + + if kms_key: + callback_context.state[auth_config.credential_key] = ( + _encrypt_auth_credential(kms_key, cred) + ) + else: + callback_context.state[auth_config.credential_key] = cred diff --git a/src/google/adk/tools/_google_credentials.py b/src/google/adk/tools/_google_credentials.py index 5f9f0a0184f..07d6ef70deb 100644 --- a/src/google/adk/tools/_google_credentials.py +++ b/src/google/adk/tools/_google_credentials.py @@ -88,6 +88,8 @@ class BaseGoogleCredentialsConfig(BaseModel): """the oauth client secret to use.""" scopes: Optional[List[str]] = None """the scopes to use.""" + kms_key_name: Optional[str] = None + """The KMS key name to encrypt sensitive credentials fields.""" _token_cache_key: Optional[str] = None """The key to cache the token in the tool context.""" @@ -95,6 +97,11 @@ class BaseGoogleCredentialsConfig(BaseModel): @model_validator(mode="after") def __post_init__(self) -> BaseGoogleCredentialsConfig: """Validate that only one of credentials, external_access_token_key or client_id/secret are provided.""" + import os + + if not self.kms_key_name: + self.kms_key_name = os.environ.get("GOOGLE_CREDENTIAL_KMS_KEY") + if self.credentials: if ( self.external_access_token_key @@ -177,13 +184,25 @@ async def get_valid_credentials( if self.credentials_config._token_cache_key else None ) - creds = ( - google.oauth2.credentials.Credentials.from_authorized_user_info( - json.loads(creds_json), self.credentials_config.scopes + if creds_json: + import json + + from ..auth.auth_credential import KmsEncryptedCredentials + + creds_data = json.loads(creds_json) + kms_key = ( + creds_data.get("kms_key_name") or self.credentials_config.kms_key_name + ) + if kms_key: + creds = KmsEncryptedCredentials.from_authorized_user_info( + creds_data, self.credentials_config.scopes ) - if creds_json - else None - ) + else: + creds = google.oauth2.credentials.Credentials.from_authorized_user_info( + creds_data, self.credentials_config.scopes + ) + else: + creds = None # If credentials are empty use the default credential if not creds: @@ -214,6 +233,22 @@ async def get_valid_credentials( if creds.valid: # Cache the refreshed credentials if token cache key is set if self.credentials_config._token_cache_key: + if self.credentials_config.kms_key_name and not isinstance( + creds, KmsEncryptedCredentials + ): + from ..auth.auth_credential import KmsEncryptedCredentials + + creds = KmsEncryptedCredentials( + token=creds.token, + refresh_token=creds.refresh_token, + id_token=creds.id_token, + token_uri=creds.token_uri, + client_id=creds.client_id, + client_secret=creds.client_secret, + scopes=creds.scopes, + expiry=creds.expiry, + kms_key_name=self.credentials_config.kms_key_name, + ) tool_context.state[self.credentials_config._token_cache_key] = ( creds.to_json() ) @@ -266,14 +301,27 @@ async def _perform_oauth_flow( if auth_response: # OAuth flow completed, create credentials - creds = google.oauth2.credentials.Credentials( - token=auth_response.oauth2.access_token, - refresh_token=auth_response.oauth2.refresh_token, - token_uri=auth_scheme.flows.authorizationCode.tokenUrl, - client_id=self.credentials_config.client_id, - client_secret=self.credentials_config.client_secret, - scopes=list(self.credentials_config.scopes), - ) + if self.credentials_config.kms_key_name: + from ..auth.auth_credential import KmsEncryptedCredentials + + creds = KmsEncryptedCredentials( + token=auth_response.oauth2.access_token, + refresh_token=auth_response.oauth2.refresh_token, + token_uri=auth_scheme.flows.authorizationCode.tokenUrl, + client_id=self.credentials_config.client_id, + client_secret=self.credentials_config.client_secret, + scopes=list(self.credentials_config.scopes), + kms_key_name=self.credentials_config.kms_key_name, + ) + else: + creds = google.oauth2.credentials.Credentials( + token=auth_response.oauth2.access_token, + refresh_token=auth_response.oauth2.refresh_token, + token_uri=auth_scheme.flows.authorizationCode.tokenUrl, + client_id=self.credentials_config.client_id, + client_secret=self.credentials_config.client_secret, + scopes=list(self.credentials_config.scopes), + ) # Cache the new credentials if token cache key is set if self.credentials_config._token_cache_key: diff --git a/tests/unittests/auth/test_kms_credentials.py b/tests/unittests/auth/test_kms_credentials.py new file mode 100644 index 00000000000..1cdf63de830 --- /dev/null +++ b/tests/unittests/auth/test_kms_credentials.py @@ -0,0 +1,264 @@ +# Copyright 2026 Google LLC +# +# Licensed 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 base64 +import json +import os +from unittest.mock import Mock + +from google.adk.auth._kms_encryptor import _DEK_FERNET_CACHE +from google.adk.auth._kms_encryptor import _KMS_KEY_DEK_CACHE +from google.adk.auth._kms_encryptor import decrypt_credentials +from google.adk.auth._kms_encryptor import encrypt_credentials +from google.adk.auth.auth_credential import AuthCredential +from google.adk.auth.auth_credential import AuthCredentialTypes +from google.adk.auth.auth_credential import HttpAuth +from google.adk.auth.auth_credential import HttpCredentials +from google.adk.auth.auth_credential import KmsEncryptedCredentials +from google.adk.auth.auth_credential import OAuth2Auth +from google.adk.auth.credential_service.session_state_credential_service import SessionStateCredentialService +from google.adk.tools._google_credentials import BaseGoogleCredentialsConfig +from google.adk.tools._google_credentials import GoogleCredentialsManager +from google.adk.tools.tool_context import ToolContext +import pytest + + +@pytest.fixture(autouse=True) +def mock_kms_client(monkeypatch): + """Mock the Google Cloud KMS client for testing.""" + + class MockKmsClient: + + def encrypt(self, request): + # Wrap the plaintext DEK by adding a prefix + ct_val = b"mock_wrapped_" + request["plaintext"] + return Mock(ciphertext=ct_val) + + def decrypt(self, request): + # Unwrap the ciphertext to retrieve original DEK + ct_val = request["ciphertext"] + assert ct_val.startswith(b"mock_wrapped_") + pt_val = ct_val[13:] + return Mock(plaintext=pt_val) + + import google.adk.auth._kms_encryptor + + monkeypatch.setattr( + google.adk.auth._kms_encryptor, + "_get_kms_client", + lambda kms_key_name: MockKmsClient(), + ) + + +def test_kms_envelope_encryption_caching_and_crypto(): + """Test that kms_encryptor properly encrypts, decrypts, and uses in-memory DEK caching.""" + key_name = ( + "projects/p1/locations/l1/keyRings/kr1/cryptoKeys/k1/cryptoKeyVersions/1" + ) + + # Reset caches + _KMS_KEY_DEK_CACHE.clear() + _DEK_FERNET_CACHE.clear() + + token = "secret_access_token" + refresh_token = "secret_refresh_token" + + # Encrypt credentials using envelope encryption + enc_token, enc_refresh, _, wrapped_dek = encrypt_credentials( + key_name, token, refresh_token, None + ) + + assert enc_token != token + assert enc_refresh != refresh_token + assert wrapped_dek is not None + + # Verify the DEK is cached + assert key_name in _KMS_KEY_DEK_CACHE + assert wrapped_dek in _DEK_FERNET_CACHE + + # Decrypt credentials + dec_token, dec_refresh, _ = decrypt_credentials( + key_name, enc_token, enc_refresh, None, wrapped_dek + ) + + assert dec_token == token + assert dec_refresh == refresh_token + + +def test_kms_encrypted_credentials_serialization(): + """Test that KmsEncryptedCredentials properly serializes to JSON with envelope encryption and deserializes back.""" + key_name = ( + "projects/p1/locations/l1/keyRings/kr1/cryptoKeys/k1/cryptoKeyVersions/1" + ) + + creds = KmsEncryptedCredentials( + token="secret_access_token", + refresh_token="secret_refresh_token", + client_id="my_client_id", + client_secret="secret_client_secret", + kms_key_name=key_name, + ) + + # Serialize to JSON (envelope encryption) + serialized = creds.to_json() + data = json.loads(serialized) + + # Ensure sensitive values are prefixed and wrapped DEK is stored + assert data["token"].startswith("kms:") + assert data["refresh_token"].startswith("kms:") + assert data["client_secret"].startswith("kms:") + assert "wrapped_dek" in data + assert data["kms_key_name"] == key_name + assert data["client_id"] == "my_client_id" + + # Deserialize back + deserialized = KmsEncryptedCredentials.from_authorized_user_info(data) + + assert deserialized.token == "secret_access_token" + assert deserialized.refresh_token == "secret_refresh_token" + assert deserialized.client_secret == "secret_client_secret" + assert deserialized.client_id == "my_client_id" + assert deserialized.kms_key_name == key_name + + +def test_kms_env_var_detection(monkeypatch): + """Test that BaseGoogleCredentialsConfig automatically detects GOOGLE_CREDENTIAL_KMS_KEY.""" + key_name = ( + "projects/p1/locations/l1/keyRings/kr1/cryptoKeys/k1/cryptoKeyVersions/1" + ) + monkeypatch.setenv("GOOGLE_CREDENTIAL_KMS_KEY", key_name) + + config = BaseGoogleCredentialsConfig( + client_id="my_client_id", + client_secret="my_client_secret", + ) + + assert config.kms_key_name == key_name + + +def test_kms_credentials_backward_compatibility(): + """Test that loading a non-encrypted credentials json works normally without crashing.""" + info = { + "token": "plaintext_token", + "refresh_token": "plaintext_refresh", + "client_id": "my_client_id", + "client_secret": "plaintext_secret", + } + + # Load via KmsEncryptedCredentials but without kms_key_name in info + creds = KmsEncryptedCredentials.from_authorized_user_info(info) + + assert creds.token == "plaintext_token" + assert creds.refresh_token == "plaintext_refresh" + assert creds.client_secret == "plaintext_secret" + assert creds.kms_key_name is None + + # to_json should not encrypt when kms_key_name is not set + serialized = creds.to_json() + data = json.loads(serialized) + assert data["token"] == "plaintext_token" + assert "kms_key_name" not in data + + +def test_session_state_credential_service_kms_encryption(monkeypatch): + """Test that SessionStateCredentialService encrypts credentials on save and decrypts on load.""" + key_name = ( + "projects/p1/locations/l1/keyRings/kr1/cryptoKeys/k1/cryptoKeyVersions/1" + ) + monkeypatch.setenv("GOOGLE_CREDENTIAL_KMS_KEY", key_name) + + # 1. Create a model with plaintext fields + cred = AuthCredential( + auth_type=AuthCredentialTypes.HTTP, + http=HttpAuth( + scheme="basic", + credentials=HttpCredentials( + username="bob", + password="secretpassword", + token="tokensecret", + ), + ), + api_key="api_key_secret", + ) + + service = SessionStateCredentialService() + + class MockConfig: + credential_key = "test_cred_key" + exchanged_auth_credential = cred + + callback_ctx = Mock() + callback_ctx.state = {} + + import asyncio + + # Save credential (should encrypt sensitive fields in state) + asyncio.run(service.save_credential(MockConfig(), callback_ctx)) + + saved_data = callback_ctx.state["test_cred_key"] + assert saved_data["apiKey"].startswith("kms:") + assert saved_data["http"]["credentials"]["password"].startswith("kms:") + assert saved_data["http"]["credentials"]["token"].startswith("kms:") + + # Load credential (should decrypt back to plaintext) + loaded = asyncio.run(service.load_credential(MockConfig(), callback_ctx)) + assert loaded.api_key == "api_key_secret" + assert loaded.http.credentials.password == "secretpassword" + assert loaded.http.credentials.token == "tokensecret" + + +def test_kms_decryption_failure_fallback(monkeypatch): + """Test that decryption failures (e.g. key destroyed) trigger fallback by returning None instead of crashing.""" + key_name = ( + "projects/p1/locations/l1/keyRings/kr1/cryptoKeys/k1/cryptoKeyVersions/1" + ) + monkeypatch.setenv("GOOGLE_CREDENTIAL_KMS_KEY", key_name) + + # Mock KMS client to raise an exception on decrypt (simulating destroyed key or permission failure) + class FailedMockKmsClient: + + def encrypt(self, request): + return Mock(ciphertext=b"mock_wrapped_" + request["plaintext"]) + + def decrypt(self, request): + raise RuntimeError("KMS key has been destroyed or IAM permission denied") + + import google.adk.auth._kms_encryptor + + monkeypatch.setattr( + google.adk.auth._kms_encryptor, + "_get_kms_client", + lambda kms_key_name: FailedMockKmsClient(), + ) + + # Encrypted dict representation of a credential + encrypted_data = { + "auth_type": "apiKey", + "api_key": "kms:some_ciphertext_base64", + } + + # Loading this via SessionStateCredentialService should return None instead of crashing the runner + service = SessionStateCredentialService() + + class MockConfig: + credential_key = "test_cred_key" + + # Create a mock callback context with the encrypted state + callback_ctx = Mock() + callback_ctx.state = {"test_cred_key": encrypted_data} + + import asyncio + + res = asyncio.run(service.load_credential(MockConfig(), callback_ctx)) + assert res is None