diff --git a/crypto/crypto.go b/crypto/crypto.go index c214b72..1c9e50b 100644 --- a/crypto/crypto.go +++ b/crypto/crypto.go @@ -1,6 +1,7 @@ package crypto import ( + "bytes" "crypto/aes" "crypto/cipher" "crypto/rand" @@ -171,3 +172,42 @@ func DecryptStream(reader io.Reader, writer io.WriteCloser, key string) error { return nil } + +// EncryptBytes encrypts a byte array and returns the encrypted bytes. +// For better performance with large data or when you have readers/writers available, +// use EncryptStream instead. +func EncryptBytes(data []byte, key string) ([]byte, error) { + reader := bytes.NewReader(data) + var buf bytes.Buffer + writer := &nopCloser{&buf} + + if err := EncryptStream(reader, writer, key); err != nil { + return nil, err + } + + return buf.Bytes(), nil +} + +// DecryptBytes decrypts a byte array and returns the decrypted bytes. +// For better performance with large data or when you have readers/writers available, +// use DecryptStream instead. +func DecryptBytes(encryptedData []byte, key string) ([]byte, error) { + reader := bytes.NewReader(encryptedData) + var buf bytes.Buffer + writer := &nopCloser{&buf} + + if err := DecryptStream(reader, writer, key); err != nil { + return nil, err + } + + return buf.Bytes(), nil +} + +// nopCloser wraps an io.Writer to implement io.WriteCloser with a no-op Close method +type nopCloser struct { + io.Writer +} + +func (nc nopCloser) Close() error { + return nil +} diff --git a/crypto/crypto_test.go b/crypto/crypto_test.go index eb97cbc..28b7a18 100644 --- a/crypto/crypto_test.go +++ b/crypto/crypto_test.go @@ -8,12 +8,6 @@ import ( "testing" ) -type nopCloser struct { - io.Writer -} - -func (nopCloser) Close() error { return nil } - func TestStreamEncryptionDecryption(t *testing.T) { tests := []struct { name string @@ -684,3 +678,48 @@ func TestStreamKeyReuse(t *testing.T) { t.Error("DecryptStream() should fail with wrong key") } } + +func TestBytesEncryptionDecryption(t *testing.T) { + tests := []struct { + name string + data []byte + key string + }{ + { + name: "simple data", + data: []byte("hello world"), + key: "test-key", + }, + { + name: "empty data", + data: []byte{}, + key: "test-key", + }, + { + name: "large data", + data: bytes.Repeat([]byte("test"), 1000), + key: "test-key", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + // Encrypt + encrypted, err := EncryptBytes(tt.data, tt.key) + if err != nil { + t.Fatalf("EncryptBytes() error = %v", err) + } + + // Decrypt + decrypted, err := DecryptBytes(encrypted, tt.key) + if err != nil { + t.Fatalf("DecryptBytes() error = %v", err) + } + + // Verify + if !bytes.Equal(decrypted, tt.data) { + t.Errorf("Decrypted data doesn't match original") + } + }) + } +}