From f2c082f73fa65f9b73f673b869f62ea1e284bdd8 Mon Sep 17 00:00:00 2001 From: Jeeva Kandasamy Date: Wed, 12 Aug 2026 15:17:06 +0530 Subject: [PATCH] fix websocket issues with policies Signed-off-by: Jeeva Kandasamy --- pkg/api/policy/mapper.go | 4 ++- pkg/api/policy/mapper_test.go | 7 +++++ pkg/service/websocket/events_listener.go | 40 ++++++++++++++++++++++-- pkg/service/websocket/handler.go | 10 +++++- pkg/service/websocket/service.go | 3 +- pkg/service/websocket/store.go | 21 +++++++------ 6 files changed, 69 insertions(+), 16 deletions(-) diff --git a/pkg/api/policy/mapper.go b/pkg/api/policy/mapper.go index 9fdf53c..55c15d3 100644 --- a/pkg/api/policy/mapper.go +++ b/pkg/api/policy/mapper.go @@ -32,7 +32,9 @@ func MapRequest(r *http.Request) RequestAccess { path = strings.TrimPrefix(path, "api/") // special non-restricted already handled by auth middleware - if path == "status" || path == "version" || strings.HasPrefix(path, "user/login") || + // /api/ws is authenticated (cookie/JWT) but not RBAC-gated: every signed-in + // user may hold a socket. Events are filtered per principal when sent. + if path == "status" || path == "version" || path == "ws" || strings.HasPrefix(path, "user/login") || strings.HasPrefix(path, "oauth/") || strings.HasPrefix(path, "plugin/gateway") { return RequestAccess{Skip: true} } diff --git a/pkg/api/policy/mapper_test.go b/pkg/api/policy/mapper_test.go index 04d36e1..b791edc 100644 --- a/pkg/api/policy/mapper_test.go +++ b/pkg/api/policy/mapper_test.go @@ -8,6 +8,13 @@ import ( policyTY "github.com/mycontroller-org/server/v2/pkg/types/policy" ) +func TestMapRequestWebsocketSkipsRBAC(t *testing.T) { + access := MapRequest(httptest.NewRequest(http.MethodGet, "/api/ws", nil)) + if !access.Skip { + t.Fatalf("websocket should skip RBAC after auth, got %+v", access) + } +} + func TestMapRequestUserProfileExact(t *testing.T) { profile := MapRequest(httptest.NewRequest(http.MethodGet, "/api/user/profile", nil)) if profile.Name != "" || profile.Kind != policyTY.ResourceUser || profile.Action != policyTY.ActionGet { diff --git a/pkg/service/websocket/events_listener.go b/pkg/service/websocket/events_listener.go index 2e73240..6a2a88f 100644 --- a/pkg/service/websocket/events_listener.go +++ b/pkg/service/websocket/events_listener.go @@ -4,8 +4,10 @@ import ( "time" ws "github.com/gorilla/websocket" + policyAPI "github.com/mycontroller-org/server/v2/pkg/api/policy" "github.com/mycontroller-org/server/v2/pkg/json" eventTY "github.com/mycontroller-org/server/v2/pkg/types/event" + policyTY "github.com/mycontroller-org/server/v2/pkg/types/policy" wsTY "github.com/mycontroller-org/server/v2/pkg/types/websocket" busTY "github.com/mycontroller-org/server/v2/plugin/bus/types" "go.uber.org/zap" @@ -68,15 +70,17 @@ func (svc *WebsocketService) processEvent(item interface{}) error { } wsClients := svc.store.getClients() - for index := range wsClients { - client := wsClients[index] + for client, subject := range wsClients { + if !svc.eventAllowed(subject, event) { + continue + } // write with write timeout err := client.SetWriteDeadline(time.Now().Add(defaultWriteTimeout)) if err != nil { svc.logger.Debug("error on setting write deadline", zap.Any("remoteAddress", client.RemoteAddr().String()), zap.Error(err)) svc.store.unregister(client) - return nil + continue } err = client.WriteMessage(ws.TextMessage, dataBytes) if err != nil { @@ -86,3 +90,33 @@ func (svc *WebsocketService) processEvent(item interface{}) error { } return nil } + +// eventAllowed reports whether this principal may see the live event. +// Quick ids are checked as the named resource; otherwise the check is +// kind:entityId. Events with neither a parseable quick id nor EntityID are denied. +func (svc *WebsocketService) eventAllowed(subject policyAPI.Subject, event *eventTY.Event) bool { + if subject.UserID == "" { + return false + } + ac := svc.api.Policy() + if ac == nil { + return false + } + resource := "" + if event.EntityQuickID != "" { + if res, err := policyAPI.ResourceFromQuickID(event.EntityQuickID); err == nil { + resource = res + } + } + if resource == "" { + if event.EntityID == "" { + return false + } + kind := policyTY.NormalizeKind(event.EntityType) + if kind == "" { + return false + } + resource = policyAPI.FormatResource(kind, event.EntityID) + } + return ac.Allowed(subject, policyTY.ActionGet, resource) == nil +} diff --git a/pkg/service/websocket/handler.go b/pkg/service/websocket/handler.go index 4db98dd..f93da5c 100644 --- a/pkg/service/websocket/handler.go +++ b/pkg/service/websocket/handler.go @@ -4,6 +4,7 @@ import ( "net/http" ws "github.com/gorilla/websocket" + middleware "github.com/mycontroller-org/server/v2/pkg/http_router/middleware" "go.uber.org/zap" ) @@ -26,6 +27,13 @@ func (svc *WebsocketService) Start() error { // this is simple example websocket // yet to implement actual version func (svc *WebsocketService) wsFunc(w http.ResponseWriter, r *http.Request) { + subject, err := middleware.SubjectFromRequest(r) + if err != nil { + svc.logger.Info("websocket rejected: no authenticated subject", zap.Error(err)) + http.Error(w, "unauthorized", http.StatusUnauthorized) + return + } + wsCon, err := upgrader.Upgrade(w, r, nil) if err != nil { svc.logger.Info("websocket upgrade error", zap.Error(err)) @@ -33,7 +41,7 @@ func (svc *WebsocketService) wsFunc(w http.ResponseWriter, r *http.Request) { } // register the new client - svc.store.register(wsCon) + svc.store.register(wsCon, subject) // NOTE: for now not serving any request, only sending the events to the listeners(ex: remote browsers) // this loop is used to close the connection immediately on remote side close diff --git a/pkg/service/websocket/service.go b/pkg/service/websocket/service.go index 8b11bf0..f04c218 100644 --- a/pkg/service/websocket/service.go +++ b/pkg/service/websocket/service.go @@ -8,6 +8,7 @@ import ( "github.com/gorilla/mux" ws "github.com/gorilla/websocket" entityAPI "github.com/mycontroller-org/server/v2/pkg/api/entities" + policyAPI "github.com/mycontroller-org/server/v2/pkg/api/policy" serviceTY "github.com/mycontroller-org/server/v2/pkg/types/service" "github.com/mycontroller-org/server/v2/pkg/types/topic" loggerUtils "github.com/mycontroller-org/server/v2/pkg/utils/logger" @@ -55,7 +56,7 @@ func New(ctx context.Context, router *mux.Router) (serviceTY.Service, error) { router: router, } - svc.store = &Store{clients: make(map[*ws.Conn]bool), mutex: sync.RWMutex{}, logger: svc.logger} + svc.store = &Store{clients: make(map[*ws.Conn]policyAPI.Subject), mutex: sync.RWMutex{}, logger: svc.logger} svc.eventsQueue = &queueUtils.QueueSpec{ Queue: queueUtils.New(svc.logger, "websocket_event_listener", defaultQueueSize, svc.processEvent, defaultWorkers), diff --git a/pkg/service/websocket/store.go b/pkg/service/websocket/store.go index 69c34d8..4d13b63 100644 --- a/pkg/service/websocket/store.go +++ b/pkg/service/websocket/store.go @@ -4,22 +4,23 @@ import ( "sync" ws "github.com/gorilla/websocket" + policyAPI "github.com/mycontroller-org/server/v2/pkg/api/policy" "go.uber.org/zap" ) type Store struct { - clients map[*ws.Conn]bool + clients map[*ws.Conn]policyAPI.Subject mutex sync.RWMutex logger *zap.Logger } // register a websocket client connection -func (s *Store) register(conn *ws.Conn) { +func (s *Store) register(conn *ws.Conn, subject policyAPI.Subject) { s.mutex.Lock() defer s.mutex.Unlock() - s.clients[conn] = true - s.logger.Debug("new websocket connection added", zap.String("remoteAddress", conn.RemoteAddr().String())) + s.clients[conn] = subject + s.logger.Debug("new websocket connection added", zap.String("remoteAddress", conn.RemoteAddr().String()), zap.String("userId", subject.UserID)) } // unregister a websocket client connection @@ -36,16 +37,16 @@ func (s *Store) unregister(conn *ws.Conn) { delete(s.clients, conn) } -// returns available websocket client connection -func (s *Store) getClients() []*ws.Conn { +// returns available websocket client connections with their access subject +func (s *Store) getClients() map[*ws.Conn]policyAPI.Subject { s.mutex.RLock() defer s.mutex.RUnlock() - wsClients := make([]*ws.Conn, 0) - for client := range s.clients { - wsClients = append(wsClients, client) + out := make(map[*ws.Conn]policyAPI.Subject, len(s.clients)) + for client, subject := range s.clients { + out[client] = subject } - return wsClients + return out } // returns the size of the client map