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" ) func makeRecipientPubs(pqPub, classicalPub crypto.RecipientPub) []crypto.RecipientPub { pubs := make([]crypto.RecipientPub, 0, 2) return append(pubs, pqPub, classicalPub) } func makeRecipientPrivs(pqPriv, classicalPriv crypto.RecipientPriv) []crypto.RecipientPriv { privs := make([]crypto.RecipientPriv, 0, 2) return append(privs, pqPriv, classicalPriv) } // --------------------------------------------------------------------------- // 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), makeRecipientPrivs(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 index := 0; index < size; index++ { plaintext[index] = byte(index) } 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), makeRecipientPubs(pqPub, classicalPub), &encrypted, rand.Reader, ); err != nil { t.Fatalf("Encrypt: %v", err) } out := &bytes.Buffer{} if err := dec.Decrypt( bytes.NewReader(encrypted.Bytes()), makeRecipientPrivs(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), makeRecipientPubs(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), makeRecipientPubs(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), makeRecipientPubs(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), makeRecipientPrivs(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}), makeRecipientPubs(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), makeRecipientPrivs(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}), makeRecipientPubs(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}), makeRecipientPubs(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, makeRecipientPrivs(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, makeRecipientPrivs(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}), makeRecipientPubs(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()), makeRecipientPrivs(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}), makeRecipientPubs(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()), makeRecipientPrivs(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), makeRecipientPubs(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]), makeRecipientPrivs(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) } func TestLimitedWriter_AllowsWritesWithinBudget(t *testing.T) { var buf bytes.Buffer lw := &limitedWriter{writer: &buf, remaining: 10} n, err := lw.Write([]byte("hello")) if err != nil { t.Fatalf("unexpected error: %v", err) } if n != 5 { t.Errorf("wrote %d bytes, want 5", n) } n, err = lw.Write([]byte("world")) if err != nil { t.Fatalf("unexpected error: %v", err) } if n != 5 { t.Errorf("wrote %d bytes, want 5", n) } if buf.String() != "helloworld" { t.Errorf("buffer = %q, want %q", buf.String(), "helloworld") } } func TestLimitedWriter_RejectsOverBudget(t *testing.T) { var buf bytes.Buffer lw := &limitedWriter{writer: &buf, remaining: 3} _, err := lw.Write([]byte("hello")) if !errors.Is(err, ErrPlaintextTooLarge) { t.Errorf("error = %v, want ErrPlaintextTooLarge", err) } if buf.Len() != 0 { t.Errorf("buffer len = %d, want 0", buf.Len()) } }