diff --git a/mfer/deserialize.go b/mfer/deserialize.go index 6da64c6..60c93b0 100644 --- a/mfer/deserialize.go +++ b/mfer/deserialize.go @@ -278,6 +278,10 @@ func (m *manifest) deserializeInner() error { return fmt.Errorf("deserialize: unmarshal inner: %w", err) } + if m.pbInner.GetVersion() != MFFile_VERSION_ONE { + return errUnknownVersion + } + // Validate inner UUID err = validateUUID(m.pbInner.GetUuid()) if err != nil { diff --git a/mfer/deserialize_path_test.go b/mfer/deserialize_path_test.go index 8572901..aec2e92 100644 --- a/mfer/deserialize_path_test.go +++ b/mfer/deserialize_path_test.go @@ -151,7 +151,9 @@ func TestDeserializeRefusesEntriesThatDecodeTooLarge(t *testing.T) { entry = protowire.AppendBytes(entry, nil) id := uuid.NewV4() - inner := protowire.AppendTag(nil, 102, protowire.BytesType) // MFFile.uuid + inner := protowire.AppendTag(nil, 100, protowire.VarintType) // MFFile.version + inner = protowire.AppendVarint(inner, uint64(MFFile_VERSION_ONE)) + inner = protowire.AppendTag(inner, 102, protowire.BytesType) // MFFile.uuid inner = protowire.AppendBytes(inner, id[:]) for range 1000 { @@ -182,7 +184,9 @@ func TestDeserializeDropsUnknownFields(t *testing.T) { entry = append(entry, unknown...) id := uuid.NewV4() - inner := protowire.AppendTag(nil, 101, protowire.BytesType) // MFFile.files + inner := protowire.AppendTag(nil, 100, protowire.VarintType) // MFFile.version + inner = protowire.AppendVarint(inner, uint64(MFFile_VERSION_ONE)) + inner = protowire.AppendTag(inner, 101, protowire.BytesType) // MFFile.files inner = protowire.AppendBytes(inner, entry) inner = protowire.AppendTag(inner, 102, protowire.BytesType) // MFFile.uuid inner = protowire.AppendBytes(inner, id[:]) diff --git a/mfer/deserialize_test.go b/mfer/deserialize_test.go index 0a8f272..878a6a8 100644 --- a/mfer/deserialize_test.go +++ b/mfer/deserialize_test.go @@ -4,11 +4,32 @@ package mfer import ( "bytes" "testing" + "uuid" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + "google.golang.org/protobuf/proto" ) +// An inner message whose version is not VERSION_ONE, whether version 0 or a +// later one, is refused with the same error as an outer message's. +func TestDeserializeRefusesUnknownInnerVersion(t *testing.T) { + t.Parallel() + + for _, version := range []MFFile_Version{MFFile_VERSION_NONE, MFFile_VERSION_ONE + 1} { + t.Run(version.String(), func(t *testing.T) { + t.Parallel() + + id := uuid.NewV4() + inner, err := proto.Marshal(&MFFile{Version: version, Uuid: id[:]}) + require.NoError(t, err) + + _, err = NewManifestFromReader(bytes.NewReader(wrapInner(t, id, inner))) + require.ErrorIs(t, err, errUnknownVersion) + }) + } +} + // TestReadAtMost gives readAtMost exactly its maximum, which it must // return whole, and twice its maximum, which it must refuse after reading // one byte past the maximum, and no more. NewManifestFromReader reads