Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 4 additions & 4 deletions cmd/broadcaster/main.go
Original file line number Diff line number Diff line change
Expand Up @@ -38,13 +38,13 @@ func NewApp(logger *zap.Logger, settings Settings) *App {

authenticator := auth.NewAuthenticator(settings.JWTSecret, settings.APIKeys)

channelIdValidator := handler.NewChannelIdValidator()
channelValidator := handler.NewChannelValidator()
registry := broadcaster.NewInMemoryRegistry(logger)

heartbeatHandler := handler.NewHeartbeatHandler()
subscribeHandler := handler.NewSubscribeHandler(channelIdValidator, registry)
unsubscribeHandler := handler.NewUnsubscribeHandler(channelIdValidator, registry)
publishHandler := handler.NewPublishHandler(channelIdValidator, registry)
subscribeHandler := handler.NewSubscribeHandler(channelValidator, registry)
unsubscribeHandler := handler.NewUnsubscribeHandler(channelValidator, registry)
publishHandler := handler.NewPublishHandler(channelValidator, registry)
authHandler := handler.NewAuthHandler(authenticator)

router := server.NewRouter(
Expand Down
20 changes: 10 additions & 10 deletions internal/auth/authenticator.go
Original file line number Diff line number Diff line change
Expand Up @@ -18,10 +18,10 @@ type Claims struct {
}

type Authentication struct {
Subject string
AuthorizedChannelsIds []string
Scope []string
IsAdmin bool
Subject string
AuthorizedChannels []string
Scope []string
IsAdmin bool
}

func (a *Authentication) IsPublisher() bool {
Expand All @@ -32,7 +32,7 @@ func (a *Authentication) IsSubscriber() bool {
return slices.Contains(a.Scope, "subscribe")
}

func (a *Authentication) IsAuthorized(channelId string) bool {
func (a *Authentication) IsAuthorized(channel string) bool {
if a.Subject == "" {
return false
}
Expand All @@ -41,7 +41,7 @@ func (a *Authentication) IsAuthorized(channelId string) bool {
return true
}

return slices.Contains(a.AuthorizedChannelsIds, channelId)
return slices.Contains(a.AuthorizedChannels, channel)
}

type contextKey string
Expand Down Expand Up @@ -104,10 +104,10 @@ func (a *Authenticator) AuthenticateJWT(tokenString string) (*Authentication, er
}

return &Authentication{
Subject: subject,
AuthorizedChannelsIds: claims.AuthorizedChannels,
Scope: claims.Scope,
IsAdmin: false,
Subject: subject,
AuthorizedChannels: claims.AuthorizedChannels,
Scope: claims.Scope,
IsAdmin: false,
}, nil
}

Expand Down
2 changes: 1 addition & 1 deletion internal/auth/authenticator_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -30,7 +30,7 @@ func TestAuthenticator_AuthenticateJWT(t *testing.T) {
assert.NoError(t, err)
assert.NotNil(t, auth)
assert.Equal(t, "test-user", auth.Subject)
assert.Equal(t, []string{"test-channel"}, auth.AuthorizedChannelsIds)
assert.Equal(t, []string{"test-channel"}, auth.AuthorizedChannels)
assert.Equal(t, []string{"subscribe"}, auth.Scope)
assert.False(t, auth.IsAdmin)
})
Expand Down
3 changes: 2 additions & 1 deletion internal/broadcaster/message.go
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@ import "time"
type Message struct {
Id string `json:"id"`
CreateTime time.Time `json:"createTime"`
ChannelId string `json:"channelId"`
Channel string `json:"channel"`
Event string `json:"event"`
Payload any `json:"payload"`
}
2 changes: 1 addition & 1 deletion internal/broadcaster/registry.go
Original file line number Diff line number Diff line change
Expand Up @@ -52,7 +52,7 @@ func (r *InMemoryRegistry) Connect(connection *Connection) error {
func (r *InMemoryRegistry) Broadcast(message Message) {
r.mu.RLock()

connectionIds, ok := r.connectionsByChannel[message.ChannelId]
connectionIds, ok := r.connectionsByChannel[message.Channel]
if !ok {
r.mu.RUnlock()

Expand Down
16 changes: 8 additions & 8 deletions internal/handler/channel.go
Original file line number Diff line number Diff line change
Expand Up @@ -7,20 +7,20 @@ import (
"github.com/goevery/broadcaster/internal/ierr"
)

type ChannelIdValidator struct {
channelIdRegex *regexp.Regexp
type ChannelValidator struct {
channelRegex *regexp.Regexp
}

func NewChannelIdValidator() *ChannelIdValidator {
return &ChannelIdValidator{
channelIdRegex: regexp.MustCompile(`^([\w-]+:?)*\w$`),
func NewChannelValidator() *ChannelValidator {
return &ChannelValidator{
channelRegex: regexp.MustCompile(`^([\w-]+:?)*\w$`),
}
}

func (v *ChannelIdValidator) Validate(channelId string) error {
valid := v.channelIdRegex.MatchString(channelId)
func (v *ChannelValidator) Validate(channel string) error {
valid := v.channelRegex.MatchString(channel)
if !valid {
return ierr.New(ierr.ErrorCodeInvalidArgument, errors.New("invalid channelId"))
return ierr.New(ierr.ErrorCodeInvalidArgument, errors.New("invalid channel"))
}

return nil
Expand Down
18 changes: 10 additions & 8 deletions internal/handler/publish.go
Original file line number Diff line number Diff line change
Expand Up @@ -12,25 +12,26 @@ import (
)

type PublishRequest struct {
ChannelId string `json:"channelId"`
Payload any `json:"payload"`
Channel string `json:"channel"`
Event string `json:"event"`
Payload any `json:"payload"`
}

type PublishHandlerInterface interface {
Handle(ctx context.Context, req PublishRequest) (broadcaster.Message, error)
}

type PublishHandler struct {
channelIdValidator *ChannelIdValidator
channelValidator *ChannelValidator
subscriptionRegistry broadcaster.Registry
}

func NewPublishHandler(
channelIdValidator *ChannelIdValidator,
channelValidator *ChannelValidator,
subscriptionRegistry broadcaster.Registry,
) *PublishHandler {
return &PublishHandler{
channelIdValidator,
channelValidator,
subscriptionRegistry,
}
}
Expand All @@ -55,20 +56,21 @@ func (h *PublishHandler) Handle(ctx context.Context, req PublishRequest) (broadc
ierr.New(ierr.ErrorCodePermissionDenied, errors.New("user not authorized to publish messages"))
}

if !authentication.IsAuthorized(req.ChannelId) {
if !authentication.IsAuthorized(req.Channel) {
return broadcaster.Message{},
ierr.New(ierr.ErrorCodePermissionDenied, errors.New("user not authorized to publish to this channel"))
}

err := h.channelIdValidator.Validate(req.ChannelId)
err := h.channelValidator.Validate(req.Channel)
if err != nil {
return broadcaster.Message{}, err
}

message := broadcaster.Message{
Id: gonanoid.Must(),
CreateTime: time.Now(),
ChannelId: req.ChannelId,
Channel: req.Channel,
Event: req.Event,
Payload: req.Payload,
}

Expand Down
14 changes: 7 additions & 7 deletions internal/handler/subscribe.go
Original file line number Diff line number Diff line change
Expand Up @@ -10,7 +10,7 @@ import (
)

type SubscribeRequest struct {
ChannelId string
Channel string `json:"channel"`
}

type SubscribeResponse struct {
Expand All @@ -23,23 +23,23 @@ type SubscribeHandlerInterface interface {
}

type SubscribeHandler struct {
channelIdValidator *ChannelIdValidator
channelValidator *ChannelValidator
subscriptionRegistry broadcaster.Registry
}

func NewSubscribeHandler(
channelIdValidator *ChannelIdValidator,
channelValidator *ChannelValidator,
subscriptionRegistry broadcaster.Registry,
) *SubscribeHandler {

return &SubscribeHandler{
channelIdValidator,
channelValidator,
subscriptionRegistry,
}
}

func (h *SubscribeHandler) Handle(ctx context.Context, req SubscribeRequest) (SubscribeResponse, error) {
err := h.channelIdValidator.Validate(req.ChannelId)
err := h.channelValidator.Validate(req.Channel)
if err != nil {
return SubscribeResponse{}, err
}
Expand All @@ -60,12 +60,12 @@ func (h *SubscribeHandler) Handle(ctx context.Context, req SubscribeRequest) (Su
ierr.New(ierr.ErrorCodePermissionDenied, errors.New("subscribe scope required to subscribe to a channel"))
}

if !connection.IsAuthorized(req.ChannelId) {
if !connection.IsAuthorized(req.Channel) {
return SubscribeResponse{},
ierr.New(ierr.ErrorCodeUnauthenticated, errors.New("user not authorized to access this channel"))
}

err = h.subscriptionRegistry.Subscribe(req.ChannelId, connection.Id)
err = h.subscriptionRegistry.Subscribe(req.Channel, connection.Id)
if err != nil {
return SubscribeResponse{}, err
}
Expand Down
12 changes: 6 additions & 6 deletions internal/handler/unsubscribe.go
Original file line number Diff line number Diff line change
Expand Up @@ -8,7 +8,7 @@ import (
)

type UnsubscribeRequest struct {
ChannelId string `json:"channelId"`
Channel string `json:"channel"`
}

type UnsubscribeResponse struct {
Expand All @@ -20,22 +20,22 @@ type UnsubscribeHandlerInterface interface {
}

type UnsubscribeHandler struct {
channelIdValidator *ChannelIdValidator
channelValidator *ChannelValidator
subscriptionRegistry broadcaster.Registry
}

func NewUnsubscribeHandler(
channelIdValidator *ChannelIdValidator,
channelValidator *ChannelValidator,
subscriptionRegistry broadcaster.Registry,
) *UnsubscribeHandler {
return &UnsubscribeHandler{
channelIdValidator,
channelValidator,
subscriptionRegistry,
}
}

func (h *UnsubscribeHandler) Handle(ctx context.Context, req UnsubscribeRequest) (UnsubscribeResponse, error) {
err := h.channelIdValidator.Validate(req.ChannelId)
err := h.channelValidator.Validate(req.Channel)
if err != nil {
return UnsubscribeResponse{}, err
}
Expand All @@ -45,7 +45,7 @@ func (h *UnsubscribeHandler) Handle(ctx context.Context, req UnsubscribeRequest)
return UnsubscribeResponse{}, errors.New("connection not found in context")
}

h.subscriptionRegistry.Unsubscribe(req.ChannelId, connection.Id)
h.subscriptionRegistry.Unsubscribe(req.Channel, connection.Id)

return UnsubscribeResponse{
Success: true,
Expand Down
10 changes: 5 additions & 5 deletions internal/server/rest_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -19,8 +19,8 @@ func TestRESTServer_Publish(t *testing.T) {
logger, _ := zap.NewDevelopment()
authenticator := auth.NewAuthenticator("test-secret", []string{"test-api-key"})
registry := broadcaster.NewMockRegistry(t)
channelIdValidator := handler.NewChannelIdValidator()
publishHandler := handler.NewPublishHandler(channelIdValidator, registry)
channelValidator := handler.NewChannelValidator()
publishHandler := handler.NewPublishHandler(channelValidator, registry)

restServer := NewRESTServer(logger, publishHandler, authenticator)

Expand All @@ -31,10 +31,10 @@ func TestRESTServer_Publish(t *testing.T) {
defer server.Close()

t.Run("valid api key", func(t *testing.T) {
body := `{"channelId":"test-channel","payload":"test-payload"}`
body := `{"channel":"test-channel","event":"test-event","payload":"test-payload"}`

registry.On("Broadcast", mock.MatchedBy(func(msg broadcaster.Message) bool {
return msg.ChannelId == "test-channel" && msg.Payload == "test-payload"
return msg.Channel == "test-channel" && msg.Event == "test-event" && msg.Payload == "test-payload"
})).Return().Once()

req, _ := http.NewRequest("POST", server.URL+"/publish", bytes.NewBuffer([]byte(body)))
Expand All @@ -49,7 +49,7 @@ func TestRESTServer_Publish(t *testing.T) {
})

t.Run("invalid api key", func(t *testing.T) {
body := `{"channelId":"test-channel","payload":"test-payload"}`
body := `{"channel":"test-channel","event":"test-event","payload":"test-payload"}`

req, _ := http.NewRequest("POST", server.URL+"/publish", bytes.NewBuffer([]byte(body)))
req.Header.Set("Authorization", "Bearer invalid-api-key")
Expand Down
24 changes: 12 additions & 12 deletions internal/server/websocket_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -21,11 +21,11 @@ func TestWebSocketServer(t *testing.T) {
logger, _ := zap.NewDevelopment()
registry := broadcaster.NewInMemoryRegistry(logger)
authenticator := auth.NewAuthenticator("test-secret", []string{"test-api-key"})
channelIdValidator := handler.NewChannelIdValidator()
channelValidator := handler.NewChannelValidator()
heartbeatHandler := handler.NewHeartbeatHandler()
subscribeHandler := handler.NewSubscribeHandler(channelIdValidator, registry)
unsubscribeHandler := handler.NewUnsubscribeHandler(channelIdValidator, registry)
publishHandler := handler.NewPublishHandler(channelIdValidator, registry)
subscribeHandler := handler.NewSubscribeHandler(channelValidator, registry)
unsubscribeHandler := handler.NewUnsubscribeHandler(channelValidator, registry)
publishHandler := handler.NewPublishHandler(channelValidator, registry)
authHandler := handler.NewAuthHandler(authenticator)

router := NewRouter(logger, heartbeatHandler, subscribeHandler, unsubscribeHandler, publishHandler, authHandler)
Expand Down Expand Up @@ -75,7 +75,7 @@ func TestWebSocketServer(t *testing.T) {
assert.Equal(t, true, authResponsePayload.Success)

// Subscribe
subscribeRequest := json.RawMessage(`{"id":2,"method":"subscribe","params":{"channelId":"test-channel"}}`)
subscribeRequest := json.RawMessage(`{"id":2,"method":"subscribe","params":{"channel":"test-channel"}}`)
err = conn.WriteJSON(subscribeRequest)
assert.NoError(t, err)

Expand All @@ -90,7 +90,7 @@ func TestWebSocketServer(t *testing.T) {
assert.NotEmpty(t, subscribeResponsePayload.SubscriptionId)

// Server sends a message
msg := broadcaster.Message{ChannelId: "test-channel", Payload: "test-payload"}
msg := broadcaster.Message{Channel: "test-channel", Payload: "test-payload"}
registry.Broadcast(msg)

var messageRequest handler.Request
Expand All @@ -102,7 +102,7 @@ func TestWebSocketServer(t *testing.T) {
var messagePayload broadcaster.Message
err = json.Unmarshal(*messageRequest.Params, &messagePayload)
assert.NoError(t, err)
assert.Equal(t, msg.ChannelId, messagePayload.ChannelId)
assert.Equal(t, msg.Channel, messagePayload.Channel)
assert.Equal(t, msg.Payload, messagePayload.Payload)

conn.Close()
Expand All @@ -128,7 +128,7 @@ func TestWebSocketServer(t *testing.T) {
assert.NoError(t, err)
defer conn.Close()

subscribeRequest := json.RawMessage(`{"id":1,"method":"subscribe","params":{"channelId":"test-channel"}}`)
subscribeRequest := json.RawMessage(`{"id":1,"method":"subscribe","params":{"channel":"test-channel"}}`)
err = conn.WriteJSON(subscribeRequest)
assert.NoError(t, err)

Expand Down Expand Up @@ -168,7 +168,7 @@ func TestWebSocketServer(t *testing.T) {
assert.NoError(t, err)

// Subscribe
subscribeRequest := json.RawMessage(`{"id":2,"method":"subscribe","params":{"channelId":"test-channel"}}`)
subscribeRequest := json.RawMessage(`{"id":2,"method":"subscribe","params":{"channel":"test-channel"}}`)
err = conn.WriteJSON(subscribeRequest)
assert.NoError(t, err)

Expand Down Expand Up @@ -213,7 +213,7 @@ func TestWebSocketServer(t *testing.T) {
assert.Equal(t, true, authResponsePayload.Success)

// Publish
publishRequest := json.RawMessage(`{"id":2,"method":"publish","params":{"channelId":"test-channel","payload":{"foo":"bar"}}}`)
publishRequest := json.RawMessage(`{"id":2,"method":"publish","params":{"channel":"test-channel","event":"test-event","payload":{"foo":"bar"}}}`)
err = conn.WriteJSON(publishRequest)
assert.NoError(t, err)

Expand Down Expand Up @@ -257,7 +257,7 @@ func TestWebSocketServer(t *testing.T) {
assert.Equal(t, true, authResponsePayload.Success)

// Subscribe should fail without subscribe scope
subscribeRequest := json.RawMessage(`{"id":2,"method":"subscribe","params":{"channelId":"test-channel"}}`)
subscribeRequest := json.RawMessage(`{"id":2,"method":"subscribe","params":{"channel":"test-channel"}}`)
err = conn.WriteJSON(subscribeRequest)
assert.NoError(t, err)

Expand Down Expand Up @@ -302,7 +302,7 @@ func TestWebSocketServer(t *testing.T) {
assert.Equal(t, true, authResponsePayload.Success)

// Publish should fail without publish scope
publishRequest := json.RawMessage(`{"id":2,"method":"publish","params":{"channelId":"test-channel","payload":{"foo":"bar"}}}`)
publishRequest := json.RawMessage(`{"id":2,"method":"publish","params":{"channel":"test-channel","event":"test-event","payload":{"foo":"bar"}}}`)
err = conn.WriteJSON(publishRequest)
assert.NoError(t, err)

Expand Down
Loading