Files
go-synapse-backupper/pkg/adapters/crypto/composite/composite_test.go
T

950 lines
28 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
package composite
import (
"bytes"
"crypto/aes"
"crypto/cipher"
"crypto/rand"
"crypto/sha256"
"encoding/base64"
"encoding/binary"
"encoding/json"
"errors"
"fmt"
"io"
"os"
"testing"
"git.tswf.io/infra/go-synapse-backupper/pkg/domain/crypto"
)
// ---------------------------------------------------------------------------
// Test KEM harness
//
// The composite format pins slot 0 = PQ (schemeID 0x0006, ciphertext length
// 1088) and slot 1 = classical (schemeID 0x0007, ciphertext length 32),
// matching the real mlkem768 + x25519 adapter contracts. Adapters' priv types
// are unexported and reject type-asserted impostors at Decapsulate, and the
// composite todo's scope forbids touching adapter packages — so tests exercise
// the composite with deterministic fake KEMs (registered under the SAME
// schemeIDs as the real adapters). The committed golden fixture uses these
// fakes; the composite production code is exercised end-to-end on the format,
// the HKDF combiner, AES-256-GCM wrapping, and chunked AEAD.
// ---------------------------------------------------------------------------
const (
fakePqSchemeID uint16 = 0x0006
fakeClassicalSchemeID uint16 = 0x0007
fakePqCtLen int = 1088 // matches crypto/mlkem EncapsulateKey768 ciphertext length
fakeClassicalCtLen int = 32 // matches X25519 ephemeral pubkey length
fakeSeedLen int = 32
)
// fakeKem derives a deterministic shared secret per pub/ct pair:
//
// ss = SHA256(pub_raw_or_priv_raw || ct)
//
// where pub.raw and priv.raw are both the random seed; randomness lives only
// in the ct (the call site's rand supplies ct bytes), so decapsulation with
// the matching priv always recovers the encryption-time ss.
type fakeKem struct {
schemeIDValue uint16
ctLenValue int
}
func newFakePqKem() crypto.KEM {
return &fakeKem{schemeIDValue: fakePqSchemeID, ctLenValue: fakePqCtLen}
}
func newFakeClassicalKem() crypto.KEM {
return &fakeKem{schemeIDValue: fakeClassicalSchemeID, ctLenValue: fakeClassicalCtLen}
}
func (k *fakeKem) SchemeID() uint16 { return k.schemeIDValue }
func (k *fakeKem) GenerateKeyPair(
rand io.Reader,
) (
crypto.RecipientPub,
crypto.RecipientPriv,
error,
) {
seed := make([]byte, fakeSeedLen)
if _, err := io.ReadFull(rand, seed); err != nil {
return nil, nil, err
}
return newFakePub(k.schemeIDValue, seed), newFakePriv(k.schemeIDValue, seed), nil
}
func (k *fakeKem) Encapsulate(
pub crypto.RecipientPub,
rand io.Reader,
) (
ciphertext []byte,
sharedSecret []byte,
err error,
) {
p, ok := pub.(*fakePub)
if !ok {
return nil, nil, errors.New("fakeKem: invalid pub type")
}
ciphertext = make([]byte, k.ctLenValue)
if _, err := io.ReadFull(rand, ciphertext); err != nil {
return nil, nil, err
}
return ciphertext, deriveFakeSS(p.raw, ciphertext), nil
}
func (k *fakeKem) Decapsulate(
priv crypto.RecipientPriv,
ciphertext []byte,
) (
sharedSecret []byte,
err error,
) {
p, ok := priv.(*fakePriv)
if !ok {
return nil, errors.New("fakeKem: invalid priv type")
}
if len(ciphertext) != k.ctLenValue {
return nil, errors.New("fakeKem: invalid ciphertext length")
}
return deriveFakeSS(p.raw, ciphertext), nil
}
func (k *fakeKem) LoadPriv(
raw []byte,
) (
crypto.RecipientPriv,
error,
) {
return newFakePriv(k.schemeIDValue, raw), nil
}
func deriveFakeSS(
raw, ciphertext []byte,
) []byte {
h := sha256.New()
h.Write(raw)
h.Write(ciphertext)
return h.Sum(nil)
}
// fakePub / fakePriv — deterministic raw-bytes-backed recipients.
type fakePub struct {
scheme uint16
raw []byte
keyID []byte
}
func newFakePub(
scheme uint16,
raw []byte,
) *fakePub {
h := sha256.Sum256(raw)
return &fakePub{
scheme: scheme,
raw: append([]byte(nil), raw...),
keyID: h[:8],
}
}
func (f *fakePub) SchemeID() uint16 { return f.scheme }
func (f *fakePub) KeyID() []byte { return f.keyID }
func (f *fakePub) Raw() []byte { return f.raw }
type fakePriv struct {
scheme uint16
raw []byte
keyID []byte
}
func newFakePriv(
scheme uint16,
raw []byte,
) *fakePriv {
h := sha256.Sum256(raw)
return &fakePriv{
scheme: scheme,
raw: append([]byte(nil), raw...),
keyID: h[:8],
}
}
func (f *fakePriv) SchemeID() uint16 { return f.scheme }
func (f *fakePriv) KeyID() []byte { return f.keyID }
func (f *fakePriv) Raw() []byte { return f.raw }
// fakeRegistry returns a Registry with the two fake KEMs registered under
// the v2 slot schemeIDs.
func fakeRegistry(
t *testing.T,
) crypto.Registry {
t.Helper()
reg := crypto.NewRegistry()
if err := reg.Register(fakePqSchemeID, newFakePqKem); err != nil {
t.Fatalf("register pq fake: %v", err)
}
if err := reg.Register(fakeClassicalSchemeID, newFakeClassicalKem); err != nil {
t.Fatalf("register classical fake: %v", err)
}
return reg
}
// generateFakeKeyPair generates (pub, priv) for slot schemeID from rand.
func generateFakeKeyPair(
t *testing.T,
reg crypto.Registry,
schemeID uint16,
rand io.Reader,
) (
crypto.RecipientPub,
crypto.RecipientPriv,
) {
t.Helper()
factory, err := reg.Lookup(schemeID)
if err != nil {
t.Fatalf("lookup 0x%04x: %v", schemeID, err)
}
pub, priv, err := factory().GenerateKeyPair(rand)
if err != nil {
t.Fatalf("generate 0x%04x: %v", schemeID, err)
}
return pub, priv
}
// standardHeaderLen returns the fixed artifact-header length given the two
// slot ciphertext lengths: 11 (magic+version+flags+nRecipients) +
// per-slot (14 + ctLen) + 72 (wrapNonce+wrappedCEK+firstPayloadNonce).
func standardHeaderLen(
pqCtLen, classicalCtLen int,
) int {
return 11 + (14 + pqCtLen) + (14 + classicalCtLen) + (12 + 48 + 12)
}
// countingReader wraps an io.Reader and counts how many bytes have been read
// — used by the adversarial-parser tests to assert the parser does NOT
// consume past the header before bailing out.
type countingReader struct {
r io.Reader
n int64
}
func (c *countingReader) Read(
p []byte,
) (int, error) {
readN, err := c.r.Read(p)
c.n += int64(readN)
return readN, err
}
// ---------------------------------------------------------------------------
// Tests (a) through (p)
// ---------------------------------------------------------------------------
// (a) Golden format fixture.
func TestGoldenFormat(
t *testing.T,
) {
goldenBytes := mustReadFile(t, "testdata/golden-1byte.pqenc")
// Header byte offsets pinned verbatim — if any of these breaks, the
// on-disk format has drifted and old .pqenc files won't decrypt.
// [0:4] magic u32 BE = 0x47535051
// [4:6] version u16 BE = 0x0002
// [6:10] flags u32 BE = 0x00000000
// [10] nRecipients u8 = 0x02
// [11:13] slot0 schemeID = 0x0006
// [13:21] slot0 keyID 8B
// [21:25] slot0 ctLen u32 = 1088
// [25:1113] slot0 ciphertext (1088 bytes)
// [1113:1115] slot1 schemeID = 0x0007
if binary.BigEndian.Uint32(goldenBytes[0:4]) != 0x47535051 {
t.Errorf("magic = 0x%08x, want 0x47535051", binary.BigEndian.Uint32(goldenBytes[0:4]))
}
if binary.BigEndian.Uint16(goldenBytes[4:6]) != 0x0002 {
t.Errorf("version = 0x%04x, want 0x0002", binary.BigEndian.Uint16(goldenBytes[4:6]))
}
if binary.BigEndian.Uint32(goldenBytes[6:10]) != 0x00000000 {
t.Errorf("flags = 0x%08x, want 0", binary.BigEndian.Uint32(goldenBytes[6:10]))
}
if goldenBytes[10] != 0x02 {
t.Errorf("nRecipients = 0x%02x, want 0x02", goldenBytes[10])
}
if binary.BigEndian.Uint16(goldenBytes[11:13]) != 0x0006 {
t.Errorf("slot0 schemeID = 0x%04x, want 0x0006", binary.BigEndian.Uint16(goldenBytes[11:13]))
}
if binary.BigEndian.Uint16(goldenBytes[1113:1115]) != 0x0007 {
t.Errorf("slot1 schemeID = 0x%04x, want 0x0007", binary.BigEndian.Uint16(goldenBytes[1113:1115]))
}
// Decrypt-equality: reconstruct privs from committed golden-keys.json
// and assert Decrypt yields the 0xAA plaintext committed via golden_generate.
pqPriv, classicalPriv := loadGoldenPrivs(t, "testdata/golden-keys.json")
dec := NewDecryptor(fakeRegistry(t))
out := &bytes.Buffer{}
if err := dec.Decrypt(
bytes.NewReader(goldenBytes),
[]crypto.RecipientPriv{pqPriv, classicalPriv},
out,
); err != nil {
t.Fatalf("Decrypt(golden) failed: %v", err)
}
if !bytes.Equal(out.Bytes(), []byte{0xAA}) {
t.Errorf("decrypted = %x, want [0xAA]", out.Bytes())
}
}
// (b) Round-trip on canonical input sizes.
func TestRoundTrip(
t *testing.T,
) {
sizes := []int{0, 1, 64*1024 - 1, 64 * 1024, 64*1024 + 1, 1 << 20}
for _, size := range sizes {
t.Run(fmt.Sprintf("size=%d", size), func(t *testing.T) {
plaintext := make([]byte, size)
for i := 0; i < size; i++ {
plaintext[i] = byte(i)
}
reg := fakeRegistry(t)
enc := NewEncryptor(reg)
dec := NewDecryptor(reg)
pqPub, pqPriv := generateFakeKeyPair(t, reg, fakePqSchemeID, rand.Reader)
classicalPub, classicalPriv := generateFakeKeyPair(t, reg, fakeClassicalSchemeID, rand.Reader)
var encrypted bytes.Buffer
if err := enc.Encrypt(
bytes.NewReader(plaintext),
[]crypto.RecipientPub{pqPub, classicalPub},
&encrypted,
rand.Reader,
); err != nil {
t.Fatalf("Encrypt: %v", err)
}
out := &bytes.Buffer{}
if err := dec.Decrypt(
bytes.NewReader(encrypted.Bytes()),
[]crypto.RecipientPriv{pqPriv, classicalPriv},
out,
); err != nil {
t.Fatalf("Decrypt: %v", err)
}
if !bytes.Equal(out.Bytes(), plaintext) {
t.Errorf("round-trip mismatch: got %d bytes, want %d", out.Len(), len(plaintext))
}
})
}
}
// (c) Empty plaintext produces exactly ONE chunk with flags=0x01 and a 16B
// (tag-only) ciphertext.
func TestEmptyPlaintextSingleFinalChunk(
t *testing.T,
) {
reg := fakeRegistry(t)
enc := NewEncryptor(reg)
pqPub, _ := generateFakeKeyPair(t, reg, fakePqSchemeID, rand.Reader)
classicalPub, _ := generateFakeKeyPair(t, reg, fakeClassicalSchemeID, rand.Reader)
var encrypted bytes.Buffer
if err := enc.Encrypt(
bytes.NewReader(nil),
[]crypto.RecipientPub{pqPub, classicalPub},
&encrypted,
rand.Reader,
); err != nil {
t.Fatalf("Encrypt: %v", err)
}
headerLen := standardHeaderLen(fakePqCtLen, fakeClassicalCtLen)
chunks := encrypted.Bytes()[headerLen:]
// Expected record: [len=16 u32 (4B)][flags=0x01 (1B)][ciphertext (16B)].
if len(chunks) != 4+1+16 {
t.Fatalf("expected 21-byte chunk record, got %d bytes", len(chunks))
}
if ctLen := binary.BigEndian.Uint32(chunks[0:4]); ctLen != 16 {
t.Errorf("ctLen = %d, want 16 (tag-only)", ctLen)
}
if chunks[4] != 0x01 {
t.Errorf("flags = 0x%02x, want 0x01", chunks[4])
}
if len(chunks[5:]) != 16 {
t.Errorf("ciphertext = %d bytes, want 16 (tag-only)", len(chunks[5:]))
}
}
// (d) Exactly-64KiB input produces TWO chunks: full body (flags=0x00) and
// zero-length final marker (flags=0x01).
func TestExactly64KiBTwoChunks(
t *testing.T,
) {
reg := fakeRegistry(t)
enc := NewEncryptor(reg)
pqPub, _ := generateFakeKeyPair(t, reg, fakePqSchemeID, rand.Reader)
classicalPub, _ := generateFakeKeyPair(t, reg, fakeClassicalSchemeID, rand.Reader)
plaintext := bytes.Repeat([]byte{0xCC}, 64*1024)
var encrypted bytes.Buffer
if err := enc.Encrypt(
bytes.NewReader(plaintext),
[]crypto.RecipientPub{pqPub, classicalPub},
&encrypted,
rand.Reader,
); err != nil {
t.Fatalf("Encrypt: %v", err)
}
headerLen := standardHeaderLen(fakePqCtLen, fakeClassicalCtLen)
chunks := encrypted.Bytes()[headerLen:]
// Chunk 1: full body. ct = 64 KiB plaintext + 16B tag.
const bodyCtLen = 64*1024 + 16
if ctLen := binary.BigEndian.Uint32(chunks[0:4]); ctLen != bodyCtLen {
t.Errorf("chunk1 ctLen = %d, want %d", ctLen, bodyCtLen)
}
if chunks[4] != 0x00 {
t.Errorf("chunk1 flags = 0x%02x, want 0x00", chunks[4])
}
// Chunk 2: zero-length final marker (ct = 16B tag), flags = 0x01.
chunk2Start := 4 + 1 + bodyCtLen
if chunk2Start+5 > len(chunks) {
t.Fatalf("file truncated before chunk 2: need offset %d, have %d", chunk2Start+5, len(chunks))
}
if ctLen := binary.BigEndian.Uint32(chunks[chunk2Start : chunk2Start+4]); ctLen != 16 {
t.Errorf("chunk2 ctLen = %d, want 16 (zero-length marker)", ctLen)
}
if chunks[chunk2Start+4] != 0x01 {
t.Errorf("chunk2 flags = 0x%02x, want 0x01", chunks[chunk2Start+4])
}
chunk3Start := chunk2Start + 4 + 1 + 16
if chunk3Start != len(chunks) {
t.Errorf("expected exactly 2 chunks; remaining = %d bytes after chunk 2", len(chunks)-chunk3Start)
}
}
// (e) Counter wrap-around: with chunkNonce[4:12]=0xFFFFFFFFFFFFFFFF, an
// attempt to encrypt a SECOND body chunk fails on counter increment and
// returns ErrNonceCounterWrapped.
func TestCounterWraparound(
t *testing.T,
) {
// firstPayloadNonce: slot 0..3 = arbitrary base; slot 4..11 = 0xFF*8.
firstPayloadNonce := make([]byte, 12)
firstPayloadNonce[0] = 0xde
firstPayloadNonce[1] = 0xad
firstPayloadNonce[2] = 0xbe
firstPayloadNonce[3] = 0xef
for i := 4; i < 12; i++ {
firstPayloadNonce[i] = 0xFF
}
cek := make([]byte, 32)
if _, err := io.ReadFull(rand.Reader, cek); err != nil {
t.Fatalf("rand: %v", err)
}
block, err := aes.NewCipher(cek)
if err != nil {
t.Fatalf("aes: %v", err)
}
gcm, err := cipher.NewGCM(block)
if err != nil {
t.Fatalf("gcm: %v", err)
}
// 64 KiB + 1 byte forces 2 body chunks; incrementing after chunk 1
// wraps to 0 and ErrNonceCounterWrapped (the second chunk's emission
// never happens).
input := make([]byte, chunkSize+1)
var out bytes.Buffer
err = encryptChunks(bytes.NewReader(input), &out, gcm, firstPayloadNonce)
if !errors.Is(err, ErrNonceCounterWrapped) {
t.Errorf("expected ErrNonceCounterWrapped, got %v", err)
}
}
// (f) Tamper 1 byte in payload → ErrTamperingDetected.
func TestTamperPayload(
t *testing.T,
) {
reg := fakeRegistry(t)
enc := NewEncryptor(reg)
dec := NewDecryptor(reg)
pqPub, pqPriv := generateFakeKeyPair(t, reg, fakePqSchemeID, rand.Reader)
classicalPub, classicalPriv := generateFakeKeyPair(t, reg, fakeClassicalSchemeID, rand.Reader)
plaintext := bytes.Repeat([]byte{0x88}, 64*1024+1) // enough to produce a body chunk + final
var encrypted bytes.Buffer
if err := enc.Encrypt(
bytes.NewReader(plaintext),
[]crypto.RecipientPub{pqPub, classicalPub},
&encrypted,
rand.Reader,
); err != nil {
t.Fatalf("Encrypt: %v", err)
}
buf := encrypted.Bytes()
headerLen := standardHeaderLen(fakePqCtLen, fakeClassicalCtLen)
tamperIdx := headerLen + 4 + 1 + 8 // into first chunk ciphertext, past len + flags
if tamperIdx >= len(buf) {
t.Fatalf("file too small to tamper: idx=%d len=%d", tamperIdx, len(buf))
}
buf[tamperIdx] ^= 0x01
out := &bytes.Buffer{}
err := dec.Decrypt(
bytes.NewReader(buf),
[]crypto.RecipientPriv{pqPriv, classicalPriv},
out,
)
if !errors.Is(err, ErrTamperingDetected) {
t.Errorf("expected ErrTamperingDetected, got %v", err)
}
}
// (g) Tamper 1 byte in wrappedCEK → ErrTamperingDetected.
func TestTamperWrappedCEK(
t *testing.T,
) {
reg := fakeRegistry(t)
enc := NewEncryptor(reg)
dec := NewDecryptor(reg)
pqPub, pqPriv := generateFakeKeyPair(t, reg, fakePqSchemeID, rand.Reader)
classicalPub, classicalPriv := generateFakeKeyPair(t, reg, fakeClassicalSchemeID, rand.Reader)
var encrypted bytes.Buffer
if err := enc.Encrypt(
bytes.NewReader([]byte{0xAA}),
[]crypto.RecipientPub{pqPub, classicalPub},
&encrypted,
rand.Reader,
); err != nil {
t.Fatalf("Encrypt: %v", err)
}
buf := encrypted.Bytes()
// wrappedCEK starts at: header-prefix (11) + slot0 (+ct) + slot1 (+ct) + wrapNonce (12).
wrapOffset := 11 + (14 + fakePqCtLen) + (14 + fakeClassicalCtLen) + 12
if wrapOffset+wrappedCekLen > len(buf) {
t.Fatalf("file too short for wrappedCEK")
}
buf[wrapOffset+5] ^= 0x01
out := &bytes.Buffer{}
err := dec.Decrypt(
bytes.NewReader(buf),
[]crypto.RecipientPriv{pqPriv, classicalPriv},
out,
)
if !errors.Is(err, ErrTamperingDetected) {
t.Errorf("expected ErrTamperingDetected, got %v", err)
}
}
// (h) Wrong priv key (swap pq.priv with another) → ErrWrongKeys.
func TestWrongPrivKey(
t *testing.T,
) {
reg := fakeRegistry(t)
enc := NewEncryptor(reg)
dec := NewDecryptor(reg)
pqPub, _ := generateFakeKeyPair(t, reg, fakePqSchemeID, rand.Reader)
classicalPub, classicalPriv := generateFakeKeyPair(t, reg, fakeClassicalSchemeID, rand.Reader)
// Different PQ priv — fresh seed, hence different KeyID.
_, wrongPqPriv := generateFakeKeyPair(t, reg, fakePqSchemeID, rand.Reader)
// Sanity: the wrong priv's keyID must not collide with the original
// pub's (otherwise this test would degrade into a keyID-collision case).
if bytes.Equal(wrongPqPriv.KeyID(), pqPub.KeyID()) {
t.Fatalf("wrongPqPriv keyID accidentally collides with pqPub keyID; reseed")
}
var encrypted bytes.Buffer
if err := enc.Encrypt(
bytes.NewReader([]byte{0x11, 0x22, 0x33}),
[]crypto.RecipientPub{pqPub, classicalPub},
&encrypted,
rand.Reader,
); err != nil {
t.Fatalf("Encrypt: %v", err)
}
out := &bytes.Buffer{}
err := dec.Decrypt(
bytes.NewReader(encrypted.Bytes()),
[]crypto.RecipientPriv{wrongPqPriv, classicalPriv},
out,
)
if !errors.Is(err, ErrWrongKeys) {
t.Errorf("expected ErrWrongKeys, got %v", err)
}
}
// (i) Format conformance: magic, version, nRecipients.
func TestFormatConformance(
t *testing.T,
) {
reg := fakeRegistry(t)
enc := NewEncryptor(reg)
pqPub, _ := generateFakeKeyPair(t, reg, fakePqSchemeID, rand.Reader)
classicalPub, _ := generateFakeKeyPair(t, reg, fakeClassicalSchemeID, rand.Reader)
var out bytes.Buffer
if err := enc.Encrypt(
bytes.NewReader([]byte{0x42}),
[]crypto.RecipientPub{pqPub, classicalPub},
&out,
rand.Reader,
); err != nil {
t.Fatalf("Encrypt: %v", err)
}
outBytes := out.Bytes()
if binary.BigEndian.Uint32(outBytes[0:4]) != 0x47535051 {
t.Errorf("magic = 0x%08x, want 0x47535051", binary.BigEndian.Uint32(outBytes[0:4]))
}
if binary.BigEndian.Uint16(outBytes[4:6]) != 0x0002 {
t.Errorf("version = 0x%04x, want 0x0002", binary.BigEndian.Uint16(outBytes[4:6]))
}
if outBytes[10] != 0x02 {
t.Errorf("nRecipients = 0x%02x, want 0x02", outBytes[10])
}
}
// (j) Adversarial parser: version==0x0001 → ErrUnsupportedVersion, with no
// GCM operations attempted (proven by the post-validation byte counter
// remaining at the prefix length — the parser does not consume past the
// header before bailing).
func TestUnsupportedVersionNoGCM(
t *testing.T,
) {
reg := fakeRegistry(t)
dec := NewDecryptor(reg)
// File: magic + version=0x0001 + flags + nRecipients=0x02 + filler.
var buf bytes.Buffer
_ = binary.Write(&buf, binary.BigEndian, uint32(0x47535051)) // magic
_ = binary.Write(&buf, binary.BigEndian, uint16(0x0001)) // version (downgrade probe)
buf.Write(make([]byte, 200)) // filler
// Dummy privs — irrelevant because the parser bails at version check,
// but Decrypt accepts the slice shape.
_, pqPriv := generateFakeKeyPair(t, reg, fakePqSchemeID, rand.Reader)
_, classicalPriv := generateFakeKeyPair(t, reg, fakeClassicalSchemeID, rand.Reader)
reader := &countingReader{r: bytes.NewReader(buf.Bytes())}
out := &bytes.Buffer{}
err := dec.Decrypt(
reader,
[]crypto.RecipientPriv{pqPriv, classicalPriv},
out,
)
if !errors.Is(err, ErrUnsupportedVersion) {
t.Errorf("expected ErrUnsupportedVersion, got %v", err)
}
// The parser consumed only the 11-byte prefix — no further bytes read,
// hence no GCM operations attempted.
if reader.n != 11 {
t.Errorf("Decrypt consumed %d bytes post-validation; expected exactly 11 (the fixed prefix)", reader.n)
}
}
// (k) nRecipients==0 → ErrMalformedHeader.
func TestZeroRecipients(
t *testing.T,
) {
var buf bytes.Buffer
_ = binary.Write(&buf, binary.BigEndian, uint32(0x47535051)) // magic
_ = binary.Write(&buf, binary.BigEndian, uint16(0x0002)) // version
_ = binary.Write(&buf, binary.BigEndian, uint32(0)) // flags
buf.WriteByte(0x00) // nRecipients = 0
buf.Write(make([]byte, 64)) // filler
dec := NewDecryptor(fakeRegistry(t))
err := dec.Decrypt(bytes.NewReader(buf.Bytes()), nil, &bytes.Buffer{})
if !errors.Is(err, ErrMalformedHeader) {
t.Errorf("expected ErrMalformedHeader, got %v", err)
}
}
// (l) nRecipients>2 OR ctLen > maxRecipientCiphertextLen →
// ErrMalformedHeader BEFORE io.ReadFull attempts to allocate the
// oversized ciphertext buffer.
func TestMalformedHeaderCtLenOverflow(
t *testing.T,
) {
t.Run("nRecipients_gt_2", func(t *testing.T) {
var buf bytes.Buffer
_ = binary.Write(&buf, binary.BigEndian, uint32(0x47535051)) // magic
_ = binary.Write(&buf, binary.BigEndian, uint16(0x0002)) // version
_ = binary.Write(&buf, binary.BigEndian, uint32(0)) // flags
buf.WriteByte(0x03) // nRecipients = 3
buf.Write(make([]byte, 200)) // filler
dec := NewDecryptor(fakeRegistry(t))
err := dec.Decrypt(bytes.NewReader(buf.Bytes()), nil, &bytes.Buffer{})
if !errors.Is(err, ErrMalformedHeader) {
t.Errorf("expected ErrMalformedHeader, got %v", err)
}
})
t.Run("ctLen_overflow", func(t *testing.T) {
var buf bytes.Buffer
_ = binary.Write(&buf, binary.BigEndian, uint32(0x47535051)) // magic
_ = binary.Write(&buf, binary.BigEndian, uint16(0x0002)) // version
_ = binary.Write(&buf, binary.BigEndian, uint32(0)) // flags
buf.WriteByte(0x02) // nRecipients = 2
// Slot 0 metadata only — schemeID, keyID, ctLen = 2 MiB (over the
// 1<<20 cap). The parser validates ctLen BEFORE allocating and
// reading per-slot ciphertext bytes, so it must reject at this
// point without attempting io.ReadFull of an oversized buffer.
_ = binary.Write(&buf, binary.BigEndian, uint16(0x0006))
buf.Write(make([]byte, 8)) // keyID
_ = binary.Write(&buf, binary.BigEndian, uint32(2*1024*1024)) // ctLen
reader := &countingReader{r: bytes.NewReader(buf.Bytes())}
dec := NewDecryptor(fakeRegistry(t))
_, pqPriv := generateFakeKeyPair(t, fakeRegistry(t), fakePqSchemeID, rand.Reader)
_, classicalPriv := generateFakeKeyPair(t, fakeRegistry(t), fakeClassicalSchemeID, rand.Reader)
err := dec.Decrypt(
reader,
[]crypto.RecipientPriv{pqPriv, classicalPriv},
&bytes.Buffer{},
)
if !errors.Is(err, ErrMalformedHeader) {
t.Errorf("expected ErrMalformedHeader, got %v", err)
}
// Consumed exactly prefix(11) + slot0 metadata(14) = 25 bytes — the
// ctLen validation fired before reading any slot1 metadata or any
// per-slot ciphertext.
if reader.n != 25 {
t.Errorf("Decrypt consumed %d bytes; expected 25 (no io.ReadFull of oversized ct)", reader.n)
}
})
}
// (m) Chunk record with length==0 AND flags&0x01==0 → ErrMalformedChunk
// (prevents an infinite-loop DoS where the parser keeps scanning zero-size
// non-final chunks).
func TestZeroLengthNonFinalChunk(
t *testing.T,
) {
reg := fakeRegistry(t)
enc := NewEncryptor(reg)
dec := NewDecryptor(reg)
pqPub, pqPriv := generateFakeKeyPair(t, reg, fakePqSchemeID, rand.Reader)
classicalPub, classicalPriv := generateFakeKeyPair(t, reg, fakeClassicalSchemeID, rand.Reader)
var encrypted bytes.Buffer
if err := enc.Encrypt(
bytes.NewReader([]byte{0xAA}),
[]crypto.RecipientPub{pqPub, classicalPub},
&encrypted,
rand.Reader,
); err != nil {
t.Fatalf("Encrypt: %v", err)
}
headerLen := standardHeaderLen(fakePqCtLen, fakeClassicalCtLen)
var corrupt bytes.Buffer
corrupt.Write(encrypted.Bytes()[:headerLen])
_ = binary.Write(&corrupt, binary.BigEndian, uint32(0)) // ctLen = 0
corrupt.WriteByte(0x00) // flags = 0x00 (NOT final)
out := &bytes.Buffer{}
err := dec.Decrypt(
bytes.NewReader(corrupt.Bytes()),
[]crypto.RecipientPriv{pqPriv, classicalPriv},
out,
)
if !errors.Is(err, ErrMalformedChunk) {
t.Errorf("expected ErrMalformedChunk, got %v", err)
}
}
// (n) Chunk record with length > 64*1024+16 → ErrMalformedChunk.
func TestOversizedChunk(
t *testing.T,
) {
reg := fakeRegistry(t)
enc := NewEncryptor(reg)
dec := NewDecryptor(reg)
pqPub, pqPriv := generateFakeKeyPair(t, reg, fakePqSchemeID, rand.Reader)
classicalPub, classicalPriv := generateFakeKeyPair(t, reg, fakeClassicalSchemeID, rand.Reader)
var encrypted bytes.Buffer
if err := enc.Encrypt(
bytes.NewReader([]byte{0xAA}),
[]crypto.RecipientPub{pqPub, classicalPub},
&encrypted,
rand.Reader,
); err != nil {
t.Fatalf("Encrypt: %v", err)
}
headerLen := standardHeaderLen(fakePqCtLen, fakeClassicalCtLen)
var corrupt bytes.Buffer
corrupt.Write(encrypted.Bytes()[:headerLen])
_ = binary.Write(&corrupt, binary.BigEndian, uint32(chunkSize+gcmTagLen+1)) // oversized
corrupt.WriteByte(0x00) // flags = 0x00
out := &bytes.Buffer{}
err := dec.Decrypt(
bytes.NewReader(corrupt.Bytes()),
[]crypto.RecipientPriv{pqPriv, classicalPriv},
out,
)
if !errors.Is(err, ErrMalformedChunk) {
t.Errorf("expected ErrMalformedChunk, got %v", err)
}
}
// (o) End-of-stream reached BEFORE any chunk with flags&0x01==1
// (truncated file after a body chunk with no final marker) →
// ErrUnexpectedEOF.
func TestPrematureEOF(
t *testing.T,
) {
reg := fakeRegistry(t)
enc := NewEncryptor(reg)
dec := NewDecryptor(reg)
pqPub, pqPriv := generateFakeKeyPair(t, reg, fakePqSchemeID, rand.Reader)
classicalPub, classicalPriv := generateFakeKeyPair(t, reg, fakeClassicalSchemeID, rand.Reader)
plaintext := bytes.Repeat([]byte{0xAB}, 64*1024+1) // 2 body chunks + final marker
var encrypted bytes.Buffer
if err := enc.Encrypt(
bytes.NewReader(plaintext),
[]crypto.RecipientPub{pqPub, classicalPub},
&encrypted,
rand.Reader,
); err != nil {
t.Fatalf("Encrypt: %v", err)
}
headerLen := standardHeaderLen(fakePqCtLen, fakeClassicalCtLen)
bodyCtLen := 64*1024 + 1 + gcmTagLen
bodyChunkRecord := 4 + 1 + bodyCtLen
truncatedLen := headerLen + bodyChunkRecord
if truncatedLen >= len(encrypted.Bytes()) {
t.Fatalf("encrypted file shorter than expected: %d vs expected truncation at %d",
len(encrypted.Bytes()), truncatedLen)
}
out := &bytes.Buffer{}
err := dec.Decrypt(
bytes.NewReader(encrypted.Bytes()[:truncatedLen]),
[]crypto.RecipientPriv{pqPriv, classicalPriv},
out,
)
if !errors.Is(err, ErrUnexpectedEOF) {
t.Errorf("expected ErrUnexpectedEOF, got %v", err)
}
}
// (p) Truncated header → ErrMalformedHeader BEFORE any recipient allocation.
func TestTruncatedHeader(
t *testing.T,
) {
dec := NewDecryptor(fakeRegistry(t))
// File is shorter than the fixed 11-byte prefix + per-recipient metadata
// (2×14=28 = 39 bytes minimum): only 30 bytes total.
var buf bytes.Buffer
_ = binary.Write(&buf, binary.BigEndian, uint32(0x47535051)) // magic
_ = binary.Write(&buf, binary.BigEndian, uint16(0x0002)) // version
_ = binary.Write(&buf, binary.BigEndian, uint32(0)) // flags
buf.WriteByte(0x02) // nRecipients = 2
buf.Write(make([]byte, 20)) // only 20 of the needed 28 metadata bytes
reader := &countingReader{r: bytes.NewReader(buf.Bytes())}
err := dec.Decrypt(reader, nil, &bytes.Buffer{})
if !errors.Is(err, ErrMalformedHeader) {
t.Errorf("expected ErrMalformedHeader, got %v", err)
}
// No recipient ct allocation happened — only the prefix (11) + partial
// metadata (20) = 31 bytes consumed; well shy of a full prefix+meta
// read that would precede any per-recipient ct allocation.
if reader.n > 39 {
t.Errorf("Decrypt consumed %d bytes; expected ≤ 39 — no recipient allocation occurred", reader.n)
}
}
// ---------------------------------------------------------------------------
// Shared test helpers
// ---------------------------------------------------------------------------
func mustReadFile(
t *testing.T,
path string,
) []byte {
t.Helper()
data, err := os.ReadFile(path)
if err != nil {
t.Fatalf("read %s: %v", path, err)
}
return data
}
func mustB64Decode(
t *testing.T,
s string,
) []byte {
t.Helper()
b, err := base64.StdEncoding.DecodeString(s)
if err != nil {
t.Fatalf("base64 decode: %v", err)
}
return b
}
// loadGoldenPrivs reconstructs the two fake privs from the committed
// golden-keys.json (the file produced by //go:build golden_generate).
type goldenKeyFile struct {
Pq string `json:"pq"`
Classical string `json:"classical"`
}
func loadGoldenPrivs(
t *testing.T,
path string,
) (
*fakePriv,
*fakePriv,
) {
t.Helper()
data := mustReadFile(t, path)
var keys goldenKeyFile
if err := json.Unmarshal(data, &keys); err != nil {
t.Fatalf("unmarshal golden keys: %v", err)
}
pqRaw := mustB64Decode(t, keys.Pq)
classicalRaw := mustB64Decode(t, keys.Classical)
if len(pqRaw) != fakeSeedLen {
t.Fatalf("pq raw len = %d, want %d", len(pqRaw), fakeSeedLen)
}
if len(classicalRaw) != fakeSeedLen {
t.Fatalf("classical raw len = %d, want %d", len(classicalRaw), fakeSeedLen)
}
return newFakePriv(fakePqSchemeID, pqRaw), newFakePriv(fakeClassicalSchemeID, classicalRaw)
}