From 148efeb269909a5ed34c1345ffe9483d4d165ae3 Mon Sep 17 00:00:00 2001 From: "open-swe[bot]" <215916821+open-swe[bot]@users.noreply.github.com> Date: Fri, 8 May 2026 14:29:23 -0700 Subject: [PATCH] feat: support TOKEN_ENCRYPTION_KEY rotation via MultiFernet [closes AB-2323] (#1275) Threat model T6: Fernet token-encryption key had no rotation path. Rotating TOKEN_ENCRYPTION_KEY immediately invalidated every github_token_encrypted value in thread metadata, forcing re-auth or bot-token re-resolution. Switch to cryptography.fernet.MultiFernet and parse TOKEN_ENCRYPTION_KEY as a comma- or newline-separated ordered list (most-recent-first). New writes encrypt under the first key; reads try every key. Deployers can prepend a new key, let active threads roll over, then drop the old key. Co-authored-by: open-swe[bot] Co-authored-by: Johannes du Plessis <51395795+johannes117@users.noreply.github.com> --- INSTALLATION.md | 26 ++++++++ agent/encryption.py | 59 ++++++++-------- tests/test_encryption.py | 141 +++++++++++++++++++++++++++++++++++++++ 3 files changed, 195 insertions(+), 31 deletions(-) create mode 100644 tests/test_encryption.py diff --git a/INSTALLATION.md b/INSTALLATION.md index 935a1ee2..3471cbad 100644 --- a/INSTALLATION.md +++ b/INSTALLATION.md @@ -418,8 +418,32 @@ DEFAULT_SANDBOX_DELETE_AFTER_STOP_SECONDS="" # Delete N seconds after stop (def # === Token Encryption === TOKEN_ENCRYPTION_KEY="" # Generate with: openssl rand -base64 32 + # Supports key rotation: see "Rotating TOKEN_ENCRYPTION_KEY" below ``` +### Rotating TOKEN_ENCRYPTION_KEY + +`TOKEN_ENCRYPTION_KEY` accepts either a single Fernet key or a comma- or +newline-separated **ordered list of keys, most-recent-first**. New writes always +encrypt under the first key; reads try every key in order. To rotate without +invalidating already-stored GitHub tokens: + +1. Generate a new key: `openssl rand -base64 32`. +2. Prepend it to `TOKEN_ENCRYPTION_KEY`, keeping the old key second: + ``` + TOKEN_ENCRYPTION_KEY="," + ``` + Restart the server. New encryptions use ``; existing ciphertexts + still decrypt against ``. +3. Let active threads cycle (each fresh OAuth flow re-encrypts under the new + key). After every active thread has re-authed, drop the old key: + ``` + TOKEN_ENCRYPTION_KEY="" + ``` + Any thread still holding ciphertext under `` will fail to decrypt + and the user will be re-prompted to authenticate — same UX as if the thread + had never authed. + ## 7. Start the server Make sure ngrok is still running from step 2, then start the LangGraph server in a second terminal: @@ -523,3 +547,5 @@ The `langgraph.json` at the project root already defines the graph entry point a - Ensure `TOKEN_ENCRYPTION_KEY` is set (generate with `openssl rand -base64 32`) - The key must be a valid 32-byte Fernet-compatible base64 string +- For key rotation, `TOKEN_ENCRYPTION_KEY` may be a comma- or newline-separated + list of keys (most-recent-first). See "Rotating TOKEN_ENCRYPTION_KEY" above. diff --git a/agent/encryption.py b/agent/encryption.py index cf169cc6..7a048eed 100644 --- a/agent/encryption.py +++ b/agent/encryption.py @@ -3,7 +3,7 @@ import logging import os -from cryptography.fernet import Fernet, InvalidToken +from cryptography.fernet import Fernet, InvalidToken, MultiFernet logger = logging.getLogger(__name__) @@ -12,59 +12,56 @@ class EncryptionKeyMissingError(ValueError): """Raised when TOKEN_ENCRYPTION_KEY environment variable is not set.""" -def _get_encryption_key() -> bytes: - """Get or derive the encryption key from environment variable. +def _parse_encryption_keys(raw: str) -> list[bytes]: + """Split TOKEN_ENCRYPTION_KEY into one or more keys (most-recent-first).""" + keys: list[bytes] = [] + for part in raw.replace("\n", ",").split(","): + stripped = part.strip() + if stripped: + keys.append(stripped.encode()) + return keys - Uses TOKEN_ENCRYPTION_KEY env var if set (must be 32 url-safe base64 bytes), - otherwise derives a key from LANGSMITH_API_KEY using SHA256. - Returns: - 32-byte Fernet-compatible key +def _get_encryption_keys() -> list[bytes]: + """Read the ordered key list from TOKEN_ENCRYPTION_KEY (most-recent-first). + + Accepts a single key or a comma/newline separated list. The first key is + used for encryption; every key is tried for decryption. Raises: - EncryptionKeyMissingError: If TOKEN_ENCRYPTION_KEY is not set + EncryptionKeyMissingError: If TOKEN_ENCRYPTION_KEY is unset or empty. """ explicit_key = os.environ.get("TOKEN_ENCRYPTION_KEY") if not explicit_key: raise EncryptionKeyMissingError - return explicit_key.encode() + keys = _parse_encryption_keys(explicit_key) + if not keys: + raise EncryptionKeyMissingError + return keys + + +def _get_fernet() -> MultiFernet: + """Build a MultiFernet from the configured key list.""" + return MultiFernet([Fernet(k) for k in _get_encryption_keys()]) def encrypt_token(token: str) -> str: - """Encrypt a token for safe storage. - - Args: - token: The plaintext token to encrypt - - Returns: - Base64-encoded encrypted token - """ + """Encrypt a token under the newest configured key.""" if not token: return "" - key = _get_encryption_key() - f = Fernet(key) - encrypted = f.encrypt(token.encode()) + encrypted = _get_fernet().encrypt(token.encode()) return encrypted.decode() def decrypt_token(encrypted_token: str) -> str: - """Decrypt an encrypted token. - - Args: - encrypted_token: The base64-encoded encrypted token - - Returns: - The plaintext token, or empty string if decryption fails - """ + """Decrypt a token, trying each configured key in order.""" if not encrypted_token: return "" try: - key = _get_encryption_key() - f = Fernet(key) - decrypted = f.decrypt(encrypted_token.encode()) + decrypted = _get_fernet().decrypt(encrypted_token.encode()) return decrypted.decode() except InvalidToken: logger.warning("Failed to decrypt token: invalid token") diff --git a/tests/test_encryption.py b/tests/test_encryption.py new file mode 100644 index 00000000..cd0b0555 --- /dev/null +++ b/tests/test_encryption.py @@ -0,0 +1,141 @@ +from __future__ import annotations + +import pytest +from cryptography.fernet import Fernet + +from agent.encryption import ( + EncryptionKeyMissingError, + _get_encryption_keys, + _parse_encryption_keys, + decrypt_token, + encrypt_token, +) + + +def _set_key(monkeypatch: pytest.MonkeyPatch, value: str) -> None: + monkeypatch.setenv("TOKEN_ENCRYPTION_KEY", value) + + +class TestParseEncryptionKeys: + def test_single_key(self) -> None: + k = Fernet.generate_key().decode() + assert _parse_encryption_keys(k) == [k.encode()] + + def test_comma_separated(self) -> None: + k1 = Fernet.generate_key().decode() + k2 = Fernet.generate_key().decode() + assert _parse_encryption_keys(f"{k1},{k2}") == [k1.encode(), k2.encode()] + + def test_newline_separated(self) -> None: + k1 = Fernet.generate_key().decode() + k2 = Fernet.generate_key().decode() + assert _parse_encryption_keys(f"{k1}\n{k2}") == [k1.encode(), k2.encode()] + + def test_strips_whitespace_and_empties(self) -> None: + k1 = Fernet.generate_key().decode() + k2 = Fernet.generate_key().decode() + raw = f" {k1} ,, \n {k2}\n,\n" + assert _parse_encryption_keys(raw) == [k1.encode(), k2.encode()] + + +class TestGetEncryptionKeys: + def test_missing_env_raises(self, monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.delenv("TOKEN_ENCRYPTION_KEY", raising=False) + with pytest.raises(EncryptionKeyMissingError): + _get_encryption_keys() + + def test_empty_env_raises(self, monkeypatch: pytest.MonkeyPatch) -> None: + _set_key(monkeypatch, "") + with pytest.raises(EncryptionKeyMissingError): + _get_encryption_keys() + + def test_whitespace_only_env_raises(self, monkeypatch: pytest.MonkeyPatch) -> None: + _set_key(monkeypatch, " ,\n ") + with pytest.raises(EncryptionKeyMissingError): + _get_encryption_keys() + + +class TestSingleKeyRoundtrip: + def test_encrypt_decrypt(self, monkeypatch: pytest.MonkeyPatch) -> None: + _set_key(monkeypatch, Fernet.generate_key().decode()) + token = "ghp_abc123" + ciphertext = encrypt_token(token) + assert ciphertext != "" + assert ciphertext != token + assert decrypt_token(ciphertext) == token + + def test_empty_token_returns_empty(self, monkeypatch: pytest.MonkeyPatch) -> None: + _set_key(monkeypatch, Fernet.generate_key().decode()) + assert encrypt_token("") == "" + assert decrypt_token("") == "" + + def test_invalid_ciphertext_returns_empty(self, monkeypatch: pytest.MonkeyPatch) -> None: + _set_key(monkeypatch, Fernet.generate_key().decode()) + assert decrypt_token("not-a-valid-fernet-token") == "" + + def test_decrypt_without_key_returns_empty(self, monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.delenv("TOKEN_ENCRYPTION_KEY", raising=False) + assert decrypt_token("anything") == "" + + +class TestMultiKeyDecrypt: + def test_decrypt_old_ciphertext_after_prepending_new_key( + self, monkeypatch: pytest.MonkeyPatch + ) -> None: + old_key = Fernet.generate_key().decode() + new_key = Fernet.generate_key().decode() + + _set_key(monkeypatch, old_key) + token = "ghp_old_secret" + old_ciphertext = encrypt_token(token) + + _set_key(monkeypatch, f"{new_key},{old_key}") + assert decrypt_token(old_ciphertext) == token + + def test_encrypts_under_first_key(self, monkeypatch: pytest.MonkeyPatch) -> None: + new_key = Fernet.generate_key().decode() + old_key = Fernet.generate_key().decode() + + _set_key(monkeypatch, f"{new_key},{old_key}") + token = "ghp_new_secret" + ciphertext = encrypt_token(token) + + assert Fernet(new_key.encode()).decrypt(ciphertext.encode()).decode() == token + + def test_decrypt_fails_when_no_key_matches(self, monkeypatch: pytest.MonkeyPatch) -> None: + old_key = Fernet.generate_key().decode() + _set_key(monkeypatch, old_key) + old_ciphertext = encrypt_token("ghp_token") + + unrelated_key = Fernet.generate_key().decode() + _set_key(monkeypatch, unrelated_key) + assert decrypt_token(old_ciphertext) == "" + + def test_newline_separated_keys(self, monkeypatch: pytest.MonkeyPatch) -> None: + old_key = Fernet.generate_key().decode() + new_key = Fernet.generate_key().decode() + + _set_key(monkeypatch, old_key) + old_ciphertext = encrypt_token("ghp_old") + + _set_key(monkeypatch, f"{new_key}\n{old_key}") + assert decrypt_token(old_ciphertext) == "ghp_old" + + +class TestRotationRoundtrip: + def test_full_rotation_lifecycle(self, monkeypatch: pytest.MonkeyPatch) -> None: + old_key = Fernet.generate_key().decode() + new_key = Fernet.generate_key().decode() + + _set_key(monkeypatch, old_key) + token = "ghp_lifecycle" + old_ciphertext = encrypt_token(token) + + _set_key(monkeypatch, f"{new_key},{old_key}") + assert decrypt_token(old_ciphertext) == token + re_encrypted = encrypt_token(decrypt_token(old_ciphertext)) + assert re_encrypted != old_ciphertext + + _set_key(monkeypatch, new_key) + assert decrypt_token(re_encrypted) == token + assert decrypt_token(old_ciphertext) == ""