242 lines
6.2 KiB
Go
242 lines
6.2 KiB
Go
package crypto
|
|
|
|
import (
|
|
"bytes"
|
|
"testing"
|
|
)
|
|
|
|
func TestDeriveKeyFromPassword(t *testing.T) {
|
|
salt, err := GenerateSalt()
|
|
if err != nil {
|
|
t.Fatalf("GenerateSalt: %v", err)
|
|
}
|
|
|
|
key1 := DeriveKeyFromPassword("hunter2", salt)
|
|
key2 := DeriveKeyFromPassword("hunter2", salt)
|
|
|
|
if !bytes.Equal(key1, key2) {
|
|
t.Fatal("same password+salt should produce same key")
|
|
}
|
|
if len(key1) != 32 {
|
|
t.Fatalf("expected 32-byte key, got %d", len(key1))
|
|
}
|
|
|
|
// Different password → different key
|
|
key3 := DeriveKeyFromPassword("different", salt)
|
|
if bytes.Equal(key1, key3) {
|
|
t.Fatal("different passwords should produce different keys")
|
|
}
|
|
|
|
// Different salt → different key
|
|
salt2, _ := GenerateSalt()
|
|
key4 := DeriveKeyFromPassword("hunter2", salt2)
|
|
if bytes.Equal(key1, key4) {
|
|
t.Fatal("different salts should produce different keys")
|
|
}
|
|
}
|
|
|
|
func TestDeriveKeyFromEnv(t *testing.T) {
|
|
key1, err := DeriveKeyFromEnv("my-secret-key-123")
|
|
if err != nil {
|
|
t.Fatalf("DeriveKeyFromEnv: %v", err)
|
|
}
|
|
if len(key1) != 32 {
|
|
t.Fatalf("expected 32-byte key, got %d", len(key1))
|
|
}
|
|
|
|
// Deterministic
|
|
key2, _ := DeriveKeyFromEnv("my-secret-key-123")
|
|
if !bytes.Equal(key1, key2) {
|
|
t.Fatal("same input should produce same key")
|
|
}
|
|
|
|
// Empty → error
|
|
_, err = DeriveKeyFromEnv("")
|
|
if err != ErrNoEncryptionKey {
|
|
t.Fatalf("expected ErrNoEncryptionKey, got %v", err)
|
|
}
|
|
}
|
|
|
|
func TestUEKRoundTrip(t *testing.T) {
|
|
uek, err := GenerateUEK()
|
|
if err != nil {
|
|
t.Fatalf("GenerateUEK: %v", err)
|
|
}
|
|
if len(uek) != 32 {
|
|
t.Fatalf("expected 32-byte UEK, got %d", len(uek))
|
|
}
|
|
|
|
salt, _ := GenerateSalt()
|
|
pdk := DeriveKeyFromPassword("my-password", salt)
|
|
|
|
ciphertext, nonce, err := WrapUEK(uek, pdk)
|
|
if err != nil {
|
|
t.Fatalf("WrapUEK: %v", err)
|
|
}
|
|
|
|
recovered, err := UnwrapUEK(ciphertext, nonce, pdk)
|
|
if err != nil {
|
|
t.Fatalf("UnwrapUEK: %v", err)
|
|
}
|
|
|
|
if !bytes.Equal(uek, recovered) {
|
|
t.Fatal("UEK round-trip mismatch")
|
|
}
|
|
}
|
|
|
|
func TestUEKWrongPassword(t *testing.T) {
|
|
uek, _ := GenerateUEK()
|
|
salt, _ := GenerateSalt()
|
|
|
|
pdk := DeriveKeyFromPassword("correct-password", salt)
|
|
ciphertext, nonce, _ := WrapUEK(uek, pdk)
|
|
|
|
wrongPDK := DeriveKeyFromPassword("wrong-password", salt)
|
|
_, err := UnwrapUEK(ciphertext, nonce, wrongPDK)
|
|
if err != ErrDecryptionFailed {
|
|
t.Fatalf("expected ErrDecryptionFailed, got %v", err)
|
|
}
|
|
}
|
|
|
|
func TestAPIKeyRoundTrip(t *testing.T) {
|
|
key, _ := DeriveKeyFromEnv("test-encryption-key")
|
|
apiKey := "sk-ant-api03-very-secret-key-1234567890"
|
|
|
|
ciphertext, nonce, err := Encrypt(apiKey, key)
|
|
if err != nil {
|
|
t.Fatalf("Encrypt: %v", err)
|
|
}
|
|
|
|
recovered, err := Decrypt(ciphertext, nonce, key)
|
|
if err != nil {
|
|
t.Fatalf("Decrypt: %v", err)
|
|
}
|
|
|
|
if recovered != apiKey {
|
|
t.Fatalf("API key round-trip mismatch: got %q", recovered)
|
|
}
|
|
}
|
|
|
|
func TestAPIKeyWrongKey(t *testing.T) {
|
|
key1, _ := DeriveKeyFromEnv("key-one")
|
|
key2, _ := DeriveKeyFromEnv("key-two")
|
|
|
|
ciphertext, nonce, _ := Encrypt("secret-api-key", key1)
|
|
|
|
_, err := Decrypt(ciphertext, nonce, key2)
|
|
if err != ErrDecryptionFailed {
|
|
t.Fatalf("expected ErrDecryptionFailed, got %v", err)
|
|
}
|
|
}
|
|
|
|
func TestAPIKeyCorruption(t *testing.T) {
|
|
key, _ := DeriveKeyFromEnv("test-key")
|
|
ciphertext, nonce, _ := Encrypt("my-key", key)
|
|
|
|
// Flip a byte in ciphertext
|
|
corrupted := make([]byte, len(ciphertext))
|
|
copy(corrupted, ciphertext)
|
|
corrupted[0] ^= 0xff
|
|
|
|
_, err := Decrypt(corrupted, nonce, key)
|
|
if err != ErrDecryptionFailed {
|
|
t.Fatalf("expected ErrDecryptionFailed for corrupted data, got %v", err)
|
|
}
|
|
}
|
|
|
|
func TestEncryptEmpty(t *testing.T) {
|
|
key, _ := DeriveKeyFromEnv("test-key")
|
|
|
|
ciphertext, nonce, err := Encrypt("", key)
|
|
if err != nil {
|
|
t.Fatalf("Encrypt empty: %v", err)
|
|
}
|
|
if ciphertext != nil || nonce != nil {
|
|
t.Fatal("encrypting empty string should return nil, nil")
|
|
}
|
|
|
|
result, err := Decrypt(nil, nil, key)
|
|
if err != nil {
|
|
t.Fatalf("Decrypt nil: %v", err)
|
|
}
|
|
if result != "" {
|
|
t.Fatalf("decrypting nil should return empty string, got %q", result)
|
|
}
|
|
}
|
|
|
|
func TestUEKUniqueness(t *testing.T) {
|
|
uek1, _ := GenerateUEK()
|
|
uek2, _ := GenerateUEK()
|
|
if bytes.Equal(uek1, uek2) {
|
|
t.Fatal("two generated UEKs should not be equal")
|
|
}
|
|
}
|
|
|
|
func TestFullPersonalKeyFlow(t *testing.T) {
|
|
// Simulate: user registers → UEK generated → UEK wrapped with password
|
|
// → user logs in → UEK unwrapped → personal API key encrypted with UEK
|
|
// → completion request decrypts API key with UEK
|
|
|
|
password := "user-secure-password-123"
|
|
apiKey := "sk-personal-byok-key-abc123"
|
|
|
|
// Registration: generate UEK and wrap with password
|
|
uek, _ := GenerateUEK()
|
|
salt, _ := GenerateSalt()
|
|
pdk := DeriveKeyFromPassword(password, salt)
|
|
wrappedUEK, uekNonce, _ := WrapUEK(uek, pdk)
|
|
|
|
// User adds a personal provider: encrypt API key with UEK
|
|
keyCiphertext, keyNonce, _ := Encrypt(apiKey, uek)
|
|
|
|
// --- simulate session boundary ---
|
|
|
|
// Login: derive PDK from password, unwrap UEK
|
|
loginPDK := DeriveKeyFromPassword(password, salt)
|
|
sessionUEK, err := UnwrapUEK(wrappedUEK, uekNonce, loginPDK)
|
|
if err != nil {
|
|
t.Fatalf("login UEK unwrap: %v", err)
|
|
}
|
|
|
|
// Completion request: decrypt API key with session UEK
|
|
recovered, err := Decrypt(keyCiphertext, keyNonce, sessionUEK)
|
|
if err != nil {
|
|
t.Fatalf("API key decrypt: %v", err)
|
|
}
|
|
|
|
if recovered != apiKey {
|
|
t.Fatalf("full flow mismatch: got %q", recovered)
|
|
}
|
|
}
|
|
|
|
func TestPasswordChangePreservesKeys(t *testing.T) {
|
|
// UEK stays the same — only the wrapping changes
|
|
oldPassword := "old-password"
|
|
newPassword := "new-password"
|
|
apiKey := "sk-my-key"
|
|
|
|
uek, _ := GenerateUEK()
|
|
salt, _ := GenerateSalt()
|
|
|
|
// Wrap with old password
|
|
oldPDK := DeriveKeyFromPassword(oldPassword, salt)
|
|
_, _, _ = WrapUEK(uek, oldPDK)
|
|
|
|
// Encrypt API key with UEK
|
|
keyCiphertext, keyNonce, _ := Encrypt(apiKey, uek)
|
|
|
|
// Password change: re-wrap UEK with new password (new salt too)
|
|
newSalt, _ := GenerateSalt()
|
|
newPDK := DeriveKeyFromPassword(newPassword, newSalt)
|
|
newWrappedUEK, newUEKNonce, _ := WrapUEK(uek, newPDK)
|
|
|
|
// Login with new password: unwrap UEK, decrypt API key
|
|
loginPDK := DeriveKeyFromPassword(newPassword, newSalt)
|
|
recoveredUEK, _ := UnwrapUEK(newWrappedUEK, newUEKNonce, loginPDK)
|
|
|
|
recovered, _ := Decrypt(keyCiphertext, keyNonce, recoveredUEK)
|
|
if recovered != apiKey {
|
|
t.Fatalf("password change broke key access: got %q", recovered)
|
|
}
|
|
}
|