mirror of
https://github.com/Sea-Haven-Industries/open-swe.git
synced 2026-09-30 05:43:14 +00:00
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:
parent
3f9dbb6597
commit
148efeb269
3 changed files with 195 additions and 31 deletions
|
|
@ -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="<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
|
||||
|
||||
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.
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
141
tests/test_encryption.py
Normal file
141
tests/test_encryption.py
Normal 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) == ""
|
||||
Loading…
Add table
Reference in a new issue