diff --git a/STATUS.md b/STATUS.md index d12262f3..a8f44d25 100644 --- a/STATUS.md +++ b/STATUS.md @@ -536,8 +536,6 @@ Session layer: - A rejected request sends STOP_SENDING with INTERNAL_ERROR (§3.3.4 SHOULD use a relevant code). - Mandatory Track Property enforcement is off unless configured (§2.5.1). -- FETCH Serialization Flags ≥ 128 are read as field bits before being rejected, - so a reset or oversized length avoids the PROTOCOL_VIOLATION (§11.4.4). - SETUP options are sorted unstably, so with more than 12 the Token order on the wire can differ from the order `heldSetupAliases` replays (§10.3.1.4). diff --git a/pkg/moqt/message/fetch_object.go b/pkg/moqt/message/fetch_object.go index f8c73329..e33e57cc 100644 --- a/pkg/moqt/message/fetch_object.go +++ b/pkg/moqt/message/fetch_object.go @@ -154,6 +154,10 @@ func (o *FetchObject) Parse(r wire.Decoder) error { return err } o.SerializationFlags = flags + // Decided on the flags alone: their bits must not be read as fields. + if err := checkFetchFlags(flags); err != nil { + return err + } // End-of-range markers: Group ID and Object ID follow (§11.4.4.2). if isEndOfRange(flags) { @@ -267,20 +271,20 @@ func (o *FetchObject) SubgroupMode() FetchSubgroupIDMode { // Validate checks the fetch object for protocol violations. func (o *FetchObject) Validate() error { - flags := o.SerializationFlags - - // End-of-range markers are always valid structurally. - if isEndOfRange(flags) { - return nil - } - - // Values >= 128 that are not end-of-range markers are PROTOCOL_VIOLATION. // Note: 0x40 with non-zero subgroup-mode LSBs stays valid — the publisher // only SHOULD zero them and the subscriber MUST ignore them (§11.4.4.1), // so rejecting the combination would itself be non-conformant. - if flags >= 128 { - return fmt.Errorf("moqt/message: fetch object has invalid serialization flags 0x%X", flags) - } + return checkFetchFlags(o.SerializationFlags) +} + +// ErrInvalidFetchFlags is a Serialization Flags value of 128 or more that is +// not an End of Range marker: "Any other value is a PROTOCOL_VIOLATION" +// (§11.4.4). [FetchObject.Parse] returns it right after the flags. +var ErrInvalidFetchFlags = errors.New("moqt/message: invalid fetch object serialization flags") +func checkFetchFlags(flags uint64) error { + if flags >= 128 && !isEndOfRange(flags) { + return fmt.Errorf("%w 0x%X", ErrInvalidFetchFlags, flags) + } return nil } diff --git a/pkg/moqt/message/fetch_test.go b/pkg/moqt/message/fetch_test.go index 32193d9f..c3291e44 100644 --- a/pkg/moqt/message/fetch_test.go +++ b/pkg/moqt/message/fetch_test.go @@ -350,6 +350,20 @@ func TestFetchObjectIsEndOfRangeCoversAllThree(t *testing.T) { } } +// TestFetchObjectParseRejectsInvalidFlags: Parse rejects Serialization Flags +// of 128 or more that are not End of Range markers right after reading them +// (§11.4.4), without reading their bits as fields: 0xFF alone, with nothing +// after it, is ErrInvalidFetchFlags rather than a truncated object. +func TestFetchObjectParseRejectsInvalidFlags(t *testing.T) { + for _, flags := range []uint64{0x80, 0xFF, 0x8D, 0x100} { + var o FetchObject + err := o.Parse(wire.NewReader(wire.AppendVarint(nil, flags))) + if !errors.Is(err, ErrInvalidFetchFlags) { + t.Errorf("Parse(flags 0x%X) = %v, want ErrInvalidFetchFlags", flags, err) + } + } +} + func TestFetchObjectValidateInvalidFlags(t *testing.T) { // Values >= 128 that are not end-of-range markers are PROTOCOL_VIOLATION. obj := &FetchObject{ diff --git a/pkg/moqt/session/datastream_in.go b/pkg/moqt/session/datastream_in.go index b3444b50..60fe4461 100644 --- a/pkg/moqt/session/datastream_in.go +++ b/pkg/moqt/session/datastream_in.go @@ -285,13 +285,13 @@ func (s *IncomingFetchStream) Cancel(code moqt.StreamResetCode) { func (s *IncomingFetchStream) ReadObject() (*message.FetchObject, error) { obj := &message.FetchObject{} if err := obj.Parse(s.rd); err != nil { + // §11.4.4: a Serialization Flags value of 128 or more that is not + // an End of Range marker "is a PROTOCOL_VIOLATION". + if errors.Is(err, message.ErrInvalidFetchFlags) { + return nil, s.sess.closeProtocolViolation(fmt.Errorf("moqt/session: fetch object: %w", err)) + } return nil, s.sess.checkFINMidObject(err) } - // §11.4.4: a Serialization Flags value of 128 or more that is not an - // End of Range marker "is a PROTOCOL_VIOLATION". - if err := obj.Validate(); err != nil { - return nil, s.sess.closeProtocolViolation(fmt.Errorf("moqt/session: fetch object: %w", err)) - } return obj, nil } diff --git a/pkg/moqt/session/fetch_test.go b/pkg/moqt/session/fetch_test.go index d30bd813..5758674e 100644 --- a/pkg/moqt/session/fetch_test.go +++ b/pkg/moqt/session/fetch_test.go @@ -260,28 +260,53 @@ func TestFetchOKEndBeforeStartClosesSession(t *testing.T) { } // TestFetchObjectInvalidFlagsCloseSession: Serialization Flags of 128 and -// above other than the End of Range values are a PROTOCOL_VIOLATION (§11.4.4). +// above other than the End of Range values are a PROTOCOL_VIOLATION (§11.4.4), +// decided on the flags alone: whatever follows them, a whole object, a reset +// or a Payload Length too large to read, is never parsed as their fields. func TestFetchObjectInvalidFlagsCloseSession(t *testing.T) { t.Parallel() - client, server := openPair(t) - go func() { - out, err := server.OpenFetchStream(message.FetchHeader{RequestID: 0}) - if err != nil { - return - } - _ = out.WriteObject(&message.FetchObject{SerializationFlags: 0x81, ObjectPayload: []byte("x")}) - _ = out.Close() - }() - ds, err := client.AcceptDataStream(t.Context()) - if err != nil { - t.Fatalf("AcceptDataStream: %v", err) - } - fs, ok := ds.(*session.IncomingFetchStream) - if !ok { - t.Fatalf("AcceptDataStream = %T, want a FETCH stream", ds) - } - if _, err := fs.ReadObject(); err == nil { - t.Fatal("ReadObject accepted Serialization Flags 0x81") + // 0xFF sets every field bit, so read as flags it would parse a Group ID + // Delta, Subgroup ID, Object ID Delta, Priority and Properties first. + flags := func(v uint64) []byte { return wire.AppendVarint(nil, v) } + for _, tc := range []struct { + name string + write func(out *session.OutgoingFetchStream) + }{ + {"then a whole object", func(out *session.OutgoingFetchStream) { + _ = out.WriteObject(&message.FetchObject{SerializationFlags: 0x81, ObjectPayload: []byte("x")}) + _ = out.Close() + }}, + {"then a reset", func(out *session.OutgoingFetchStream) { + _, _ = out.Write(flags(0xFF)) + out.Cancel(moqt.StreamResetCancelled) + }}, + {"then an oversized Payload Length", func(out *session.OutgoingFetchStream) { + _, _ = out.Write(wire.AppendVarint(flags(0x81), 1<<62-1)) + _ = out.Close() + }}, + } { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + client, server := openPair(t) + go func() { + out, err := server.OpenFetchStream(message.FetchHeader{RequestID: 0}) + if err != nil { + return + } + tc.write(out) + }() + ds, err := client.AcceptDataStream(t.Context()) + if err != nil { + t.Fatalf("AcceptDataStream: %v", err) + } + fs, ok := ds.(*session.IncomingFetchStream) + if !ok { + t.Fatalf("AcceptDataStream = %T, want a FETCH stream", ds) + } + if _, err := fs.ReadObject(); err == nil { + t.Fatal("ReadObject accepted invalid Serialization Flags") + } + requireClosedProtocolViolation(t, client) + }) } - requireClosedProtocolViolation(t, client) }