mirror of
https://github.com/Sea-Haven-Industries/open-swe.git
synced 2026-09-30 08:03:15 +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 ===
|
||||||
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.
|
||||||
|
|
|
||||||
|
|
@ -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
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