diff --git a/STATUS.md b/STATUS.md index a8f44d25..3938b8ca 100644 --- a/STATUS.md +++ b/STATUS.md @@ -196,7 +196,7 @@ By package, bottom-up along the dependency stack: | 10.9 | REQUEST_UPDATE | 0x02 | DONE | A REQUEST_UPDATE opening a request stream closes the session with PROTOCOL_VIOLATION (`ErrUnexpectedRequestUpdate`). | | 10.10 | PUBLISH_STATE_NOTIFY | 0x22 | DONE | Only the publisher may send it; enforced by brokers and the relay. | | 10.11 | PUBLISH | 0x1D | DONE | | -| 10.12 | PUBLISH_DONE | 0x0B | DONE | Sent once every stream of the subscription has closed and no datagram send is in progress, with the exact Stream Count; written on its own goroutine, so subscribers do not wait on each other. When a track's last upstream ends, its PUBLISH_DONE code reaches subscribers if it is about the track (TRACK_ENDED, MALFORMED_TRACK); codes about the relay's own upstream subscription become INTERNAL_ERROR. | +| 10.12 | PUBLISH_DONE | 0x0B | DONE | Sent once every stream of the subscription has closed and no datagram send is in progress, with the exact Stream Count; written on its own goroutine, so subscribers do not wait on each other. When a track's last upstream ends, its PUBLISH_DONE code reaches subscribers if it is about the track (TRACK_ENDED, MALFORMED_TRACK); codes about the relay's own upstream subscription become INTERNAL_ERROR. Session `Publication.Done` resets the subgroups still open with CANCELLED, refuses later opens and writes (`ErrPublicationEnded`), and counts every subgroup opened, however the opens race it. | | 10.13 | FETCH | 0x16 | DONE | Standalone, the only kind in draft-20. From the cache, a Location is non-existent only on a signal: a Prior Group or Object ID Gap, a Group's or the Track's end, or an upstream's FETCH. Other uncached Locations are FETCHed from a fetch-capable upstream in one span, within FILL_TIMEOUT, or else marked End of Unknown (or Timed-Out) Range. | | 10.14 | FETCH_OK | 0x18 | DONE | An End Location before the FETCH's Start closes the session. A Start relative to the Largest Object is compared through End ≤ Largest; an End of {0,0} is let through, as it cannot be told apart from "no content yet". | | 10.15 | TRACK_STATUS | 0x0D | DONE | Reply via REQUEST_OK, then FIN; any follow-up from the requester closes the session. | @@ -527,17 +527,6 @@ Session layer: SUBSCRIBE_NAMESPACE or SUBSCRIBE_TRACKS, where §10.19/§10.20 make it a PROTOCOL_VIOLATION) is an error rather than a legal message checked against §10.4. -- `Publication`'s automatic PUBLISH_DONE UPDATE_FAILED is sent while its subgroup - streams are open, `WriteObject` still succeeds after `Done`, and a subgroup - opened concurrently with `Done` is missing from the Stream Count (§10.12). -- `Publication`'s REQUEST_UPDATE_OK carries LARGEST_OBJECT only for Objects it - wrote itself, not the one its SUBSCRIBE_OK or PUBLISH reported (§10.2.17, - §10.9.1). -- 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). -- 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). Relay: diff --git a/pkg/moqt/session/datastream_out.go b/pkg/moqt/session/datastream_out.go index 2547ada6..d38f5e49 100644 --- a/pkg/moqt/session/datastream_out.go +++ b/pkg/moqt/session/datastream_out.go @@ -86,9 +86,13 @@ type OutgoingSubgroupStream struct { encHavePrev bool // Set by [Publication.OpenSubgroup]: onObject is told each written - // object's Location, and paused reports a Forward State of 0 (§11.4.3). + // object's Location, paused reports a Forward State of 0 (§11.4.3), ended + // reports that the publication ended (Done), and onEnd is told when the + // stream is FINished or reset. onObject func(group, object uint64) paused func() bool + ended func() bool + onEnd func() } // WithDeliveryTimeouts returns a shallow copy of s configured with the §8 @@ -168,9 +172,13 @@ func (s *OutgoingSubgroupStream) WriteObjectReceivedAt( } s.sawFirstObject = true + // §10.12: Done reset the stream before PUBLISH_DONE. + if s.ended != nil && s.ended() { + return ErrPublicationEnded + } // §5.1: no Objects while the Forward State is 0; §11.4.3: reset. if s.paused != nil && s.paused() { - s.dst.CancelWrite(uint64(moqt.StreamResetCancelled)) + s.Cancel(moqt.StreamResetCancelled) return ErrForwardPaused } if err := s.checkObjectTimeout(receivedAt); err != nil { @@ -247,7 +255,7 @@ func (s *OutgoingSubgroupStream) checkObjectTimeout(receivedAt time.Time) error } elapsed := time.Since(receivedAt) if elapsed > s.objectTimeout { - s.dst.CancelWrite(uint64(moqt.StreamResetDeliveryTimeout)) + s.Cancel(moqt.StreamResetDeliveryTimeout) return fmt.Errorf("%w (elapsed %s, limit %s)", ErrDeliveryTimeout, elapsed, s.objectTimeout) } @@ -262,6 +270,9 @@ func (s *OutgoingSubgroupStream) checkObjectTimeout(receivedAt time.Time) error // stream if the peer has not acknowledged all data within the timeout (§8). func (s *OutgoingSubgroupStream) Close() error { err := s.dst.Close() + if s.onEnd != nil { + s.onEnd() + } tracked, ok := s.dst.(DeliveryTrackingSendStream) if s.subgroupTimeout > 0 && ok { finished := tracked.Finished() @@ -285,6 +296,9 @@ func (s *OutgoingSubgroupStream) Close() error { // Cancel resets the stream with the given application code (§3.3.4). func (s *OutgoingSubgroupStream) Cancel(code moqt.StreamResetCode) { s.dst.CancelWrite(uint64(code)) + if s.onEnd != nil { + s.onEnd() + } } // SetSendPriority forwards the composite §7.2 scheduling key to the underlying @@ -367,25 +381,45 @@ func (s *Session) OpenSubgroupContext( ctx context.Context, h message.SubgroupHeader, ) (*OutgoingSubgroupStream, error) { - dst, err := s.conn.OpenUniStream() + sg, reset, err := s.openSubgroup(ctx, h, false) if err != nil { return nil, err } + if reset { + // ctx was cancelled just after the header write went through, and + // the stream is reset. + return nil, fmt.Errorf("moqt/session: write SUBGROUP_HEADER: %w", ctx.Err()) + } + return sg, nil +} + +// openSubgroup opens a subgroup stream and writes its header, which +// cancelling ctx interrupts by resetting the stream, first marking what was +// written reliable when markReliable is set (§11.4.3). It returns the stream +// whenever the whole header was written, since the peer can then attribute +// the stream to its track, with reset reporting that ctx reset it just after. +func (s *Session) openSubgroup( + ctx context.Context, + h message.SubgroupHeader, + markReliable bool, +) (sg *OutgoingSubgroupStream, reset bool, err error) { + dst, err := s.conn.OpenUniStream() + if err != nil { + return nil, false, err + } stop := context.AfterFunc(ctx, func() { + if r, ok := dst.(ReliableResetStream); ok && markReliable { + r.SetReliableBoundary() + } dst.CancelWrite(uint64(moqt.StreamResetCancelled)) }) if err := message.WriteSubgroupHeader(dst, h); err != nil { stop() dst.CancelWrite(uint64(moqt.StreamResetInternalError)) if ctx.Err() != nil { - return nil, fmt.Errorf("moqt/session: write SUBGROUP_HEADER: %w", ctx.Err()) + return nil, false, fmt.Errorf("moqt/session: write SUBGROUP_HEADER: %w", ctx.Err()) } - return nil, fmt.Errorf("moqt/session: write SUBGROUP_HEADER: %w", err) - } - if !stop() { - // The AfterFunc already ran: ctx was cancelled while (or just - // after) the header write went through — the stream is reset. - return nil, fmt.Errorf("moqt/session: write SUBGROUP_HEADER: %w", ctx.Err()) + return nil, false, fmt.Errorf("moqt/session: write SUBGROUP_HEADER: %w", err) } - return &OutgoingSubgroupStream{header: h, dst: dst}, nil + return &OutgoingSubgroupStream{header: h, dst: dst}, !stop(), nil } diff --git a/pkg/moqt/session/export_test.go b/pkg/moqt/session/export_test.go index 198e1411..2ae2cf7b 100644 --- a/pkg/moqt/session/export_test.go +++ b/pkg/moqt/session/export_test.go @@ -37,3 +37,11 @@ func OpenRequestForTest(s *Session, first message.Message) (Stream, error) { func WithSetupOptionForTest(kv wire.KVPair) Option { return func(c *config) { c.setupOptions = append(c.setupOptions, kv) } } + +// OpenSubgroupsForTest reports how many subgroups p tracks as open, which Done +// would reset. +func (p *Publication) OpenSubgroupsForTest() int { + p.subMu.Lock() + defer p.subMu.Unlock() + return len(p.open) +} diff --git a/pkg/moqt/session/options.go b/pkg/moqt/session/options.go index 5edb3e92..150f9776 100644 --- a/pkg/moqt/session/options.go +++ b/pkg/moqt/session/options.go @@ -166,13 +166,10 @@ func WithGrease() Option { // set, it returns *ErrUnsupportedMandatoryTrackProperty; [Request.AcceptPublish] // refuses such a PUBLISH with UNSUPPORTED_EXTENSION. // -// If this option is never called, enforcement is disabled and all properties -// pass through. Leave it unset only when the application checks the -// properties itself (§2.5.1). -// -// End subscribers that interpret track data should call this option to opt -// in to enforcement. Pass an empty (non-nil) map to reject all mandatory -// properties, or populate the map with the types you support. +// If this option is never called, or types is empty or nil, no Mandatory Track +// Property is known and every one is refused: an endpoint that does not +// understand one "MUST NOT process or forward that track" (§2.5.1). List the +// types this endpoint understands to accept them. func WithKnownMandatoryTrackProperties(types map[message.PropertyType]struct{}) Option { return func(c *config) { c.knownMandatoryTrackProperties = types diff --git a/pkg/moqt/session/publication_done_test.go b/pkg/moqt/session/publication_done_test.go new file mode 100644 index 00000000..2c49d821 --- /dev/null +++ b/pkg/moqt/session/publication_done_test.go @@ -0,0 +1,311 @@ +package session_test + +import ( + "context" + "errors" + "io" + "maps" + "slices" + "sync" + "testing" + "time" + + "github.com/floatdrop/moq-go/pkg/moqt" + "github.com/floatdrop/moq-go/pkg/moqt/message" + "github.com/floatdrop/moq-go/pkg/moqt/session" + "github.com/floatdrop/moq-go/pkg/moqt/session/sessiontest" +) + +// drainDataStreams accepts sess's data streams for a second and reads each to +// its end, reporting per Group ID whether the stream ended cleanly (FIN) or +// not. A stream that has not ended a second later is missing from the report. +func drainDataStreams(t *testing.T, sess *session.Session) <-chan map[uint64]bool { + t.Helper() + out := make(chan map[uint64]bool, 1) + go func() { + var ( + mu sync.Mutex + clean = map[uint64]bool{} + wg sync.WaitGroup + ) + ctx, cancel := context.WithTimeout(t.Context(), time.Second) + defer cancel() + for { + ds, err := sess.AcceptDataStream(ctx) + if err != nil { + break + } + sg, ok := ds.(*session.IncomingSubgroupStream) + if !ok { + continue + } + wg.Go(func() { + for { + if _, err := sg.ReadDecoded(); err != nil { + mu.Lock() + clean[sg.Header.GroupID] = errors.Is(err, io.EOF) + mu.Unlock() + return + } + } + }) + } + finished := make(chan struct{}) + go func() { wg.Wait(); close(finished) }() + select { + case <-finished: + case <-time.After(time.Second): + } + mu.Lock() + out <- maps.Clone(clean) + mu.Unlock() + }() + return out +} + +// publishDoneOn serves sub's broker and returns the PUBLISH_DONE it reads. +func publishDoneOn(t *testing.T, sub *session.Subscription) <-chan *message.PublishDone { + t.Helper() + got := make(chan *message.PublishDone, 1) + b := sub.Broker() + go func() { + _ = b.Serve(t.Context(), func(m message.Message) bool { + if d, ok := m.(*message.PublishDone); ok { + got <- d + return false + } + return true + }) + }() + return got +} + +func subgroupHeader(group uint64) message.SubgroupHeader { + return message.SubgroupHeader{SubgroupIDMode: message.SubgroupIDImplicitZero, GroupID: group} +} + +// TestPublicationDoneResetsOpenSubgroups: Done resets the subgroups still open +// with CANCELLED, since PUBLISH_DONE MUST NOT be sent "until it has closed all +// streams it will ever open" (§10.12), leaves a FINished one alone, counts both +// in the Stream Count, and refuses further writes. +func TestPublicationDoneResetsOpenSubgroups(t *testing.T) { + cli, srv := openPair(t) + sub, pub := subscribePair(t, cli, srv) + done := publishDoneOn(t, sub) + streams := drainDataStreams(t, cli) + + open, err := pub.OpenSubgroup(subgroupHeader(1)) + if err != nil { + t.Fatalf("OpenSubgroup: %v", err) + } + if err := open.WriteObject(&message.SubgroupObject{Payload: []byte("a")}); err != nil { + t.Fatalf("WriteObject: %v", err) + } + finished, err := pub.OpenSubgroup(subgroupHeader(2)) + if err != nil { + t.Fatalf("OpenSubgroup: %v", err) + } + if err := finished.WriteObject(&message.SubgroupObject{Payload: []byte("b")}); err != nil { + t.Fatalf("WriteObject: %v", err) + } + if err := finished.Close(); err != nil { + t.Fatalf("Close: %v", err) + } + + if err := pub.Done(moqt.PublishDoneTrackEnded, ""); err != nil { + t.Fatalf("Done: %v", err) + } + if err := open.WriteObject( + &message.SubgroupObject{Payload: []byte("c")}, + ); !errors.Is( + err, + session.ErrPublicationEnded, + ) { + t.Errorf("WriteObject after Done = %v, want ErrPublicationEnded", err) + } + select { + case d := <-done: + if d.StreamCount != 2 { + t.Errorf("PUBLISH_DONE Stream Count = %d, want 2", d.StreamCount) + } + case <-time.After(2 * time.Second): + t.Fatal("no PUBLISH_DONE") + } + clean := <-streams + if ok, seen := clean[1]; !seen || ok { + t.Errorf("open subgroup: seen %v, ended cleanly %v; want reset", seen, ok) + } + if ok, seen := clean[2]; !seen || !ok { + t.Errorf("finished subgroup: seen %v, ended cleanly %v; want its FIN", seen, ok) + } +} + +// TestPublicationDoneKeepsFINOfDeliveryTimeoutCopy: a subgroup FINished +// through the copy WithDeliveryTimeouts returns is as finished as one closed +// directly, so Done does not reset it: a publisher that delivered every Object +// "MUST close the stream with a FIN" (§11.4.3). +func TestPublicationDoneKeepsFINOfDeliveryTimeoutCopy(t *testing.T) { + cli, srv := openPair(t) + sub, pub := subscribePair(t, cli, srv) + done := publishDoneOn(t, sub) + streams := drainDataStreams(t, cli) + + sg, err := pub.OpenSubgroup(subgroupHeader(1)) + if err != nil { + t.Fatalf("OpenSubgroup: %v", err) + } + timed := sg.WithDeliveryTimeouts(message.DeliveryTimeouts{}, message.DeliveryTimeouts{}) + if err := timed.WriteObject(&message.SubgroupObject{Payload: []byte("a")}); err != nil { + t.Fatalf("WriteObject: %v", err) + } + if err := timed.Close(); err != nil { + t.Fatalf("Close: %v", err) + } + if n := pub.OpenSubgroupsForTest(); n != 0 { + t.Errorf("%d subgroups still tracked as open after their FIN; Done would reset them", n) + } + if ok, seen := (<-streams)[1]; !seen || !ok { + t.Fatalf("subgroup: seen %v, ended cleanly %v; want its FIN", seen, ok) + } + if err := pub.Done(moqt.PublishDoneTrackEnded, ""); err != nil { + t.Fatalf("Done: %v", err) + } + select { + case d := <-done: + if d.StreamCount != 1 { + t.Errorf("PUBLISH_DONE Stream Count = %d, want 1", d.StreamCount) + } + case <-time.After(2 * time.Second): + t.Fatal("no PUBLISH_DONE") + } +} + +// TestPublicationDoneStreamCountExact: however OpenSubgroup races Done, the +// Stream Count is "the number of data streams the publisher opened" (§10.12): +// every OpenSubgroup that returned a stream. The subscriber never sees more; +// it can see fewer, when Done's reset loses a header it already wrote, which +// §10.12 has subscribers allow for. +func TestPublicationDoneStreamCountExact(t *testing.T) { + cli, srv := openPair(t) + sub, pub := subscribePair(t, cli, srv) + done := publishDoneOn(t, sub) + streams := drainDataStreams(t, cli) + + var ( + mu sync.Mutex + opened uint64 + wg sync.WaitGroup + ) + start := make(chan struct{}) + for w := range uint64(8) { + wg.Go(func() { + <-start + for i := uint64(0); ; i++ { + sg, err := pub.OpenSubgroup(subgroupHeader(w*1000 + i)) + if errors.Is(err, session.ErrPublicationEnded) { + return + } + if errors.Is(err, session.ErrNoStreamCredit) { + continue // the subscriber has not read enough streams yet + } + if err != nil { + t.Errorf("OpenSubgroup: %v", err) + return + } + mu.Lock() + opened++ + mu.Unlock() + _ = sg.Close() + } + }) + } + close(start) + time.Sleep(5 * time.Millisecond) + if err := pub.Done(moqt.PublishDoneTrackEnded, ""); err != nil { + t.Fatalf("Done: %v", err) + } + wg.Wait() + select { + case d := <-done: + if d.StreamCount != opened { + t.Errorf("PUBLISH_DONE Stream Count = %d, but %d subgroups opened", d.StreamCount, opened) + } + case <-time.After(2 * time.Second): + t.Fatal("no PUBLISH_DONE") + } + if got := uint64(len(<-streams)); got > opened { + t.Errorf("subscriber saw %d subgroup streams, more than the %d opened", got, opened) + } +} + +// reliableRecConn hands out uni streams that implement +// [session.ReliableResetStream] and record, in order, each reliable boundary +// and reset they see. +type reliableRecConn struct { + session.Conn + + mu sync.Mutex + events []string +} + +func (c *reliableRecConn) OpenUniStream() (session.SendStream, error) { + s, err := c.Conn.OpenUniStream() + if err != nil { + return nil, err + } + return &reliableRecStream{SendStream: s, c: c}, nil +} + +func (c *reliableRecConn) record(e string) { + c.mu.Lock() + c.events = append(c.events, e) + c.mu.Unlock() +} + +type reliableRecStream struct { + session.SendStream + + c *reliableRecConn +} + +func (s *reliableRecStream) SetReliableBoundary() { s.c.record("reliable") } + +func (s *reliableRecStream) CancelWrite(code uint64) { + s.c.record("reset") + s.SendStream.CancelWrite(code) +} + +// TestPublicationDoneResetKeepsHeaderReliable: before Done resets an open +// subgroup it marks what was written as reliable, so over RESET_STREAM_AT the +// header still reaches the subscriber, which can then "accurately account for +// reset data streams when handling PUBLISH_DONE" (§11.4.3). +func TestPublicationDoneResetKeepsHeaderReliable(t *testing.T) { + cliConn, rawSrv := sessiontest.NewConnPair() + srvConn := &reliableRecConn{Conn: rawSrv} + cli, srv := openSessions(t, cliConn, srvConn, nil, nil) + sub, pub := subscribePair(t, cli, srv) + done := publishDoneOn(t, sub) + drainDataStreams(t, cli) + + sg, err := pub.OpenSubgroup(subgroupHeader(1)) + if err != nil { + t.Fatalf("OpenSubgroup: %v", err) + } + if err := sg.WriteObject(&message.SubgroupObject{Payload: []byte("a")}); err != nil { + t.Fatalf("WriteObject: %v", err) + } + if err := pub.Done(moqt.PublishDoneTrackEnded, ""); err != nil { + t.Fatalf("Done: %v", err) + } + select { + case <-done: + case <-time.After(2 * time.Second): + t.Fatal("no PUBLISH_DONE") + } + srvConn.mu.Lock() + events := slices.Clone(srvConn.events) + srvConn.mu.Unlock() + if want := []string{"reliable", "reset"}; !slices.Equal(events, want) { + t.Fatalf("subgroup stream saw %v, want %v", events, want) + } +} diff --git a/pkg/moqt/session/publication_largest_test.go b/pkg/moqt/session/publication_largest_test.go new file mode 100644 index 00000000..c7f38e1d --- /dev/null +++ b/pkg/moqt/session/publication_largest_test.go @@ -0,0 +1,91 @@ +package session_test + +import ( + "testing" + + "github.com/floatdrop/moq-go/pkg/moqt/message" + "github.com/floatdrop/moq-go/pkg/moqt/session" +) + +// TestPublicationUpdateOKCarriesAnnouncedLargest: a Publication's +// REQUEST_UPDATE_OK reports LARGEST_OBJECT from what its SUBSCRIBE_OK or +// PUBLISH already reported, not only from Objects it wrote itself: "If Objects +// have been published on this Track the Publisher MUST include this parameter" +// (§10.2.17). A larger Object written later wins; a smaller one does not. +func TestPublicationUpdateOKCarriesAnnouncedLargest(t *testing.T) { + check := func(t *testing.T, ok *message.RequestOK, err error, group, object uint64) { + t.Helper() + must(t, err) + p, found := ok.Parameters.Find(message.ParamLargestObject) + if !found || p.Group != group || p.Object != object { + t.Fatalf("REQUEST_UPDATE_OK LARGEST_OBJECT = %+v (found %v), want {%d, %d}", p, found, group, object) + } + } + + t.Run("SUBSCRIBE_OK", func(t *testing.T) { + client, server := openPair(t) + pubs := make(chan *session.Publication, 1) + go func() { + r, err := server.AcceptRequest(t.Context()) + if err != nil { + return + } + p, err := r.AcceptSubscribe(&message.SubscribeOK{ + TrackAlias: 1, + Parameters: message.Parameters{message.LargestObjectParam(5, 9)}, + }) + if err == nil { + pubs <- p + } + }() + sub, err := client.Subscribe(t.Context(), &message.Subscribe{Name: []byte("t")}) + must(t, err) + pub := <-pubs + b := pub.Broker() + go func() { _ = b.Serve(t.Context(), nil) }() + + ok, err := sub.Update(t.Context(), message.Parameters{message.ForwardParam(true)}) + check(t, ok, err, 5, 9) + + go drainOneSubgroup(t, client) + sg, err := pub.OpenSubgroup(message.SubgroupHeader{GroupID: 6, SubgroupIDMode: message.SubgroupIDExplicit}) + must(t, err) + must(t, sg.WriteObjectAt(0, &message.SubgroupObject{Payload: []byte("a")})) + must(t, sg.Close()) + ok, err = sub.Update(t.Context(), message.Parameters{message.ForwardParam(true)}) + check(t, ok, err, 6, 0) + }) + + t.Run("PUBLISH", func(t *testing.T) { + client, server := openPair(t) + pubs := make(chan *session.Publication, 1) + go func() { + p, err := server.Publish(t.Context(), &message.Publish{ + Name: []byte("t"), TrackAlias: 7, + Parameters: message.Parameters{message.LargestObjectParam(3, 1)}, + }) + if err == nil { + pubs <- p + } + }() + r, err := client.AcceptRequest(t.Context()) + must(t, err) + in, err := r.AcceptPublish() + must(t, err) + pub := <-pubs + b := pub.Broker() + go func() { _ = b.Serve(t.Context(), nil) }() + + ok, err := in.Update(t.Context(), message.Parameters{message.ForwardParam(true)}) + check(t, ok, err, 3, 1) + + // An Object written below the announced one leaves it the largest. + go drainOneSubgroup(t, client) + sg, err := pub.OpenSubgroup(message.SubgroupHeader{GroupID: 2, SubgroupIDMode: message.SubgroupIDExplicit}) + must(t, err) + must(t, sg.WriteObjectAt(5, &message.SubgroupObject{Payload: []byte("a")})) + must(t, sg.Close()) + ok, err = in.Update(t.Context(), message.Parameters{message.ForwardParam(true)}) + check(t, ok, err, 3, 1) + }) +} diff --git a/pkg/moqt/session/publish.go b/pkg/moqt/session/publish.go index fe464ee4..017ff27b 100644 --- a/pkg/moqt/session/publish.go +++ b/pkg/moqt/session/publish.go @@ -31,21 +31,32 @@ type Publication struct { alias uint64 - // subgroupCount counts subgroup streams opened via OpenSubgroup, used as - // the §10.12 Stream Count when Done sends PUBLISH_DONE. - subgroupCount atomic.Uint64 + // The subgroups opened via OpenSubgroup, for Done (§10.12). subMu + // orders each OpenSubgroup against Done, which waits for the opens in + // flight (opening) and cancels their header writes (endCtx), so every + // subgroup opened is counted in subgroupCount, and the ones still open + // (open) are reset before PUBLISH_DONE. A subgroup's onEnd closes over + // the pointer registered here, which its copies (WithDeliveryTimeouts) + // share. + subMu sync.Mutex + subgroupCount uint64 + open map[*OutgoingSubgroupStream]struct{} + opening sync.WaitGroup + endCtx context.Context + endCancel context.CancelFunc // paused is the inverse of the §5.1 Forward State. paused atomic.Bool - // largest is the largest Location written through OpenSubgroup streams, - // reported as LARGEST_OBJECT in REQUEST_UPDATE_OK (§10.9.1). + // largest is the largest Location this side announced in SUBSCRIBE_OK + // or PUBLISH, or wrote through OpenSubgroup streams since, reported as + // LARGEST_OBJECT in REQUEST_UPDATE_OK (§10.9.1, §10.2.17). largestMu sync.Mutex largest message.Location hasLargest bool - // ended is latched by the first Done, so PUBLISH_DONE is sent once and no - // subgroup opens after it. + // ended is latched by the first Done, under subMu, so PUBLISH_DONE is + // sent once and no subgroup opens or writes after it. ended atomic.Bool brokerInit sync.Once @@ -57,23 +68,36 @@ type Publication struct { // resumes. var ErrForwardPaused = errors.New("moqt/session: Forward State is 0; not sending objects") -// ErrPublicationEnded is returned by [Publication.OpenSubgroup] once the -// publication has sent PUBLISH_DONE — by [Publication.Done], or automatically -// after a declined REQUEST_UPDATE (§10.9.1). +// ErrPublicationEnded is returned by [Publication.OpenSubgroup], and by the +// WriteObject methods of a subgroup it opened, once the publication has ended +// — by [Publication.Done], or automatically after a declined REQUEST_UPDATE +// (§10.9.1). var ErrPublicationEnded = errors.New("moqt/session: publication ended (PUBLISH_DONE sent)") // newPublication builds a Publication whose initial Forward State is the -// establishing message's FORWARD (§5.1), or 1 when omitted (§10.2.18). -func newPublication(s *Session, stream Stream, requestID, alias uint64, establishing message.Parameters) *Publication { +// establishing message's FORWARD (§5.1), or 1 when omitted (§10.2.18), and +// whose Largest Object starts at the LARGEST_OBJECT this side announced in +// its SUBSCRIBE_OK or PUBLISH (§10.2.17), if any. +func newPublication( + s *Session, + stream Stream, + requestID, alias uint64, + establishing, announced message.Parameters, +) *Publication { // The subscriber may send REQUEST_UPDATE (§10.9) but not // PUBLISH_STATE_NOTIFY (§10.10). p := &Publication{ Stream: stream, s: s, requestID: requestID, alias: alias, peerUpdate: true, updateScope: message.ScopeUpdateFromSubscriber, + open: make(map[*OutgoingSubgroupStream]struct{}), } + p.endCtx, p.endCancel = context.WithCancel(context.Background()) if f, ok := establishing.Find(message.ParamForward); ok { p.paused.Store(f.Byte == 0) } + if lo, ok := announced.Find(message.ParamLargestObject); ok { + p.largest, p.hasLargest = message.Location{Group: lo.Group, Object: lo.Object}, true + } return p } @@ -98,9 +122,10 @@ func (p *Publication) Broker() *RequestBroker { // Serve callback still sees the update). Any other parameter declines the // whole update with NOT_SUPPORTED. // -// The REQUEST_UPDATE_OK carries LARGEST_OBJECT (§10.9.1) once objects have -// been written through [Publication.OpenSubgroup] streams; other objects are -// not seen. +// The REQUEST_UPDATE_OK carries LARGEST_OBJECT (§10.9.1, §10.2.17): the +// larger of the one this side's SUBSCRIBE_OK or PUBLISH reported and the +// largest Object written through [Publication.OpenSubgroup] streams; other +// objects are not seen. func (p *Publication) ApplyUpdate(upd *message.RequestUpdate) (*message.RequestOK, error) { forward, setForward := false, false for _, prm := range upd.Parameters { @@ -161,41 +186,92 @@ func (p *Publication) TrackAlias() uint64 { return p.alias } // publication's track, filling in the Track Alias automatically — h.TrackAlias // is ignored and overwritten. It is otherwise identical to // [Session.OpenSubgroup]: the caller MUST Close the returned stream to FIN it -// once all objects are written, or Cancel to reset. +// once all objects are written, or Cancel to reset. After [Publication.Done] +// it opens nothing, and the WriteObject methods of a subgroup it opened fail, +// both with [ErrPublicationEnded]. func (p *Publication) OpenSubgroup(h message.SubgroupHeader) (*OutgoingSubgroupStream, error) { + p.subMu.Lock() if p.ended.Load() { + p.subMu.Unlock() return nil, ErrPublicationEnded } if p.paused.Load() { + p.subMu.Unlock() return nil, ErrForwardPaused } + p.opening.Add(1) + p.subMu.Unlock() + defer p.opening.Done() + h.TrackAlias = p.alias - sg, err := p.s.OpenSubgroup(h) + // Once its header is written the peer can attribute the stream, so it + // counts, even if Done reset it just after; its writes then fail. + sg, _, err := p.s.openSubgroup(p.endCtx, h, true) if err != nil { + if p.endCtx.Err() != nil { + return nil, ErrPublicationEnded // Done reset the header write + } return nil, err } - p.subgroupCount.Add(1) sg.onObject = p.noteObject sg.paused = p.paused.Load + sg.ended = p.ended.Load + sg.onEnd = func() { p.forget(sg) } + p.subMu.Lock() + p.subgroupCount++ + p.open[sg] = struct{}{} + p.subMu.Unlock() return sg, nil } -// Done ends the publication (§10.12): it writes a PUBLISH_DONE with the given -// status code and reason, then FINs the request stream. The §10.12 Stream Count -// is set to the number of subgroup streams opened via [Publication.OpenSubgroup] -// so a subscriber knows how many data streams to expect; this is exact only when -// every subgroup was opened through this handle (subgroups opened via -// [Session.OpenSubgroup] directly are not counted — send PUBLISH_DONE yourself -// via message.Marshal if you need a different count). +// forget drops a subgroup that was FINished or reset from the ones Done +// resets. +func (p *Publication) forget(sg *OutgoingSubgroupStream) { + p.subMu.Lock() + delete(p.open, sg) + p.subMu.Unlock() +} + +// Done ends the publication (§10.12): "A sender MUST NOT send PUBLISH_DONE +// until it has closed all streams it will ever open", so Done stops new +// subgroups, resets the ones still open with CANCELLED, then writes a +// PUBLISH_DONE with the given status code and reason and FINs the request +// stream. It does not wait for subgroups to drain: finish them with Close +// first to deliver their objects. A subgroup's WriteObject after Done fails +// with [ErrPublicationEnded]. +// +// The Stream Count is the number of subgroup streams opened via +// [Publication.OpenSubgroup], exact however those opens race Done. Subgroups +// opened via [Session.OpenSubgroup] directly are not counted — send +// PUBLISH_DONE yourself via message.Marshal if you need a different count. // // Only the first call sends; later ones return nil. func (p *Publication) Done(code moqt.PublishDoneCode, reason string) error { - if !p.ended.CompareAndSwap(false, true) { + p.subMu.Lock() + if p.ended.Load() { + p.subMu.Unlock() return nil } + p.ended.Store(true) + p.subMu.Unlock() + // Opens in flight finish promptly: their header writes are cancelled. + p.endCancel() + p.opening.Wait() + + p.subMu.Lock() + open, count := p.open, p.subgroupCount + p.open = nil + p.subMu.Unlock() + // §11.4.3: ending the subscription early resets the subgroups it cut + // short, keeping what was written, header first, reliable so the + // subscriber can attribute each reset stream when handling PUBLISH_DONE. + for sg := range open { + sg.MarkReliable() + sg.dst.CancelWrite(uint64(moqt.StreamResetCancelled)) + } if err := p.writeThenClose(&message.PublishDone{ StatusCode: code, - StreamCount: p.subgroupCount.Load(), + StreamCount: count, ErrorReason: reason, }); err != nil { return fmt.Errorf("moqt/session: write PUBLISH_DONE: %w", err) @@ -250,7 +326,7 @@ func (s *Session) Publish(ctx context.Context, m *message.Publish) (*Publication func(stream Stream, _ *message.RequestOK) (*Publication, error) { // The PUBLISH sets the initial Forward State (§5.1); PUBLISH_OK // carries no subscription parameters. - p := newPublication(s, stream, m.RequestID, m.TrackAlias, m.Parameters) + p := newPublication(s, stream, m.RequestID, m.TrackAlias, m.Parameters, m.Parameters) p.answered = message.TypePublish return p, nil }) diff --git a/pkg/moqt/session/request.go b/pkg/moqt/session/request.go index 535e4d76..86d688a1 100644 --- a/pkg/moqt/session/request.go +++ b/pkg/moqt/session/request.go @@ -875,9 +875,10 @@ func (r *Request) Reply(msg message.Message) error { } // RejectError writes a REQUEST_ERROR with Retry Interval 0 (§10.6.2), stops -// reading and FINs the stream (§3.3.3). If the REQUEST_ERROR cannot be -// written, the stream is reset instead so the requester is not left waiting. -// Use [Request.Reject] to invite a retry. +// reading with CANCELLED (§3.3.4) and FINs the stream (§3.3.3). If the +// REQUEST_ERROR cannot be written, the stream is reset with INTERNAL_ERROR +// instead so the requester is not left waiting. Use [Request.Reject] to invite +// a retry. func (r *Request) RejectError(code moqt.RequestErrorCode, reason string) error { return r.Reject(&RequestRejectedError{Code: code, Reason: reason}) } @@ -907,7 +908,9 @@ func (r *Request) Reject(rej *RequestRejectedError) error { resetStream(r.Stream) return err } - r.Stream.CancelRead(uint64(moqt.StreamResetInternalError)) + // §3.3.4: "SHOULD use a relevant error code". The request ended, so this + // side stopped reading because it was cancelled, not because it failed. + r.Stream.CancelRead(uint64(moqt.StreamResetCancelled)) return r.Stream.Close() } @@ -931,7 +934,7 @@ func (r *Request) AcceptSubscribe(ok *message.SubscribeOK) (*Publication, error) if err := message.Marshal(r.Stream, ok); err != nil { return nil, fmt.Errorf("moqt/session: write SUBSCRIBE_OK: %w", err) } - return newPublication(r.s, r.Stream, sub.RequestID, ok.TrackAlias, sub.Parameters), nil + return newPublication(r.s, r.Stream, sub.RequestID, ok.TrackAlias, sub.Parameters, ok.Parameters), nil } // AcceptPublish accepts an inbound PUBLISH (§10.11): it registers the Track diff --git a/pkg/moqt/session/request_reject_internal_test.go b/pkg/moqt/session/request_reject_internal_test.go new file mode 100644 index 00000000..c7173490 --- /dev/null +++ b/pkg/moqt/session/request_reject_internal_test.go @@ -0,0 +1,34 @@ +package session + +import ( + "slices" + "testing" + + "github.com/floatdrop/moq-go/pkg/moqt" + "github.com/floatdrop/moq-go/pkg/moqt/message" +) + +// cancelRecordingStream is a Stream that accepts every write and records +// each CancelRead's code. The tests using it are synchronous. +type cancelRecordingStream struct { + brokerStubStream + + readCodes []uint64 +} + +func (s *cancelRecordingStream) CancelRead(code uint64) { s.readCodes = append(s.readCodes, code) } + +// TestRejectStopsSendingWithCancelled: once its REQUEST_ERROR is written, a +// rejected request stops reading with CANCELLED, the relevant code (§3.3.4: +// "The stream was cancelled by either endpoint"), not INTERNAL_ERROR. +func TestRejectStopsSendingWithCancelled(t *testing.T) { + t.Parallel() + stream := &cancelRecordingStream{} + r := &Request{Stream: stream, First: &message.Subscribe{}} + if err := r.RejectError(moqt.RequestDoesNotExist, "no"); err != nil { + t.Fatalf("RejectError: %v", err) + } + if want := []uint64{uint64(moqt.StreamResetCancelled)}; !slices.Equal(stream.readCodes, want) { + t.Fatalf("STOP_SENDING codes = %v, want %v (CANCELLED)", stream.readCodes, want) + } +} diff --git a/pkg/moqt/session/setup_token_test.go b/pkg/moqt/session/setup_token_test.go index c5ca384d..22e7816b 100644 --- a/pkg/moqt/session/setup_token_test.go +++ b/pkg/moqt/session/setup_token_test.go @@ -205,6 +205,38 @@ func TestSetupTokensSent(t *testing.T) { } } +// TestSetupTokensHeldMatchesPeerWithManyTokens: with more SETUP tokens than +// a small sort keeps in order, among other options, the peer still receives +// them in the order given, so the REGISTERs the sender believes held are those +// the peer's cache holds (§10.3.1.3, §10.3.1.4). Varied sizes make which ones +// fit depend on that order. +func TestSetupTokensHeldMatchesPeerWithManyTokens(t *testing.T) { + // Other options interleaved as a caller might pass them, so the SETUP + // encoder's sort by Type has to move the tokens past them. + others := []session.Option{ + session.WithMaxAuthTokenCacheSize(1), session.WithAuthority("relay.example"), + session.WithMaxFilterRanges(4), session.WithMaxRequestUpdates(8), + } + var opts []session.Option + for i := range uint64(24) { + if i%6 == 0 { + opts = append(opts, others[i/6]) + } + opts = append(opts, session.WithSetupToken(register(i+1, strings.Repeat("x", int(i%5)*20)))) + } + client, server := openPairWithOpts(t, opts, []session.Option{session.WithMaxAuthTokenCacheSize(300)}) + + held := client.SetupTokenAliases() + for alias := range uint64(24) { + alias++ + _, _, err := server.TokenCache().Resolve(alias) + if want := slices.Contains(held, alias); (err == nil) != want { + t.Errorf("server cache holds alias %d: %v, but the client believes %v (held %v)", + alias, err == nil, want, held) + } + } +} + // TestSetupTokensDefaultCacheIsZero: a peer that advertises no cache size has // the default 0, so no REGISTER is held. func TestSetupTokensDefaultCacheIsZero(t *testing.T) { diff --git a/pkg/moqt/session/track_properties.go b/pkg/moqt/session/track_properties.go index 4fdea66e..f15ec151 100644 --- a/pkg/moqt/session/track_properties.go +++ b/pkg/moqt/session/track_properties.go @@ -34,9 +34,10 @@ func (e *ErrUnsupportedMandatoryTrackProperty) Error() string { } // ErrMalformedTrackProperties is wrapped by the error [ValidateTrackProperties] -// returns when raw Track Properties do not parse (§2.5). Assumption: the draft -// does not cover this, and rejecting with MALFORMED_TRACK (§10.6 defines it -// only for FETCH) is this package's choice. +// returns when raw Track Properties do not parse: a Key-Value-Pair that +// "cannot be parsed" makes the track malformed (§12.7, §2.4.2). Answering +// with MALFORMED_TRACK, which §10.6 defines only for FETCH, is this package's +// choice. var ErrMalformedTrackProperties = errors.New("moqt/session: malformed track properties") // ErrTrackPropertiesNotAllowed is returned, and nothing sent, when asked to @@ -48,9 +49,9 @@ var ErrTrackPropertiesNotAllowed = errors.New("moqt/session: track properties no // // knownMandatory is the set of Mandatory Track Property types this endpoint // supports. Every mandatory property found in raw that is not in this set -// causes *ErrUnsupportedMandatoryTrackProperty to be returned. An empty -// (non-nil) map means "I support no mandatory extensions" — any mandatory -// property will be rejected. +// causes *ErrUnsupportedMandatoryTrackProperty to be returned. An empty or nil +// map means "I support no mandatory extensions" — any mandatory property will +// be rejected. // // Returns the parsed pairs on success. context is used in the error message // to identify the source message (e.g. "SUBSCRIBE_OK"). @@ -83,7 +84,7 @@ func ValidateTrackProperties( // *ErrUnsupportedMandatoryTrackProperty or an error wrapping // [ErrMalformedTrackProperties]; see [TrackPropertiesRejectCode]. It is for // callers that bypass [Request.AcceptPublish] and the outbound openers, which -// already check. Without that option only the values are checked. +// already check. Without that option no Mandatory Track Property is known. func (s *Session) CheckTrackProperties(raw []byte, context string) error { return s.validateTrackProperties(raw, context) } @@ -92,16 +93,13 @@ func (s *Session) CheckTrackProperties(raw []byte, context string) error { // the draft makes session-fatal (see [message.CheckTrackPropertyValues]) // closes the session with PROTOCOL_VIOLATION. Then, against the session's // configured set of known mandatory track property types, an unknown one is -// refused (§2.5.1); if WithKnownMandatoryTrackProperties was never called (the -// map is nil), that check is skipped, for endpoints that pass Track -// Properties through. +// refused (§2.5.1). Without WithKnownMandatoryTrackProperties none is known, +// so every Mandatory Track Property is refused: an endpoint that does not +// understand one "MUST NOT process or forward that track". func (s *Session) validateTrackProperties(raw []byte, context string) error { if err := s.checkTrackPropertyValues(raw, context); err != nil { return err } - if s.knownMandatoryTrackProperties == nil { - return nil // not configured — skip enforcement - } _, err := ValidateTrackProperties(raw, s.knownMandatoryTrackProperties, context) return err } diff --git a/pkg/moqt/session/track_properties_test.go b/pkg/moqt/session/track_properties_test.go index 6321e724..9595e71e 100644 --- a/pkg/moqt/session/track_properties_test.go +++ b/pkg/moqt/session/track_properties_test.go @@ -141,10 +141,10 @@ func TestErrUnsupportedMandatoryTrackPropertyError(t *testing.T) { // TestSubscribeMandatoryTrackPropertyRejected verifies that Subscribe() // returns *ErrUnsupportedMandatoryTrackProperty when the server sends -// SUBSCRIBE_OK with an unknown mandatory track property and the client has -// opted in to enforcement via WithKnownMandatoryTrackProperties. +// SUBSCRIBE_OK with an unknown mandatory track property and the client knows +// none, here through an explicitly empty WithKnownMandatoryTrackProperties. func TestSubscribeMandatoryTrackPropertyRejected(t *testing.T) { - // Client opts in with empty known set → all mandatory are unknown → reject. + // An empty known set → all mandatory are unknown → reject. cli, srv := openPairWithOpts(t, []session.Option{session.WithKnownMandatoryTrackProperties(map[message.PropertyType]struct{}{})}, nil, @@ -579,168 +579,65 @@ func TestInboundPublishMandatoryTrackPropertyValidation(t *testing.T) { } // --------------------------------------------------------------------------- -// Default (no option): enforcement disabled — mandatory properties pass through +// Default (no option): no Mandatory Track Property is known, so each is refused // --------------------------------------------------------------------------- -// TestSubscribeDefaultNoEnforcement verifies that when -// WithKnownMandatoryTrackProperties is NOT called, mandatory track properties -// in SUBSCRIBE_OK are silently accepted. This is the correct default for -// relays and forwarding endpoints. -func TestSubscribeDefaultNoEnforcement(t *testing.T) { - // No WithKnownMandatoryTrackProperties → enforcement disabled. - cli, srv := openPairWithOpts(t, nil, nil) - ctx := t.Context() - - var ( - wg sync.WaitGroup - serverErr error - clientErr error - gotOK *message.SubscribeOK - ) - - wg.Go(func() { - r, err := srv.AcceptRequest(ctx) - if err != nil { - serverErr = err - return - } - serverErr = r.Reply(&message.SubscribeOK{ - TrackAlias: 7, - TrackProperties: mandatoryTrackProps(0x5000, 42), - }) - }) - - wg.Go(func() { - stream, err := cli.Subscribe(ctx, &message.Subscribe{ - Namespace: wire.TrackNamespace{[]byte("ns")}, - Name: []byte("track"), - }) - if err != nil { - clientErr = err - return - } - defer stream.Close() - gotOK = stream.OK - }) - - wg.Wait() - - if serverErr != nil { - t.Fatalf("server: %v", serverErr) - } - if clientErr != nil { - t.Fatalf("client Subscribe: %v", clientErr) - } - if gotOK == nil { - t.Fatal("gotOK is nil") - } - if gotOK.TrackAlias != 7 { - t.Errorf("TrackAlias = %d, want 7", gotOK.TrackAlias) - } -} - -// TestFetchDefaultNoEnforcement verifies that when -// WithKnownMandatoryTrackProperties is NOT called, mandatory track properties -// in FETCH_OK are silently accepted. -func TestFetchDefaultNoEnforcement(t *testing.T) { - cli, srv := openPairWithOpts(t, nil, nil) - ctx := t.Context() - - var ( - wg sync.WaitGroup - serverErr error - clientErr error - gotOK *message.FetchOK - ) - - wg.Go(func() { - r, err := srv.AcceptRequest(ctx) - if err != nil { - serverErr = err - return - } - serverErr = r.Reply(&message.FetchOK{ - EndOfTrack: true, - EndLocation: message.Location{Group: 1, Object: 5}, - TrackProperties: mandatoryTrackProps(0x6000, 99), - }) - }) - - wg.Go(func() { - stream, err := cli.Fetch(ctx, &message.Fetch{ - Namespace: wire.TrackNamespace{[]byte("ns")}, - Name: []byte("track"), +// TestDefaultRefusesUnknownMandatoryTrackProperty: without +// WithKnownMandatoryTrackProperties no Mandatory Track Property is understood, +// so one in SUBSCRIBE_OK, FETCH_OK or TRACK_STATUS_OK fails the request with +// *ErrUnsupportedMandatoryTrackProperty (§2.5.1: "MUST NOT process or forward +// that track"), and one in PUBLISH is refused with UNSUPPORTED_EXTENSION. +func TestDefaultRefusesUnknownMandatoryTrackProperty(t *testing.T) { + props := mandatoryTrackProps(0x5000, 42) + ns := wire.TrackNamespace{[]byte("ns")} + for _, tc := range []struct { + name string + reply message.Message + open func(t *testing.T, cli *session.Session) error + }{ + {"SUBSCRIBE_OK", &message.SubscribeOK{TrackAlias: 7, TrackProperties: props}, + func(t *testing.T, cli *session.Session) error { + _, err := cli.Subscribe(t.Context(), &message.Subscribe{Namespace: ns, Name: []byte("t")}) + return err + }}, + {"FETCH_OK", &message.FetchOK{EndLocation: message.Location{Group: 1}, TrackProperties: props}, + func(t *testing.T, cli *session.Session) error { + _, err := cli.Fetch(t.Context(), &message.Fetch{Namespace: ns, Name: []byte("t")}) + return err + }}, + {"TRACK_STATUS_OK", &message.RequestOK{TrackProperties: props}, + func(t *testing.T, cli *session.Session) error { + _, err := cli.TrackStatus(t.Context(), &message.TrackStatus{Namespace: ns, Name: []byte("t")}) + return err + }}, + } { + t.Run(tc.name, func(t *testing.T) { + cli, srv := openPairWithOpts(t, nil, nil) + go func() { + if r, err := srv.AcceptRequest(t.Context()); err == nil { + _ = r.Reply(tc.reply) + } + }() + err := tc.open(t, cli) + u, ok := errors.AsType[*session.ErrUnsupportedMandatoryTrackProperty](err) + if !ok || u.PropertyType != 0x5000 { + t.Fatalf("error = %v, want *ErrUnsupportedMandatoryTrackProperty for 0x5000", err) + } }) - if err != nil { - clientErr = err - return - } - defer stream.Close() - gotOK = stream.OK - }) - - wg.Wait() - - if serverErr != nil { - t.Fatalf("server: %v", serverErr) - } - if clientErr != nil { - t.Fatalf("client Fetch: %v", clientErr) - } - if gotOK == nil { - t.Fatal("gotOK is nil") } -} - -// TestTrackStatusDefaultNoEnforcement verifies that when -// WithKnownMandatoryTrackProperties is NOT called, mandatory track properties -// in TRACK_STATUS_OK are silently accepted. -func TestTrackStatusDefaultNoEnforcement(t *testing.T) { - cli, srv := openPairWithOpts(t, nil, nil) - ctx := t.Context() - - var ( - wg sync.WaitGroup - serverErr error - clientErr error - gotOK *message.TrackStatusOK - ) - - wg.Go(func() { - r, err := srv.AcceptRequest(ctx) - if err != nil { - serverErr = err - return - } - serverErr = r.Reply(&message.RequestOK{ - TrackProperties: mandatoryTrackProps(0x7000, 1), - }) - }) - - wg.Go(func() { - ts, err := cli.TrackStatus(ctx, &message.TrackStatus{ - Namespace: wire.TrackNamespace{[]byte("ns")}, - Name: []byte("track"), - }) - if err != nil { - clientErr = err - return + t.Run("PUBLISH", func(t *testing.T) { + cli, srv := openPairWithOpts(t, nil, nil) + go func() { + if r, err := srv.AcceptRequest(t.Context()); err == nil { + _, _ = r.AcceptPublish() + } + }() + _, err := cli.Publish(t.Context(), &message.Publish{Namespace: ns, Name: []byte("t"), TrackProperties: props}) + rej, ok := errors.AsType[*session.RequestRejectedError](err) + if !ok || rej.Code != moqt.RequestUnsupportedExtension { + t.Fatalf("Publish = %v, want REQUEST_ERROR UNSUPPORTED_EXTENSION", err) } - defer ts.Close() - gotOK = ts.OK }) - - wg.Wait() - - if serverErr != nil { - t.Fatalf("server: %v", serverErr) - } - if clientErr != nil { - t.Fatalf("client TrackStatus: %v", clientErr) - } - if gotOK == nil { - t.Fatal("gotOK is nil") - } } // --------------------------------------------------------------------------- diff --git a/pkg/moqt/session/track_status.go b/pkg/moqt/session/track_status.go index fb02e26c..bb5e5e50 100644 --- a/pkg/moqt/session/track_status.go +++ b/pkg/moqt/session/track_status.go @@ -101,8 +101,9 @@ func (e *readErrRecorder) Read(p []byte) (int, error) { func (s *Session) TrackStatus(ctx context.Context, m *message.TrackStatus) (*TrackStatusRequest, error) { return awaitRequestResponse(ctx, s, m, func(stream Stream, ok *message.RequestOK) (*TrackStatusRequest, error) { - // §2.5.1: reject tracks with unknown mandatory track properties. - // TRACK_STATUS_OK carries the same Track Properties as SUBSCRIBE_OK. + // Assumption: §2.5.1 lists only PUBLISH, SUBSCRIBE_OK and FETCH_OK, + // but TRACK_STATUS_OK carries the Track Properties "it would have + // set in a SUBSCRIBE_OK" (§10.15), so the SUBSCRIBE_OK rule applies. if err := s.validateTrackProperties(ok.TrackProperties, "TRACK_STATUS_OK"); err != nil { _ = stream.Close() return nil, err diff --git a/pkg/moqt/wire/kv.go b/pkg/moqt/wire/kv.go index 32803e29..98bad754 100644 --- a/pkg/moqt/wire/kv.go +++ b/pkg/moqt/wire/kv.go @@ -39,8 +39,11 @@ func (w *Writer) KVPair(p KVPair, prev uint64) uint64 { // KVPairs appends a list of KVPairs with delta encoding starting from prev=0. // Pairs are sorted by Type before encoding so callers do not need to order them. +// The sort is stable: pairs of one Type keep the caller's order, which a +// repeated option can depend on: which SETUP REGISTERs fit the peer's cache +// depends on their order (§10.3.1.3, §10.3.1.4). func (w *Writer) KVPairs(pairs []KVPair) { - slices.SortFunc(pairs, func(a, b KVPair) int { return cmp.Compare(a.Type, b.Type) }) + slices.SortStableFunc(pairs, func(a, b KVPair) int { return cmp.Compare(a.Type, b.Type) }) var prev uint64 for _, p := range pairs { prev = w.KVPair(p, prev) diff --git a/pkg/moqt/wire/wire_test.go b/pkg/moqt/wire/wire_test.go index 04568aea..218b0234 100644 --- a/pkg/moqt/wire/wire_test.go +++ b/pkg/moqt/wire/wire_test.go @@ -2,9 +2,11 @@ package wire import ( "bytes" + "cmp" "errors" "io" "math" + "slices" "strings" "testing" "unicode/utf8" @@ -208,6 +210,31 @@ func TestTrackNamespaceTooManyFieldsRejected(t *testing.T) { } } +// TestKVPairsKeepOrderWithinAType: KVPairs sorts by Type for the delta +// encoding but keeps pairs of one Type in the caller's order, however many +// there are. Which of a SETUP's AUTHORIZATION TOKEN REGISTERs fit the peer's +// cache depends on their order (§10.3.1.3, §10.3.1.4), so it must survive. +func TestKVPairsKeepOrderWithinAType(t *testing.T) { + var pairs []KVPair + for i := range 40 { + pairs = append(pairs, KVPair{Type: 4, IntVal: uint64(i)}) + if i%3 == 0 { + pairs = append(pairs, KVPair{Type: 2, IntVal: uint64(100 + i)}) + } + } + w := NewWriter(nil) + w.KVPairs(slices.Clone(pairs)) + got, err := NewReader(w.Bytes()).KVPairsRemaining() + if err != nil { + t.Fatalf("KVPairsRemaining: %v", err) + } + want := slices.Clone(pairs) + slices.SortStableFunc(want, func(a, b KVPair) int { return cmp.Compare(a.Type, b.Type) }) + if !slices.EqualFunc(got, want, func(a, b KVPair) bool { return a.Type == b.Type && a.IntVal == b.IntVal }) { + t.Fatalf("decoded %v, want %v", got, want) + } +} + func TestKVPairsRoundTrip(t *testing.T) { pairs := []KVPair{ {Type: 2, IntVal: 42}, // even -> varint diff --git a/pkg/relay/handler_subscribe.go b/pkg/relay/handler_subscribe.go index bb9560ec..832e6103 100644 --- a/pkg/relay/handler_subscribe.go +++ b/pkg/relay/handler_subscribe.go @@ -646,8 +646,8 @@ func includeProperties(ps message.Parameters) bool { // upstream SUBSCRIBE failed with err. // // An unknown Mandatory Track Property is UNSUPPORTED_EXTENSION (§2.5.1); -// unparseable Track Properties are MALFORMED_TRACK (an interpretation: the -// draft does not cover them). An upstream REQUEST_ERROR code about the track +// unparseable Track Properties make the track malformed (§12.7, §2.4.2), and +// MALFORMED_TRACK answers them (an interpretation: §10.6 defines it for FETCH). An upstream REQUEST_ERROR code about the track // or the publisher's load passes through with its Retry Interval (§10.6.2), // MALFORMED_TRACK included though §10.6.2 scopes it to FETCH; one about the // relay's own hop or its Next Object filter, or an unknown one, becomes diff --git a/pkg/relay/internal/registry/downstream_write_test.go b/pkg/relay/internal/registry/downstream_write_test.go index 66c74535..947d9791 100644 --- a/pkg/relay/internal/registry/downstream_write_test.go +++ b/pkg/relay/internal/registry/downstream_write_test.go @@ -2,6 +2,8 @@ package registry_test import ( "bytes" + "errors" + "slices" "sync" "sync/atomic" "testing" @@ -59,10 +61,17 @@ type recordingStream struct { mu sync.Mutex buf []byte + readCodes []uint64 // each CancelRead's code closeOnce sync.Once closed chan struct{} } +func (s *recordingStream) CancelRead(code uint64) { + s.mu.Lock() + s.readCodes = append(s.readCodes, code) + s.mu.Unlock() +} + // newRecordingStream returns an open recordingStream. func newRecordingStream() *recordingStream { return &recordingStream{closed: make(chan struct{})} @@ -156,6 +165,57 @@ func TestDownstreamSub_TerminateBeforeOKAnswersWithRequestError(t *testing.T) { if re.ErrorCode != moqt.RequestDoesNotExist { t.Errorf("REQUEST_ERROR code = 0x%X, want DOES_NOT_EXIST", uint64(re.ErrorCode)) } + // §3.3.4: the request ended by REQUEST_ERROR is cancelled, not failed. + stream.mu.Lock() + codes := stream.readCodes + stream.mu.Unlock() + if want := []uint64{uint64(moqt.StreamResetCancelled)}; !slices.Equal(codes, want) { + t.Errorf("STOP_SENDING codes = %v, want %v (CANCELLED)", codes, want) + } +} + +// failingWriteStream is a recordingStream whose writes fail, recording each +// CancelWrite's code too. +type failingWriteStream struct { + *recordingStream + + writeCodes []uint64 +} + +func (s *failingWriteStream) Write([]byte) (int, error) { return 0, errors.New("transport gone") } + +func (s *failingWriteStream) CancelWrite(code uint64) { + s.mu.Lock() + s.writeCodes = append(s.writeCodes, code) + s.mu.Unlock() +} + +// TestDownstreamSub_TerminateBeforeOKWriteFailureResets: when the REQUEST_ERROR +// cannot be written, the stream is reset with INTERNAL_ERROR, a genuine +// failure (§3.3.4), as [session.Request.RejectError] does, not cancelled. +func TestDownstreamSub_TerminateBeforeOKWriteFailureResets(t *testing.T) { + t.Parallel() + stream := &failingWriteStream{recordingStream: newRecordingStream()} + sub := registry.NewDownstreamSub(1, nil, stream, 7) + sub.TerminateWithPublishDone(moqt.PublishDoneTrackEnded, "upstream gone") + + want := []uint64{uint64(moqt.StreamResetInternalError)} + deadline := time.Now().Add(2 * time.Second) + for { + stream.mu.Lock() + reads, writes := slices.Clone(stream.readCodes), slices.Clone(stream.writeCodes) + stream.mu.Unlock() + if len(reads) > 0 && len(writes) > 0 { + if !slices.Equal(reads, want) || !slices.Equal(writes, want) { + t.Fatalf("STOP_SENDING %v, RESET_STREAM %v, want both %v (INTERNAL_ERROR)", reads, writes, want) + } + return + } + if time.Now().After(deadline) { + t.Fatalf("STOP_SENDING %v, RESET_STREAM %v, want both %v (INTERNAL_ERROR)", reads, writes, want) + } + time.Sleep(10 * time.Millisecond) + } } // TestDownstreamSub_TerminateAfterOKSendsPublishDone: once SUBSCRIBE_OK is out, diff --git a/pkg/relay/internal/registry/subscription.go b/pkg/relay/internal/registry/subscription.go index 601879a3..499eca4d 100644 --- a/pkg/relay/internal/registry/subscription.go +++ b/pkg/relay/internal/registry/subscription.go @@ -829,13 +829,19 @@ func (d *DownstreamSub) sendPublishDone(done *pendingPublishDone, streamCount ui d.writeMu.Lock() defer d.writeMu.Unlock() if !d.okSent { - _ = message.Marshal(d.Stream, &message.RequestError{ + // Mirror [session.Request.RejectError] (§3.3.4): an answer that + // cannot be written resets the stream as a failure; once written, + // nothing reads this stream any more, so the read side is + // cancelled. + if err := message.Marshal(d.Stream, &message.RequestError{ ErrorCode: moqt.RequestDoesNotExist, ErrorReason: done.reason, - }) - // Mirror [session.Request.RejectError]: nothing reads this - // stream any more, so cancel the read side too. - d.Stream.CancelRead(uint64(moqt.StreamResetInternalError)) + }); err != nil { + d.Stream.CancelRead(uint64(moqt.StreamResetInternalError)) + d.Stream.CancelWrite(uint64(moqt.StreamResetInternalError)) + return + } + d.Stream.CancelRead(uint64(moqt.StreamResetCancelled)) } else { _ = message.Marshal(d.Stream, &message.PublishDone{ StatusCode: done.code, diff --git a/pkg/relay/relay.go b/pkg/relay/relay.go index a2d4b2d8..ab204eba 100644 --- a/pkg/relay/relay.go +++ b/pkg/relay/relay.go @@ -79,7 +79,7 @@ type Config struct { // refused with UNSUPPORTED_EXTENSION (§2.5.1). Empty (the default) // refuses every Mandatory Track Property. Set it here rather than with // session.WithKnownMandatoryTrackProperties in SessionOptions, which - // overrides this field (and turns the check off with a nil map). + // overrides this field. KnownMandatoryTrackProperties []message.PropertyType // Logger is used for relay-level events (accept loop start/stop, @@ -345,7 +345,6 @@ func New(listener Listener, cfg Config) *Relay { // Prepended, so it is the SETUP budget unless the caller states one — and // stated twice it is advertised twice, which is why [Config.MaxFilterRanges] // is the way to change it rather than another WithMaxFilterRanges here. - // Always non-nil, so every session enforces §2.5.1. knownMandatory := make(map[message.PropertyType]struct{}, len(cfg.KnownMandatoryTrackProperties)) for _, t := range cfg.KnownMandatoryTrackProperties { knownMandatory[t] = struct{}{}