diff --git a/pkg/service/agentservice.go b/pkg/service/agentservice.go index f60d8dfd1..b5128cb9e 100644 --- a/pkg/service/agentservice.go +++ b/pkg/service/agentservice.go @@ -211,8 +211,10 @@ func (s *AgentService) ServeHTTP(w http.ResponseWriter, r *http.Request) { if s.signalMessageSizeLimit > 0 { conn.SetReadLimit(s.signalMessageSizeLimit) } - s.HandleConnection(r.Context(), NewWSSignalConnection(conn, s.signalMessageSizeLimit), registration) - conn.Close() + sigConn := NewWSSignalConnection(conn, s.signalMessageSizeLimit) + defer sigConn.Close() + + s.HandleConnection(r.Context(), sigConn, registration) } } diff --git a/pkg/service/rtcservice.go b/pkg/service/rtcservice.go index b42bb16ab..520bec369 100644 --- a/pkg/service/rtcservice.go +++ b/pkg/service/rtcservice.go @@ -467,6 +467,7 @@ func (s *RTCService) serve(w http.ResponseWriter, r *http.Request, needsJoinRequ // websocket established sigConn := NewWSSignalConnection(conn, s.limits.SignalMessageSizeLimit) + defer sigConn.CloseWithReason("") pLogger.Debugw("sending initial response", "response", logger.Proto(initialResponse)) count, err := sigConn.WriteResponse(initialResponse) if err != nil { @@ -495,9 +496,7 @@ func (s *RTCService) serve(w http.ResponseWriter, r *http.Request, needsJoinRequ defer func() { // when the source is terminated, this means Participant.Close had been called and RTC connection is done // we would terminate the signal connection as well - closeMsg := websocket.FormatCloseMessage(websocket.CloseNormalClosure, "") - _ = conn.WriteControl(websocket.CloseMessage, closeMsg, time.Now().Add(time.Second)) - _ = conn.Close() + sigConn.CloseWithReason("") }() defer func() { if r := rtc.Recover(pLogger); r != nil { diff --git a/pkg/service/wsprotocol.go b/pkg/service/wsprotocol.go index 407063976..58ca514aa 100644 --- a/pkg/service/wsprotocol.go +++ b/pkg/service/wsprotocol.go @@ -22,6 +22,7 @@ import ( "sync" "time" + "github.com/frostbyte73/core" "github.com/gorilla/websocket" "google.golang.org/protobuf/proto" @@ -45,6 +46,8 @@ type WSSignalConnection struct { // maximum size (in bytes) of a single decompressed message; 0 disables the limit messageSizeLimit int64 + + closed core.Fuse } func NewWSSignalConnection(conn types.WebsocketClient, messageSizeLimit int64) *WSSignalConnection { @@ -88,10 +91,14 @@ func (c *WSSignalConnection) readMessage() (int, []byte, error) { } func (c *WSSignalConnection) Close() error { + c.closed.Break() + return c.conn.Close() } func (c *WSSignalConnection) CloseWithReason(reason string) error { + c.closed.Break() + msg := websocket.FormatCloseMessage(websocket.CloseNormalClosure, reason) _ = c.conn.WriteControl(websocket.CloseMessage, msg, time.Now().Add(closeWriteTimeout)) return c.conn.Close() @@ -213,10 +220,16 @@ func (c *WSSignalConnection) pingWorker() { ticker := time.NewTicker(pingFrequency) defer ticker.Stop() - for range ticker.C { - err := c.conn.WriteControl(websocket.PingMessage, []byte(""), time.Now().Add(pingTimeout)) - if err != nil { + for { + select { + case <-c.closed.Watch(): return + + case <-ticker.C: + err := c.conn.WriteControl(websocket.PingMessage, []byte(""), time.Now().Add(pingTimeout)) + if err != nil { + return + } } } }