PLT-6226 fixing race in IsAuth (#7296)

* Fixing race in isAuth function

* PLT-6226 fixing race in IsAuth

* Moving int64 to top so it's aligned

* Adding comment and fixing asymmetric call
Этот коммит содержится в:
Corey Hulen
2017-09-05 07:58:47 -07:00
коммит произвёл GitHub
родитель 7843dc3cfa
Коммит d6383643cb
4 изменённых файлов: 71 добавлений и 31 удалений

Просмотреть файл

@@ -5,6 +5,7 @@ package app
import ( import (
"fmt" "fmt"
"sync/atomic"
"time" "time"
"github.com/mattermost/platform/einterfaces" "github.com/mattermost/platform/einterfaces"
@@ -28,11 +29,11 @@ const (
) )
type WebConn struct { type WebConn struct {
sessionExpiresAt int64 // This should stay at the top for 64-bit alignment of 64-bit words accessed atomically
WebSocket *websocket.Conn WebSocket *websocket.Conn
Send chan model.WebSocketMessage Send chan model.WebSocketMessage
SessionToken string sessionToken atomic.Value
SessionExpiresAt int64 session atomic.Value
Session *model.Session
UserId string UserId string
T goi18n.TranslateFunc T goi18n.TranslateFunc
Locale string Locale string
@@ -49,15 +50,47 @@ func NewWebConn(ws *websocket.Conn, session model.Session, t goi18n.TranslateFun
}() }()
} }
return &WebConn{ wc := &WebConn{
Send: make(chan model.WebSocketMessage, SEND_QUEUE_SIZE), Send: make(chan model.WebSocketMessage, SEND_QUEUE_SIZE),
WebSocket: ws, WebSocket: ws,
UserId: session.UserId, UserId: session.UserId,
SessionToken: session.Token, T: t,
SessionExpiresAt: session.ExpiresAt, Locale: locale,
T: t,
Locale: locale,
} }
wc.SetSession(&session)
wc.SetSessionToken(session.Token)
wc.SetSessionExpiresAt(session.ExpiresAt)
return wc
}
func (c *WebConn) GetSessionExpiresAt() int64 {
return atomic.LoadInt64(&c.sessionExpiresAt)
}
func (c *WebConn) SetSessionExpiresAt(v int64) {
atomic.StoreInt64(&c.sessionExpiresAt, v)
}
func (c *WebConn) GetSessionToken() string {
return c.sessionToken.Load().(string)
}
func (c *WebConn) SetSessionToken(v string) {
c.sessionToken.Store(v)
}
func (c *WebConn) GetSession() *model.Session {
return c.session.Load().(*model.Session)
}
func (c *WebConn) SetSession(v *model.Session) {
if v != nil {
v = v.DeepCopy()
}
c.session.Store(v)
} }
func (c *WebConn) ReadPump() { func (c *WebConn) ReadPump() {
@@ -175,7 +208,7 @@ func (c *WebConn) WritePump() {
} }
case <-authTicker.C: case <-authTicker.C:
if c.SessionToken == "" { if c.GetSessionToken() == "" {
l4g.Debug(fmt.Sprintf("websocket.authTicker: did not authenticate ip=%v", c.WebSocket.RemoteAddr())) l4g.Debug(fmt.Sprintf("websocket.authTicker: did not authenticate ip=%v", c.WebSocket.RemoteAddr()))
return return
} }
@@ -187,29 +220,28 @@ func (c *WebConn) WritePump() {
func (webCon *WebConn) InvalidateCache() { func (webCon *WebConn) InvalidateCache() {
webCon.AllChannelMembers = nil webCon.AllChannelMembers = nil
webCon.LastAllChannelMembersTime = 0 webCon.LastAllChannelMembersTime = 0
webCon.SessionExpiresAt = 0 webCon.SetSession(nil)
webCon.Session = nil webCon.SetSessionExpiresAt(0)
} }
func (webCon *WebConn) IsAuthenticated() bool { func (webCon *WebConn) IsAuthenticated() bool {
// Check the expiry to see if we need to check for a new session // Check the expiry to see if we need to check for a new session
if webCon.SessionExpiresAt < model.GetMillis() { if webCon.GetSessionExpiresAt() < model.GetMillis() {
if webCon.SessionToken == "" { if webCon.GetSessionToken() == "" {
return false return false
} }
session, err := GetSession(webCon.SessionToken) session, err := GetSession(webCon.GetSessionToken())
if err != nil { if err != nil {
l4g.Error(utils.T("api.websocket.invalid_session.error"), err.Error()) l4g.Error(utils.T("api.websocket.invalid_session.error"), err.Error())
webCon.SessionToken = "" webCon.SetSessionToken("")
webCon.SessionExpiresAt = 0 webCon.SetSession(nil)
webCon.Session = nil webCon.SetSessionExpiresAt(0)
return false return false
} }
webCon.SessionToken = session.Token webCon.SetSession(session)
webCon.SessionExpiresAt = session.ExpiresAt webCon.SetSessionExpiresAt(session.ExpiresAt)
webCon.Session = session
} }
return true return true
@@ -278,18 +310,20 @@ func (webCon *WebConn) ShouldSendEvent(msg *model.WebSocketEvent) bool {
func (webCon *WebConn) IsMemberOfTeam(teamId string) bool { func (webCon *WebConn) IsMemberOfTeam(teamId string) bool {
if webCon.Session == nil { currentSession := webCon.GetSession()
session, err := GetSession(webCon.SessionToken)
if currentSession == nil || len(currentSession.Token) == 0 {
session, err := GetSession(webCon.GetSessionToken())
if err != nil { if err != nil {
l4g.Error(utils.T("api.websocket.invalid_session.error"), err.Error()) l4g.Error(utils.T("api.websocket.invalid_session.error"), err.Error())
return false return false
} else { } else {
webCon.Session = session webCon.SetSession(session)
currentSession = session
} }
} }
member := webCon.Session.GetTeamByTeamId(teamId) member := currentSession.GetTeamByTeamId(teamId)
if member != nil { if member != nil {
return true return true

Просмотреть файл

@@ -43,7 +43,7 @@ func (wr *WebSocketRouter) ServeWebSocket(conn *WebConn, r *model.WebSocketReque
} }
if r.Action == model.WEBSOCKET_AUTHENTICATION_CHALLENGE { if r.Action == model.WEBSOCKET_AUTHENTICATION_CHALLENGE {
if conn.SessionToken != "" { if conn.GetSessionToken() != "" {
return return
} }
@@ -63,7 +63,8 @@ func (wr *WebSocketRouter) ServeWebSocket(conn *WebConn, r *model.WebSocketReque
UpdateLastActivityAtIfNeeded(*session) UpdateLastActivityAtIfNeeded(*session)
}() }()
conn.SessionToken = session.Token conn.SetSession(session)
conn.SetSessionToken(session.Token)
conn.UserId = session.UserId conn.UserId = session.UserId
HubRegister(conn) HubRegister(conn)

Просмотреть файл

@@ -37,6 +37,11 @@ type Session struct {
TeamMembers []*TeamMember `json:"team_members" db:"-"` TeamMembers []*TeamMember `json:"team_members" db:"-"`
} }
func (me *Session) DeepCopy() *Session {
copy := *me
return &copy
}
func (me *Session) ToJson() string { func (me *Session) ToJson() string {
b, err := json.Marshal(me) b, err := json.Marshal(me)
if err != nil { if err != nil {

Просмотреть файл

@@ -23,7 +23,7 @@ type webSocketHandler struct {
func (wh webSocketHandler) ServeWebSocket(conn *app.WebConn, r *model.WebSocketRequest) { func (wh webSocketHandler) ServeWebSocket(conn *app.WebConn, r *model.WebSocketRequest) {
l4g.Debug("/api/v3/users/websocket:%s", r.Action) l4g.Debug("/api/v3/users/websocket:%s", r.Action)
session, sessionErr := app.GetSession(conn.SessionToken) session, sessionErr := app.GetSession(conn.GetSessionToken())
if sessionErr != nil { if sessionErr != nil {
l4g.Error(utils.T("api.web_socket_handler.log.error"), "/api/v3/users/websocket", r.Action, r.Seq, conn.UserId, sessionErr.SystemMessage(utils.T), sessionErr.Error()) l4g.Error(utils.T("api.web_socket_handler.log.error"), "/api/v3/users/websocket", r.Action, r.Seq, conn.UserId, sessionErr.SystemMessage(utils.T), sessionErr.Error())
sessionErr.DetailedError = "" sessionErr.DetailedError = ""