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] <open-swe@users.noreply.github.com>
Co-authored-by: Johannes du Plessis <51395795+johannes117@users.noreply.github.com>
This commit is contained in:
open-swe[bot] 2026-05-08 14:29:23 -07:00 • committed by GitHub
parent 3f9dbb6597
commit 148efeb269
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 195 additions and 31 deletions

View file

@ -418,8 +418,32 @@ DEFAULT_SANDBOX_DELETE_AFTER_STOP_SECONDS="" # Delete N seconds after stop (def
# === Token Encryption === # === Token Encryption ===
TOKEN_ENCRYPTION_KEY="" # Generate with: openssl rand -base64 32 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="<new_key>,<old_key>"
```
Restart the server. New encryptions use `<new_key>`; existing ciphertexts
still decrypt against `<old_key>`.
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="<new_key>"
```
Any thread still holding ciphertext under `<old_key>` 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 ## 7. Start the server
Make sure ngrok is still running from step 2, then start the LangGraph server in a second terminal: 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`) - Ensure `TOKEN_ENCRYPTION_KEY` is set (generate with `openssl rand -base64 32`)
- The key must be a valid 32-byte Fernet-compatible base64 string - 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.

View file

@ -3,7 +3,7 @@
import logging import logging
import os import os
from cryptography.fernet import Fernet, InvalidToken from cryptography.fernet import Fernet, InvalidToken, MultiFernet
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@ -12,59 +12,56 @@ class EncryptionKeyMissingError(ValueError):
"""Raised when TOKEN_ENCRYPTION_KEY environment variable is not set.""" """Raised when TOKEN_ENCRYPTION_KEY environment variable is not set."""
def _get_encryption_key() -> bytes: def _parse_encryption_keys(raw: str) -> list[bytes]:
"""Get or derive the encryption key from environment variable. """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: def _get_encryption_keys() -> list[bytes]:
32-byte Fernet-compatible key """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: 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") explicit_key = os.environ.get("TOKEN_ENCRYPTION_KEY")
if not explicit_key: if not explicit_key:
raise EncryptionKeyMissingError 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: def encrypt_token(token: str) -> str:
"""Encrypt a token for safe storage. """Encrypt a token under the newest configured key."""
Args:
token: The plaintext token to encrypt
Returns:
Base64-encoded encrypted token
"""
if not token: if not token:
return "" return ""
key = _get_encryption_key() encrypted = _get_fernet().encrypt(token.encode())
f = Fernet(key)
encrypted = f.encrypt(token.encode())
return encrypted.decode() return encrypted.decode()
def decrypt_token(encrypted_token: str) -> str: def decrypt_token(encrypted_token: str) -> str:
"""Decrypt an encrypted token. """Decrypt a token, trying each configured key in order."""
Args:
encrypted_token: The base64-encoded encrypted token
Returns:
The plaintext token, or empty string if decryption fails
"""
if not encrypted_token: if not encrypted_token:
return "" return ""
try: try:
key = _get_encryption_key() decrypted = _get_fernet().decrypt(encrypted_token.encode())
f = Fernet(key)
decrypted = f.decrypt(encrypted_token.encode())
return decrypted.decode() return decrypted.decode()
except InvalidToken: except InvalidToken:
logger.warning("Failed to decrypt token: invalid token") logger.warning("Failed to decrypt token: invalid token")

141
tests/test_encryption.py Normal file
View file

@ -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) == ""