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) } }