Simplified the websockets handler code (#30)

* Simplified the websockets handler code

* created function validateAuthKey

* Using go default error interface.
This commit is contained in:
Claudemiro
2016-08-13 12:25:38 -03:00
committed by GitHub
parent 0ac6c7a55c
commit 9ad14daeeb
5 changed files with 229 additions and 200 deletions
-35
View File
@@ -4,41 +4,6 @@
package ipe
// Error Codes
const (
// 4000 - 4099
// Indicates an error resulting in the connection being closed by Pusher,
// and that attempting to reconnect using the same parameters will not succeed.
applicationOnlyAcceptsSSL = 4000
applicationDoesNotExists = 4001
applicationDisabled = 4003
applicationIsOverConnectionQuota = 4004 // Not Implemented
pathNotFound = 4005 // Not Implemented
invalidVersionStringFormat = 4006
unsupportedProtocolVersion = 4007
noProtocolVersionSupplied = 4008
// 4100 - 4199
// Indicates an error resulting in the connection being closed by Pusher,
// and the client may reconnect after 1s or more
overCapacity = 4100 // Not Implemented
// 4200 - 4299
// Indicate an error resulting in the connection being closed by Pusher,
// and the client my reconnect immediately
genericReconnectImmediately = 4200
pongReplyNotReceived = 4201 // Ping was sent to the client, but no reply was received; Not Implemented
closedAfterInactivity = 4202 // Client has been inactive for a long time (24 hours) and client does not suppot ping.; Not Implemented
// 4300 - 4399
// Any other type of error
clientRejectedDueToRateLimit = 4301 // Not Implemented
// Pusher send null, This app use this error code to send the null value
// see ErrorEvent
otherError = 0
)
// Only this version is supported
const supportedProtocolVersion = 7
+14 -8
View File
@@ -4,6 +4,8 @@
package ipe
import "fmt"
// Base interface
type websocketError interface {
GetCode() int
@@ -24,6 +26,10 @@ func (e baseWebsocketError) GetMsg() string {
return e.Msg
}
func (e baseWebsocketError) Error() string {
return fmt.Sprintf("%d: %s", e.Code, e.Msg)
}
// Unsupprted protocol version
type unsupportedProtocolVersionError struct {
baseWebsocketError
@@ -31,7 +37,7 @@ type unsupportedProtocolVersionError struct {
func newUnsupportedProtocolVersionError() unsupportedProtocolVersionError {
return unsupportedProtocolVersionError{
baseWebsocketError{Code: unsupportedProtocolVersion, Msg: "Unsupported protocol version"},
baseWebsocketError{Code: 4007, Msg: "Unsupported protocol version"},
}
}
@@ -43,7 +49,7 @@ type applicationDoesNotExistsError struct {
func newApplicationDoesNotExistsError() applicationDoesNotExistsError {
return applicationDoesNotExistsError{
baseWebsocketError{Code: applicationDoesNotExists, Msg: "Could not found an app with the given key"},
baseWebsocketError{Code: 4001, Msg: "Could not found an app with the given key"},
}
}
@@ -54,7 +60,7 @@ type noProtocolVersionSuppliedError struct {
func newNoProtocolVersionSuppliedError() noProtocolVersionSuppliedError {
return noProtocolVersionSuppliedError{
baseWebsocketError{Code: noProtocolVersionSupplied, Msg: "No protocol version supplied"},
baseWebsocketError{Code: 4008, Msg: "No protocol version supplied"},
}
}
@@ -66,7 +72,7 @@ type applicationDisabledError struct {
func newApplicationDisabledError() noProtocolVersionSuppliedError {
return noProtocolVersionSuppliedError{
baseWebsocketError{Code: applicationDisabled, Msg: "Application disabled"},
baseWebsocketError{Code: 4003, Msg: "Application disabled"},
}
}
@@ -77,7 +83,7 @@ type applicationOnlyAccepsSSLError struct {
func newApplicationOnlyAccepsSSLError() applicationOnlyAccepsSSLError {
return applicationOnlyAccepsSSLError{
baseWebsocketError{Code: applicationOnlyAcceptsSSL, Msg: "Application only accepts SSL connections, reconnect using wss://"},
baseWebsocketError{Code: 4000, Msg: "Application only accepts SSL connections, reconnect using wss://"},
}
}
@@ -88,7 +94,7 @@ type invalidVersionStringFormatError struct {
func newInvalidVersionStringFormatError() invalidVersionStringFormatError {
return invalidVersionStringFormatError{
baseWebsocketError{Code: invalidVersionStringFormat, Msg: "Invalid version string format"},
baseWebsocketError{Code: 4006, Msg: "Invalid version string format"},
}
}
@@ -101,7 +107,7 @@ type genericReconnectImmediatelyError struct {
func newGenericReconnectImmediatelyError() genericReconnectImmediatelyError {
return genericReconnectImmediatelyError{
baseWebsocketError{Code: genericReconnectImmediately, Msg: "Generic reconnect immediately"},
baseWebsocketError{Code: 4200, Msg: "Generic reconnect immediately"},
}
}
@@ -113,6 +119,6 @@ type genericError struct {
func newGenericError(msg string) genericError {
return genericError{
baseWebsocketError{Code: otherError, Msg: msg},
baseWebsocketError{Code: 0, Msg: msg},
}
}
+12 -17
View File
@@ -160,26 +160,21 @@ type errorEvent struct {
// Create a new error event
// Pusher protocol is very strange in some parts
// It send null in some errors.
// So I created this GENERIC_ERROR thing, just to verify if the json must have null on the error code
func newErrorEvent(code int, message string) errorEvent {
var data interface{}
if code == otherError {
data = struct {
Code *int `json:"code"`
Message string `json:"message"`
}{
nil,
message,
}
type dataErrorEvent struct {
Code *int `json:"code"`
Message string `json:"message"`
}
var data = dataErrorEvent{
Message: message,
}
if code == 0 {
data.Code = nil
} else {
data = struct {
Code int `json:"code"`
Message string `json:"message"`
}{
code,
message,
}
data.Code = &code
}
return errorEvent{Event: "pusher:error", Data: data}
+29
View File
@@ -0,0 +1,29 @@
package ipe
import (
"bytes"
"encoding/json"
"testing"
)
func Test_newErrorEvent_with_invalid_code(t *testing.T) {
event := newErrorEvent(0, "The error message")
data, _ := json.Marshal(event)
expected := `{"event":"pusher:error","data":{"code":null,"message":"The error message"}}`
if bytes.Compare(data, []byte(expected)) != 0 {
t.Errorf("%s != %s", string(data), expected)
}
}
func Test_newErrorEvent_with_valid_code(t *testing.T) {
event := newErrorEvent(4007, "Unsupported protocol version")
data, _ := json.Marshal(event)
expected := `{"event":"pusher:error","data":{"code":4007,"message":"Unsupported protocol version"}}`
if bytes.Compare(data, []byte(expected)) != 0 {
t.Errorf("%s != %s", string(data), expected)
}
}
+174 -140
View File
@@ -29,8 +29,59 @@ var upgrader = websocket.Upgrader{
CheckOrigin: func(r *http.Request) bool { return true },
}
// Handle open Subscriber.
func onOpen(conn *websocket.Conn, w http.ResponseWriter, r *http.Request, sessionID string, app *app) websocketError {
func handleMessages(
conn *websocket.Conn, w http.ResponseWriter,
r *http.Request, sessionID string, app *app) {
var event struct {
Event string `json:"event"`
}
for {
_, message, err := conn.ReadMessage()
if err != nil {
handleError(conn, sessionID, app, err)
return
}
if err := json.Unmarshal(message, &event); err != nil {
emitWSError(newGenericReconnectImmediatelyError(), conn)
return
}
log.Infof("websockets: Handling %s event", event.Event)
switch event.Event {
case "pusher:ping":
onPing(conn)
case "pusher:subscribe":
onSubscribe(conn, sessionID, app, message)
case "pusher:unsubscribe":
onUnsubscribe(conn, sessionID, app, message)
default:
if utils.IsClientEvent(event.Event) {
onClientEvent(conn, sessionID, app, message)
}
}
} // For
}
func handleError(conn *websocket.Conn, sessionID string, app *app, err error) {
log.Errorf("%+v", err)
if err == io.EOF {
onClose(sessionID, app)
} else if _, ok := err.(*websocket.CloseError); ok {
onClose(sessionID, app)
} else {
emitWSError(newGenericReconnectImmediatelyError(), conn)
}
}
func onOpen(
conn *websocket.Conn, w http.ResponseWriter,
r *http.Request, sessionID string, app *app) error {
params := r.URL.Query()
p := params.Get("protocol")
@@ -65,150 +116,143 @@ func onOpen(conn *websocket.Conn, w http.ResponseWriter, r *http.Request, sessio
return nil
}
// Handle the close event
func onClose(sessionID string, app *app) {
app.Disconnect(sessionID)
}
// Handle messages
//
// If there is an unrecoverable error then break the loop,
// otherwise just keep going.
func onMessage(conn *websocket.Conn, w http.ResponseWriter, r *http.Request, sessionID string, app *app) {
var event struct {
Event string `json:"event"`
func onPing(conn *websocket.Conn) {
if err := conn.WriteJSON(newPongEvent()); err != nil {
emitWSError(newGenericReconnectImmediatelyError(), conn)
}
}
func onClientEvent(
conn *websocket.Conn, sessionID string, app *app, message []byte) {
if !app.UserEvents {
emitWSError(newGenericError("To send client events, you must enable this feature in the Settings."), conn)
}
for {
_, message, err := conn.ReadMessage()
clientEvent := rawEvent{}
if err != nil {
log.Errorf("%+v", err)
if err == io.EOF {
onClose(sessionID, app)
} else if _, ok := err.(*websocket.CloseError); ok {
onClose(sessionID, app)
} else {
emitWSError(newGenericReconnectImmediatelyError(), conn)
}
break
if err := json.Unmarshal(message, &clientEvent); err != nil {
log.Error(err)
emitWSError(newGenericReconnectImmediatelyError(), conn)
return
}
channel, err := app.FindChannelByChannelID(clientEvent.Channel)
if err != nil {
emitWSError(newGenericError(fmt.Sprintf("Could not find a channel with the id %s", clientEvent.Channel)), conn)
}
if !channel.IsPresenceOrPrivate() {
emitWSError(newGenericError("Client event rejected - only supported on private and presence channels"), conn)
return
}
if err := app.Publish(channel, clientEvent, sessionID); err != nil {
log.Error(err)
emitWSError(newGenericReconnectImmediatelyError(), conn)
return
}
}
func onUnsubscribe(
conn *websocket.Conn, sessionID string, app *app, message []byte) {
unsubscribeEvent := unsubscribeEvent{}
if err := json.Unmarshal(message, &unsubscribeEvent); err != nil {
emitWSError(newGenericReconnectImmediatelyError(), conn)
}
connection, err := app.FindConnection(sessionID)
if err != nil {
emitWSError(newGenericError(fmt.Sprintf("Could not find a connection with the id %s", sessionID)), conn)
}
channel, err := app.FindChannelByChannelID(unsubscribeEvent.Data.Channel)
if err != nil {
emitWSError(newGenericError(fmt.Sprintf("Could not find a channel with the id %s", unsubscribeEvent.Data.Channel)), conn)
}
if err := app.Unsubscribe(channel, connection); err != nil {
emitWSError(newGenericReconnectImmediatelyError(), conn)
return
}
}
func onSubscribe(
conn *websocket.Conn, sessionID string, app *app, message []byte) {
subscribeEvent := subscribeEvent{}
if err := json.Unmarshal(message, &subscribeEvent); err != nil {
emitWSError(newGenericReconnectImmediatelyError(), conn)
return
}
connection, err := app.FindConnection(sessionID)
if err != nil {
emitWSError(newGenericReconnectImmediatelyError(), conn)
return
}
channelName := strings.TrimSpace(subscribeEvent.Data.Channel)
if !utils.IsChannelNameValid(channelName) {
emitWSError(newGenericError(fmt.Sprintf("This channel name is not valid")), conn)
return
}
isPresence := utils.IsPresenceChannel(channelName)
isPrivate := utils.IsPrivateChannel(channelName)
if isPresence || isPrivate {
toSign := []string{connection.SocketID, channelName}
if isPresence || len(subscribeEvent.Data.ChannelData) > 0 {
toSign = append(toSign, subscribeEvent.Data.ChannelData)
}
if err := json.Unmarshal(message, &event); err != nil {
emitWSError(newGenericReconnectImmediatelyError(), conn)
break
if !validateAuthKey(subscribeEvent.Data.Auth, toSign, app) {
emitWSError(newGenericError(fmt.Sprintf("Auth value for subscription to %s is invalid", channelName)), conn)
return
}
}
log.Infof("websockets: Handling %s event", event.Event)
channel := app.FindOrCreateChannelByChannelID(channelName)
log.Info(subscribeEvent.Data.ChannelData)
switch event.Event {
case "pusher:ping":
if err := conn.WriteJSON(newPongEvent()); err != nil {
emitWSError(newGenericReconnectImmediatelyError(), conn)
}
case "pusher:subscribe":
subscribeEvent := subscribeEvent{}
if err := app.Subscribe(channel, connection, subscribeEvent.Data.ChannelData); err != nil {
emitWSError(newGenericReconnectImmediatelyError(), conn)
}
}
if err := json.Unmarshal(message, &subscribeEvent); err != nil {
emitWSError(newGenericReconnectImmediatelyError(), conn)
break
}
func validateAuthKey(givenAuthKey string, toSign []string, app *app) bool {
expectedAuthKey := fmt.Sprintf("%s:%s", app.Key, utils.HashMAC([]byte(strings.Join(toSign, ":")), []byte(app.Secret)))
return givenAuthKey == expectedAuthKey
}
connection, err := app.FindConnection(sessionID)
// Emit an Websocket ErrorEvent
func emitWSError(err error, conn *websocket.Conn) {
e, ok := err.(websocketError)
if err != nil {
emitWSError(newGenericReconnectImmediatelyError(), conn)
break
}
if !ok {
log.Error(err)
return
}
channelName := strings.TrimSpace(subscribeEvent.Data.Channel)
event := newErrorEvent(e.GetCode(), e.GetMsg())
if !utils.IsChannelNameValid(channelName) {
emitWSError(newGenericError(fmt.Sprintf("This channel name is not valid")), conn)
break
}
isPresence := utils.IsPresenceChannel(channelName)
isPrivate := utils.IsPrivateChannel(channelName)
if isPresence || isPrivate {
toSign := []string{connection.SocketID, channelName}
if isPresence || len(subscribeEvent.Data.ChannelData) > 0 {
toSign = append(toSign, subscribeEvent.Data.ChannelData)
}
expectedAuthKey := fmt.Sprintf("%s:%s", app.Key, utils.HashMAC([]byte(strings.Join(toSign, ":")), []byte(app.Secret)))
if subscribeEvent.Data.Auth != expectedAuthKey {
emitWSError(newGenericError(fmt.Sprintf("Auth value for subscription to %s is invalid", channelName)), conn)
continue
}
}
channel := app.FindOrCreateChannelByChannelID(channelName)
log.Info(subscribeEvent.Data.ChannelData)
if err := app.Subscribe(channel, connection, subscribeEvent.Data.ChannelData); err != nil {
emitWSError(newGenericReconnectImmediatelyError(), conn)
}
case "pusher:unsubscribe":
unsubscribeEvent := unsubscribeEvent{}
if err := json.Unmarshal(message, &unsubscribeEvent); err != nil {
emitWSError(newGenericReconnectImmediatelyError(), conn)
}
connection, err := app.FindConnection(sessionID)
if err != nil {
emitWSError(newGenericError(fmt.Sprintf("Could not find a connection with the id %s", sessionID)), conn)
}
channel, err := app.FindChannelByChannelID(unsubscribeEvent.Data.Channel)
if err != nil {
emitWSError(newGenericError(fmt.Sprintf("Could not find a channel with the id %s", unsubscribeEvent.Data.Channel)), conn)
}
if err := app.Unsubscribe(channel, connection); err != nil {
emitWSError(newGenericReconnectImmediatelyError(), conn)
break
}
default: // CLient Events ??
// see http://pusher.com/docs/client_api_guide/client_events#trigger-events
if utils.IsClientEvent(event.Event) {
if !app.UserEvents {
emitWSError(newGenericError("To send client events, you must enable this feature in the Settings."), conn)
}
clientEvent := rawEvent{}
if err := json.Unmarshal(message, &clientEvent); err != nil {
log.Error(err)
emitWSError(newGenericReconnectImmediatelyError(), conn)
break
}
channel, err := app.FindChannelByChannelID(clientEvent.Channel)
if err != nil {
emitWSError(newGenericError(fmt.Sprintf("Could not find a channel with the id %s", clientEvent.Channel)), conn)
}
if !channel.IsPresenceOrPrivate() {
emitWSError(newGenericError("Client event rejected - only supported on private and presence channels"), conn)
break
}
if err := app.Publish(channel, clientEvent, sessionID); err != nil {
log.Error(err)
emitWSError(newGenericReconnectImmediatelyError(), conn)
break
}
}
} // switch
} // For
if err := conn.WriteJSON(event); err != nil {
log.Error(err)
}
}
func newWebsocketHandler(DB db) goji.Handler {
@@ -248,15 +292,5 @@ func (h *websocketHandler) ServeHTTPC(ctx context.Context, w http.ResponseWriter
return
}
onMessage(conn, w, r, sessionID, app)
}
// Emit an Websocket ErrorEvent
func emitWSError(err websocketError, conn *websocket.Conn) {
event := newErrorEvent(err.GetCode(), err.GetMsg())
if err := conn.WriteJSON(event); err != nil {
log.Error(err)
}
handleMessages(conn, w, r, sessionID, app)
}