Refactor & harden websocket implementation (data race issues)
This commit is contained in:
@@ -109,8 +109,8 @@ func (ws *server) Send(t string, payload interface{}, userIDs ...uint64) error {
|
||||
func (ws *server) StoreSession(s *session) {
|
||||
ws.l.Lock()
|
||||
defer ws.l.Unlock()
|
||||
if s.identity != nil {
|
||||
ws.storeSession(s, s.identity.Identity(), s.id)
|
||||
if i := s.Identity(); i != nil {
|
||||
ws.storeSession(s, i.Identity(), s.id)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -126,8 +126,8 @@ func (ws *server) storeSession(w io.Writer, uid, sid uint64) {
|
||||
func (ws *server) RemoveSession(s *session) {
|
||||
ws.l.Lock()
|
||||
defer ws.l.Unlock()
|
||||
if s.identity != nil {
|
||||
uid := s.identity.Identity()
|
||||
if i := s.Identity(); i != nil {
|
||||
uid := i.Identity()
|
||||
delete(ws.sessions[uid], s.id)
|
||||
|
||||
if len(ws.sessions[uid]) == 0 {
|
||||
|
||||
@@ -25,10 +25,22 @@ var (
|
||||
)
|
||||
|
||||
type (
|
||||
conection interface {
|
||||
Close() error
|
||||
RemoteAddr() net.Addr
|
||||
WriteMessage(messageType int, data []byte) error
|
||||
SetWriteDeadline(t time.Time) error
|
||||
ReadMessage() (messageType int, p []byte, err error)
|
||||
SetReadDeadline(t time.Time) error
|
||||
SetPongHandler(h func(appData string) error)
|
||||
}
|
||||
|
||||
session struct {
|
||||
l sync.RWMutex
|
||||
|
||||
id uint64
|
||||
once sync.Once
|
||||
conn *websocket.Conn
|
||||
conn conection
|
||||
|
||||
ctx context.Context
|
||||
ctxCancel context.CancelFunc
|
||||
@@ -69,34 +81,16 @@ func Session(ctx context.Context, ws *server, conn *websocket.Conn) *session {
|
||||
return s
|
||||
}
|
||||
|
||||
func (s *session) connected() (err error) {
|
||||
s.logger.Info("connected", zap.String("remoteAddr", s.conn.RemoteAddr().String()))
|
||||
|
||||
//// Tell everyone that user has connected
|
||||
//if err = s.sendPresence("connected"); err != nil {
|
||||
// return
|
||||
//}
|
||||
//
|
||||
//
|
||||
//// Create a heartbeat every minute for this user
|
||||
//go func() {
|
||||
// defer sentry.Recover()
|
||||
//
|
||||
// t := time.NewTicker(time.Second * 60)
|
||||
// for {
|
||||
// select {
|
||||
// case <-s.ctx.Done():
|
||||
// return
|
||||
// case <-t.C:
|
||||
// _ = s.sendPresence("active")
|
||||
// }
|
||||
// }
|
||||
//}()
|
||||
|
||||
return nil
|
||||
func (s *session) Identity() auth.Identifiable {
|
||||
s.l.RLock()
|
||||
defer s.l.RUnlock()
|
||||
return s.identity
|
||||
}
|
||||
|
||||
func (s *session) disconnected() {
|
||||
func (s *session) disconnect() {
|
||||
s.l.Lock()
|
||||
defer s.l.Unlock()
|
||||
|
||||
// Cancel context
|
||||
s.ctxCancel()
|
||||
|
||||
@@ -104,46 +98,40 @@ func (s *session) disconnected() {
|
||||
|
||||
// Close connection
|
||||
_ = s.conn.Close()
|
||||
|
||||
close(s.send)
|
||||
close(s.stop)
|
||||
s.conn = nil
|
||||
}
|
||||
|
||||
//func (s *session) sendPresence(_ string) error {
|
||||
// return nil
|
||||
//}
|
||||
|
||||
func (s *session) Handle() (err error) {
|
||||
if err = s.connected(); err != nil {
|
||||
s.Close()
|
||||
return
|
||||
}
|
||||
|
||||
func (s *session) Handle() error {
|
||||
go func() {
|
||||
// Close unidentified connections in 5sec
|
||||
<-time.NewTimer(time.Second * 5).C
|
||||
if s.identity == nil {
|
||||
s.Write([]byte(closingUnidentifiedConn))
|
||||
if s.Identity() == nil {
|
||||
_, _ = s.Write(closingUnidentifiedConn)
|
||||
s.logger.Info("closing unidentified connection")
|
||||
s.Close()
|
||||
}
|
||||
}()
|
||||
|
||||
go func() {
|
||||
if err = s.readLoop(); err != nil {
|
||||
if err := s.readLoop(); err != nil {
|
||||
s.logger.Error("read failure", zap.Error(err))
|
||||
}
|
||||
s.Close()
|
||||
}()
|
||||
|
||||
if err = s.writeLoop(); err != nil {
|
||||
if err := s.writeLoop(); err != nil {
|
||||
s.logger.Error("write failure", zap.Error(err))
|
||||
}
|
||||
|
||||
return
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *session) Close() {
|
||||
s.once.Do(func() {
|
||||
s.disconnected()
|
||||
s.disconnect()
|
||||
s.server.RemoveSession(s)
|
||||
})
|
||||
}
|
||||
@@ -164,12 +152,12 @@ func (s *session) readLoop() (err error) {
|
||||
)
|
||||
|
||||
for {
|
||||
if s.conn == nil {
|
||||
return nil
|
||||
if raw, err = s.read(); err != nil {
|
||||
return errHandler("read failed", err)
|
||||
}
|
||||
|
||||
if _, raw, err = s.conn.ReadMessage(); err != nil {
|
||||
return errHandler("read failed", err)
|
||||
if raw == nil {
|
||||
continue
|
||||
}
|
||||
|
||||
if err = s.procRawMessage(raw); err != nil {
|
||||
@@ -178,6 +166,23 @@ func (s *session) readLoop() (err error) {
|
||||
}
|
||||
}
|
||||
|
||||
func (s *session) read() (raw []byte, err error) {
|
||||
defer func() {
|
||||
if recovered := recover(); recovered != nil {
|
||||
s.logger.Debug("recovering from websocket read panic", zap.Any("recovered-error", recovered))
|
||||
}
|
||||
}()
|
||||
|
||||
s.l.RLock()
|
||||
defer s.l.RUnlock()
|
||||
|
||||
if _, raw, err = s.conn.ReadMessage(); err != nil {
|
||||
return nil, errHandler("websocket read failed", err)
|
||||
}
|
||||
|
||||
return
|
||||
}
|
||||
|
||||
func (s *session) procRawMessage(raw []byte) (err error) {
|
||||
pw := payloadWrap{}
|
||||
if err = json.Unmarshal(raw, &pw); err != nil {
|
||||
@@ -194,10 +199,11 @@ func (s *session) procRawMessage(raw []byte) (err error) {
|
||||
return fmt.Errorf("unauthorized: %w", err)
|
||||
}
|
||||
|
||||
i := s.Identity()
|
||||
s.logger.Debug(
|
||||
"authenticated",
|
||||
zap.Uint64("userID", s.identity.Identity()),
|
||||
zap.Uint64s("roles", s.identity.Roles()),
|
||||
zap.Uint64("userID", i.Identity()),
|
||||
zap.Uint64s("roles", i.Roles()),
|
||||
)
|
||||
|
||||
s.server.StoreSession(s)
|
||||
@@ -206,7 +212,7 @@ func (s *session) procRawMessage(raw []byte) (err error) {
|
||||
return
|
||||
}
|
||||
|
||||
if s.identity == nil {
|
||||
if s.Identity() == nil {
|
||||
return fmt.Errorf("unauthenticated session")
|
||||
}
|
||||
|
||||
@@ -214,65 +220,14 @@ func (s *session) procRawMessage(raw []byte) (err error) {
|
||||
return fmt.Errorf("unknown message type '%s'", pw.Type)
|
||||
}
|
||||
|
||||
// reads send & stop channels and sends received messages to websocket connection via write fn()
|
||||
func (s *session) writeLoop() error {
|
||||
ticker := time.NewTicker(s.config.PingPeriod)
|
||||
|
||||
defer func() {
|
||||
ticker.Stop()
|
||||
s.Close() // break readLoop
|
||||
}()
|
||||
|
||||
write := func(msg []byte) (err error) {
|
||||
defer func() {
|
||||
if recErr := recover(); recErr != nil {
|
||||
s.logger.Debug("recovering from write panic", zap.Error(err))
|
||||
}
|
||||
}()
|
||||
|
||||
if s.conn == nil {
|
||||
// Connection closed, nowhere to write
|
||||
return
|
||||
}
|
||||
|
||||
if err = s.conn.SetWriteDeadline(time.Now().Add(s.config.Timeout)); err != nil {
|
||||
return fmt.Errorf("deadline error: %w", err)
|
||||
}
|
||||
|
||||
if msg != nil && s.conn != nil {
|
||||
return s.conn.WriteMessage(websocket.TextMessage, msg)
|
||||
}
|
||||
|
||||
return
|
||||
}
|
||||
|
||||
ping := func() (err error) {
|
||||
defer func() {
|
||||
if recErr := recover(); recErr != nil {
|
||||
s.logger.Debug("recovering from ping panic", zap.Error(err))
|
||||
}
|
||||
}()
|
||||
|
||||
if s.conn == nil {
|
||||
// Connection closed, nothing to ping
|
||||
return
|
||||
}
|
||||
|
||||
if err = s.conn.SetWriteDeadline(time.Now().Add(s.config.Timeout)); err != nil {
|
||||
return
|
||||
}
|
||||
|
||||
if s.conn != nil {
|
||||
return s.conn.WriteMessage(websocket.PingMessage, nil)
|
||||
}
|
||||
|
||||
return
|
||||
}
|
||||
defer ticker.Stop()
|
||||
defer s.Close() // break readLoop
|
||||
|
||||
for {
|
||||
if s.conn == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
select {
|
||||
case msg, ok := <-s.send:
|
||||
if !ok {
|
||||
@@ -280,31 +235,62 @@ func (s *session) writeLoop() error {
|
||||
return nil
|
||||
}
|
||||
|
||||
if err := errHandler("send failed", write(msg)); err != nil {
|
||||
if err := s.write(websocket.TextMessage, msg); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
case msg := <-s.stop:
|
||||
// continue with wait & write/ping loop
|
||||
|
||||
case msg, ok := <-s.stop:
|
||||
if !ok {
|
||||
// channel closed
|
||||
return nil
|
||||
}
|
||||
|
||||
// Shutdown requested, don't care if the message is delivered
|
||||
_ = write(msg)
|
||||
if err := s.write(websocket.TextMessage, msg); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// stopping, break the loop.
|
||||
return nil
|
||||
|
||||
case <-ticker.C:
|
||||
if err := ping(); err != nil {
|
||||
return errHandler("ping failed", err)
|
||||
if err := s.write(websocket.PingMessage, nil); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// continue with wait & write/ping loop
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// writes messages to websocket connection
|
||||
func (s *session) write(t int, msg []byte) (err error) {
|
||||
s.l.RLock()
|
||||
defer s.l.RUnlock()
|
||||
|
||||
defer func() {
|
||||
if recovered := recover(); recovered != nil {
|
||||
s.logger.Debug("recovering from websocket write panic", zap.Any("recovered-error", recovered))
|
||||
}
|
||||
}()
|
||||
|
||||
if err = s.conn.SetWriteDeadline(time.Now().Add(s.config.Timeout)); err != nil {
|
||||
return fmt.Errorf("deadline error: %w", err)
|
||||
}
|
||||
|
||||
return errHandler("websocket write failed", s.conn.WriteMessage(t, msg))
|
||||
}
|
||||
|
||||
func (s *session) authenticate(p *payloadAuth) error {
|
||||
identity, err := s.server.tokenValidator(s.ctx, p.AccessToken)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if s.identity != nil {
|
||||
if s.identity.Identity() != identity.Identity() {
|
||||
if i := s.Identity(); i != nil {
|
||||
if i.Identity() != identity.Identity() {
|
||||
return fmt.Errorf("identity does not match")
|
||||
}
|
||||
}
|
||||
@@ -313,6 +299,9 @@ func (s *session) authenticate(p *payloadAuth) error {
|
||||
return fmt.Errorf("invalid identity")
|
||||
}
|
||||
|
||||
s.l.Lock()
|
||||
defer s.l.Unlock()
|
||||
|
||||
s.identity = identity
|
||||
_, _ = s.Write(ok)
|
||||
return nil
|
||||
@@ -320,6 +309,12 @@ func (s *session) authenticate(p *payloadAuth) error {
|
||||
|
||||
// sendBytes sends byte to channel or timeout
|
||||
func (s *session) Write(p []byte) (int, error) {
|
||||
defer func() {
|
||||
if recovered := recover(); recovered != nil {
|
||||
s.logger.Debug("recovering from websocket write panic", zap.Any("recovered-error", recovered))
|
||||
}
|
||||
}()
|
||||
|
||||
select {
|
||||
case s.send <- p:
|
||||
return len(p), nil
|
||||
@@ -328,7 +323,7 @@ func (s *session) Write(p []byte) (int, error) {
|
||||
}
|
||||
}
|
||||
|
||||
func errHandler(wrap string, err error) error {
|
||||
func errHandler(prefix string, err error) error {
|
||||
if err == nil {
|
||||
return nil
|
||||
}
|
||||
@@ -342,5 +337,6 @@ func errHandler(wrap string, err error) error {
|
||||
// suppress errors when reading/writing from/to a closed connection
|
||||
return nil
|
||||
}
|
||||
return fmt.Errorf(wrap+": %w", err)
|
||||
|
||||
return fmt.Errorf(prefix+": %w", err)
|
||||
}
|
||||
|
||||
@@ -12,6 +12,40 @@ import (
|
||||
"go.uber.org/zap"
|
||||
)
|
||||
|
||||
//type (
|
||||
// mockConn struct {
|
||||
// close func() error
|
||||
// remoteAddr func() net.Addr
|
||||
// writeMessage func(messageType int, data []byte) error
|
||||
// setWriteDeadline func(t time.Time) error
|
||||
// readMessage func() (messageType int, p []byte, err error)
|
||||
// setReadDeadline func(t time.Time) error
|
||||
// setPongHandler func(h func(appData string) error)
|
||||
// }
|
||||
//)
|
||||
//
|
||||
//func MockConn() *mockConn {
|
||||
// return &mockConn{
|
||||
// close: func() (err error) { return },
|
||||
// remoteAddr: func() (addr net.Addr) { return &net.IPAddr{IP: net.IPv4(0, 0, 0, 0)} },
|
||||
// writeMessage: func(messageType int, data []byte) (err error) { return },
|
||||
// setWriteDeadline: func(t time.Time) (err error) { return },
|
||||
// readMessage: func() (messageType int, p []byte, err error) { return },
|
||||
// setReadDeadline: func(t time.Time) (err error) { return },
|
||||
// setPongHandler: func(h func(appData string) error) {},
|
||||
// }
|
||||
//}
|
||||
//
|
||||
//func (c *mockConn) Close() error { return c.close() }
|
||||
//func (c *mockConn) RemoteAddr() net.Addr { return c.remoteAddr() }
|
||||
//func (c *mockConn) WriteMessage(messageType int, data []byte) error {
|
||||
// return c.writeMessage(messageType, data)
|
||||
//}
|
||||
//func (c *mockConn) SetWriteDeadline(t time.Time) error { return c.setWriteDeadline(t) }
|
||||
//func (c *mockConn) ReadMessage() (messageType int, p []byte, err error) { return c.readMessage() }
|
||||
//func (c *mockConn) SetReadDeadline(t time.Time) error { return c.setReadDeadline(t) }
|
||||
//func (c *mockConn) SetPongHandler(h func(appData string) error) { c.setPongHandler(h) }
|
||||
|
||||
func TestSession_procRawMessage(t *testing.T) {
|
||||
var (
|
||||
req = require.New(t)
|
||||
@@ -24,11 +58,6 @@ func TestSession_procRawMessage(t *testing.T) {
|
||||
nil,
|
||||
options.WebsocketOpt{},
|
||||
func(ctx context.Context, accessToken string) (auth.Identifiable, error) {
|
||||
//token, err := jwt.Parse([]byte(accessToken))
|
||||
//if err != nil {
|
||||
// return nil, err
|
||||
//}
|
||||
//return auth.IdentityFromToken(token), nil
|
||||
switch accessToken {
|
||||
case "one":
|
||||
return identity1, nil
|
||||
@@ -59,23 +88,23 @@ func TestSession_procRawMessage(t *testing.T) {
|
||||
}
|
||||
|
||||
req.EqualError(s.procRawMessage([]byte("{}")), "unauthenticated session")
|
||||
req.Nil(s.identity)
|
||||
req.Nil(s.Identity())
|
||||
|
||||
req.EqualError(s.procRawMessage(mockResponse(nil)), "unauthorized: failed to parse token: EOF")
|
||||
req.Nil(s.identity)
|
||||
req.Nil(s.Identity())
|
||||
|
||||
token = []byte("one")
|
||||
req.NoError(s.procRawMessage(mockResponse(token)))
|
||||
req.NotNil(s.identity)
|
||||
req.Equal(identity1.Identity(), s.identity.Identity())
|
||||
req.NotNil(s.Identity())
|
||||
req.Equal(identity1.Identity(), s.Identity().Identity())
|
||||
|
||||
req.EqualError(s.procRawMessage([]byte("{}")), "unknown message type ''")
|
||||
req.Equal(identity1.Identity(), s.identity.Identity())
|
||||
req.Equal(identity1.Identity(), s.Identity().Identity())
|
||||
|
||||
token = []byte("one")
|
||||
req.NoError(s.procRawMessage(mockResponse(token)))
|
||||
req.NotNil(s.identity)
|
||||
req.Equal(identity1.Identity(), s.identity.Identity())
|
||||
req.NotNil(s.Identity())
|
||||
req.Equal(identity1.Identity(), s.Identity().Identity())
|
||||
|
||||
token = []byte("two")
|
||||
req.EqualError(s.procRawMessage(mockResponse(token)), "unauthorized: identity does not match")
|
||||
|
||||
Reference in New Issue
Block a user