Code cleanup

This commit is contained in:
claudemiro
2016-03-06 20:21:04 -03:00
parent e71294de18
commit 5ce17c8856
10 changed files with 49 additions and 48 deletions
+3 -3
View File
@@ -27,10 +27,10 @@ type app struct {
WebHooks bool WebHooks bool
URLWebHook string URLWebHook string
Channels map[string]*channel `json:"-"` Channels map[string]*channel
Connections map[string]*connection `json:"-"` Connections map[string]*connection
Stats *expvar.Map `json:"-"` Stats *expvar.Map
} }
func newApp(name, appID, key, secret string, onlySSL, disabled, userEvents, webHooks bool, webHookURL string) *app { func newApp(name, appID, key, secret string, onlySSL, disabled, userEvents, webHooks bool, webHookURL string) *app {
+5 -5
View File
@@ -22,7 +22,7 @@ func newTestApp() *app {
func TestConnect(t *testing.T) { func TestConnect(t *testing.T) {
app := newTestApp() app := newTestApp()
app.Connect(newConnection("socketID", nil)) app.Connect(newConnection("socketID", mockSocket{}))
if len(app.Connections) != 1 { if len(app.Connections) != 1 {
t.Errorf("len(app.Connections) == %d, wants %d", len(app.Connections), 1) t.Errorf("len(app.Connections) == %d, wants %d", len(app.Connections), 1)
@@ -33,7 +33,7 @@ func TestConnect(t *testing.T) {
func TestDisconnect(t *testing.T) { func TestDisconnect(t *testing.T) {
app := newTestApp() app := newTestApp()
app.Connect(newConnection("socketID", nil)) app.Connect(newConnection("socketID", mockSocket{}))
app.Disconnect("socketID") app.Disconnect("socketID")
if len(app.Connections) != 0 { if len(app.Connections) != 0 {
@@ -45,7 +45,7 @@ func TestDisconnect(t *testing.T) {
func TestFindConnection(t *testing.T) { func TestFindConnection(t *testing.T) {
app := newTestApp() app := newTestApp()
app.Connect(newConnection("socketID", nil)) app.Connect(newConnection("socketID", mockSocket{}))
if _, err := app.FindConnection("socketID"); err != nil { if _, err := app.FindConnection("socketID"); err != nil {
t.Errorf("app.FindConnection('socketID') == _, %q, wants %v", err, nil) t.Errorf("app.FindConnection('socketID') == _, %q, wants %v", err, nil)
@@ -165,7 +165,7 @@ func Test_New_Subscriber(t *testing.T) {
t.Errorf("len(app.Connections) == %d, wants %d", len(app.Connections), 0) t.Errorf("len(app.Connections) == %d, wants %d", len(app.Connections), 0)
} }
conn := newConnection("1", nil) conn := newConnection("1", mockSocket{})
app.Connect(conn) app.Connect(conn)
if len(app.Connections) != 1 { if len(app.Connections) != 1 {
@@ -175,7 +175,7 @@ func Test_New_Subscriber(t *testing.T) {
func Test_find_subscriber(t *testing.T) { func Test_find_subscriber(t *testing.T) {
app := newTestApp() app := newTestApp()
conn := newConnection("1", nil) conn := newConnection("1", mockSocket{})
app.Connect(conn) app.Connect(conn)
conn, err := app.FindConnection("1") conn, err := app.FindConnection("1")
+4 -4
View File
@@ -13,7 +13,7 @@ func TestIsOccupied(t *testing.T) {
t.Errorf("c.IsOccupied() == %t, wants %t", c.IsOccupied(), false) t.Errorf("c.IsOccupied() == %t, wants %t", c.IsOccupied(), false)
} }
c.Subscriptions["ID"] = newSubscription(newConnection("ID", nil), "") c.Subscriptions["ID"] = newSubscription(newConnection("ID", mockSocket{}), "")
if !c.IsOccupied() { if !c.IsOccupied() {
t.Errorf("c.IsOccupied() == %t, wants %t", c.IsOccupied(), true) t.Errorf("c.IsOccupied() == %t, wants %t", c.IsOccupied(), true)
@@ -69,8 +69,8 @@ func TestTotalSubscriptions(t *testing.T) {
func TestTotalUsers(t *testing.T) { func TestTotalUsers(t *testing.T) {
c := newChannel("ID") c := newChannel("ID")
c.Subscriptions["1"] = newSubscription(newConnection("ID", nil), "") c.Subscriptions["1"] = newSubscription(newConnection("ID", mockSocket{}), "")
c.Subscriptions["2"] = newSubscription(newConnection("ID", nil), "") c.Subscriptions["2"] = newSubscription(newConnection("ID", mockSocket{}), "")
if c.TotalSubscriptions() != len(c.Subscriptions) { if c.TotalSubscriptions() != len(c.Subscriptions) {
t.Errorf("c.TotalSubscriptions() == %d, wants %d", c.TotalSubscriptions(), len(c.Subscriptions)) t.Errorf("c.TotalSubscriptions() == %d, wants %d", c.TotalSubscriptions(), len(c.Subscriptions))
@@ -84,7 +84,7 @@ func TestTotalUsers(t *testing.T) {
func TestIsSubscribed(t *testing.T) { func TestIsSubscribed(t *testing.T) {
c := newChannel("ID") c := newChannel("ID")
conn := newConnection("ID", nil) conn := newConnection("ID", mockSocket{})
if c.IsSubscribed(conn) { if c.IsSubscribed(conn) {
t.Errorf("c.IsSubscribed(%q) == %t, wants %t", conn, c.IsSubscribed(conn), false) t.Errorf("c.IsSubscribed(%q) == %t, wants %t", conn, c.IsSubscribed(conn), false)
+15 -8
View File
@@ -8,18 +8,30 @@ import (
"time" "time"
log "github.com/golang/glog" log "github.com/golang/glog"
"github.com/gorilla/websocket"
) )
// socket interface to write to the client
type socket interface {
WriteJSON(interface{}) error
}
// mockSocket is a mock implementation of socket
// used in the test suite
type mockSocket struct{}
func (s mockSocket) WriteJSON(i interface{}) error {
return nil
}
// An User Connection // An User Connection
type connection struct { type connection struct {
SocketID string SocketID string
Socket *websocket.Conn Socket socket
CreatedAt time.Time CreatedAt time.Time
} }
// Create a new Subscriber // Create a new Subscriber
func newConnection(socketID string, s *websocket.Conn) *connection { func newConnection(socketID string, s socket) *connection {
log.Infof("Creating a new Subscriber %+v", socketID) log.Infof("Creating a new Subscriber %+v", socketID)
return &connection{SocketID: socketID, Socket: s, CreatedAt: time.Now()} return &connection{SocketID: socketID, Socket: s, CreatedAt: time.Now()}
@@ -28,11 +40,6 @@ func newConnection(socketID string, s *websocket.Conn) *connection {
// Publish the message to websocket atached to this client // Publish the message to websocket atached to this client
func (conn *connection) Publish(m interface{}) { func (conn *connection) Publish(m interface{}) {
go func() { go func() {
if conn.Socket == nil {
log.Info("Socket is nil. Maybe you are testing this app?")
return
}
if err := conn.Socket.WriteJSON(m); err != nil { if err := conn.Socket.WriteJSON(m); err != nil {
log.Errorf("Error publishing message to connection %+v, %s", conn, err) log.Errorf("Error publishing message to connection %+v, %s", conn, err)
} }
+2 -6
View File
@@ -4,15 +4,11 @@
package ipe package ipe
import ( import "testing"
"testing"
"github.com/gorilla/websocket"
)
func TestNewConnection(t *testing.T) { func TestNewConnection(t *testing.T) {
expectedSocketID := "socketID" expectedSocketID := "socketID"
expectedSocket := &websocket.Conn{} expectedSocket := mockSocket{}
c := newConnection(expectedSocketID, expectedSocket) c := newConnection(expectedSocketID, expectedSocket)
+5 -5
View File
@@ -17,13 +17,13 @@ func (p params) Get(key string) string {
return p[key] return p[key]
} }
// A handlerHTTPC responds to an HTTP request with custom application context. // A contextHandler responds to an HTTP request with custom application context.
type handlerHTTPC interface { type contextHandler interface {
ServeHTTPC(ctx *applicationContext, p params, w http.ResponseWriter, r *http.Request) ServeWithContext(ctx *applicationContext, p params, w http.ResponseWriter, r *http.Request)
} }
type handlerHTTPCFunc func(ctx *applicationContext, p params, w http.ResponseWriter, r *http.Request) type contextHandlerFunc func(ctx *applicationContext, p params, w http.ResponseWriter, r *http.Request)
func (c handlerHTTPCFunc) ServeHTTPC(ctx *applicationContext, p params, w http.ResponseWriter, r *http.Request) { func (c contextHandlerFunc) ServeWithContext(ctx *applicationContext, p params, w http.ResponseWriter, r *http.Request) {
c(ctx, p, w, r) c(ctx, p, w, r)
} }
+7 -7
View File
@@ -46,8 +46,8 @@ func prepareQueryString(params url.Values) string {
// * The request path (e.g. /some/resource) // * The request path (e.g. /some/resource)
// * The query parameters sorted by key, with keys converted to lowercase, then joined as in the query string. // * The query parameters sorted by key, with keys converted to lowercase, then joined as in the query string.
// Note that the string must not be url escaped (e.g. given the keys auth_key: foo, Name: Something else, you get auth_key=foo&name=Something else) // Note that the string must not be url escaped (e.g. given the keys auth_key: foo, Name: Something else, you get auth_key=foo&name=Something else)
func restAuthenticationHandler(ctx *applicationContext, h handlerHTTPC) handlerHTTPC { func restAuthenticationHandler(ctx *applicationContext, h contextHandler) contextHandler {
return handlerHTTPCFunc(func(ctx *applicationContext, p params, w http.ResponseWriter, r *http.Request) { return contextHandlerFunc(func(ctx *applicationContext, p params, w http.ResponseWriter, r *http.Request) {
appID := p.Get("app_id") appID := p.Get("app_id")
app, err := ctx.DB.GetAppByAppID(appID) app, err := ctx.DB.GetAppByAppID(appID)
@@ -68,7 +68,7 @@ func restAuthenticationHandler(ctx *applicationContext, h handlerHTTPC) handlerH
toSign := strings.ToUpper(r.Method) + "\n" + r.URL.Path + "\n" + queryString toSign := strings.ToUpper(r.Method) + "\n" + r.URL.Path + "\n" + queryString
if utils.HashMAC([]byte(toSign), []byte(app.Secret)) == signature { if utils.HashMAC([]byte(toSign), []byte(app.Secret)) == signature {
h.ServeHTTPC(ctx, p, w, r) h.ServeWithContext(ctx, p, w, r)
} else { } else {
log.Error("Not authorized") log.Error("Not authorized")
http.Error(w, "Not authorized", http.StatusUnauthorized) http.Error(w, "Not authorized", http.StatusUnauthorized)
@@ -77,8 +77,8 @@ func restAuthenticationHandler(ctx *applicationContext, h handlerHTTPC) handlerH
} }
// Check if the application is disabled // Check if the application is disabled
func restCheckAppDisabledHandler(ctx *applicationContext, h handlerHTTPC) handlerHTTPC { func restCheckAppDisabledHandler(ctx *applicationContext, h contextHandler) contextHandler {
return handlerHTTPCFunc(func(ctx *applicationContext, p params, w http.ResponseWriter, r *http.Request) { return contextHandlerFunc(func(ctx *applicationContext, p params, w http.ResponseWriter, r *http.Request) {
appID := p.Get("app_id") appID := p.Get("app_id")
currentApp, err := ctx.DB.GetAppByAppID(appID) currentApp, err := ctx.DB.GetAppByAppID(appID)
@@ -93,12 +93,12 @@ func restCheckAppDisabledHandler(ctx *applicationContext, h handlerHTTPC) handle
return return
} }
h.ServeHTTPC(ctx, p, w, r) h.ServeWithContext(ctx, p, w, r)
}) })
} }
// commonHandlers combine restCheckAppDisabledHandler and restAuthenticationHandler handlers // commonHandlers combine restCheckAppDisabledHandler and restAuthenticationHandler handlers
func commonHandlers(ctx *applicationContext, h handlerHTTPCFunc) handlerHTTPC { func commonHandlers(ctx *applicationContext, h contextHandlerFunc) contextHandler {
return restCheckAppDisabledHandler(ctx, restAuthenticationHandler(ctx, h)) return restCheckAppDisabledHandler(ctx, restAuthenticationHandler(ctx, h))
} }
+2 -2
View File
@@ -21,10 +21,10 @@ func init() {
testApp.AddChannel(newChannel("c2")) testApp.AddChannel(newChannel("c2"))
testApp.AddChannel(newChannel("private-c3")) testApp.AddChannel(newChannel("private-c3"))
conn := newConnection("123.456", nil) conn := newConnection("123.456", mockSocket{})
testApp.Subscribe(channel, conn, "{}") testApp.Subscribe(channel, conn, "{}")
conn = newConnection("321.654", nil) conn = newConnection("321.654", mockSocket{})
testApp.Subscribe(channel, conn, "{}") testApp.Subscribe(channel, conn, "{}")
db := newMemdb() db := newMemdb()
+1 -1
View File
@@ -53,7 +53,7 @@ func Start(filename string) {
router.GET("/apps/{app_id}/channels/{channel_name}/users", commonHandlers(ctx, getChannelUsers)) router.GET("/apps/{app_id}/channels/{channel_name}/users", commonHandlers(ctx, getChannelUsers))
router.GET("/app/{key}", handlerHTTPCFunc(wsHandler)) router.GET("/app/{key}", contextHandlerFunc(wsHandler))
if conf.SSL { if conf.SSL {
go func() { go func() {
+5 -7
View File
@@ -13,7 +13,7 @@ import (
type router struct { type router struct {
ctx *applicationContext ctx *applicationContext
mux *mux.Router mux *mux.Router
routes map[string]handlerHTTPC routes map[string]contextHandler
} }
func newRouter(ctx *applicationContext) *router { func newRouter(ctx *applicationContext) *router {
@@ -23,19 +23,17 @@ func newRouter(ctx *applicationContext) *router {
} }
} }
func (a *router) GET(path string, handler handlerHTTPC) { func (a *router) GET(path string, handler contextHandler) {
a.Handle("GET", path, handler) a.Handle("GET", path, handler)
} }
func (a *router) POST(path string, handler handlerHTTPC) { func (a *router) POST(path string, handler contextHandler) {
a.Handle("POST", path, handler) a.Handle("POST", path, handler)
} }
func (a *router) Handle(method, path string, handler handlerHTTPC) { func (a *router) Handle(method, path string, handler contextHandler) {
a.mux.Methods(method).Path(path).HandlerFunc(func(w http.ResponseWriter, r *http.Request) { a.mux.Methods(method).Path(path).HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
p := params(mux.Vars(r)) handler.ServeWithContext(a.ctx, params(mux.Vars(r)), w, r)
handler.ServeHTTPC(a.ctx, p, w, r)
}) })
} }