diff --git a/cmd/broadcaster/main.go b/cmd/broadcaster/main.go index 899bfbc..8ceebd4 100644 --- a/cmd/broadcaster/main.go +++ b/cmd/broadcaster/main.go @@ -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( diff --git a/internal/auth/authenticator.go b/internal/auth/authenticator.go index f896715..fa4a3b9 100644 --- a/internal/auth/authenticator.go +++ b/internal/auth/authenticator.go @@ -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 { @@ -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 } @@ -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 @@ -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 } diff --git a/internal/auth/authenticator_test.go b/internal/auth/authenticator_test.go index 8a22401..4ec26e6 100644 --- a/internal/auth/authenticator_test.go +++ b/internal/auth/authenticator_test.go @@ -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) }) diff --git a/internal/broadcaster/message.go b/internal/broadcaster/message.go index 6953b79..79282af 100644 --- a/internal/broadcaster/message.go +++ b/internal/broadcaster/message.go @@ -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"` } diff --git a/internal/broadcaster/registry.go b/internal/broadcaster/registry.go index b4d9076..2c15276 100644 --- a/internal/broadcaster/registry.go +++ b/internal/broadcaster/registry.go @@ -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() diff --git a/internal/handler/channel.go b/internal/handler/channel.go index 0030299..a0a114b 100644 --- a/internal/handler/channel.go +++ b/internal/handler/channel.go @@ -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 diff --git a/internal/handler/publish.go b/internal/handler/publish.go index 86ba7bf..2f6142f 100644 --- a/internal/handler/publish.go +++ b/internal/handler/publish.go @@ -12,8 +12,9 @@ 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 { @@ -21,16 +22,16 @@ type PublishHandlerInterface interface { } 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, } } @@ -55,12 +56,12 @@ 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 } @@ -68,7 +69,8 @@ func (h *PublishHandler) Handle(ctx context.Context, req PublishRequest) (broadc message := broadcaster.Message{ Id: gonanoid.Must(), CreateTime: time.Now(), - ChannelId: req.ChannelId, + Channel: req.Channel, + Event: req.Event, Payload: req.Payload, } diff --git a/internal/handler/subscribe.go b/internal/handler/subscribe.go index 4fc50a2..99dfbd0 100644 --- a/internal/handler/subscribe.go +++ b/internal/handler/subscribe.go @@ -10,7 +10,7 @@ import ( ) type SubscribeRequest struct { - ChannelId string + Channel string `json:"channel"` } type SubscribeResponse struct { @@ -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 } @@ -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 } diff --git a/internal/handler/unsubscribe.go b/internal/handler/unsubscribe.go index 55ca957..388d0e2 100644 --- a/internal/handler/unsubscribe.go +++ b/internal/handler/unsubscribe.go @@ -8,7 +8,7 @@ import ( ) type UnsubscribeRequest struct { - ChannelId string `json:"channelId"` + Channel string `json:"channel"` } type UnsubscribeResponse struct { @@ -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 } @@ -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, diff --git a/internal/server/rest_test.go b/internal/server/rest_test.go index ab9dd64..1d5521c 100644 --- a/internal/server/rest_test.go +++ b/internal/server/rest_test.go @@ -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) @@ -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))) @@ -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") diff --git a/internal/server/websocket_test.go b/internal/server/websocket_test.go index 9dca89b..464f642 100644 --- a/internal/server/websocket_test.go +++ b/internal/server/websocket_test.go @@ -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) @@ -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) @@ -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 @@ -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() @@ -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) @@ -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) @@ -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) @@ -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) @@ -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) diff --git a/test/broadcaster.html b/test/broadcaster.html index acaf6b3..9ab4aaf 100644 --- a/test/broadcaster.html +++ b/test/broadcaster.html @@ -34,12 +34,23 @@