Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
38 changes: 38 additions & 0 deletions sdk/manifest.go
Original file line number Diff line number Diff line change
@@ -1,5 +1,13 @@
package sdk

import "fmt"

// Segment describes one chunk of the payload.
//
// Size and EncryptedSize are optional in the wire format.
// If absent, use the default sizes.
// Since our JSON parser doesn't distinguish an omitted key from an
// explicit 0, always check both (EncryptedSize is never 0).
type Segment struct {
Hash string `json:"hash"`
Size int64 `json:"segmentSize"`
Expand All @@ -19,6 +27,36 @@ type IntegrityInformation struct {
Segments []Segment `json:"segments"`
}

// resolveSegmentSizes returns the plaintext and ciphertext sizes of seg in bytes,
// substituting the manifest-level default for whichever field the writer
// omitted.
//
// EncryptedSize is never ambiguous on its own: ciphertext is never
// legitimately zero-length (there is always at least a nonce and a tag), so
// a raw 0 always means the key was left out because it equals
// DefaultEncryptedSegSize.
//
// Size is ambiguous on its own. For example, web-sdk decides emits Size and
// EncryptedSize only when they are not the default size (128 and 128+28 for AES-GCM-256).
// This determines the correct plaintext and ciphertext based on that understanding.
func (i IntegrityInformation) resolveSegmentSizes(seg Segment) (int64, int64, error) {
encryptedSize := seg.EncryptedSize
if encryptedSize == 0 {
encryptedSize = i.DefaultEncryptedSegSize
}

size := seg.Size
if size == 0 && encryptedSize == i.DefaultEncryptedSegSize {
size = i.DefaultSegmentSize
}

if size < 0 || encryptedSize <= 0 {
return 0, 0, fmt.Errorf("%w: segmentSize=%d encryptedSegmentSize=%d", ErrSegSizeUnresolved, size, encryptedSize)
}

return size, encryptedSize, nil
}

type KeyAccess struct {
KeyType string `json:"type"`
KasURL string `json:"url"`
Expand Down
49 changes: 33 additions & 16 deletions sdk/tdf.go
Original file line number Diff line number Diff line change
Expand Up @@ -907,7 +907,14 @@ func (s SDK) LoadTDF(reader io.ReadSeeker, opts ...TDFReaderOption) (*Reader, er

var payloadSize int64
for _, seg := range manifestObj.Segments {
payloadSize += seg.Size
// Sizes the writer left to the manifest-level default have to be
// filled in here too: without it the payload looks shorter than it
// is, and every read bounded by payloadSize comes up short.
size, _, err := manifestObj.resolveSegmentSizes(seg)
if err != nil {
return nil, err
}
payloadSize += size
}

return &Reader{
Expand Down Expand Up @@ -984,18 +991,23 @@ func (r *Reader) WriteTo(writer io.Writer) (int64, error) {
var payloadReadOffset int64
var decryptedDataOffset int64
for _, seg := range r.manifest.Segments {
if decryptedDataOffset+seg.Size < r.cursor {
decryptedDataOffset += seg.Size
payloadReadOffset += seg.EncryptedSize
segSize, encryptedSegSize, err := r.manifest.resolveSegmentSizes(seg)
if err != nil {
return totalBytes, err
}

if decryptedDataOffset+segSize < r.cursor {
decryptedDataOffset += segSize
payloadReadOffset += encryptedSegSize
continue
}

readBuf, err := r.tdfReader.ReadPayload(payloadReadOffset, seg.EncryptedSize)
readBuf, err := r.tdfReader.ReadPayload(payloadReadOffset, encryptedSegSize)
if err != nil {
return totalBytes, fmt.Errorf("TDFReader.ReadPayload failed: %w", err)
}

if int64(len(readBuf)) != seg.EncryptedSize {
if int64(len(readBuf)) != encryptedSegSize {
return totalBytes, ErrSegSizeMismatch
}

Expand Down Expand Up @@ -1034,9 +1046,9 @@ func (r *Reader) WriteTo(writer io.Writer) (int64, error) {
return totalBytes, errWriteFailed
}

payloadReadOffset += seg.EncryptedSize
payloadReadOffset += encryptedSegSize
r.cursor += int64(n)
decryptedDataOffset += seg.Size
decryptedDataOffset += segSize
}

return totalBytes, nil
Expand Down Expand Up @@ -1084,6 +1096,11 @@ func (r *Reader) ReadAt(buf []byte, offset int64) (int, error) { //nolint:funlen
var segStart int64 // plaintext offset of seg
startIndex := int64(-1) // offset of the request within decryptedBuf
for _, seg := range r.manifest.Segments {
segSize, encryptedSegSize, err := r.manifest.resolveSegmentSizes(seg)
if err != nil {
return 0, err
}

// Segment.Size positions every plaintext offset derived below --
// including for the segments this request skips over -- but nothing
// authenticates it: the root signature aggregates only Segment.Hash.
Expand All @@ -1093,18 +1110,18 @@ func (r *Reader) ReadAt(buf []byte, offset int64) (int, error) { //nolint:funlen
// is the per-segment form of the check doPayloadKeyUnwrap already
// applies to the manifest defaults. Deriving Size from EncryptedSize
// rather than the reverse keeps the arithmetic from overflowing.
if seg.EncryptedSize < gcmIvSize+aesBlockSize || seg.Size != seg.EncryptedSize-(gcmIvSize+aesBlockSize) {
if encryptedSegSize < gcmIvSize+aesBlockSize || segSize != encryptedSegSize-(gcmIvSize+aesBlockSize) {
return 0, fmt.Errorf("%w: segment declares size %d with encrypted size %d",
ErrSegSizeMismatch, seg.Size, seg.EncryptedSize)
ErrSegSizeMismatch, segSize, encryptedSegSize)
}

segEnd := segStart + seg.Size
segEnd := segStart + segSize

// Wholly before the request. The comparison is <= rather than < so
// that a request starting exactly on a segment boundary, or a
// zero-length request, does not pull in the preceding segment.
if segEnd <= offset {
payloadReadOffset += seg.EncryptedSize
payloadReadOffset += encryptedSegSize
segStart = segEnd
continue
}
Expand All @@ -1118,12 +1135,12 @@ func (r *Reader) ReadAt(buf []byte, offset int64) (int, error) { //nolint:funlen
startIndex = offset - segStart
}

readBuf, err := r.tdfReader.ReadPayload(payloadReadOffset, seg.EncryptedSize)
readBuf, err := r.tdfReader.ReadPayload(payloadReadOffset, encryptedSegSize)
if err != nil {
return 0, fmt.Errorf("TDFReader.ReadPayload failed: %w", err)
}

if int64(len(readBuf)) != seg.EncryptedSize {
if int64(len(readBuf)) != encryptedSegSize {
return 0, ErrSegSizeMismatch
}

Expand Down Expand Up @@ -1156,7 +1173,7 @@ func (r *Reader) ReadAt(buf []byte, offset int64) (int, error) { //nolint:funlen
return 0, errWriteFailed
}

payloadReadOffset += seg.EncryptedSize
payloadReadOffset += encryptedSegSize
segStart = segEnd
}

Expand Down Expand Up @@ -1581,7 +1598,7 @@ func calculateSignature(data []byte, secret []byte, alg IntegrityAlgorithm, isLe
return string(hmac), nil
}
if kGMACPayloadLength > len(data) {
return "", errors.New("fail to create gmac signature")
return "", fmt.Errorf("%w: ciphertext length=%d", ErrGMACSignatureFailed, len(data))
}

if isLegacyTDF {
Expand Down
17 changes: 10 additions & 7 deletions sdk/tdf_readat_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -227,21 +227,24 @@ func TestReaderReadAtNonUniformEdges(t *testing.T) {
// its own to catch it.
func TestReaderReadAtDeclaredSizeMismatch(t *testing.T) {
for _, tc := range []struct {
name string
mutate func(segments []Segment)
name string
mutate func(segments []Segment)
wantErr error
}{
// Understating the first segment shifts every later segment down by
// five bytes. The read below starts past that segment, so it is skipped
// and never decrypted.
{"understated", func(segments []Segment) { segments[0].Size = 5 }},
{"overstated", func(segments []Segment) { segments[0].Size = 40 }},
{"understated", func(segments []Segment) { segments[0].Size = 5 }, ErrSegSizeMismatch},
{"overstated", func(segments []Segment) { segments[0].Size = 40 }, ErrSegSizeMismatch},
// Sizes that sum back to something plausible: payloadSize is the sum of
// every Size, so a pair that overflows to a small positive total gets
// past the range check on offset and reaches the segment walk.
// past the range check on offset and reaches the segment walk. A
// negative declared Size is caught by resolveSegmentSizes itself,
// before the arithmetic consistency check below it ever runs.
{"negative", func(segments []Segment) {
segments[0].Size = math.MinInt64 + 1
segments[1].Size = math.MinInt64 + 7
}},
}, ErrSegSizeUnresolved},
} {
t.Run(tc.name, func(t *testing.T) {
reader, _ := newNonUniformReader(t, []int{10, 10, 10})
Expand All @@ -258,7 +261,7 @@ func TestReaderReadAtDeclaredSizeMismatch(t *testing.T) {
// reader that trusted Size would report a full 20 bytes of shifted
// plaintext rather than an error.
n, err := reader.ReadAt(make([]byte, 20), 5)
require.ErrorIs(t, err, ErrSegSizeMismatch)
require.ErrorIs(t, err, tc.wantErr)
assert.Zero(t, n)
})
}
Expand Down
Loading
Loading