diff --git a/rtpconn/whipclient.go b/rtpconn/whipclient.go index 6bc2e52..fe8e9a2 100644 --- a/rtpconn/whipclient.go +++ b/rtpconn/whipclient.go @@ -8,6 +8,9 @@ import ( "github.com/jech/galene/conn" "github.com/jech/galene/group" + "github.com/jech/galene/sdpfrag" + + "github.com/pion/sdp/v3" "github.com/pion/webrtc/v4" ) @@ -242,3 +245,57 @@ func (c *WhipClient) GotICECandidate(init webrtc.ICECandidateInit) error { } return c.connection.addICECandidate(&init) } + +func (c *WhipClient) Restart(ctx context.Context, frag sdpfrag.SDPFrag) (sdpfrag.SDPFrag, error) { + c.mu.Lock() + conn := c.connection + c.mu.Unlock() + if conn == nil { + return sdpfrag.SDPFrag{}, errors.New("no connection") + } + + offer := conn.pc.RemoteDescription() + var sdpOffer sdp.SessionDescription + err := sdpOffer.Unmarshal([]byte(offer.SDP)) + if err != nil { + return sdpfrag.SDPFrag{}, nil + } + sdpOffer2, _ := sdpfrag.PatchSDP(sdpOffer, frag) + offer2, err := sdpOffer2.Marshal() + if err != nil { + return sdpfrag.SDPFrag{}, nil + } + err = conn.pc.SetRemoteDescription(webrtc.SessionDescription{ + Type: webrtc.SDPTypeOffer, + SDP: string(offer2), + }) + if err != nil { + return sdpfrag.SDPFrag{}, err + } + + answer, err := conn.pc.CreateAnswer(nil) + if err != nil { + return sdpfrag.SDPFrag{}, err + } + + gatherComplete := webrtc.GatheringCompletePromise(conn.pc) + + err = conn.pc.SetLocalDescription(answer) + if err != nil { + return sdpfrag.SDPFrag{}, err + } + + select { + case <-ctx.Done(): + return sdpfrag.SDPFrag{}, ctx.Err() + case <-gatherComplete: + } + + sdpAnswer2 := conn.pc.LocalDescription() + var answer2 sdp.SessionDescription + err = answer2.Unmarshal([]byte(sdpAnswer2.SDP)) + if err != nil { + return sdpfrag.SDPFrag{}, err + } + return sdpfrag.FromSDP(answer2), nil +} diff --git a/sdpfrag/sdpfrag.go b/sdpfrag/sdpfrag.go index b8afe0e..f4fe4a4 100644 --- a/sdpfrag/sdpfrag.go +++ b/sdpfrag/sdpfrag.go @@ -226,3 +226,50 @@ func PatchSDP(s sdp.SessionDescription, f SDPFrag) (sdp.SessionDescription, bool } return s, override } + +func getCandidates(as []sdp.Attribute, mid *string, index *uint16, ufrag string) []webrtc.ICECandidateInit { + var cs []webrtc.ICECandidateInit + for _, a := range as { + if a.Key != "candidate" { + continue + } + c := webrtc.ICECandidateInit{ + Candidate: a.Value, + SDPMid: mid, + SDPMLineIndex: index, + } + if ufrag != "" { + c.UsernameFragment = &ufrag + } + cs = append(cs, c) + } + return cs + +} + +// FromSDP extracts the ICE information from an [sdp.SessionDescription] +func FromSDP(s sdp.SessionDescription) SDPFrag { + var f SDPFrag + ufrag, _ := s.Attribute("ice-ufrag") + f.UsernameFragment = ufrag + pwd, _ := s.Attribute("ice-pwd") + f.Password = pwd + f.Candidates = getCandidates(s.Attributes, nil, nil, ufrag) + + for index, m := range s.MediaDescriptions { + var mm MediaDescription + mm.MLine = m.MediaName.String() + mid, _ := m.Attribute("mid") + mm.Mid = mid + ufrag, _ := m.Attribute("ice-ufrag") + mm.UsernameFragment = ufrag + pwd, _ := m.Attribute("ice-pwd") + mm.Password = pwd + i := uint16(index) + mm.Candidates = getCandidates( + m.Attributes, &mid, &i, ufrag, + ) + f.MediaDescriptions = append(f.MediaDescriptions, mm) + } + return f +} diff --git a/sdpfrag/sdpfrag_test.go b/sdpfrag/sdpfrag_test.go index 96dc92a..a970c65 100644 --- a/sdpfrag/sdpfrag_test.go +++ b/sdpfrag/sdpfrag_test.go @@ -300,3 +300,29 @@ func TestPatchSDP(t *testing.T) { testPatchSDP(t, sdpOffer, sdpfragRestart, "HWmk", 0) }) } + +func TestFromSDP(t *testing.T) { + frag := `m=audio 9 UDP/TLS/RTP/SAVPF 111 63 9 0 8 13 110 126 +a=ice-ufrag:qiKa +a=ice-pwd:bcfs93hb/+ZLuUE2K50HVbkr +a=mid:0 +m=video 9 UDP/TLS/RTP/SAVPF 96 97 103 104 107 108 109 114 115 116 117 118 39 40 45 46 98 99 100 101 119 120 121 +a=ice-ufrag:qiKa +a=ice-pwd:bcfs93hb/+ZLuUE2K50HVbkr +a=mid:1 +` + var f SDPFrag + err := f.Unmarshal([]byte(frag)) + if err != nil { + t.Fatalf("Unmarshal: %v", err) + } + var s sdp.SessionDescription + err = s.Unmarshal([]byte(sdpOffer)) + if err != nil { + t.Fatalf("Unmarshal: %v", err) + } + ff := FromSDP(s) + if !reflect.DeepEqual(ff, f) { + t.Errorf("Not equal: %v %v", ff, f) + } +} diff --git a/webserver/whip.go b/webserver/whip.go index b2346d3..a1bdd8b 100644 --- a/webserver/whip.go +++ b/webserver/whip.go @@ -381,9 +381,28 @@ func whipResourceHandler(w http.ResponseWriter, r *http.Request) { } uu, pp := frag.UFragPwd() if uu != u || pp != p { - http.Error(w, "ICE restarts are not supported yet", - http.StatusUnprocessableEntity, + frag2, err := c.Restart(r.Context(), frag) + if err != nil { + log.Printf("WHIP restart: %v", err) + http.Error(w, "internal server error", + http.StatusInternalServerError, + ) + return + } + c.SetETag("\"" + newId() + "\"") + f2, err := frag2.Marshal() + if err != nil { + log.Printf("WHIP marshal frag: %v", err) + http.Error(w, "internal server error", + http.StatusInternalServerError, + ) + return + } + w.Header().Set( + "Content-Type", "application/trickle-ice-sdpfrag", ) + w.Header().Set("ETag", c.ETag()) + w.Write(f2) return } for _, init := range frag.AllCandidates() {