diff --git a/CONTRIBUTING.md b/CONTRIBUTING.md index ed2ece9..d124503 100644 --- a/CONTRIBUTING.md +++ b/CONTRIBUTING.md @@ -44,9 +44,9 @@ The module builds with `CGO_ENABLED=0`. | Path | What it is | | --- | --- | | `.` (`rtc`) | `Client`, `Call`, publisher, subscriber, media engine, stats, event fan-out. `call.go` is the most important file. | -| `coordinator/` | Coordinator REST + event websocket. Exactly one endpoint is called: `POST /api/v2/video/call/{type}/{id}/join`. | +| `coordinator/` | Coordinator REST + event websocket. Joins call `POST /api/v2/video/call/{type}/{id}/fast_join` (candidate SFUs, `fastjoin.go`), or `.../join` on the legacy flow, for reconnects and migrations, and where `fast_join` is missing. | | `coordinator/models/` | Generated from the public OpenAPI spec. Do not hand-edit. | -| `signal/` | The SFU signalling client: the protobuf websocket plus the twirp `SignalServer` RPCs. | +| `signal/` | The SFU signalling client: the protobuf websocket plus the twirp `SignalServer` and `FastJoinServer` RPCs. | | `pc/` | `pc.Transport`, the peer-connection wrapper. | | `track/` | Local tracks driven by a `SampleProvider`, simulcast layers, RED transcoding for audio. | | `audio/` | `audio.PCM`, resampling, chunking, WAV, G.711. | diff --git a/README.md b/README.md index 785d186..40167db 100644 --- a/README.md +++ b/README.md @@ -104,8 +104,17 @@ if _, err := call.AddTrack(track.TrackInfo(), track); err != nil { return writer.Write(audio.FromInt16(samples, 24000, 1)) ``` +To publish from the start, pass the track to `Join` instead: `rtc.WithTrack(track.TrackInfo(), track)`. +Its offer is then built while the coordinator request is in flight and answered by the SFU's +join itself, where `AddTrack` after `Join` costs a renegotiation. + ## Join latency +`Join` takes the fast flow: one coordinator request returns candidate SFUs, and one request +to the SFU joins, answers the publisher offer and returns the subscriber offer. Where the +deployment has no fast join, it joins the legacy way; `call.JoinFlow()` says which ran, and +`rtc.WithJoinFlow(rtc.JoinFlowLegacy)` asks for the legacy one. + A call records how its first join went: every network step, timer and local step, what it waited for, and each step in round trips to its peer. `String()` draws it as a DAG with the critical path marked; it marshals to JSON. diff --git a/call.go b/call.go index 21fee4e..299a387 100644 --- a/call.go +++ b/call.go @@ -65,6 +65,10 @@ type joinOptions struct { location string migratingFrom string + flow JoinFlow + // tracks are published from the start: see WithTrack. + tracks []trackWithInfo + // used by integration tests to induce specific behaviour beforeSubscriberSendAnswer func(*signal_rpc.SendAnswerRequest) error } @@ -76,6 +80,7 @@ func defaultJoinOptions() joinOptions { }, subscriber: SubscriberFunc(func(OnTrackReceived) {}), create: true, + flow: JoinFlowFast, } } @@ -283,6 +288,10 @@ type Call struct { coordinatorState atomic.Pointer[CallState] // trace records the first join's steps; see JoinTrace. trace joinTracer + // joinFlow is the JoinFlow the first join took. + joinFlow atomic.Value + // attached is closed once a fast join's websocket attach has finished, well or not. + attached atomic.Pointer[chan struct{}] GetCred GetCredentialsFunc cred atomic.Pointer[models.Credentials] @@ -754,6 +763,23 @@ func (c *Call) Client() *signal.Client { return c.peer.Load().client.Load() } +// credentials are the SFU credentials in use, or none yet: a fast join builds its peer +// connections before the coordinator has said which SFU to use. +func (c *Call) credentials() models.Credentials { + if cred := c.cred.Load(); cred != nil { + return *cred + } + return models.Credentials{} +} + +func iceServers(servers []models.ICEServerResponse) []webrtc.ICEServer { + var out []webrtc.ICEServer + for _, s := range servers { + out = append(out, webrtc.ICEServer{URLs: s.Urls, Username: s.Username, Credential: s.Password}) + } + return out +} + // unifiedSessionID returns an ID that stays the same for the lifetime of this // Call, across every reconnect and migration, unlike SessionID which a rejoin // rotates. It is what lets server-side stats be stitched back together into one @@ -805,11 +831,13 @@ func humanizeSdkType(sdkType sfu_models.SdkType) string { return sdkType.String() } -// Join performs the coordinator join, opens the SFU websocket, sends the SFU -// JoinRequest and brings up the publisher and subscriber peer connections. +// Join joins the call: the coordinator request, the SFU join and both peer +// connections, returning once the SFU has accepted the client. Media starts flowing +// shortly after, as ICE and DTLS complete. // -// The coordinator half runs only on the first call; the reconnect and migration -// paths re-enter Join to rebuild the SFU side against fresh credentials. +// By default it takes the fast join (JoinFlowFast): see WithJoinFlow. The reconnect +// and migration paths re-enter Join to rebuild the SFU side against fresh +// credentials, always through the SFU websocket's JoinRequest. func (c *Call) Join(ctx context.Context, opts ...JoinOption) (*sfu_events.JoinResponse, error) { options := defaultJoinOptions() for _, o := range opts { @@ -821,17 +849,40 @@ func (c *Call) Join(ctx context.Context, opts ...JoinOption) (*sfu_events.JoinRe c.cc.claimConnectTrace(rec) } - if err := c.joinCoordinator(ctx, options); err != nil { - return nil, err + if !reconnecting && c.GetCred == nil && options.flow == JoinFlowFast { + resp, err := c.fastJoin(ctx, opts, options, rec) + if !errors.Is(err, errFastJoinUnavailable) { + return resp, err + } + c.logger.WithField("err", err).Warn("fast join unavailable, joining the legacy way") } + resp, err := c.legacyJoin(ctx, opts, options, rec, reconnecting) + if err != nil || reconnecting { + return resp, err + } + c.joinFlow.Store(JoinFlowLegacy) + for _, t := range options.tracks { + if _, err := c.AddTrack(t.info, t.tracks[0]); err != nil { + return nil, xerr.Wrap(err) + } + } + return resp, nil +} + +// setSessionID picks the session ID of a join: the one asked for, or else the current +// one, or else a new one. +func (c *Call) setSessionID(options joinOptions) { if options.sessionID != "" { // always overwrite if sessionID is provided c.SessionID.Store(options.sessionID) } else if c.SessionID.Load() == "" { c.SessionID.Store(uuid.New().String()) } +} +// rememberJoinOptions keeps the first join's options, which the reconnects reuse. +func (c *Call) rememberJoinOptions(opts []JoinOption) { if c.joinOptions == nil { if opts == nil { // nil vs empty slice, todo better approach @@ -839,7 +890,65 @@ func (c *Call) Join(ctx context.Context, opts ...JoinOption) (*sfu_events.JoinRe } c.joinOptions = opts } +} + +// joinRequest is the SFU websocket's JoinRequest for the current session and +// credentials. +func (c *Call) joinRequest(options joinOptions, publisherSDP, subscriberSDP string) *sfu_events.JoinRequest { + return &sfu_events.JoinRequest{ + Token: c.credentials().Token, + SessionId: c.SessionID.Load(), + PublisherSdp: publisherSDP, + ClientDetails: c.sfuClientDetails(), + ReconnectDetails: options.reconnectDetails, + Source: c.cc.source.toSfuParticipantSource(), + SubscriberSdp: subscriberSDP, + // SessionId rotates on every rejoin, so without this the server cannot + // tell that the sessions before and after an outage were the same call + // from the same client, and its stats are split across them. + UnifiedSessionId: c.unifiedSessionID(), + Capabilities: options.clientCapabilities(), + PreferredPublishOptions: options.preferredPublishOptions, + } +} + +func (c *Call) sfuClientDetails() *sfu_models.ClientDetails { + return &sfu_models.ClientDetails{ + Sdk: &sfu_models.Sdk{ + Type: c.cc.clientDetails.sdkType(), + Major: c.cc.clientDetails.SDKVersion.Major, + Minor: c.cc.clientDetails.SDKVersion.Minor, + Patch: c.cc.clientDetails.SDKVersion.Patch, + }, + Os: &sfu_models.OS{ + Name: c.cc.clientDetails.OSName, + }, + Browser: &sfu_models.Browser{ + Name: c.cc.clientDetails.BrowserName, + }, + } +} + +// startTracing starts the call's stats trace buffers, when the call reports stats. +func (c *Call) startTracing() { + if !c.statsEnabled() { + return + } + attempt := int64(c.reconnectAttempt.Load()) - 1 + c.Tracing.Store(rtcstats.NewCallTraceBuffer("", attempt, c.credentials().Server.EdgeName)) + c.Client().Tracing.Store(c.Tracing.Load()) +} +// legacyJoin is the join through the coordinator's join and the SFU websocket's +// JoinRequest, which every reconnect takes. +func (c *Call) legacyJoin( + ctx context.Context, opts []JoinOption, options joinOptions, rec *jointrace.Recorder, reconnecting bool, +) (*sfu_events.JoinResponse, error) { + if err := c.joinCoordinator(ctx, options); err != nil { + return nil, err + } + c.setSessionID(options) + c.rememberJoinOptions(opts) c.externalRTCP = options.externalRTCP publisherSDP := "" @@ -856,44 +965,8 @@ func (c *Call) Join(ctx context.Context, opts ...JoinOption) (*sfu_events.JoinRe } } - req := &sfu_events.JoinRequest{ - Token: c.cred.Load().Token, - SessionId: c.SessionID.Load(), - PublisherSdp: publisherSDP, - ClientDetails: &sfu_models.ClientDetails{ - Sdk: &sfu_models.Sdk{ - Type: c.cc.clientDetails.sdkType(), - Major: c.cc.clientDetails.SDKVersion.Major, - Minor: c.cc.clientDetails.SDKVersion.Minor, - Patch: c.cc.clientDetails.SDKVersion.Patch, - }, - Os: &sfu_models.OS{ - Name: c.cc.clientDetails.OSName, - }, - Browser: &sfu_models.Browser{ - Name: c.cc.clientDetails.BrowserName, - }, - }, - ReconnectDetails: options.reconnectDetails, - Source: c.cc.source.toSfuParticipantSource(), - SubscriberSdp: subscriberSDP, - // SessionId rotates on every rejoin, so without this the server cannot - // tell that the sessions before and after an outage were the same call - // from the same client, and its stats are split across them. - UnifiedSessionId: c.unifiedSessionID(), - Capabilities: options.clientCapabilities(), - PreferredPublishOptions: options.preferredPublishOptions, - } - - attempt := int64(c.reconnectAttempt.Load()) - 1 - sfuid := c.cred.Load().Server.EdgeName - - if c.statsEnabled() { - // Update Tracer for active call - c.Tracing.Store(rtcstats.NewCallTraceBuffer("", attempt, sfuid)) - // Update Tracer for signal client - c.Client().Tracing.Store(c.Tracing.Load()) - } + req := c.joinRequest(options, publisherSDP, subscriberSDP) + c.startTracing() pcsStart := time.Now() if err := c.initPubAndSub(options); err != nil { @@ -1146,6 +1219,13 @@ func (c *Call) restoreICE(ctx context.Context) error { } func (c *Call) Leave(reason string) error { + // A fast join returns before its websocket has attached; the leave goes on it. + if attached := c.attached.Load(); attached != nil { + select { + case <-*attached: + case <-time.After(fastAttachTimeout): + } + } // Report rtcstats for the last time before leaving if err := c.reportRtcStats(c.callCtx); err != nil { c.logger.WithField("err", err).Error("failed to report rtcstats to the SFU") @@ -1252,13 +1332,20 @@ func (c *Call) OnIceTrickle(trickle *sfu_events.SfuEvent_IceTrickle) { transport.AddICECandidate(candidate) } -func (c *Call) AddSimulcastTracks(trackInfo *sfu_models.TrackInfo, tracks ...webrtc.TrackLocal) (*webrtc.RTPTransceiver, error) { +// recordPublishedTrack adds a track to what the call publishes, which is what a +// reconnect restores, and returns its info with the mid it will be sent on. +func (c *Call) recordPublishedTrack(trackInfo *sfu_models.TrackInfo, tracks ...webrtc.TrackLocal) *sfu_models.TrackInfo { trackInfo = proto.Clone(trackInfo).(*sfu_models.TrackInfo) c.publishedTracksMu.Lock() + defer c.publishedTracksMu.Unlock() // todo, this logic needs to change once we support remove track trackInfo.Mid = strconv.Itoa(len(c.publishedTracks)) c.publishedTracks = append(c.publishedTracks, trackWithInfo{tracks: tracks, info: trackInfo}) - c.publishedTracksMu.Unlock() + return trackInfo +} + +func (c *Call) AddSimulcastTracks(trackInfo *sfu_models.TrackInfo, tracks ...webrtc.TrackLocal) (*webrtc.RTPTransceiver, error) { + trackInfo = c.recordPublishedTrack(trackInfo, tracks...) peerPub := c.publisherPeer() if peerPub == nil { return nil, xerr.Wrap(fmt.Errorf("add simulcast tracks: %w", errNoPeerConnection)) @@ -1272,12 +1359,7 @@ func (c *Call) AddSimulcastTracks(trackInfo *sfu_models.TrackInfo, tracks ...web } func (c *Call) AddTrack(trackInfo *sfu_models.TrackInfo, track webrtc.TrackLocal) (*webrtc.RTPTransceiver, error) { - trackInfo = proto.Clone(trackInfo).(*sfu_models.TrackInfo) - c.publishedTracksMu.Lock() - // todo, this logic needs to change once we support remove track - trackInfo.Mid = strconv.Itoa(len(c.publishedTracks)) - c.publishedTracks = append(c.publishedTracks, trackWithInfo{tracks: []webrtc.TrackLocal{track}, info: trackInfo}) - c.publishedTracksMu.Unlock() + trackInfo = c.recordPublishedTrack(trackInfo, track) peerPub := c.publisherPeer() if peerPub == nil { return nil, xerr.Wrap(fmt.Errorf("add track: %w", errNoPeerConnection)) @@ -1321,7 +1403,7 @@ func (c *Call) sendSubscriptions(ctx context.Context, trackDetails []*signal_rpc return xerr.Wrap(err) } rec.Add(jointrace.Span{ - Name: jointrace.SubSubscribe, After: []string{jointrace.SFUJoin}, + Name: jointrace.SubSubscribe, After: afterFirst(rec, jointrace.SFUJoin, jointrace.SFUFastJoin), Start: start, End: time.Now(), Kind: jointrace.KindNet, Peer: jointrace.PeerSFU, }) diff --git a/client.go b/client.go index f493914..b2e9f15 100644 --- a/client.go +++ b/client.go @@ -741,6 +741,20 @@ func (c *Client) connectWithRetries( _type, id string, joinCallRequest models.JoinCallRequest, ) (*models.JoinCallResponse, error) { + return retryJoin(ctx, c, joinCallRequest, func(ctx context.Context) (models.JoinCallResponse, error) { + return c.CoordinatorClientInterface.JoinCall(ctx, _type, id, joinCallRequest) + }) +} + +// retryJoin runs a coordinator join until it succeeds, retrying what IsRetryableError +// allows with a backoff. A first-ever join of a user the coordinator does not know yet +// waits once for the websocket, which is what creates the user. +func retryJoin[T any]( + ctx context.Context, + c *Client, + joinCallRequest models.JoinCallRequest, + join func(context.Context) (T, error), +) (*T, error) { backoff := 100 * time.Millisecond var lastError error waitedForUser := false @@ -756,7 +770,7 @@ func (c *Client) connectWithRetries( } c.Tracing.Load().Emit(rtcstats.CoordinatorJoinCallEvent, joinCallRequest) - result, err := c.CoordinatorClientInterface.JoinCall(ctx, _type, id, joinCallRequest) + result, err := join(ctx) if err == nil { c.Tracing.Load().Emit(rtcstats.CoordinatorJoinCallResponseEvent, result) return &result, nil @@ -856,9 +870,20 @@ func (c *Client) joinCoordinator( c.shareRTTs(rec) c.Tracing.Load().Emit(rtcstats.CoordinatorConnectedEvent, result) - getCred := func(forceReload bool, excludeSFUID string) (models.Credentials, error) { + return result, c.legacyCredentials(ctx, callType, id, joinCallRequest, result.Credentials), nil +} + +// legacyCredentials is the GetCredentialsFunc of a joined call: first the credentials the +// join handed out, then, for a reconnect or a migration, a fresh coordinator join. +func (c *Client) legacyCredentials( + ctx context.Context, + callType, id string, + joinCallRequest models.JoinCallRequest, + first models.Credentials, +) GetCredentialsFunc { + return func(forceReload bool, excludeSFUID string) (models.Credentials, error) { if !forceReload && excludeSFUID == "" { - return result.Credentials, nil + return first, nil } req := joinCallRequest if excludeSFUID != "" { @@ -874,8 +899,40 @@ func (c *Client) joinCoordinator( c.Tracing.Load().Emit(rtcstats.CoordinatorConnectedEvent, retryResult) return retryResult.Credentials, nil } +} + +// fastJoinCoordinator is joinCoordinator for the fast join: it POSTs to fast_join and +// returns the candidate SFUs instead of credentials for one. +func (c *Client) fastJoinCoordinator( + ctx context.Context, + callType, id string, + joinCallRequest models.JoinCallRequest, + rec *jointrace.Recorder, +) (*models.FastJoinCallResponse, error) { + if joinCallRequest.Location == "" { + joinCallRequest.Location = LocationAuto + } - return result, getCred, nil + c.Tracing.Load().Emit(rtcstats.CoordinatorConnectEvent, joinCallRequest) + start := time.Now() + joinCtx := jointrace.WithStep(ctx, rec, jointrace.CoordFastJoin, jointrace.PeerCoordinator) + result, err := retryJoin(joinCtx, c, joinCallRequest, func(ctx context.Context) (models.FastJoinCallResponse, error) { + return c.CoordinatorClientInterface.FastJoinCall(ctx, callType, id, joinCallRequest) + }) + if err != nil { + return nil, err + } + note := "" + if jointrace.Reused(joinCtx) { + note = "reused connection" + } + rec.Add(jointrace.Span{ + Name: jointrace.CoordFastJoin, + Start: start, End: time.Now(), Kind: jointrace.KindNet, Peer: jointrace.PeerCoordinator, Note: note, + }) + c.shareRTTs(rec) + c.Tracing.Load().Emit(rtcstats.CoordinatorConnectedEvent, result) + return result, nil } func (c *Client) StatsReportingInterval() time.Duration { diff --git a/client_ws_test.go b/client_ws_test.go index 95f0523..ca043c2 100644 --- a/client_ws_test.go +++ b/client_ws_test.go @@ -3,6 +3,7 @@ package rtc import ( "context" "encoding/json" + "fmt" "net/http" "net/http/httptest" "net/url" @@ -32,21 +33,47 @@ type fakeCoordinator struct { watches chan url.Values events chan string known atomic.Bool + + // fastJoins receives the query of every fast_join. Until serveFastJoin, fast_join + // is a 404, as on a coordinator from before it. + fastJoins chan url.Values + candidates atomic.Pointer[[]models.SFUCandidate] + token string +} + +// serveFastJoin makes fast_join answer with a candidate per SFU, in order. +func (f *fakeCoordinator) serveFastJoin(sfus ...*testutil.FakeSFU) { + candidates := make([]models.SFUCandidate, len(sfus)) + for i, sfu := range sfus { + cred := fakeSFUCredentials(sfu, fmt.Sprintf("sfu-fake-%d", i+1), f.token) + candidates[i] = models.SFUCandidate{Server: cred.Server, Token: cred.Token, SetupGrant: fmt.Sprintf("grant-%d", i+1)} + } + f.candidates.Store(&candidates) } func newFakeCoordinator(t *testing.T, wsDelay time.Duration, unknownUsers bool) *fakeCoordinator { t.Helper() f := &fakeCoordinator{ - sfu: testutil.NewFakeSFU(), - joins: make(chan url.Values, 4), - watches: make(chan url.Values, 4), - events: make(chan string, 4), + sfu: testutil.NewFakeSFU(), + joins: make(chan url.Values, 4), + watches: make(chan url.Values, 4), + events: make(chan string, 4), + fastJoins: make(chan url.Values, 4), } t.Cleanup(f.sfu.Close) f.known.Store(!unknownUsers) token, err := testutil.GenerateToken("test-api-key", "test-api-secret", "ws-user", time.Hour) require.NoError(t, err) + f.token = token.Token + unknownUser := func(w http.ResponseWriter) bool { + if f.known.Load() { + return false + } + w.WriteHeader(http.StatusNotFound) + _, _ = w.Write([]byte(`{"code":16,"message":"JoinCall failed with error: \"the user ws-user does not exist\"","StatusCode":404}`)) + return true + } mux := http.NewServeMux() mux.HandleFunc("GET /api/v2/connect", func(w http.ResponseWriter, r *http.Request) { @@ -88,13 +115,25 @@ func newFakeCoordinator(t *testing.T, wsDelay time.Duration, unknownUsers bool) mux.HandleFunc("POST /api/v2/video/call/{type}/{id}/join", func(w http.ResponseWriter, r *http.Request) { f.joins <- r.URL.Query() w.Header().Set("Content-Type", "application/json") - if !f.known.Load() { - w.WriteHeader(http.StatusNotFound) - _, _ = w.Write([]byte(`{"code":16,"message":"JoinCall failed with error: \"the user ws-user does not exist\"","StatusCode":404}`)) + if unknownUser(w) { return } _ = json.NewEncoder(w).Encode(models.JoinCallResponse{Credentials: fakeSFUCredentials(f.sfu, "sfu-fake", token.Token)}) }) + mux.HandleFunc("POST /api/v2/video/call/{type}/{id}/fast_join", func(w http.ResponseWriter, r *http.Request) { + candidates := f.candidates.Load() + if candidates == nil { + http.NotFound(w, r) + return + } + f.fastJoins <- r.URL.Query() + w.Header().Set("Content-Type", "application/json") + w.Header().Set("Server-Timing", "fastjoin;dur=1.5") + if unknownUser(w) { + return + } + _ = json.NewEncoder(w).Encode(models.FastJoinCallResponse{Candidates: *candidates}) + }) mux.HandleFunc("GET /api/v2/video/call/{type}/{id}", func(w http.ResponseWriter, r *http.Request) { f.watches <- r.URL.Query() w.Header().Set("Content-Type", "application/json") @@ -105,16 +144,16 @@ func newFakeCoordinator(t *testing.T, wsDelay time.Duration, unknownUsers bool) return f } -func (f *fakeCoordinator) client(t *testing.T) *Client { +func (f *fakeCoordinator) client(t *testing.T, opts ...Option) *Client { t.Helper() token, err := testutil.GenerateToken("test-api-key", "test-api-secret", "ws-user", time.Hour) require.NoError(t, err) client, err := NewClient(token.APIKey, User{ID: "ws-user"}, StaticToken(token.Token), - WithCoordinatorOptions( + append([]Option{WithCoordinatorOptions( coordinator.ApiURL(f.srv.URL), coordinator.WithWsURL("ws"+strings.TrimPrefix(f.srv.URL, "http")+"/api/v2/connect"), - )) + )}, opts...)...) require.NoError(t, err) t.Cleanup(func() { _ = client.Close() }) return client diff --git a/cmd/joinbench/bench.go b/cmd/joinbench/bench.go index a1833c0..416f3b0 100644 --- a/cmd/joinbench/bench.go +++ b/cmd/joinbench/bench.go @@ -90,9 +90,15 @@ type joined struct { received atomic.Int64 } -func (b *bench) join(ctx context.Context, client *rtc.Client, user, callID string) (*joined, error) { +// join joins user to the call, publishing audio if asked: with the join on the fast +// flow, which is what it is for, and right after it on the legacy one. +func (b *bench) join(ctx context.Context, client *rtc.Client, user, callID string, publish bool) (*joined, error) { j := &joined{user: user, call: client.Call(b.cfg.CallType, callID)} - opts := []rtc.JoinOption{rtc.WithOnTrack(rtc.SubscriberFunc(func(remote rtc.OnTrackReceived) { + flow := rtc.JoinFlowLegacy + if b.cfg.Flow == flowFast { + flow = rtc.JoinFlowFast + } + opts := []rtc.JoinOption{rtc.WithJoinFlow(flow), rtc.WithOnTrack(rtc.SubscriberFunc(func(remote rtc.OnTrackReceived) { // The first-packet stamp is taken on read: keep reading. go func() { for { @@ -106,14 +112,35 @@ func (b *bench) join(ctx context.Context, client *rtc.Client, user, callID strin if b.cfg.Location != "" { opts = append(opts, rtc.WithLocation(b.cfg.Location)) } + var info *sfu_models.TrackInfo + var audio webrtc.TrackLocal + if publish { + var err error + if info, audio, err = silentAudio(); err != nil { + return nil, err + } + if flow == rtc.JoinFlowFast { + opts = append(opts, rtc.WithTrack(info, audio)) + } + } var err error if j.resp, err = j.call.Join(ctx, opts...); err != nil { return nil, fmt.Errorf("%s join: %w", user, err) } + if got := j.call.JoinFlow(); got != flow { + _ = j.call.Leave("joinbench: wrong flow") + return nil, fmt.Errorf("%s: asked for the %s flow, the join took %s", user, flow, got) + } if b.cfg.SFU != "" && !samePin(b.cfg.SFU, j.sfu()) { _ = j.call.Leave("joinbench: wrong SFU") return nil, fmt.Errorf("%s: asked for SFU %s, the coordinator returned %s", user, b.cfg.SFU, j.sfu()) } + if publish && flow == rtc.JoinFlowLegacy { + if _, err := j.call.AddTrack(info, audio); err != nil { + _ = j.call.Leave("joinbench: publish failed") + return nil, fmt.Errorf("%s publish: %w", j.user, err) + } + } return j, nil } @@ -148,17 +175,14 @@ func (*silence) NextSample(ctx context.Context) (media.Sample, error) { func (*silence) CurrentAudioLevel() uint8 { return 127 } -func (j *joined) publishAudio() error { +func silentAudio() (*sfu_models.TrackInfo, webrtc.TrackLocal, error) { info := &sfu_models.TrackInfo{TrackId: uuid.NewString(), TrackType: sfu_models.TrackType_TRACK_TYPE_AUDIO} audio, err := track.NewAudioTrack(info, &silence{}, webrtc.RTPCodecCapability{MimeType: webrtc.MimeTypeOpus, ClockRate: 48000, Channels: 2}) if err != nil { - return err - } - if _, err := j.call.AddTrack(info, audio); err != nil { - return fmt.Errorf("%s publish: %w", j.user, err) + return nil, nil, err } - return nil + return info, audio, nil } // subscribeTo subscribes to user's audio, as an app does with the call state from the @@ -233,15 +257,12 @@ func (b *bench) scenario(ctx context.Context, mode, scenario string, cl *clients return err } } - alice, err := b.join(ctx, cl.alice, aliceID, r.CallID) + alice, err := b.join(ctx, cl.alice, aliceID, r.CallID, true) if err != nil { return err } defer alice.leave() r.SFU = alice.sfu() - if err := alice.publishAudio(); err != nil { - return err - } aliceTrace, err := alice.await(ctx, jointrace.PubRTP) if scenario == scenarioPubSub { r.addTrace(rolePublisher, aliceID, aliceTrace) @@ -256,7 +277,7 @@ func (b *bench) scenario(ctx context.Context, mode, scenario string, cl *clients return err } } - bob, err := b.join(ctx, cl.bob, bobID, r.CallID) + bob, err := b.join(ctx, cl.bob, bobID, r.CallID, scenario == scenarioOneToOne) if err != nil { return err } @@ -266,9 +287,6 @@ func (b *bench) scenario(ctx context.Context, mode, scenario string, cl *clients } steps := []string{jointrace.SubRTP} if scenario == scenarioOneToOne { - if err := bob.publishAudio(); err != nil { - return err - } steps = append(steps, jointrace.PubRTP) } if err := bob.subscribeTo(ctx, aliceID); err != nil { diff --git a/cmd/joinbench/config.go b/cmd/joinbench/config.go index a58cd37..d00ed29 100644 --- a/cmd/joinbench/config.go +++ b/cmd/joinbench/config.go @@ -68,7 +68,7 @@ func parseConfig(args []string, getenv func(string) string, output io.Writer) (c fs := flag.NewFlagSet("joinbench", flag.ContinueOnError) fs.SetOutput(output) fs.StringVar(&c.Env, "env", envLocal, "local (the T03 stack) or staging (STREAM_* from the environment)") - fs.StringVar(&c.Flow, "flow", flowLegacy, "join path: legacy (fast lands with T23)") + fs.StringVar(&c.Flow, "flow", flowLegacy, "join path: legacy (coordinator join, SFU websocket join) or fast (fast_join, FastJoin); a fast join that falls back to legacy fails the run") fs.StringVar(&modes, "mode", "cold,warm", "comma-separated: cold (new Client per run), warm (one Client, one discarded join)") fs.StringVar(&scenarios, "scenario", scenarioPubSub, "comma-separated: pubsub (alice publishes, bob subscribes), one-to-one (bob publishes and subscribes, timed)") fs.DurationVar(&rtt, "rtt", rtt, "round trip injected with WithNetworkDelay (default 100ms for local, 0 for staging)") @@ -90,11 +90,9 @@ func parseConfig(args []string, getenv func(string) string, output io.Writer) (c } switch c.Flow { - case flowLegacy: - case flowFast: - return config{}, errors.New("-flow fast is not implemented yet: the fast join client lands with T23") + case flowLegacy, flowFast: default: - return config{}, fmt.Errorf("-flow %q: want legacy", c.Flow) + return config{}, fmt.Errorf("-flow %q: want legacy or fast", c.Flow) } var err error if c.Modes, err = list("mode", modes, modeCold, modeWarm); err != nil { diff --git a/cmd/joinbench/config_test.go b/cmd/joinbench/config_test.go index 7e9abc6..35e5c4e 100644 --- a/cmd/joinbench/config_test.go +++ b/cmd/joinbench/config_test.go @@ -81,8 +81,7 @@ func TestParseConfigRejects(t *testing.T) { env func(string) string want string }{ - "fast flow": {[]string{"-flow", "fast"}, localEnv, "T23"}, - "unknown flow": {[]string{"-flow", "quick"}, localEnv, "want legacy"}, + "unknown flow": {[]string{"-flow", "quick"}, localEnv, "want legacy or fast"}, "unknown mode": {[]string{"-mode", "cold,hot"}, localEnv, `"hot"`}, "unknown scenario": {[]string{"-scenario", "mesh"}, localEnv, `"mesh"`}, "no runs": {[]string{"-runs", "0"}, localEnv, "-runs"}, diff --git a/cmd/joinbench/live_test.go b/cmd/joinbench/live_test.go new file mode 100644 index 0000000..a8a37ac --- /dev/null +++ b/cmd/joinbench/live_test.go @@ -0,0 +1,76 @@ +package main + +import ( + "bufio" + "encoding/json" + "io" + "os" + "path/filepath" + "slices" + "testing" + "time" + + "github.com/stretchr/testify/require" +) + +// fastJoinInterimBudget is the warm fast join's time to media against the T21 coordinator +// and T22 SFU, in round trips: 9.7 to 11.1 RTT median on the local stack on 2026-10-01. +// The SFU's candidates still come on the websocket, and ICE and DTLS take 6 RTT; T30-T33 +// take it to 3.5. +const fastJoinInterimBudget = 11.5 + +// TestFastJoinWithinTheInterimBudget runs warm fast joins against the local stack (T03) +// with 100 ms injected and checks the median time to media both ways, within +// fastJoinInterimBudget round trips plus 30 ms. It needs the stack running with a +// fast_join coordinator and a FastJoin SFU: +// +// eval "$(~/src/video-sfu/.factory/3rtt/tools/local-stack.sh env)" +// JOINBENCH_LIVE=local go test -run WithinTheInterimBudget -v ./cmd/joinbench +func TestFastJoinWithinTheInterimBudget(t *testing.T) { + if os.Getenv("JOINBENCH_LIVE") != envLocal { + t.Skip("set JOINBENCH_LIVE=local, with the local stack's STREAM_* in the environment") + } + + out := filepath.Join(t.TempDir(), "fast.jsonl") + cfg, err := parseConfig([]string{ + "-env", envLocal, "-flow", flowFast, "-mode", modeWarm, "-scenario", "pubsub,one-to-one", + "-runs", "5", "-out", out, + }, os.Getenv, io.Discard) + require.NoError(t, err) + require.NoError(t, run(cfg, testWriter{t})) + + f, err := os.Open(out) + require.NoError(t, err) + defer f.Close() + var publish, subscribe []float64 + lines := bufio.NewScanner(f) + for lines.Scan() { + var r runResult + require.NoError(t, json.Unmarshal(lines.Bytes(), &r)) + require.Empty(t, r.Error, "run %d of %s", r.Run, r.Scenario) + require.Equal(t, flowFast, r.Flow) + require.NotNil(t, r.Subscribe) + subscribe = append(subscribe, r.Subscribe.Ms) + if r.Publish != nil { + publish = append(publish, r.Publish.Ms) + } + } + require.NoError(t, lines.Err()) + require.Len(t, subscribe, 10) + require.Len(t, publish, 10) + + bound := fastJoinInterimBudget*float64(cfg.RTT/time.Millisecond) + 30 + for name, times := range map[string][]float64{"publish": publish, "subscribe": subscribe} { + slices.Sort(times) + median := (times[len(times)/2-1] + times[len(times)/2]) / 2 + t.Logf("%s to media: median %.0f ms, bound %.0f ms, runs %v", name, median, bound, times) + require.LessOrEqual(t, median, bound, "%s to media", name) + } +} + +type testWriter struct{ t *testing.T } + +func (w testWriter) Write(p []byte) (int, error) { + w.t.Log(string(p)) + return len(p), nil +} diff --git a/coordinator/client.go b/coordinator/client.go index 60fe940..52c8864 100644 --- a/coordinator/client.go +++ b/coordinator/client.go @@ -119,6 +119,13 @@ type CoordinatorClientInterface interface { joinCallRequest models.JoinCallRequest, ) (models.JoinCallResponse, error) + FastJoinCall( + ctx context.Context, + _type string, + id string, + joinCallRequest models.JoinCallRequest, + ) (models.FastJoinCallResponse, error) + WatchCall(ctx context.Context, _type, id, connectionID string) error Connect( @@ -209,19 +216,43 @@ func (c *Client) JoinCall( joinCallRequest models.JoinCallRequest, ) (models.JoinCallResponse, error) { var response models.JoinCallResponse - query := map[string]any{} - for k := range c.joinQuery { - query[k] = c.joinQuery.Get(k) - } err := c.makeRequest(ctx, http.MethodPost, "/api/v2/video/call/{type}/{id}/join", map[string]any{ "type": _type, "id": id, }, - query, joinCallRequest, &response) + c.joinQueryParams(), joinCallRequest, &response) return response, xerr.Wrapf(err, "join call %s:%s", _type, id) } +// FastJoinCall is JoinCall for the fast join: instead of credentials for one SFU the +// coordinator has already set the call up on, it returns candidate SFUs, each with a +// setup grant that lets it create the call itself. A coordinator without the endpoint +// answers 404, which IsNotFound reports. +func (c *Client) FastJoinCall( + ctx context.Context, + _type string, + id string, + joinCallRequest models.JoinCallRequest, +) (models.FastJoinCallResponse, error) { + var response models.FastJoinCallResponse + err := c.makeRequest(ctx, http.MethodPost, "/api/v2/video/call/{type}/{id}/fast_join", + map[string]any{ + "type": _type, + "id": id, + }, + c.joinQueryParams(), joinCallRequest, &response) + return response, xerr.Wrapf(err, "fast join call %s:%s", _type, id) +} + +func (c *Client) joinQueryParams() map[string]any { + query := map[string]any{} + for k := range c.joinQuery { + query[k] = c.joinQuery.Get(k) + } + return query +} + // WatchCall subscribes the websocket connection connectionID to the call's // events. The coordinator's GetCall does that for a client-side request that // carries a connection_id; the call state it returns is discarded. @@ -310,10 +341,14 @@ func statusError(status int, body []byte) *Error { retry := status == http.StatusTooManyRequests || status/100 == 5 var reported models.APIError + var e *Error if err := json.Unmarshal(body, &reported); err == nil && reported.Message != "" { - return NewError(int(reported.Code), reported.Message, retry) + e = NewError(int(reported.Code), reported.Message, retry) + } else { + e = NewError(0, fmt.Sprintf("unexpected status code %d: %s", status, body), retry) } - return NewError(0, fmt.Sprintf("unexpected status code %d: %s", status, body), retry) + e.Status = status + return e } // RawHandler forwards an event to the interceptors without going through the diff --git a/coordinator/client_test.go b/coordinator/client_test.go index 6ea2c27..8fc214d 100644 --- a/coordinator/client_test.go +++ b/coordinator/client_test.go @@ -123,6 +123,78 @@ func TestJoinCallRequestPath(t *testing.T) { }}, resp.Credentials.IceServers) } +// TestFastJoinCallRequestPath is TestJoinCallRequestPath for fast_join, which answers +// with candidate SFUs, each with its own token, ICE servers and grant, in order. +func TestFastJoinCallRequestPath(t *testing.T) { + t.Parallel() + + var path, auth string + var query url.Values + var body models.JoinCallRequest + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + path, auth, query = r.URL.Path, r.Header.Get("authorization"), r.URL.Query() + raw, err := io.ReadAll(r.Body) + require.NoError(t, err) + require.NoError(t, json.Unmarshal(raw, &body)) + w.Header().Set("Content-Type", "application/json") + _, err = io.WriteString(w, `{ + "call": {"id": "the-call", "type": "default", "cid": "default:the-call"}, + "duration": "1ms", + "candidates": [ + { + "server": {"url": "https://sfu-1.example.com/twirp", "ws_endpoint": "wss://sfu-1.example.com/ws", "edge_name": "sfu-1"}, + "token": "sfu-1-token", + "ice_servers": [{"urls": ["turn:turn-1.example.com:3478"], "username": "u1", "password": "p1"}], + "setup_grant": "grant-1" + }, + { + "server": {"url": "https://sfu-2.example.com/twirp", "ws_endpoint": "wss://sfu-2.example.com/ws", "edge_name": "sfu-2"}, + "token": "sfu-2-token", + "setup_grant": "grant-2" + } + ] + }`) + require.NoError(t, err) + })) + defer srv.Close() + + client := newTestClient(t, coordinator.ApiURL(srv.URL), + coordinator.WithJoinQuery(map[string][]string{"sfu_id": {"sfu-2"}})) + resp, err := client.FastJoinCall(context.Background(), "default", "the-call", models.JoinCallRequest{Location: "auto"}) + require.NoError(t, err) + + require.Equal(t, "/api/v2/video/call/default/the-call/fast_join", path) + require.Equal(t, "jwt-token", auth) + require.Equal(t, "sfu-2", query.Get("sfu_id"), "pinned like join") + require.NotContains(t, query, "connection_id") + require.Equal(t, "auto", body.Location) + + require.Equal(t, "the-call", resp.Call.ID) + require.Len(t, resp.Candidates, 2) + first := resp.Candidates[0].Credentials() + require.Equal(t, "https://sfu-1.example.com/twirp", first.Server.URL) + require.Equal(t, "wss://sfu-1.example.com/ws", first.Server.WsEndpoint) + require.Equal(t, "sfu-1-token", first.Token) + require.Equal(t, []string{"turn:turn-1.example.com:3478"}, first.IceServers[0].Urls) + require.Equal(t, "grant-1", resp.Candidates[0].SetupGrant) + require.Equal(t, "sfu-2-token", resp.Candidates[1].Token) +} + +// A coordinator without fast_join answers it with a plain 404: the join goes the +// legacy way, which an unknown user's 404 must not be mistaken for. +func TestFastJoinCallNotFound(t *testing.T) { + t.Parallel() + + srv := httptest.NewServer(http.NotFoundHandler()) + defer srv.Close() + + client := newTestClient(t, coordinator.ApiURL(srv.URL)) + _, err := client.FastJoinCall(context.Background(), "default", "the-call", models.JoinCallRequest{}) + require.Error(t, err) + require.True(t, coordinator.IsNotFound(err)) + require.False(t, coordinator.IsUnknownUser(err)) +} + func TestJoinCallCarriesTheJoinQuery(t *testing.T) { t.Parallel() diff --git a/coordinator/error.go b/coordinator/error.go index 7047119..27c12bd 100644 --- a/coordinator/error.go +++ b/coordinator/error.go @@ -3,6 +3,7 @@ package coordinator import ( "errors" "fmt" + "net/http" "strings" sfumodels "github.com/GetStream/protocol/protobuf/video/sfu/models" @@ -14,6 +15,8 @@ type Error struct { Code int Message string ShouldRetry bool + // Status is the HTTP status of the response the error came from, or zero. + Status int } func NewError(code int, message string, shouldRetry bool) *Error { @@ -43,6 +46,13 @@ func IsUnknownUser(err error) bool { return ok && strings.HasSuffix(rest, ` does not exist"`) } +// IsNotFound reports whether err is an HTTP 404: for an endpoint, that the coordinator +// or an edge in front of it does not have it. +func IsNotFound(err error) bool { + coordErr := &Error{} + return errors.As(err, &coordErr) && coordErr.Status == http.StatusNotFound +} + // IsRetryableError reports whether err is worth retrying. Errors the // coordinator did not classify -- network failures, timeouts -- are retryable. func IsRetryableError(err error) bool { diff --git a/coordinator/mocks/mock_coordinator_client.go b/coordinator/mocks/mock_coordinator_client.go index 3e13663..62674fe 100644 --- a/coordinator/mocks/mock_coordinator_client.go +++ b/coordinator/mocks/mock_coordinator_client.go @@ -27,6 +27,9 @@ var _ coordinator.CoordinatorClientInterface = &CoordinatorClientInterfaceMock{} // ConnectFunc: func(ctx context.Context, joinRequest *models.WSAuthMessage) (*models.ConnectedEvent, error) { // panic("mock out the Connect method") // }, +// FastJoinCallFunc: func(ctx context.Context, _type string, id string, joinCallRequest models.JoinCallRequest) (models.FastJoinCallResponse, error) { +// panic("mock out the FastJoinCall method") +// }, // GetInterceptorFunc: func() *event.Store[models.WebsocketEvent] { // panic("mock out the GetInterceptor method") // }, @@ -49,6 +52,9 @@ type CoordinatorClientInterfaceMock struct { // ConnectFunc mocks the Connect method. ConnectFunc func(ctx context.Context, joinRequest *models.WSAuthMessage) (*models.ConnectedEvent, error) + // FastJoinCallFunc mocks the FastJoinCall method. + FastJoinCallFunc func(ctx context.Context, _type string, id string, joinCallRequest models.JoinCallRequest) (models.FastJoinCallResponse, error) + // GetInterceptorFunc mocks the GetInterceptor method. GetInterceptorFunc func() *event.Store[models.WebsocketEvent] @@ -70,6 +76,17 @@ type CoordinatorClientInterfaceMock struct { // JoinRequest is the joinRequest argument value. JoinRequest *models.WSAuthMessage } + // FastJoinCall holds details about calls to the FastJoinCall method. + FastJoinCall []struct { + // Ctx is the ctx argument value. + Ctx context.Context + // _type is the _type argument value. + _type string + // ID is the id argument value. + ID string + // JoinCallRequest is the joinCallRequest argument value. + JoinCallRequest models.JoinCallRequest + } // GetInterceptor holds details about calls to the GetInterceptor method. GetInterceptor []struct { } @@ -98,6 +115,7 @@ type CoordinatorClientInterfaceMock struct { } lockClose sync.RWMutex lockConnect sync.RWMutex + lockFastJoinCall sync.RWMutex lockGetInterceptor sync.RWMutex lockJoinCall sync.RWMutex lockWatchCall sync.RWMutex @@ -166,6 +184,50 @@ func (mock *CoordinatorClientInterfaceMock) ConnectCalls() []struct { return calls } +// FastJoinCall calls FastJoinCallFunc. +func (mock *CoordinatorClientInterfaceMock) FastJoinCall(ctx context.Context, _type string, id string, joinCallRequest models.JoinCallRequest) (models.FastJoinCallResponse, error) { + if mock.FastJoinCallFunc == nil { + panic("CoordinatorClientInterfaceMock.FastJoinCallFunc: method is nil but CoordinatorClientInterface.FastJoinCall was just called") + } + callInfo := struct { + Ctx context.Context + _type string + ID string + JoinCallRequest models.JoinCallRequest + }{ + Ctx: ctx, + _type: _type, + ID: id, + JoinCallRequest: joinCallRequest, + } + mock.lockFastJoinCall.Lock() + mock.calls.FastJoinCall = append(mock.calls.FastJoinCall, callInfo) + mock.lockFastJoinCall.Unlock() + return mock.FastJoinCallFunc(ctx, _type, id, joinCallRequest) +} + +// FastJoinCallCalls gets all the calls that were made to FastJoinCall. +// Check the length with: +// +// len(mockedCoordinatorClientInterface.FastJoinCallCalls()) +func (mock *CoordinatorClientInterfaceMock) FastJoinCallCalls() []struct { + Ctx context.Context + _type string + ID string + JoinCallRequest models.JoinCallRequest +} { + var calls []struct { + Ctx context.Context + _type string + ID string + JoinCallRequest models.JoinCallRequest + } + mock.lockFastJoinCall.RLock() + calls = mock.calls.FastJoinCall + mock.lockFastJoinCall.RUnlock() + return calls +} + // GetInterceptor calls GetInterceptorFunc. func (mock *CoordinatorClientInterfaceMock) GetInterceptor() *event.Store[models.WebsocketEvent] { if mock.GetInterceptorFunc == nil { diff --git a/coordinator/models/fastjoin.go b/coordinator/models/fastjoin.go new file mode 100644 index 0000000..f783858 --- /dev/null +++ b/coordinator/models/fastjoin.go @@ -0,0 +1,31 @@ +package models + +// FastJoinCallResponse is the coordinator's answer to fast_join: the call state join +// returns, plus the SFUs the client may join, in the order to try them. The coordinator +// makes no call to any SFU: the first SFU the client reaches creates the call from the +// candidate's setup grant. +type FastJoinCallResponse struct { + Call CallResponse `json:"call"` + Created bool `json:"created"` + Duration string `json:"duration"` + Members []MemberResponse `json:"members"` + Membership *MemberResponse `json:"membership,omitempty"` + OwnCapabilities []OwnCapability `json:"own_capabilities"` + StatsOptions StatsOptions `json:"stats_options"` + Candidates []SFUCandidate `json:"candidates"` +} + +// SFUCandidate is one SFU a fast join may go to. Its token is bound to that SFU, so a +// candidate's token, ICE servers and grant are only ever used with its own server. +type SFUCandidate struct { + Server SFUResponse `json:"server"` + Token string `json:"token"` + IceServers []ICEServerResponse `json:"ice_servers"` + // SetupGrant is the coordinator-signed call setup the SFU verifies to create the call. + SetupGrant string `json:"setup_grant"` +} + +// Credentials returns the candidate as the credentials the rest of the SDK takes. +func (c SFUCandidate) Credentials() Credentials { + return Credentials{IceServers: c.IceServers, Server: c.Server, Token: c.Token} +} diff --git a/fastjoin.go b/fastjoin.go new file mode 100644 index 0000000..e51a086 --- /dev/null +++ b/fastjoin.go @@ -0,0 +1,506 @@ +package rtc + +import ( + "context" + "errors" + "fmt" + "strconv" + "strings" + "time" + + sfu_events "github.com/GetStream/protocol/protobuf/video/sfu/event" + sfu_models "github.com/GetStream/protocol/protobuf/video/sfu/models" + "github.com/GetStream/protocol/protobuf/video/sfu/signal_rpc" + "github.com/pion/webrtc/v4" + "github.com/twitchtv/twirp" + + "github.com/GetStream/getstream-go-webrtc/coordinator" + "github.com/GetStream/getstream-go-webrtc/coordinator/models" + "github.com/GetStream/getstream-go-webrtc/internal/xerr" + "github.com/GetStream/getstream-go-webrtc/jointrace" + "github.com/GetStream/getstream-go-webrtc/signal" +) + +// JoinFlow is how Join reaches the SFU. +type JoinFlow string + +const ( + // JoinFlowFast is the default. The coordinator's fast_join returns candidate SFUs + // without contacting any; the peer connections and the publisher offer are built + // while it is in flight. One FastJoin request to the first SFU that takes the client + // creates the call there if needed, answers the publisher offer and returns the + // subscriber offer, whose answer goes out without waiting. The SFU websocket + // attaches alongside and carries only events. + // + // A deployment without fast join -- a coordinator or edge that does not know + // fast_join, SFUs without FastJoin, or an SFU that cannot create the call itself -- + // is joined with JoinFlowLegacy instead; Call.JoinFlow tells which ran. + JoinFlowFast JoinFlow = "fast" + // JoinFlowLegacy is the coordinator's join, which sets the call up on one SFU, then + // the SFU websocket's JoinRequest, SetPublisher and SendAnswer, each waited for. It + // is kept for benchmarking the fast join against. + JoinFlowLegacy JoinFlow = "legacy" +) + +// WithJoinFlow picks how Join reaches the SFU. The default is JoinFlowFast. +func WithJoinFlow(flow JoinFlow) JoinOption { + return func(o *joinOptions) { + o.flow = flow + } +} + +// WithTrack publishes track from the start of the call, as AddTrack would after Join. +// On the fast join its offer is built while the coordinator request is in flight and +// is answered by the SFU's join itself, so it costs no signalling of its own. A later +// AddTrack costs a renegotiation. +func WithTrack(info *sfu_models.TrackInfo, track webrtc.TrackLocal) JoinOption { + return func(o *joinOptions) { + o.tracks = append(o.tracks, trackWithInfo{tracks: []webrtc.TrackLocal{track}, info: info}) + } +} + +// JoinFlow is the flow the call's first join took, or "" before it has joined. +func (c *Call) JoinFlow() JoinFlow { + flow, _ := c.joinFlow.Load().(JoinFlow) + return flow +} + +// errFastJoinUnavailable means the deployment cannot be fast joined, and the join +// goes the legacy way. +var errFastJoinUnavailable = errors.New("fast join unavailable") + +const ( + // fastJoinRounds is how many times fast_join is asked for candidates when none of + // them took the client: the grants of the first set may have expired. + fastJoinRounds = 2 + // fastAttachTimeout bounds the websocket attach. The SFU drops a fast-joined + // participant whose websocket has not attached after 10 s. + fastAttachTimeout = 10 * time.Second +) + +// fastJoinLocal is what the fast join prepares on the client while fast_join is in +// flight. +type fastJoinLocal struct { + offer publisherOffer + tracks []*sfu_models.TrackInfo + subscriberSDP string +} + +// fastJoin runs the fast join, returning errFastJoinUnavailable, with nothing left +// behind, when the deployment does not support it. +func (c *Call) fastJoin(ctx context.Context, opts []JoinOption, options joinOptions, rec *jointrace.Recorder) (*sfu_events.JoinResponse, error) { + c.setSessionID(options) + c.rememberJoinOptions(opts) + c.externalRTCP = options.externalRTCP + c.trace.setFast(true) + + type coordResult struct { + resp *models.FastJoinCallResponse + err error + } + coordDone := make(chan coordResult, 1) + go func() { + resp, err := c.cc.fastJoinCoordinator(ctx, c.Type, c.Id, options.coordinatorRequest(), rec) + coordDone <- coordResult{resp, err} + }() + + local, err := c.prepareFastJoin(ctx, options, rec) + coord := <-coordDone + if err == nil { + err = coord.err + } + if err == nil { + var resp *sfu_events.JoinResponse + if resp, err = c.fastJoinSFU(ctx, options, local, coord.resp, rec); err == nil { + return resp, nil + } + } + + c.abandonFastJoin(rec) + if coordinator.IsNotFound(err) && !coordinator.IsUnknownUser(err) { + return nil, fmt.Errorf("%w: the coordinator has no fast_join: %w", errFastJoinUnavailable, err) + } + return nil, err +} + +// prepareFastJoin builds the peer connections, adds the join's tracks and creates the +// publisher offer, and the receive-only SDP the SFU builds the subscriber offer from. +func (c *Call) prepareFastJoin(ctx context.Context, options joinOptions, rec *jointrace.Recorder) (fastJoinLocal, error) { + var local fastJoinLocal + start := time.Now() + if err := c.initPubAndSub(options); err != nil { + return local, xerr.Wrap(err) + } + pub, sub := c.publisherPeer(), c.subscriberPeer() + pub.fastPath.Store(true) + sub.fastPath.Store(true) + + sdp, err := subscriberJoinSDPFromMediaEngine(options.subscriberPeerConfig) + if err != nil { + return local, xerr.Wrap(err) + } + local.subscriberSDP = sdp + + if len(options.tracks) > 0 { + for _, t := range options.tracks { + if _, err := pub.AddTrack(c.recordPublishedTrack(t.info, t.tracks...), t.tracks[0]); err != nil { + return local, xerr.Wrap(err) + } + } + offerCtx, cancel := context.WithTimeout(ctx, 5*time.Second) + defer cancel() + if local.offer, err = pub.joinOfferNow(offerCtx); err != nil { + return local, err + } + local.tracks = pub.trackInfos() + } + rec.Add(jointrace.Span{ + Name: jointrace.PCsCreate, Start: start, End: time.Now(), + Kind: jointrace.KindLocal, Peer: jointrace.PeerLocal, + Note: "includes the publisher offer, in parallel with coord.fastjoin", + }) + return local, nil +} + +// abandonFastJoin undoes a fast join that did not work out, so a legacy join starts +// from a call that was never joined. +func (c *Call) abandonFastJoin(rec *jointrace.Recorder) { + c.releaseOldPubSub(0) + c.publishedTracksMu.Lock() + c.publishedTracks = nil + c.publishedTracksMu.Unlock() + c.trace.setFast(false) + rec.Remove(jointrace.PCsCreate) +} + +// fastJoinSFU joins the first candidate SFU that takes the client, asking the +// coordinator for new candidates once if none does. +func (c *Call) fastJoinSFU( + ctx context.Context, options joinOptions, local fastJoinLocal, coord *models.FastJoinCallResponse, rec *jointrace.Recorder, +) (*sfu_events.JoinResponse, error) { + var err error + for round := range fastJoinRounds { + if round > 0 { + if coord, err = c.cc.fastJoinCoordinator(ctx, c.Type, c.Id, options.coordinatorRequest(), nil); err != nil { + return nil, err + } + } + c.applyFastJoinCoordinator(options, coord) + var resp *sfu_events.JoinResponse + resp, err = c.joinCandidates(ctx, options, local, coord.Candidates, rec) + if err == nil || errors.Is(err, errFastJoinUnavailable) || errors.Is(err, errFastJoinFatal) || ctx.Err() != nil { + return resp, err + } + c.logger.WithField("err", err).Warn("no fast join candidate took the client") + } + return nil, err +} + +// applyFastJoinCoordinator records the call state fast_join returned, as joinCoordinator +// does for join. +func (c *Call) applyFastJoinCoordinator(options joinOptions, resp *models.FastJoinCallResponse) { + state := &CallState{ + JoinCallRequest: ptrTo(options.coordinatorRequest()), + CallResponse: resp.Call, + Members: resp.Members, + OwnCapabilities: resp.OwnCapabilities, + StatsOptions: resp.StatsOptions, + Membership: resp.Membership, + } + if cred := c.cred.Load(); cred != nil { + state.Url, state.Token, state.WebsocketUrl, state.EdgeName = + cred.Server.URL, cred.Token, cred.Server.WsEndpoint, cred.Server.EdgeName + } + c.coordinatorState.Store(state) +} + +// errFastJoinFatal marks an SFU answer no other candidate would change. +var errFastJoinFatal = errors.New("fast join refused") + +// joinCandidates tries the candidates in order. What happens after a failed one +// follows the SFU's error: see fastJoinOutcome. +func (c *Call) joinCandidates( + ctx context.Context, options joinOptions, local fastJoinLocal, candidates []models.SFUCandidate, rec *jointrace.Recorder, +) (*sfu_events.JoinResponse, error) { + if len(candidates) == 0 { + return nil, xerr.Error("fast_join returned no SFU candidates") + } + var errs []error + unavailable := 0 + start := time.Now() + for i, candidate := range candidates { + if err := ctx.Err(); err != nil { + return nil, err + } + cred := candidate.Credentials() + client := signal.NewClient(cred, c, c.signalOptions()...) + c.getPeer().client.Store(client) + c.SetCredentials(cred) + c.startTracing() + + // A candidate's spans are only the join's if it takes the client. + attempt := rec.Scratch() + attach := c.dialAttach(client, attempt) + stepCtx := jointrace.WithStep(ctx, attempt, jointrace.SFUFastJoin, jointrace.PeerSFU) + resp, err := client.FastJoin(stepCtx, c.fastJoinRequest(options, local, candidate)) + outcome, err := fastJoinOutcome(resp, err) + if outcome == fastJoinJoined { + note := "" + if i > 0 { + note = fmt.Sprintf("candidate %d of %d", i+1, len(candidates)) + } + attempt.Add(jointrace.Span{ + Name: jointrace.SFUFastJoin, After: []string{jointrace.CoordFastJoin, jointrace.PCsCreate}, + Start: start, End: time.Now(), Kind: jointrace.KindNet, Peer: jointrace.PeerSFU, Note: note, + }) + serverTimings(stepCtx, resp.GetServerTimings()) + rec.Merge(attempt) + return c.fastJoined(options, local, resp, attach, rec) + } + attach.cancel() + err = xerr.Wrapf(err, "fast join %s", cred.Server.EdgeName) + c.logger.WithField("err", err).Warn("fast join candidate failed") + switch outcome { + case fastJoinLegacy: + return nil, fmt.Errorf("%w: %w", errFastJoinUnavailable, err) + case fastJoinRefused: + return nil, fmt.Errorf("%w: %w", errFastJoinFatal, err) + case fastJoinUnavailable: + unavailable++ + } + errs = append(errs, err) + } + if unavailable == len(candidates) { + return nil, fmt.Errorf("%w: no SFU has FastJoin: %w", errFastJoinUnavailable, errors.Join(errs...)) + } + return nil, errors.Join(errs...) +} + +func (c *Call) fastJoinRequest(options joinOptions, local fastJoinLocal, candidate models.SFUCandidate) *signal_rpc.FastJoinRequest { + return &signal_rpc.FastJoinRequest{ + Token: candidate.Token, + SetupGrant: candidate.SetupGrant, + SessionId: c.SessionID.Load(), + UnifiedSessionId: c.unifiedSessionID(), + PublisherSdp: local.offer.sdp.SDP, + Tracks: local.tracks, + SubscriberSdp: local.subscriberSDP, + ClientDetails: c.sfuClientDetails(), + Capabilities: options.clientCapabilities(), + Source: c.cc.source.toSfuParticipantSource(), + PreferredPublishOptions: options.preferredPublishOptions, + } +} + +type fastJoinResult int + +const ( + fastJoinJoined fastJoinResult = iota + // fastJoinNext: try the next candidate. + fastJoinNext + // fastJoinUnavailable: this SFU has no FastJoin; try the next, and join the legacy + // way if none has. + fastJoinUnavailable + // fastJoinLegacy: the SFU cannot create the call itself; join the legacy way. + fastJoinLegacy + // fastJoinRefused: no SFU would take the client. + fastJoinRefused +) + +// fastJoinOutcome classifies a FastJoin answer. The SFU reports its own errors in the +// response with a twirp success; a twirp or transport error is the path to it failing. +func fastJoinOutcome(resp *signal_rpc.FastJoinResponse, err error) (fastJoinResult, error) { + if err != nil { + var twirpErr twirp.Error + if errors.As(err, &twirpErr) && (twirpErr.Code() == twirp.BadRoute || twirpErr.Code() == twirp.Unimplemented) { + return fastJoinUnavailable, err + } + return fastJoinNext, err + } + sfuErr := resp.GetError() + if sfuErr == nil { + return fastJoinJoined, nil + } + err = signal.NewError(sfuErr.GetCode(), sfuErr.GetMessage(), sfuErr.GetShouldRetry()) + switch sfuErr.GetCode() { + case sfu_models.ErrorCode_ERROR_CODE_CALL_PARTICIPANT_LIMIT_REACHED: + return fastJoinRefused, err + case sfu_models.ErrorCode_ERROR_CODE_INTERNAL_SERVER_ERROR: + if strings.Contains(sfuErr.GetMessage(), "join through the coordinator") { + return fastJoinLegacy, err + } + } + // SFU_FULL and SFU_SHUTTING_DOWN created nothing; UNAUTHENTICATED is about this + // candidate's token or grant, and the next has its own. + return fastJoinNext, err +} + +// serverTimings records the SFU's own account of the FastJoin, as a Server-Timing header +// is recorded for the coordinator. +func serverTimings(ctx context.Context, timings []*signal_rpc.ServerTiming) { + var total, longest float64 + parts := make([]string, 0, len(timings)) + for _, t := range timings { + ms := t.GetDurationMs() + parts = append(parts, t.GetName()+"="+strconv.FormatFloat(ms, 'f', -1, 64)) + if t.GetName() == "total" { + total = ms + } + longest = max(longest, ms) + } + if total == 0 { + total = longest + } + jointrace.ServerDuration(ctx, time.Duration(total*float64(time.Millisecond)), "server_timings: "+strings.Join(parts, " ")) +} + +// fastJoined applies a successful FastJoin: the call state, the publisher answer and +// the subscriber offer. The websocket attach finishes on its own; the health monitor +// starts once it has. +func (c *Call) fastJoined( + options joinOptions, local fastJoinLocal, resp *signal_rpc.FastJoinResponse, attach *fastAttach, rec *jointrace.Recorder, +) (*sfu_events.JoinResponse, error) { + now := time.Now() + c.joinFlow.Store(JoinFlowFast) + cred := c.credentials() + // Reconnects and migrations go through the coordinator's join. + c.GetCred = c.cc.legacyCredentials(c.callCtx, c.Type, c.Id, options.coordinatorRequest(), cred) + c.cc.watchCall(c.callCtx, c.Type, c.Id) + c.trace.markJoined() + + joinResp := &sfu_events.JoinResponse{ + CallState: resp.GetCallState(), + FastReconnectDeadlineSeconds: resp.GetFastReconnectDeadlineSeconds(), + PublishOptions: resp.GetPublishOptions(), + } + // Before the subscriber offer: its tracks are looked up in the store. + c.store.Store(NewParticipantStore(c, joinResp.GetCallState())) + c.applyJoinResponse(joinResp) + + pub, sub := c.publisherPeer(), c.subscriberPeer() + // The peer connections were built before the SFU was chosen: they only use its + // TURN servers for gathering after an ICE restart. + for _, t := range []*webrtc.PeerConnection{pub.PC, sub.PC} { + setICEServers(t, cred.IceServers, c) + } + pub.startTracing() + sub.startTracing() + if sdp := resp.GetPublisherSdp(); sdp != "" { + pub.HandleRemoteDescriptionWithNegotiationID(webrtc.SessionDescription{ + Type: webrtc.SDPTypeAnswer, SDP: sdp, + }, local.offer.negotiationID) + } + if sdp := resp.GetSubscriberSdp(); sdp != "" { + c.trace.mu.Lock() + c.trace.subOfferAt = now + c.trace.mu.Unlock() + sub.HandleRemoteDescriptionWithNegotiationID(webrtc.SessionDescription{ + Type: webrtc.SDPTypeOffer, SDP: sdp, + }, resp.GetSubscriberNegotiationId()) + } + + attached := make(chan struct{}) + c.attached.Store(&attached) + go func() { + defer close(attached) + err := attach.join(c.joinRequestAttach(options), rec) + if err != nil { + c.logger.WithField("err", err).Warn("fast join websocket attach failed") + } + c.mu.Lock() + left := c.nextReconnectStrategy == sfu_models.WebsocketReconnectStrategy_WEBSOCKET_RECONNECT_STRATEGY_DISCONNECT + c.mu.Unlock() + if left { + _ = attach.client.Disconnect(true) + return + } + // A failed attach leaves the call without a websocket, which the health + // monitor repairs with a rejoin. + c.onceConnect.Do(func() { + go c.webrtcStatsWorker() + go c.monitorHealth() + }) + }() + return joinResp, nil +} + +func (c *Call) joinRequestAttach(options joinOptions) *sfu_events.JoinRequest { + req := c.joinRequest(options, "", "") + req.AttachFastJoin = true + return req +} + +func setICEServers(pc *webrtc.PeerConnection, servers []models.ICEServerResponse, c *Call) { + if len(servers) == 0 { + return + } + cfg := pc.GetConfiguration() + cfg.ICEServers = iceServers(servers) + if err := pc.SetConfiguration(cfg); err != nil { + c.logger.WithField("err", err).Warn("could not set the SFU's ICE servers") + } +} + +// fastAttach is the websocket of a fast join, dialled while the FastJoin is in flight. +// Its JoinRequest waits for the FastJoin to succeed: an attach reaching the SFU before +// the FastJoin finds no participant and fails. +type fastAttach struct { + client *signal.Client + start time.Time + dialed chan struct{} + conn *signal.Dialed + dialErr error + ctx context.Context + stop context.CancelFunc + // attempt holds the dial's spans until the candidate has taken the client. + attempt *jointrace.Recorder +} + +func (c *Call) dialAttach(client *signal.Client, attempt *jointrace.Recorder) *fastAttach { + ctx, stop := context.WithTimeout(c.callCtx, fastAttachTimeout) + a := &fastAttach{ + client: client, start: time.Now(), dialed: make(chan struct{}), ctx: ctx, stop: stop, attempt: attempt, + } + go func() { + defer close(a.dialed) + a.conn, a.dialErr = client.Dial(jointrace.WithStep(ctx, attempt, jointrace.SFUWSDial, jointrace.PeerSFU)) + if a.dialErr == nil { + attempt.Add(jointrace.Span{ + Name: jointrace.SFUWSDial, After: []string{jointrace.CoordFastJoin}, + Start: a.start, End: time.Now(), Kind: jointrace.KindNet, Peer: jointrace.PeerSFU, + }) + } + }() + return a +} + +// cancel drops the websocket of a candidate that did not take the client. +func (a *fastAttach) cancel() { + a.stop() + go func() { + <-a.dialed + if a.conn != nil { + a.conn.Close() + } + }() +} + +// join attaches the websocket to the participant the FastJoin created. +func (a *fastAttach) join(req *sfu_events.JoinRequest, rec *jointrace.Recorder) error { + defer a.stop() + <-a.dialed + rec.Merge(a.attempt) + if a.dialErr != nil { + return a.dialErr + } + start := time.Now() + if _, err := a.client.Join(a.ctx, a.conn, req); err != nil { + return err + } + rec.Add(jointrace.Span{ + Name: jointrace.SFUWS, After: []string{jointrace.SFUWSDial, jointrace.SFUFastJoin}, + Start: start, End: time.Now(), Kind: jointrace.KindNet, Peer: jointrace.PeerSFU, + Note: "attach; nothing waits for it but the SFU's candidates", + }) + return nil +} diff --git a/fastjoin_test.go b/fastjoin_test.go new file mode 100644 index 0000000..c12a0f4 --- /dev/null +++ b/fastjoin_test.go @@ -0,0 +1,495 @@ +package rtc + +import ( + "context" + "sync/atomic" + "testing" + "time" + + sfu_events "github.com/GetStream/protocol/protobuf/video/sfu/event" + sfu_models "github.com/GetStream/protocol/protobuf/video/sfu/models" + "github.com/GetStream/protocol/protobuf/video/sfu/signal_rpc" + "github.com/pion/webrtc/v4" + "github.com/stretchr/testify/require" + "github.com/twitchtv/twirp" + + "github.com/GetStream/getstream-go-webrtc/internal/testutil" + "github.com/GetStream/getstream-go-webrtc/jointrace" +) + +// fastJoinCall is a call on a client of f, with no health monitor: nothing here +// reconnects. +func fastJoinCall(t *testing.T, f *fakeCoordinator, id string, opts ...Option) *Call { + t.Helper() + + call := f.client(t, opts...).Call(testutil.DefaultCallType, id) + call.onceConnect.Do(func() {}) + return call +} + +func joinFast(t *testing.T, call *Call, opts ...JoinOption) error { + t.Helper() + + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cancel() + _, err := call.Join(ctx, opts...) + if err == nil { + t.Cleanup(func() { _ = call.Leave("test over") }) + } + return err +} + +func sfuError(code sfu_models.ErrorCode, message string) func(context.Context, *signal_rpc.FastJoinRequest) (*signal_rpc.FastJoinResponse, error) { + return func(context.Context, *signal_rpc.FastJoinRequest) (*signal_rpc.FastJoinResponse, error) { + return &signal_rpc.FastJoinResponse{Error: &sfu_models.Error{Code: code, Message: message}}, nil + } +} + +// rpcsOf drains what the fake SFU recorded so far and returns the requests of type T. +func rpcsOf[T any](f *testutil.FakeSFU) []T { + var got []T + for { + select { + case req := <-f.RPCRequests: + if r, ok := req.(T); ok { + got = append(got, r) + } + default: + return got + } + } +} + +func TestFastJoinOutcome(t *testing.T) { + t.Parallel() + + answer := func(code sfu_models.ErrorCode, message string) *signal_rpc.FastJoinResponse { + return &signal_rpc.FastJoinResponse{Error: &sfu_models.Error{Code: code, Message: message}} + } + for name, tc := range map[string]struct { + resp *signal_rpc.FastJoinResponse + err error + want fastJoinResult + }{ + "joined": {resp: &signal_rpc.FastJoinResponse{}, want: fastJoinJoined}, + "full": {resp: answer(sfu_models.ErrorCode_ERROR_CODE_SFU_FULL, "full"), want: fastJoinNext}, + "shutting down": {resp: answer(sfu_models.ErrorCode_ERROR_CODE_SFU_SHUTTING_DOWN, "bye"), want: fastJoinNext}, + "unauthenticated": {resp: answer(sfu_models.ErrorCode_ERROR_CODE_UNAUTHENTICATED, "grant"), want: fastJoinNext}, + "call is full": {resp: answer(sfu_models.ErrorCode_ERROR_CODE_CALL_PARTICIPANT_LIMIT_REACHED, "limit"), want: fastJoinRefused}, + "no bus credential": { + resp: answer(sfu_models.ErrorCode_ERROR_CODE_INTERNAL_SERVER_ERROR, "this server cannot create the call; join through the coordinator"), + want: fastJoinLegacy, + }, + "other internal error": {resp: answer(sfu_models.ErrorCode_ERROR_CODE_INTERNAL_SERVER_ERROR, "boom"), want: fastJoinNext}, + "no FastJoinServer": {err: twirp.NewError(twirp.BadRoute, "no such route"), want: fastJoinUnavailable}, + "unimplemented": {err: twirp.NewError(twirp.Unimplemented, "no"), want: fastJoinUnavailable}, + "transport error": {err: twirp.NewError(twirp.Unavailable, "connection refused"), want: fastJoinNext}, + } { + t.Run(name, func(t *testing.T) { + t.Parallel() + got, err := fastJoinOutcome(tc.resp, tc.err) + require.Equal(t, tc.want, got) + if tc.want == fastJoinJoined { + require.NoError(t, err) + } else { + require.Error(t, err) + } + }) + } +} + +// TestFastJoin runs the whole fast join against an SFU with real peer connections: +// one FastJoin answers the publisher and offers alice's audio, the answer to it goes +// out without Join waiting for it, the websocket attaches after the FastJoin, and +// media flows both ways without a trickled candidate or a SetPublisher. +func TestFastJoin(t *testing.T) { + t.Parallel() + + m := newMediaSFU(t) + release := make(chan struct{}) + m.holdAnswers.Store(&release) + f := newFakeCoordinator(t, 0, false) + f.serveFastJoin(m.fake) + call := fastJoinCall(t, f, "fast-join") + traces := make(chan jointrace.Trace, 1) + call.OnJoinTrace(func(tr jointrace.Trace) { traces <- tr }) + + ctx, cancel := context.WithCancel(context.Background()) + t.Cleanup(cancel) + audio, err := webrtc.NewTrackLocalStaticSample( + webrtc.RTPCodecCapability{MimeType: webrtc.MimeTypeOpus}, "audio", "prefix:TRACK_TYPE_AUDIO") + require.NoError(t, err) + var gotAlice atomic.Bool + require.NoError(t, joinFast(t, call, + WithTrack(&sfu_models.TrackInfo{TrackId: "published-audio", TrackType: sfu_models.TrackType_TRACK_TYPE_AUDIO}, audio), + WithOnTrack(SubscriberFunc(func(OnTrackReceived) { gotAlice.Store(true) })), + WithPublisherPeerConfiguration(loopbackPeerConfig()), + WithSubscriberPeerConfiguration(loopbackPeerConfig()))) + require.Equal(t, JoinFlowFast, call.JoinFlow()) + require.Empty(t, f.joins, "no coordinator join") + require.NotContains(t, <-f.fastJoins, "connection_id") + + req, err := testutil.NextRPCRequest[*signal_rpc.FastJoinRequest](m.fake, time.Second) + require.NoError(t, err) + require.Equal(t, "Bearer "+f.token, <-m.fake.Authorizations, "the candidate's token") + require.Equal(t, f.token, req.GetToken()) + require.Equal(t, "grant-1", req.GetSetupGrant()) + require.Equal(t, call.SessionID.Load(), req.GetSessionId()) + require.Contains(t, req.GetPublisherSdp(), "m=audio", "the offer was built before the FastJoin") + require.Len(t, req.GetTracks(), 1) + require.NotEmpty(t, req.GetTracks()[0].GetMid(), "with the offer's mid") + require.Contains(t, req.GetSubscriberSdp(), "a=recvonly") + + // Join has returned while the SFU still holds the subscriber's answer. + answer, err := testutil.NextRPCRequest[*signal_rpc.SendAnswerRequest](m.fake, iceTimeout) + require.NoError(t, err) + require.Equal(t, sfu_models.PeerType_PEER_TYPE_SUBSCRIBER, answer.GetPeerType()) + require.Equal(t, uint32(1), answer.GetNegotiationId(), "the FastJoin's subscriber_negotiation_id") + close(release) + + // The fake answers an attach that comes before the FastJoin with an error. + attach, err := testutil.NextRequestOf[*sfu_events.SfuRequest_JoinRequest](m.fake, iceTimeout) + require.NoError(t, err) + require.True(t, attach.JoinRequest.GetAttachFastJoin()) + require.Equal(t, req.GetSessionId(), attach.JoinRequest.GetSessionId()) + + go writeSamplesUntilDone(ctx, audio) + go writeSamplesUntilDone(ctx, m.aliceAudio) + var trace jointrace.Trace + select { + case trace = <-traces: + case <-time.After(iceTimeout): + t.Fatalf("no media both ways; recorded so far:\n%s", call.JoinTrace()) + } + t.Logf("\n%s", trace) + require.True(t, gotAlice.Load(), "alice's audio reached OnTrack") + + for _, rpc := range rpcsOf[any](m.fake) { + switch rpc.(type) { + case *sfu_models.ICETrickle: + t.Fatal("the fast join trickles no candidates") + case *signal_rpc.SetPublisherRequest: + t.Fatal("the FastJoin answered the publisher") + } + } + + for _, name := range []string{ + jointrace.CoordFastJoin, jointrace.CoordFastJoin + jointrace.DetailServer, + jointrace.PCsCreate, jointrace.SFUFastJoin, jointrace.SFUFastJoin + jointrace.DetailServer, + jointrace.SFUWSDial, jointrace.SFUWS, jointrace.PubSFUCandidates, jointrace.SubSFUCandidates, + jointrace.SubAnswer, jointrace.SubSendAnswer, + jointrace.PubICE, jointrace.PubDTLS, jointrace.PubRTP, jointrace.SubICE, jointrace.SubDTLS, jointrace.SubRTP, + } { + _, ok := trace.Span(name) + require.True(t, ok, "no %s span", name) + } + for _, name := range []string{ + jointrace.CoordJoin, jointrace.SFUJoin, jointrace.PubSetPublisher, jointrace.PubTrickleOut, + jointrace.PubDebounce, jointrace.SubDebounce, jointrace.SubOffer, + } { + _, ok := trace.Span(name) + require.False(t, ok, "a %s span on the fast join", name) + } + pcs, _ := trace.Span(jointrace.PCsCreate) + require.Empty(t, pcs.After, "pcs.create runs alongside coord.fastjoin") + fast, _ := trace.Span(jointrace.SFUFastJoin) + require.ElementsMatch(t, []string{jointrace.CoordFastJoin, jointrace.PCsCreate}, fast.After) + server, _ := trace.Span(jointrace.SFUFastJoin + jointrace.DetailServer) + require.Contains(t, server.Note, "total=2") + for _, s := range trace.Spans { + require.NotContains(t, s.After, jointrace.SubSendAnswer, "%s waits for SendAnswer", s.Name) + } + names := trace.CriticalPath().Names() + require.NotContains(t, names, jointrace.SFUWSDial, "the websocket dial runs alongside the FastJoin") + require.NotContains(t, names, jointrace.SubSendAnswer) +} + +// TestFastJoinWithNetworkDelay joins over a simulated 100 ms network with nothing warm: +// the signalling steps are whole round trips, the dial and the FastJoin overlap, and +// the attach costs one round trip after both. +func TestFastJoinWithNetworkDelay(t *testing.T) { + t.Parallel() + + const rtt = 100 * time.Millisecond + m := newMediaSFU(t) + f := newFakeCoordinator(t, 0, false) + f.serveFastJoin(m.fake) + call := fastJoinCall(t, f, "fast-join-delay", WithNetworkDelay(rtt)) + traces := make(chan jointrace.Trace, 1) + call.OnJoinTrace(func(tr jointrace.Trace) { traces <- tr }) + + ctx, cancel := context.WithCancel(context.Background()) + t.Cleanup(cancel) + audio, err := webrtc.NewTrackLocalStaticSample( + webrtc.RTPCodecCapability{MimeType: webrtc.MimeTypeOpus}, "audio", "prefix:TRACK_TYPE_AUDIO") + require.NoError(t, err) + go writeSamplesUntilDone(ctx, audio) + go writeSamplesUntilDone(ctx, m.aliceAudio) + require.NoError(t, joinFast(t, call, + WithTrack(&sfu_models.TrackInfo{TrackId: "published-audio", TrackType: sfu_models.TrackType_TRACK_TYPE_AUDIO}, audio), + WithOnTrack(SubscriberFunc(func(OnTrackReceived) {})), + WithPublisherPeerConfiguration(loopbackPeerConfig()), + WithSubscriberPeerConfiguration(loopbackPeerConfig()))) + var trace jointrace.Trace + select { + case trace = <-traces: + case <-time.After(iceTimeout): + t.Fatalf("no media both ways; recorded so far:\n%s", call.JoinTrace()) + } + t.Logf("\n%s", trace) + + span := func(name string) jointrace.Span { + t.Helper() + s, ok := trace.Span(name) + require.True(t, ok, "no %s span", name) + return s + } + within := func(name string, want float64, got time.Duration) { + t.Helper() + exact := time.Duration(want * float64(rtt)) + require.InDelta(t, float64(exact), float64(got), float64(max(rtt, exact))/10, + "%s: want %.1f RTT (%s), got %s", name, want, exact, got) + } + // Each a new connection: TCP, then the request. + within(jointrace.CoordFastJoin, 2, span(jointrace.CoordFastJoin).Duration()) + within(jointrace.SFUFastJoin, 2, span(jointrace.SFUFastJoin).Duration()) + within(jointrace.SFUWSDial, 2, span(jointrace.SFUWSDial).Duration()) + within(jointrace.SFUWS, 1, span(jointrace.SFUWS).Duration()) + fast, dial := span(jointrace.SFUFastJoin), span(jointrace.SFUWSDial) + within("dial alongside FastJoin", 0, dial.Start.Sub(fast.Start)) + within("candidates after the attach", 5, span(jointrace.SubSFUCandidates).End.Sub(trace.JoinAt)) + require.Less(t, span(jointrace.PCsCreate).Duration(), rtt/2, "local, under coord.fastjoin") +} + +// TestFastJoinTriesTheNextCandidate: a candidate that refuses the client, or cannot be +// reached, is followed by the next, and gets no websocket attach. +func TestFastJoinTriesTheNextCandidate(t *testing.T) { + t.Parallel() + + for name, first := range map[string]func(t *testing.T) *testutil.FakeSFU{ + "full": func(*testing.T) *testutil.FakeSFU { + return testutil.NewFakeSFU(testutil.WithSignalRPC(testutil.SignalRPC{ + FastJoin: sfuError(sfu_models.ErrorCode_ERROR_CODE_SFU_FULL, "sfu is full"), + })) + }, + "shutting down": func(*testing.T) *testutil.FakeSFU { + return testutil.NewFakeSFU(testutil.WithSignalRPC(testutil.SignalRPC{ + FastJoin: sfuError(sfu_models.ErrorCode_ERROR_CODE_SFU_SHUTTING_DOWN, "sfu is shutting down"), + })) + }, + "unauthenticated": func(*testing.T) *testutil.FakeSFU { + return testutil.NewFakeSFU(testutil.WithSignalRPC(testutil.SignalRPC{ + FastJoin: sfuError(sfu_models.ErrorCode_ERROR_CODE_UNAUTHENTICATED, "grant for another sfu"), + })) + }, + "unreachable": func(*testing.T) *testutil.FakeSFU { + sfu := testutil.NewFakeSFU() + sfu.Close() + return sfu + }, + } { + t.Run(name, func(t *testing.T) { + t.Parallel() + + sfu1, sfu2 := first(t), testutil.NewFakeSFU() + t.Cleanup(sfu1.Close) + t.Cleanup(sfu2.Close) + f := newFakeCoordinator(t, 0, false) + f.serveFastJoin(sfu1, sfu2) + call := fastJoinCall(t, f, "next-candidate") + require.NoError(t, joinFast(t, call)) + require.Equal(t, JoinFlowFast, call.JoinFlow()) + + req, err := testutil.NextRPCRequest[*signal_rpc.FastJoinRequest](sfu2, time.Second) + require.NoError(t, err) + require.Equal(t, "grant-2", req.GetSetupGrant(), "with its own grant") + attach, err := testutil.NextRequestOf[*sfu_events.SfuRequest_JoinRequest](sfu2, 5*time.Second) + require.NoError(t, err) + require.True(t, attach.JoinRequest.GetAttachFastJoin()) + _, err = testutil.NextRequestOf[*sfu_events.SfuRequest_JoinRequest](sfu1, 200*time.Millisecond) + require.Error(t, err, "the refusing candidate got an attach") + + require.Eventually(t, func() bool { + _, ok := call.JoinTrace().Span(jointrace.SFUWS) + return ok + }, 5*time.Second, 10*time.Millisecond) + trace := call.JoinTrace() + fast, _ := trace.Span(jointrace.SFUFastJoin) + require.Equal(t, "candidate 2 of 2", fast.Note) + dial, ok := trace.Span(jointrace.SFUWSDial) + require.True(t, ok) + ws, _ := trace.Span(jointrace.SFUWS) + require.False(t, ws.Start.Before(dial.End), "the dial is the attached websocket's") + require.Empty(t, f.joins) + }) + } +} + +// TestFastJoinAsksForNewCandidatesOnce: when every candidate refuses the grant, the +// grants may have expired, and fast_join is asked once more before Join gives up. +func TestFastJoinAsksForNewCandidatesOnce(t *testing.T) { + t.Parallel() + + t.Run("second round joins", func(t *testing.T) { + t.Parallel() + + var calls atomic.Int32 + sfu := testutil.NewFakeSFU(testutil.WithSignalRPC(testutil.SignalRPC{ + FastJoin: func(ctx context.Context, req *signal_rpc.FastJoinRequest) (*signal_rpc.FastJoinResponse, error) { + if calls.Add(1) == 1 { + return sfuError(sfu_models.ErrorCode_ERROR_CODE_UNAUTHENTICATED, "grant expired")(ctx, req) + } + return &signal_rpc.FastJoinResponse{}, nil + }, + })) + t.Cleanup(sfu.Close) + f := newFakeCoordinator(t, 0, false) + f.serveFastJoin(sfu) + call := fastJoinCall(t, f, "second-round") + require.NoError(t, joinFast(t, call)) + require.Equal(t, JoinFlowFast, call.JoinFlow()) + require.Len(t, f.fastJoins, 2) + require.EqualValues(t, 2, calls.Load()) + }) + + t.Run("gives up", func(t *testing.T) { + t.Parallel() + + sfu := testutil.NewFakeSFU(testutil.WithSignalRPC(testutil.SignalRPC{ + FastJoin: sfuError(sfu_models.ErrorCode_ERROR_CODE_UNAUTHENTICATED, "bad token"), + })) + t.Cleanup(sfu.Close) + f := newFakeCoordinator(t, 0, false) + f.serveFastJoin(sfu) + call := fastJoinCall(t, f, "gives-up") + err := joinFast(t, call) + require.ErrorContains(t, err, "bad token") + require.Len(t, f.fastJoins, fastJoinRounds) + require.Empty(t, f.joins, "a refused grant is not a reason for the legacy join") + }) +} + +// TestFastJoinRefusedByTheCall: a full call is full on every SFU. +func TestFastJoinRefusedByTheCall(t *testing.T) { + t.Parallel() + + sfu1 := testutil.NewFakeSFU(testutil.WithSignalRPC(testutil.SignalRPC{ + FastJoin: sfuError(sfu_models.ErrorCode_ERROR_CODE_CALL_PARTICIPANT_LIMIT_REACHED, "call is full"), + })) + sfu2 := testutil.NewFakeSFU() + t.Cleanup(sfu1.Close) + t.Cleanup(sfu2.Close) + f := newFakeCoordinator(t, 0, false) + f.serveFastJoin(sfu1, sfu2) + call := fastJoinCall(t, f, "call-full") + require.ErrorContains(t, joinFast(t, call), "call is full") + require.Empty(t, rpcsOf[*signal_rpc.FastJoinRequest](sfu2)) + require.Empty(t, f.joins) + require.Len(t, f.fastJoins, 1) +} + +// TestFastJoinFallsBackToTheLegacyJoin: where the deployment cannot fast join, Join +// takes the legacy flow from a call that was never joined. +func TestFastJoinFallsBackToTheLegacyJoin(t *testing.T) { + t.Parallel() + + for name, candidates := range map[string]func(t *testing.T) []*testutil.FakeSFU{ + "coordinator without fast_join": nil, + "SFUs without FastJoin": func(t *testing.T) []*testutil.FakeSFU { + sfus := []*testutil.FakeSFU{testutil.NewFakeSFU(testutil.WithoutFastJoin()), testutil.NewFakeSFU(testutil.WithoutFastJoin())} + for _, sfu := range sfus { + t.Cleanup(sfu.Close) + } + return sfus + }, + "SFU that cannot create the call": func(t *testing.T) []*testutil.FakeSFU { + sfu := testutil.NewFakeSFU(testutil.WithSignalRPC(testutil.SignalRPC{ + FastJoin: sfuError(sfu_models.ErrorCode_ERROR_CODE_INTERNAL_SERVER_ERROR, + "this server cannot create the call; join through the coordinator"), + })) + t.Cleanup(sfu.Close) + return []*testutil.FakeSFU{sfu} + }, + } { + t.Run(name, func(t *testing.T) { + t.Parallel() + + f := newFakeCoordinator(t, 0, false) + if candidates != nil { + f.serveFastJoin(candidates(t)...) + } + call := fastJoinCall(t, f, "fallback") + require.NoError(t, joinFast(t, call)) + require.Equal(t, JoinFlowLegacy, call.JoinFlow()) + require.Len(t, f.joins, 1) + + join, err := testutil.NextRequestOf[*sfu_events.SfuRequest_JoinRequest](f.sfu, time.Second) + require.NoError(t, err) + require.False(t, join.JoinRequest.GetAttachFastJoin()) + trace := call.JoinTrace() + for _, name := range []string{jointrace.CoordJoin, jointrace.PCsCreate, jointrace.SFUWSDial, jointrace.SFUJoin} { + _, ok := trace.Span(name) + require.True(t, ok, "no %s span", name) + } + _, ok := trace.Span(jointrace.SFUFastJoin) + require.False(t, ok) + dial, _ := trace.Span(jointrace.SFUWSDial) + require.Equal(t, []string{jointrace.PCsCreate}, dial.After, "the legacy join's dial") + pcs, _ := trace.Span(jointrace.PCsCreate) + coord, _ := trace.Span(jointrace.CoordJoin) + require.False(t, pcs.Start.Before(coord.End), "the legacy join's pcs.create, not the abandoned fast one") + }) + } +} + +// TestFastJoinDoesNotWaitForTheWebsocket: an attach answered late holds up neither Join +// nor anything it set up, and Leave waits for it to say goodbye on it. +func TestFastJoinDoesNotWaitForTheWebsocket(t *testing.T) { + t.Parallel() + + const attachDelay = time.Second + sfu := testutil.NewFakeSFU(testutil.WithJoinHandler(func(*sfu_events.JoinRequest) *sfu_events.SfuEvent { + time.Sleep(attachDelay) + return &sfu_events.SfuEvent{EventPayload: &sfu_events.SfuEvent_JoinResponse{JoinResponse: &sfu_events.JoinResponse{}}} + })) + t.Cleanup(sfu.Close) + f := newFakeCoordinator(t, 0, false) + f.serveFastJoin(sfu) + call := fastJoinCall(t, f, "late-attach") + + started := time.Now() + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cancel() + _, err := call.Join(ctx) + require.NoError(t, err) + require.Less(t, time.Since(started), attachDelay, "Join waited for the attach") + _, attached := call.JoinTrace().Span(jointrace.SFUWS) + require.False(t, attached) + + require.NoError(t, call.Leave("test over")) + require.GreaterOrEqual(t, time.Since(started), attachDelay, "Leave waits for the attach") + leave, err := testutil.NextRequestOf[*sfu_events.SfuRequest_LeaveCallRequest](sfu, time.Second) + if err == nil { + require.Equal(t, "test over", leave.LeaveCallRequest.GetReason()) + } + _, attached = call.JoinTrace().Span(jointrace.SFUWS) + require.True(t, attached) +} + +// TestFastJoinUnknownUser: the coordinator knows a user only once its websocket is up, +// on fast_join as on join. +func TestFastJoinUnknownUser(t *testing.T) { + t.Parallel() + + const wsDelay = 200 * time.Millisecond + sfu := testutil.NewFakeSFU() + t.Cleanup(sfu.Close) + f := newFakeCoordinator(t, wsDelay, true) + f.serveFastJoin(sfu) + started := time.Now() + call := fastJoinCall(t, f, "fast-unknown-user") + require.NoError(t, joinFast(t, call)) + require.GreaterOrEqual(t, time.Since(started), wsDelay) + require.Equal(t, JoinFlowFast, call.JoinFlow(), "an unknown user is not a coordinator without fast_join") + require.Len(t, f.fastJoins, 2, "refused, then joined once the websocket is up") +} diff --git a/go.mod b/go.mod index 7a3b485..9c4a238 100644 --- a/go.mod +++ b/go.mod @@ -4,7 +4,7 @@ go 1.27.0 require ( github.com/GetStream/getstream-go/v5 v5.2.0 - github.com/GetStream/protocol v1.49.0 + github.com/GetStream/protocol v1.50.0-3rtt.1 github.com/gammazero/deque v1.1.0 github.com/gobwas/ws v1.4.0 github.com/golang-jwt/jwt/v5 v5.3.0 diff --git a/go.sum b/go.sum index 43c8137..e46c594 100644 --- a/go.sum +++ b/go.sum @@ -2,6 +2,8 @@ github.com/GetStream/getstream-go/v5 v5.2.0 h1:nrZkIL0mzHtByMeAeNFD4jpGIT5XPa0P3 github.com/GetStream/getstream-go/v5 v5.2.0/go.mod h1:qnod4MVOSwxfpqtXLpxjbMo7lDcepOlg8X2GrpBs/cw= github.com/GetStream/protocol v1.49.0 h1:AzTRH6+PRNHjFjBlcymbjZsVE+wZBp/wPG6w+/9FCw4= github.com/GetStream/protocol v1.49.0/go.mod h1:jZALur2IoWnpcLnFzLtnFomgItt4SBHj+U1TjPYC114= +github.com/GetStream/protocol v1.50.0-3rtt.1 h1:kqsE1ZBf3gpO47lc2JLkyNVhVR9JNGFRYUtn83nLgIo= +github.com/GetStream/protocol v1.50.0-3rtt.1/go.mod h1:jZALur2IoWnpcLnFzLtnFomgItt4SBHj+U1TjPYC114= github.com/RaveNoX/go-jsoncommentstrip v1.0.0/go.mod h1:78ihd09MekBnJnxpICcwzCMzGrKSKYe4AqU6PDYYpjk= github.com/apapsch/go-jsonmerge/v2 v2.0.0 h1:axGnT1gRIfimI7gJifB699GoE/oq+F2MU7Dml6nw9rQ= github.com/apapsch/go-jsonmerge/v2 v2.0.0/go.mod h1:lvDnEdqiQrp0O42VQGgmlKpxL1AP2+08jFMw88y4klk= diff --git a/internal/testutil/fakesfu.go b/internal/testutil/fakesfu.go index b072f85..3486da4 100644 --- a/internal/testutil/fakesfu.go +++ b/internal/testutil/fakesfu.go @@ -56,10 +56,16 @@ type FakeSFU struct { mu sync.Mutex conn *websocket.Connection[sfu_events.SfuRequest, sfu_events.SfuEvent] - - connected chan struct{} - once sync.Once - tls bool + // fastJoined are the sessions a FastJoin created, which a websocket may attach to. + fastJoined map[string]bool + fastJoinSeen bool + noFastJoin bool + + connected chan struct{} + once sync.Once + attached chan struct{} + attachedOnce sync.Once + tls bool } // FakeSFUOption configures a FakeSFU. Options are applied before the server @@ -90,6 +96,14 @@ func WithTLS() FakeSFUOption { } } +// WithoutFastJoin serves no FastJoinServer, as an SFU from before the fast join: FastJoin +// requests get twirp's bad_route. +func WithoutFastJoin() FakeSFUOption { + return func(f *FakeSFU) { + f.noFastJoin = true + } +} + // WithSignalRPC overrides the twirp signalling RPC answers. func WithSignalRPC(rpc SignalRPC) FakeSFUOption { return func(f *FakeSFU) { @@ -111,6 +125,10 @@ type SignalRPC struct { SendMetrics func(context.Context, *sfu_signal_rpc.SendMetricsRequest) (*sfu_signal_rpc.SendMetricsResponse, error) StartNoiseCancellation func(context.Context, *sfu_signal_rpc.StartNoiseCancellationRequest) (*sfu_signal_rpc.StartNoiseCancellationResponse, error) StopNoiseCancellation func(context.Context, *sfu_signal_rpc.StopNoiseCancellationRequest) (*sfu_signal_rpc.StopNoiseCancellationResponse, error) + + // FastJoin answers the FastJoinServer's one RPC. By default it accepts with the + // JoinResponse's call state and no SDPs: nothing to answer and nothing to offer. + FastJoin func(context.Context, *sfu_signal_rpc.FastJoinRequest) (*sfu_signal_rpc.FastJoinResponse, error) } // NewFakeSFU starts a fake SFU. Close it when the test is done. @@ -123,6 +141,8 @@ func NewFakeSFU(opts ...FakeSFUOption) *FakeSFU { Authorizations: make(chan string, 64), joinResponse: &sfu_events.JoinResponse{}, connected: make(chan struct{}), + fastJoined: map[string]bool{}, + attached: make(chan struct{}), } for _, opt := range opts { opt(f) @@ -135,6 +155,10 @@ func NewFakeSFU(opts ...FakeSFUOption) *FakeSFU { mux.HandleFunc(wsPath, f.serve) mux.Handle("/", f.recordAuthorization(sfu_signal_rpc.NewSignalServerServer( &signalRPCService{f: f}, twirp.WithServerPathPrefix("")))) + if !f.noFastJoin { + fastJoin := sfu_signal_rpc.NewFastJoinServerServer(&fastJoinService{f: f}, twirp.WithServerPathPrefix("")) + mux.Handle(fastJoin.PathPrefix(), f.recordAuthorization(fastJoin)) + } f.srv = httptest.NewUnstartedServer(mux) if f.tls { f.srv.StartTLS() @@ -194,12 +218,26 @@ func (f *FakeSFU) CloseConnection() error { // It writes to the most recent connection. After a reconnect, wait for that // connection's JoinRequest before sending: the fake records a request only once // the connection carrying it is installed. +// +// Once a FastJoin has started, Send also waits for a websocket to attach to it, as the +// SFU sends a fast-joined participant nothing before. func (f *FakeSFU) Send(event *sfu_events.SfuEvent, timeout time.Duration) error { + deadline := time.After(timeout) select { case <-f.connected: - case <-time.After(timeout): + case <-deadline: return xerr.Error("no client connected to the fake sfu") } + f.mu.Lock() + fastJoining := f.fastJoinSeen + f.mu.Unlock() + if fastJoining { + select { + case <-f.attached: + case <-deadline: + return xerr.Error("no websocket attached to the fast join") + } + } f.mu.Lock() conn := f.conn @@ -307,6 +345,9 @@ func (f *FakeSFU) serve(w http.ResponseWriter, r *http.Request) { continue } err = conn.Write(answer) + if err == nil && payload.JoinRequest.GetAttachFastJoin() && answer.GetJoinResponse() != nil { + f.attachedOnce.Do(func() { close(f.attached) }) + } case *sfu_events.SfuRequest_HealthCheckRequest: err = conn.Write(&sfu_events.SfuEvent{ EventPayload: &sfu_events.SfuEvent_HealthCheckResponse{ @@ -326,6 +367,21 @@ func (f *FakeSFU) joinAnswer(req *sfu_events.JoinRequest) *sfu_events.SfuEvent { if f.onJoinRequest != nil { return f.onJoinRequest(req) } + if req.GetAttachFastJoin() { + f.mu.Lock() + joined := f.fastJoined[req.GetSessionId()] + f.mu.Unlock() + if !joined { + // As the SFU answers an attach with no FastJoin before it. + return &sfu_events.SfuEvent{EventPayload: &sfu_events.SfuEvent_Error{Error: &sfu_events.Error{ + Error: &sfu_models.Error{ + Code: sfu_models.ErrorCode_ERROR_CODE_PARTICIPANT_NOT_FOUND, + Message: "participant not found", + }, + ReconnectStrategy: sfu_models.WebsocketReconnectStrategy_WEBSOCKET_RECONNECT_STRATEGY_REJOIN, + }}} + } + } return &sfu_events.SfuEvent{ EventPayload: &sfu_events.SfuEvent_JoinResponse{JoinResponse: f.joinResponse}, } @@ -348,6 +404,32 @@ func (f *FakeSFU) recordRPC(req any) { } } +// fastJoinService serves the twirp FastJoinServer. A FastJoin that succeeds is what a +// websocket attach needs. +type fastJoinService struct { + f *FakeSFU +} + +var _ sfu_signal_rpc.FastJoinServer = (*fastJoinService)(nil) + +func (s *fastJoinService) FastJoin(ctx context.Context, req *sfu_signal_rpc.FastJoinRequest) (*sfu_signal_rpc.FastJoinResponse, error) { + s.f.recordRPC(req) + s.f.mu.Lock() + s.f.fastJoinSeen = true + s.f.mu.Unlock() + resp := &sfu_signal_rpc.FastJoinResponse{CallState: s.f.joinResponse.GetCallState()} + var err error + if h := s.f.rpc.FastJoin; h != nil { + resp, err = h(ctx, req) + } + if err == nil && resp.GetError() == nil { + s.f.mu.Lock() + s.f.fastJoined[req.GetSessionId()] = true + s.f.mu.Unlock() + } + return resp, err +} + // signalRPCService serves the twirp SignalServer, recording every request and // delegating to the test's overrides. type signalRPCService struct { diff --git a/join_trace.go b/join_trace.go index 02055f8..5e11d07 100644 --- a/join_trace.go +++ b/join_trace.go @@ -23,6 +23,10 @@ type joinTracer struct { timer *time.Timer handler func(jointrace.Trace) + // fast is set when the first join takes the fast path, whose steps depend on each + // other differently. + fast bool + // Moments the spans are built from that are not spans of their own. pubSignalSent time.Time subOfferAt time.Time @@ -62,6 +66,18 @@ func (j *joinTracer) markJoined() { j.joined = true } +func (j *joinTracer) setFast(fast bool) { + j.mu.Lock() + defer j.mu.Unlock() + j.fast = fast +} + +func (j *joinTracer) isFast() bool { + j.mu.Lock() + defer j.mu.Unlock() + return j.fast +} + func (j *joinTracer) snapshot() jointrace.Trace { j.mu.Lock() rec, sample := j.rec, j.udpRTT @@ -118,6 +134,10 @@ func (c *Call) peerSpans(publisher bool, t pc.Timing) { if rec == nil { return } + if c.trace.isFast() { + fastPeerSpans(rec, publisher, t) + return + } if publisher { rec.Add(jointrace.Span{ Name: jointrace.PubDebounce, After: []string{jointrace.SFUJoin}, @@ -157,6 +177,49 @@ func (c *Call) peerSpans(publisher bool, t pc.Timing) { }) } +// fastPeerSpans is peerSpans for a fast join. The offer and answer went with the +// FastJoin, but the SFU's candidates still come on the websocket, so ICE waits for both. +func fastPeerSpans(rec *jointrace.Recorder, publisher bool, t pc.Timing) { + candidates, signalled := jointrace.SubSFUCandidates, jointrace.SubAnswer + ice, dtls := jointrace.SubICE, jointrace.SubDTLS + if publisher { + candidates, signalled = jointrace.PubSFUCandidates, jointrace.SFUFastJoin + ice, dtls = jointrace.PubICE, jointrace.PubDTLS + } + if attached, ok := rec.Get(jointrace.SFUWS); ok { + rec.Add(jointrace.Span{ + Name: candidates, After: []string{jointrace.SFUWS}, + Start: attached.End, End: t.FirstRemoteCandidate, + Kind: jointrace.KindNet, Peer: jointrace.PeerSFU, + Note: "trickled on the websocket once it attaches", + }) + } + rec.Add(jointrace.Span{ + Name: ice, After: []string{signalled, candidates}, + Start: t.ICEChecking, End: t.ICEConnected, + Kind: jointrace.KindNet, Peer: jointrace.PeerUDP, + }) + rec.Add(jointrace.Span{ + Name: dtls, After: []string{ice}, + Start: t.ICEConnected, End: t.DTLSConnected, + Kind: jointrace.KindNet, Peer: jointrace.PeerUDP, + }) +} + +// subscriberAnswered records the client's side of the fast join's subscriber offer: +// from the FastJoin response that carried it to the answer being ready to send. +func (c *Call) subscriberAnswered(at time.Time) { + rec := c.trace.recorder() + c.trace.mu.Lock() + offerAt := c.trace.subOfferAt + c.trace.mu.Unlock() + rec.Add(jointrace.Span{ + Name: jointrace.SubAnswer, After: []string{jointrace.SFUFastJoin}, + Start: offerAt, End: at, Kind: jointrace.KindLocal, Peer: jointrace.PeerLocal, + Note: "sent without waiting", + }) +} + // firstRTP records the first RTP packet sent by the publisher or received by the // subscriber. func (c *Call) firstRTP(publisher bool, dtlsConnected, at time.Time) { diff --git a/join_trace_test.go b/join_trace_test.go index 84caba8..cb88d64 100644 --- a/join_trace_test.go +++ b/join_trace_test.go @@ -31,6 +31,8 @@ type mediaSFU struct { fake *testutil.FakeSFU pub, sub *sfuWebRTCPeer aliceAudio *webrtc.TrackLocalStaticSample + // holdAnswers, when set, holds every SendAnswer until it is closed. + holdAnswers atomic.Pointer[chan struct{}] } func newMediaSFU(t *testing.T, opts ...testutil.FakeSFUOption) *mediaSFU { @@ -38,16 +40,37 @@ func newMediaSFU(t *testing.T, opts ...testutil.FakeSFUOption) *mediaSFU { m := &mediaSFU{} var pub, sub atomic.Pointer[sfuWebRTCPeer] + callState := &sfu_models.CallState{ + Participants: []*sfu_models.Participant{ + {UserId: "alice", SessionId: "session-a", TrackLookupPrefix: "prefix-a"}, + }, + ParticipantCount: &sfu_models.ParticipantCount{Total: 2}, + } opts = append([]testutil.FakeSFUOption{ - testutil.WithJoinResponse(&sfu_events.JoinResponse{ - CallState: &sfu_models.CallState{ - Participants: []*sfu_models.Participant{ - {UserId: "alice", SessionId: "session-a", TrackLookupPrefix: "prefix-a"}, - }, - ParticipantCount: &sfu_models.ParticipantCount{Total: 2}, - }, - }), + testutil.WithJoinResponse(&sfu_events.JoinResponse{CallState: callState}), testutil.WithSignalRPC(testutil.SignalRPC{ + // As the SFU's: the publisher answered, and the subscriber offered alice's + // audio, which a fast-joined participant is subscribed to. + FastJoin: func(_ context.Context, req *signal_rpc.FastJoinRequest) (*signal_rpc.FastJoinResponse, error) { + resp := &signal_rpc.FastJoinResponse{ + CallState: callState, + SubscriberNegotiationId: 1, + ServerTimings: []*signal_rpc.ServerTiming{{Name: "total", DurationMs: 2}}, + } + if req.GetPublisherSdp() != "" { + answer, err := pub.Load().Answer(req.GetPublisherSdp()) + if err != nil { + return nil, err + } + resp.PublisherSdp = answer + } + offer, err := sub.Load().Offer() + if err != nil { + return nil, err + } + resp.SubscriberSdp = offer + return resp, nil + }, SetPublisher: func(_ context.Context, req *signal_rpc.SetPublisherRequest) (*signal_rpc.SetPublisherResponse, error) { answer, err := pub.Load().Answer(req.GetSdp()) if err != nil { @@ -55,7 +78,14 @@ func newMediaSFU(t *testing.T, opts ...testutil.FakeSFUOption) *mediaSFU { } return &signal_rpc.SetPublisherResponse{Sdp: answer}, nil }, - SendAnswer: func(_ context.Context, req *signal_rpc.SendAnswerRequest) (*signal_rpc.SendAnswerResponse, error) { + SendAnswer: func(ctx context.Context, req *signal_rpc.SendAnswerRequest) (*signal_rpc.SendAnswerResponse, error) { + if hold := m.holdAnswers.Load(); hold != nil { + select { + case <-*hold: + case <-ctx.Done(): + return nil, ctx.Err() + } + } if err := sub.Load().AcceptAnswer(req.GetSdp()); err != nil { return nil, err } @@ -297,7 +327,7 @@ func TestJoinTraceOfASecondJoinOnTheSameClient(t *testing.T) { call.onceConnect.Do(func() {}) ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) defer cancel() - _, err := call.Join(ctx) + _, err := call.Join(ctx, WithJoinFlow(JoinFlowLegacy)) require.NoError(t, err) trace := call.JoinTrace() require.NoError(t, call.Leave("test over")) diff --git a/jointrace/httptrace.go b/jointrace/httptrace.go index 44af3a7..9e67d47 100644 --- a/jointrace/httptrace.go +++ b/jointrace/httptrace.go @@ -169,14 +169,23 @@ func (s *step) clientTrace() *httptrace.ClientTrace { // Server-Timing response header, as a detail span centred in the wait for the first // byte. It uses the "total" metric when the header has one, otherwise the longest. func ServerTiming(ctx context.Context, header string) { - s, _ := ctx.Value(stepKey{}).(*step) - if s == nil || header == "" { + if header == "" { return } dur, ok := parseServerTiming(header) if !ok { return } + ServerDuration(ctx, dur, "server-timing: "+header) +} + +// ServerDuration records dur, the server's own time for the context's step as the +// server reported it in the response body, as ServerTiming does from a header. +func ServerDuration(ctx context.Context, dur time.Duration, note string) { + s, _ := ctx.Value(stepKey{}).(*step) + if s == nil || dur <= 0 { + return + } s.mu.Lock() wrote, first := s.wroteAt, s.firstByte s.mu.Unlock() @@ -188,7 +197,7 @@ func ServerTiming(ctx context.Context, header string) { dur = wait } start := wrote.Add((wait - dur) / 2) - s.detail(DetailServer, start, start.Add(dur), KindLocal, "server-timing") + s.detail(DetailServer, start, start.Add(dur), KindLocal, note) } func parseServerTiming(header string) (time.Duration, bool) { diff --git a/jointrace/recorder.go b/jointrace/recorder.go index 90a0e95..0bb5ecb 100644 --- a/jointrace/recorder.go +++ b/jointrace/recorder.go @@ -75,6 +75,62 @@ func (r *Recorder) Extend(s Span) { r.addLocked(s) } +// Remove drops the named spans and their detail spans, for steps that are redone from +// scratch and must be recorded again. +func (r *Recorder) Remove(names ...string) { + if r == nil { + return + } + r.mu.Lock() + defer r.mu.Unlock() + if r.sealed { + return + } + drop := make(map[string]bool, len(names)) + for _, name := range names { + drop[name] = true + } + kept := r.order[:0] + for _, name := range r.order { + if s := r.spans[name]; drop[name] || drop[s.Parent] { + delete(r.spans, name) + continue + } + kept = append(kept, name) + } + r.order = kept +} + +// Scratch returns an empty recorder with r's origin, for an attempt whose spans belong +// in r only if it works out: Merge them then. It is nil when r is. +func (r *Recorder) Scratch() *Recorder { + if r == nil { + return nil + } + return NewRecorder(r.origin) +} + +// Merge records o's spans and round-trip times into r, as Add and SetRTT would. +func (r *Recorder) Merge(o *Recorder) { + if r == nil || o == nil || r == o { + return + } + t := o.Trace() + r.mu.Lock() + defer r.mu.Unlock() + for _, s := range t.Spans { + r.addLocked(s) + } + if r.sealed { + return + } + for p, d := range t.RTT { + if _, ok := r.rtt[p]; !ok { + r.rtt[p] = d + } + } +} + // Has reports whether a span of that name was recorded. func (r *Recorder) Has(name string) bool { if r == nil { diff --git a/jointrace/span.go b/jointrace/span.go index 1385c0b..c955074 100644 --- a/jointrace/span.go +++ b/jointrace/span.go @@ -57,14 +57,18 @@ const ( SubRTP = "sub.rtp" ) -// The steps of the fast join path. +// The steps of the fast join path. It shares pcs.create, sfu.ws.dial, pub.sfu.candidates, +// the ICE, DTLS and RTP steps and sub.sendanswer with the legacy path; until the SFU puts +// its candidates in the SDPs they arrive on the websocket, which is why sfu.ws and the +// *.sfu.candidates steps are still there. const ( - CoordFastJoin = "coord.fastjoin" - SFUFastJoin = "sfu.fastjoin" - SFUWS = "sfu.ws" - SubAnswer = "sub.answer" - PubICEDTLS = "pub.ice+dtls" - SubICEDTLS = "sub.ice+dtls" + CoordFastJoin = "coord.fastjoin" + SFUFastJoin = "sfu.fastjoin" + SFUWS = "sfu.ws" + SubAnswer = "sub.answer" + SubSFUCandidates = "sub.sfu.candidates" + PubICEDTLS = "pub.ice+dtls" + SubICEDTLS = "sub.ice+dtls" ) // Suffixes of the detail spans a network step is split into. A detail span is named diff --git a/jointrace/trace_test.go b/jointrace/trace_test.go index 8d99f4f..9227c1c 100644 --- a/jointrace/trace_test.go +++ b/jointrace/trace_test.go @@ -116,6 +116,36 @@ func TestRecorderKeepsTheFirstRecording(t *testing.T) { require.Empty(t, nilRec.Trace().Spans) } +// TestScratchSpansCountOnlyOnceMerged is a fast join's candidates: each records into its +// own scratch recorder, and only the one that took the client is merged. +func TestScratchSpansCountOnlyOnceMerged(t *testing.T) { + rec := NewRecorder(t0) + rec.Add(Span{Name: CoordFastJoin, Start: at(0), End: at(100)}) + failed, won := rec.Scratch(), rec.Scratch() + failed.Add(Span{Name: SFUWSDial, Start: at(100), End: at(150)}) + failed.SetRTT(PeerSFU, 30*time.Millisecond) + won.Add(Span{Name: SFUWSDial, Start: at(160), End: at(360)}) + won.Add(Span{Name: SFUWSDial + DetailTCP, Parent: SFUWSDial, Start: at(160), End: at(260)}) + won.SetRTT(PeerSFU, 100*time.Millisecond) + won.Add(Span{Name: CoordFastJoin, Start: at(1), End: at(2)}) + rec.Merge(won) + + tr := rec.Trace() + require.Len(t, tr.Spans, 3) + dial, _ := tr.Span(SFUWSDial) + require.Equal(t, at(160), dial.Start, "the candidate that took the client") + coord, _ := tr.Span(CoordFastJoin) + require.Equal(t, at(100), coord.End, "a merge keeps what was recorded first") + require.Equal(t, 100*time.Millisecond, tr.RTT[PeerSFU]) + + rec.Remove(SFUWSDial) + require.Len(t, rec.Trace().Spans, 1, "with its detail spans") + + var nilRec *Recorder + require.Nil(t, nilRec.Scratch()) + nilRec.Merge(won) +} + func TestReportRoundTripsAsJSON(t *testing.T) { tr := todayJoin() raw, err := json.Marshal(tr) diff --git a/location_test.go b/location_test.go index 1b50dff..7ebfa26 100644 --- a/location_test.go +++ b/location_test.go @@ -64,7 +64,7 @@ func TestJoinSendsLocationAuto(t *testing.T) { call.onceConnect.Do(func() {}) ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) defer cancel() - _, err = call.Join(ctx, tc.opts...) + _, err = call.Join(ctx, append(tc.opts, WithJoinFlow(JoinFlowLegacy))...) require.NoError(t, err) t.Cleanup(func() { _ = call.Leave("test over") }) diff --git a/pc/pc.go b/pc/pc.go index 2e8c1de..1eb53de 100644 --- a/pc/pc.go +++ b/pc/pc.go @@ -416,6 +416,12 @@ func (t *Transport) HandleRemoteDescriptionWithNegotiationID(sd webrtc.SessionDe }) } +// ReportFailure hands err, from work done off the event loop on its behalf, to the +// loop, which reports it through Handler.OnNegotiationFailed like its own failures. +func (t *Transport) ReportFailure(name string, err error) { + t.enqueue(name, func() error { return err }) +} + // AddICECandidate adds a candidate trickled by the SFU. func (t *Transport) AddICECandidate(candidate webrtc.ICECandidateInit) { t.mu.Lock() diff --git a/publisher.go b/publisher.go index 5300f15..feb4f08 100644 --- a/publisher.go +++ b/publisher.go @@ -40,9 +40,21 @@ type publisher struct { iceRecovery *iceRecovery + // fastPath marks a publisher of a fast join: its first offer goes out with the + // FastJoin, and it sends the ICE-lite SFU no candidates. + fastPath atomic.Bool + // joinOffer, while set, takes the next offer instead of SetPublisher. + joinOffer atomic.Pointer[chan publisherOffer] + Tracing atomic.Pointer[rtcstats.TraceBuffer] } +// publisherOffer is an offer the fast join sends to the SFU itself. +type publisherOffer struct { + sdp webrtc.SessionDescription + negotiationID uint32 +} + type TrackDetails struct { Info *sfu_models.TrackInfo Tracks []webrtc.TrackLocal @@ -63,11 +75,7 @@ func newPublisher(c *Call, peerConfig pc.PeerConfig) (*publisher, error) { }), c: c, } - attempt := int64(c.reconnectAttempt.Load()) - 1 - sfuid := c.cred.Load().Server.EdgeName - if c.statsEnabled() { - pub.Tracing.Store(rtcstats.NewPubTraceBuffer("", attempt, sfuid)) - } + pub.startTracing() if peerConfig.Registry == nil { peerConfig.Registry = &interceptor.Registry{} @@ -104,15 +112,8 @@ func newPublisher(c *Call, peerConfig pc.PeerConfig) (*publisher, error) { c.firstRTP(true, dtls, at) }, nil)) - cred := c.cred.Load() if peerConfig.Config.ICEServers == nil { - for _, iceServer := range cred.IceServers { - peerConfig.Config.ICEServers = append(peerConfig.Config.ICEServers, webrtc.ICEServer{ - URLs: iceServer.Urls, - Username: iceServer.Username, - Credential: iceServer.Password, - }) - } + peerConfig.Config.ICEServers = iceServers(c.credentials().IceServers) } pub.Tracing.Load().Emit(rtcstats.PeerCreateEvent, peerConfig.Config) @@ -138,6 +139,40 @@ func newPublisher(c *Call, peerConfig pc.PeerConfig) (*publisher, error) { return pub, nil } +// startTracing gives the publisher its stats trace buffer, when the call reports stats. +func (p *publisher) startTracing() { + if p.c.statsEnabled() { + attempt := int64(p.c.reconnectAttempt.Load()) - 1 + p.Tracing.Store(rtcstats.NewPubTraceBuffer("", attempt, p.c.credentials().Server.EdgeName)) + } +} + +// joinOfferNow creates the publisher offer for a fast join, straight away rather than +// after the negotiation debounce, and returns it instead of sending it with +// SetPublisher. The SFU's answer comes back in the FastJoin response. +func (p *publisher) joinOfferNow(ctx context.Context) (publisherOffer, error) { + offers := make(chan publisherOffer, 1) + p.joinOffer.Store(&offers) + p.Negotiate(true) + select { + case offer := <-offers: + return offer, nil + case <-ctx.Done(): + p.joinOffer.Store(nil) + return publisherOffer{}, xerr.Wrapf(ctx.Err(), "create the publisher offer") + } +} + +// trackInfos are the tracks the publisher sends, as the SFU is told about them. +func (p *publisher) trackInfos() []*sfu_models.TrackInfo { + var tracks []*sfu_models.TrackInfo + p.tracks.Range(func(value *TrackDetails) bool { + tracks = append(tracks, value.Info) + return true + }) + return tracks +} + func (p *publisher) AddTrack(info *sfu_models.TrackInfo, t webrtc.TrackLocal) (*webrtc.RTPTransceiver, error) { p.mu.Lock() defer p.mu.Unlock() @@ -210,7 +245,7 @@ func (p *publisher) AddSimulcastTracks(trackInfo *sfu_models.TrackInfo, tracks . } func (p *publisher) OnICECandidateSender(c *webrtc.ICECandidate, target sfu_models.PeerType) error { - if c == nil { + if c == nil || p.fastPath.Load() { return nil } ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) @@ -277,6 +312,10 @@ func (p *publisher) OnTrack(t *webrtc.TrackRemote, _ *webrtc.RTPReceiver) { } func (p *publisher) OnOffer(sd webrtc.SessionDescription, negotiationID uint32) error { + if offers := p.joinOffer.Swap(nil); offers != nil { + *offers <- publisherOffer{sdp: sd, negotiationID: negotiationID} + return nil + } rec := p.c.trace.recorder() rec.Add(jointrace.Span{ Name: jointrace.PubOffer, After: []string{jointrace.PubDebounce}, @@ -285,17 +324,10 @@ func (p *publisher) OnOffer(sd webrtc.SessionDescription, negotiationID uint32) ctx, cancel := context.WithTimeout(context.Background(), time.Second*3) defer cancel() - var tracks []*sfu_models.TrackInfo - - p.tracks.Range(func(value *TrackDetails) bool { - tracks = append(tracks, value.Info) - return true - }) - req := &sfu_signal_rpc.SetPublisherRequest{ Sdp: sd.SDP, SessionId: p.c.SessionID.Load(), - Tracks: tracks, + Tracks: p.trackInfos(), } sent := time.Now() p.c.trace.mu.Lock() diff --git a/signal/client.go b/signal/client.go index 1c91393..bf5369c 100644 --- a/signal/client.go +++ b/signal/client.go @@ -128,7 +128,8 @@ var defaultOptions = options{ var _ sfu_signal_rpc.SignalServer = (*Client)(nil) type Client struct { - rpc atomic.Value + rpc atomic.Value + fastRPC atomic.Value options signalEventStore *event.Store[*sfu_events.SfuEvent] @@ -250,8 +251,7 @@ func NewClient(cred models.Credentials, handler Handler, opts ...Option) *Client if o.withTracing { client.Tracing.Store(rtcstats.NewClientTraceBuffer("")) } - client.setRPC(client.getSignalRPCClient(cred)) - client.cred.Store(&cred) + client.SetCredentials(cred) return client } @@ -260,21 +260,45 @@ func (c *Client) getSignalRPCClient(cred models.Credentials) sfu_signal_rpc.Sign if c.format == websocket.FormatText { clientCreator = sfu_signal_rpc.NewSignalServerJSONClient } - rpcClient := clientCreator( - cred.Server.URL, - &http.Client{ - Timeout: 5 * time.Second, - Transport: rtretry.NewRoundTripperRetryer(c.transport), - }, + return clientCreator(cred.Server.URL, c.rpcHTTPClient(), c.rpcClientOptions(cred)...) +} + +// getFastJoinRPCClient is the FastJoinServer twin of getSignalRPCClient: same base URL, +// HTTP client and authorization. +func (c *Client) getFastJoinRPCClient(cred models.Credentials) sfu_signal_rpc.FastJoinServer { + clientCreator := sfu_signal_rpc.NewFastJoinServerProtobufClient + if c.format == websocket.FormatText { + clientCreator = sfu_signal_rpc.NewFastJoinServerJSONClient + } + return clientCreator(cred.Server.URL, c.rpcHTTPClient(), c.rpcClientOptions(cred)...) +} + +func (c *Client) rpcHTTPClient() *http.Client { + return &http.Client{ + Timeout: 5 * time.Second, + Transport: rtretry.NewRoundTripperRetryer(c.transport), + } +} + +func (c *Client) rpcClientOptions(cred models.Credentials) []twirp.ClientOption { + return []twirp.ClientOption{ twirp.WithClientPathPrefix(""), twirp.WithClientInterceptors(twirpAuthInterceptor(cred.Token)), - ) - return rpcClient + } } func (c *Client) SetCredentials(cred models.Credentials) { c.cred.Store(&cred) c.setRPC(c.getSignalRPCClient(cred)) + c.fastRPC.Store(c.getFastJoinRPCClient(cred)) +} + +// FastJoin joins the SFU in one request: it creates the call from the request's setup +// grant if needed, answers the publisher offer and returns the subscriber offer. The +// websocket then attaches to the participant it created: Connect with a JoinRequest +// that has AttachFastJoin set. +func (c *Client) FastJoin(ctx context.Context, request *sfu_signal_rpc.FastJoinRequest) (*sfu_signal_rpc.FastJoinResponse, error) { + return c.fastRPC.Load().(sfu_signal_rpc.FastJoinServer).FastJoin(ctx, request) } // DialedAt is when the last websocket to the SFU finished opening, or the zero time if @@ -288,6 +312,28 @@ func (c *Client) DialedAt() time.Time { } func (c *Client) Connect(ctx context.Context, joinRequest *sfu_events.JoinRequest) (*sfu_events.JoinResponse, error) { + dialed, err := c.Dial(ctx) + if err != nil { + return nil, err + } + return c.Join(ctx, dialed, joinRequest) +} + +// Dialed is a websocket to the SFU that has not sent its JoinRequest yet. +type Dialed struct { + conn *websocket.Connection[sfu_events.SfuEvent, sfu_events.SfuRequest] + endpoint string +} + +// Close drops a websocket that will not be joined. +func (d *Dialed) Close() { + _ = d.conn.Close() +} + +// Dial opens the websocket to the SFU without joining: Join sends the JoinRequest. A +// fast join dials while its FastJoin is in flight, and attaches once the SFU has +// created the participant. +func (c *Client) Dial(ctx context.Context) (*Dialed, error) { endpoint := c.cred.Load().Server.WsEndpoint wsConn, err := wsdial.Dial(ctx, endpoint, c.dial, c.tlsConfig) if err != nil { @@ -313,6 +359,13 @@ func (c *Client) Connect(ctx context.Context, joinRequest *sfu_events.JoinReques }) conn := websocket.NewConnection[sfu_events.SfuEvent, sfu_events.SfuRequest](wsConn, true, websocket.FormatBinary, codec) + return &Dialed{conn: conn, endpoint: endpoint}, nil +} + +// Join sends joinRequest on a dialed websocket and waits for the SFU's JoinResponse. +// It takes ownership of dialed: on any error the websocket is closed. +func (c *Client) Join(ctx context.Context, dialed *Dialed, joinRequest *sfu_events.JoinRequest) (*sfu_events.JoinResponse, error) { + conn, endpoint := dialed.conn, dialed.endpoint // Every path out of here other than a JoinResponse abandons conn: without // this the websocket stays open with nothing reading it, and a first join diff --git a/subscriber.go b/subscriber.go index 72dc38d..1b1bdc9 100644 --- a/subscriber.go +++ b/subscriber.go @@ -5,6 +5,7 @@ import ( "encoding/json" "errors" "fmt" + "sync" "sync/atomic" "time" @@ -69,6 +70,14 @@ type subscriber struct { iceRecovery *iceRecovery + // fastPath marks a subscriber of a fast join: it sends the ICE-lite SFU no + // candidates, and sends its answers without holding up the peer connection. + fastPath atomic.Bool + // answerMu guards lastAnswer, the send of the previous answer, which the next one + // waits for so the SFU gets them in order. + answerMu sync.Mutex + lastAnswer chan struct{} + Tracing atomic.Pointer[rtcstats.TraceBuffer] } @@ -82,11 +91,7 @@ func newSubscriber(c *Call, s Subscriber, peerConfig pc.PeerConfig, beforeSendAn return a.SSRC < b.SSRC }), } - attempt := int64(c.reconnectAttempt.Load()) - 1 - sfuid := c.cred.Load().Server.EdgeName - if c.statsEnabled() { - sub.Tracing.Store(rtcstats.NewSubTraceBuffer("", attempt, sfuid)) - } + sub.startTracing() if peerConfig.MediaEngine == nil { peerConfig.MediaEngine = &webrtc.MediaEngine{} @@ -157,6 +162,14 @@ func newSubscriber(c *Call, s Subscriber, peerConfig pc.PeerConfig, beforeSendAn return sub, nil } +// startTracing gives the subscriber its stats trace buffer, when the call reports stats. +func (s *subscriber) startTracing() { + if s.c.statsEnabled() { + attempt := int64(s.c.reconnectAttempt.Load()) - 1 + s.Tracing.Store(rtcstats.NewSubTraceBuffer("", attempt, s.c.credentials().Server.EdgeName)) + } +} + // requestICERestart asks the SFU to restart ICE on the subscriber peer // connection. The SFU responds with a new offer over the websocket, which // OnSubscriberOffer applies. @@ -189,7 +202,7 @@ func (s *subscriber) Unbind() { } func (s *subscriber) OnICECandidateSender(c *webrtc.ICECandidate, target sfu_models.PeerType) error { - if c == nil { + if c == nil || s.fastPath.Load() { return nil } ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second) @@ -281,6 +294,35 @@ func (s *subscriber) OnAnswer(sd webrtc.SessionDescription, negotiationId uint32 SessionId: s.c.SessionID.Load(), NegotiationId: negotiationId, } + if s.fastPath.Load() { + s.c.subscriberAnswered(time.Now()) + s.sendAnswerAsync(req) + return nil + } + return s.sendAnswer(req) +} + +// sendAnswerAsync sends req without waiting for the SFU, so ICE can start on the +// candidates queued behind it. Answers still reach the SFU one at a time and in order, +// and a failed one is reported as a negotiation failure. +func (s *subscriber) sendAnswerAsync(req *signal_rpc.SendAnswerRequest) { + done := make(chan struct{}) + s.answerMu.Lock() + prev := s.lastAnswer + s.lastAnswer = done + s.answerMu.Unlock() + go func() { + defer close(done) + if prev != nil { + <-prev + } + if err := s.sendAnswer(req); err != nil { + s.ReportFailure("send answer", err) + } + }() +} + +func (s *subscriber) sendAnswer(req *signal_rpc.SendAnswerRequest) error { if s.beforeSendAnswer != nil { if err := s.beforeSendAnswer(req); err != nil { return xerr.Wrap(err) @@ -292,14 +334,22 @@ func (s *subscriber) OnAnswer(sd webrtc.SessionDescription, negotiationId uint32 if err != nil { return xerr.Wrap(err) } - s.c.trace.mu.Lock() - offerAt := s.c.trace.subOfferAt - s.c.trace.mu.Unlock() - rec.Add(jointrace.Span{ - Name: jointrace.SubSendAnswer, After: []string{jointrace.SubOffer}, - Start: offerAt, End: time.Now(), Kind: jointrace.KindNet, Peer: jointrace.PeerSFU, - Note: "includes creating the answer", - }) + if answered, ok := rec.Get(jointrace.SubAnswer); ok { + rec.Add(jointrace.Span{ + Name: jointrace.SubSendAnswer, After: []string{jointrace.SubAnswer}, + Start: answered.End, End: time.Now(), Kind: jointrace.KindNet, Peer: jointrace.PeerSFU, + Note: "not awaited", + }) + } else { + s.c.trace.mu.Lock() + offerAt := s.c.trace.subOfferAt + s.c.trace.mu.Unlock() + rec.Add(jointrace.Span{ + Name: jointrace.SubSendAnswer, After: []string{jointrace.SubOffer}, + Start: offerAt, End: time.Now(), Kind: jointrace.KindNet, Peer: jointrace.PeerSFU, + Note: "includes creating the answer", + }) + } // ICE may have finished while the RPC was in flight. s.c.peerSpans(false, s.Timing()) if err := answer.GetError(); err != nil {