174b5f9048
Ultraworked with [Sisyphus](https://github.com/code-yeongyu/oh-my-openagent) Co-authored-by: Sisyphus <clio-agent@sisyphuslabs.ai>
998 lines
30 KiB
Go
998 lines
30 KiB
Go
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())
|
||
}
|
||
}
|