package encryption import ( "bytes" "encoding/base64" "strings" "testing" ) func TestCipherRoundtrip(t *testing.T) { cipher, err := New(bytes.Repeat([]byte{1}, keySize)) if err != nil { t.Fatalf("New() error = %v", err) } plaintext := []byte("secret-proxmox-token") ciphertext, err := cipher.Encrypt(plaintext) if err != nil { t.Fatalf("Encrypt() error = %v", err) } if bytes.Contains(ciphertext, plaintext) { t.Fatal("ciphertext contains plaintext") } if ciphertext[0] != currentVersion { t.Fatalf("version = %d, want %d", ciphertext[0], currentVersion) } got, err := cipher.Decrypt(ciphertext) if err != nil { t.Fatalf("Decrypt() error = %v", err) } if string(got) != string(plaintext) { t.Fatalf("plaintext = %q, want %q", got, plaintext) } } func TestDecryptRejectsWrongKey(t *testing.T) { cipherA, err := New(bytes.Repeat([]byte{1}, keySize)) if err != nil { t.Fatalf("New() error = %v", err) } cipherB, err := New(bytes.Repeat([]byte{2}, keySize)) if err != nil { t.Fatalf("New() error = %v", err) } ciphertext, err := cipherA.Encrypt([]byte("secret-proxmox-token")) if err != nil { t.Fatalf("Encrypt() error = %v", err) } if _, err := cipherB.Decrypt(ciphertext); err == nil { t.Fatal("Decrypt() error = nil, want error") } } func TestDecryptRejectsUnknownVersion(t *testing.T) { cipher, err := New(bytes.Repeat([]byte{1}, keySize)) if err != nil { t.Fatalf("New() error = %v", err) } ciphertext, err := cipher.Encrypt([]byte("secret-proxmox-token")) if err != nil { t.Fatalf("Encrypt() error = %v", err) } ciphertext[0] = 99 if _, err := cipher.Decrypt(ciphertext); err == nil { t.Fatal("Decrypt() error = nil, want error") } } func TestNewRejectsInvalidKeySize(t *testing.T) { if _, err := New(bytes.Repeat([]byte{1}, keySize-1)); err == nil { t.Fatal("New() error = nil, want error") } } func TestNewFromBase64(t *testing.T) { encoded := base64.StdEncoding.EncodeToString(bytes.Repeat([]byte{1}, keySize)) if _, err := NewFromBase64(encoded); err != nil { t.Fatalf("NewFromBase64() error = %v", err) } } func TestNewFromBase64RejectsMissingKey(t *testing.T) { _, err := NewFromBase64("") if err == nil { t.Fatal("NewFromBase64() error = nil, want error") } if !strings.Contains(err.Error(), "MASTER_KEY_BASE64") { t.Fatalf("error = %q, want MASTER_KEY_BASE64 hint", err) } }