From 50512ca41504d85a68ef90b57bcb0551ad66be3a Mon Sep 17 00:00:00 2001 From: yushen Date: Mon, 13 Apr 2026 11:30:07 +0800 Subject: [PATCH] fix(security): add CSP header & move WebSocket auth to initial message Tokens in WebSocket URL query params leak into server/proxy logs and browser history. Move authentication to a post-connect JSON message so the token never appears in the URL. Also add Content-Security-Policy middleware as a browser-level XSS defense layer. Co-Authored-By: Claude Opus 4.6 (1M context) --- packages/core/api/ws-client.ts | 37 ++++--- server/cmd/server/integration_test.go | 16 ++- server/cmd/server/router.go | 1 + server/internal/middleware/csp.go | 11 +++ server/internal/realtime/hub.go | 136 ++++++++++++++++---------- server/internal/realtime/hub_test.go | 17 +++- 6 files changed, 152 insertions(+), 66 deletions(-) create mode 100644 server/internal/middleware/csp.go diff --git a/packages/core/api/ws-client.ts b/packages/core/api/ws-client.ts index 23f68435f..238750107 100644 --- a/packages/core/api/ws-client.ts +++ b/packages/core/api/ws-client.ts @@ -27,28 +27,41 @@ export class WSClient { connect() { const url = new URL(this.baseUrl); - if (this.token) url.searchParams.set("token", this.token); if (this.workspaceId) url.searchParams.set("workspace_id", this.workspaceId); this.ws = new WebSocket(url.toString()); this.ws.onopen = () => { - this.logger.info("connected"); - if (this.hasConnectedBefore) { - for (const cb of this.onReconnectCallbacks) { - try { - cb(); - } catch { - // ignore reconnect callback errors - } - } + if (this.token) { + this.ws!.send(JSON.stringify({ type: "auth", token: this.token })); } - this.hasConnectedBefore = true; }; this.ws.onmessage = (event) => { - const msg = JSON.parse(event.data as string) as WSMessage; + const data = JSON.parse(event.data as string); + + if (data.type === "auth_ok") { + this.logger.info("connected"); + if (this.hasConnectedBefore) { + for (const cb of this.onReconnectCallbacks) { + try { + cb(); + } catch { + // ignore reconnect callback errors + } + } + } + this.hasConnectedBefore = true; + return; + } + + if (data.type === "auth_error") { + this.logger.error("ws auth failed:", data.error); + return; + } + + const msg = data as WSMessage; this.logger.debug("received", msg.type); const eventHandlers = this.handlers.get(msg.type); if (eventHandlers) { diff --git a/server/cmd/server/integration_test.go b/server/cmd/server/integration_test.go index 5bd1549d0..3d08e590b 100644 --- a/server/cmd/server/integration_test.go +++ b/server/cmd/server/integration_test.go @@ -748,13 +748,27 @@ func TestInvalidRequestBodies(t *testing.T) { func TestWebSocketIntegration(t *testing.T) { // Connect WebSocket client - wsURL := "ws" + strings.TrimPrefix(testServer.URL, "http") + "/ws?token=" + testToken + "&workspace_id=" + testWorkspaceID + wsURL := "ws" + strings.TrimPrefix(testServer.URL, "http") + "/ws?workspace_id=" + testWorkspaceID conn, _, err := websocket.DefaultDialer.Dial(wsURL, nil) if err != nil { t.Fatalf("WebSocket connection failed: %v", err) } defer conn.Close() + // Authenticate via initial message + if err := conn.WriteMessage(websocket.TextMessage, []byte(`{"type":"auth","token":"`+testToken+`"}`)); err != nil { + t.Fatalf("failed to send auth message: %v", err) + } + conn.SetReadDeadline(time.Now().Add(2 * time.Second)) + _, authResp, err := conn.ReadMessage() + if err != nil { + t.Fatalf("failed to read auth response: %v", err) + } + if string(authResp) != `{"type":"auth_ok"}` { + t.Fatalf("expected auth_ok, got: %s", authResp) + } + conn.SetReadDeadline(time.Time{}) + // Allow Hub goroutine to process the register and add client to room time.Sleep(100 * time.Millisecond) diff --git a/server/cmd/server/router.go b/server/cmd/server/router.go index 9b69a5512..9ba24a3b5 100644 --- a/server/cmd/server/router.go +++ b/server/cmd/server/router.go @@ -78,6 +78,7 @@ func NewRouter(pool *pgxpool.Pool, hub *realtime.Hub, bus *events.Bus) chi.Route r.Use(chimw.RequestID) r.Use(middleware.RequestLogger) r.Use(chimw.Recoverer) + r.Use(middleware.ContentSecurityPolicy) r.Use(cors.Handler(cors.Options{ AllowedOrigins: allowedOrigins(), AllowedMethods: []string{"GET", "POST", "PUT", "PATCH", "DELETE", "OPTIONS"}, diff --git a/server/internal/middleware/csp.go b/server/internal/middleware/csp.go new file mode 100644 index 000000000..6d1c32632 --- /dev/null +++ b/server/internal/middleware/csp.go @@ -0,0 +1,11 @@ +package middleware + +import "net/http" + +func ContentSecurityPolicy(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Security-Policy", + "default-src 'self'; script-src 'self'; style-src 'self' 'unsafe-inline'; img-src 'self' https: data:; connect-src 'self' wss:; frame-ancestors 'none'; object-src 'none'") + next.ServeHTTP(w, r) + }) +} diff --git a/server/internal/realtime/hub.go b/server/internal/realtime/hub.go index ae797b2dd..4041a4966 100644 --- a/server/internal/realtime/hub.go +++ b/server/internal/realtime/hub.go @@ -2,10 +2,12 @@ package realtime import ( "context" + "encoding/json" "log/slog" "net/http" "strings" "sync" + "time" "github.com/golang-jwt/jwt/v5" "github.com/gorilla/websocket" @@ -220,60 +222,21 @@ func (h *Hub) Broadcast(message []byte) { h.broadcast <- message } -// HandleWebSocket upgrades an HTTP connection to WebSocket with JWT or PAT auth. +const authTimeout = 5 * time.Second + +type authMessage struct { + Type string `json:"type"` + Token string `json:"token"` +} + +// HandleWebSocket upgrades an HTTP connection to WebSocket. +// Authentication is performed via the first message after connection establishment, +// keeping the token out of URL query parameters (and thus out of server logs, +// proxy logs, and browser history). func HandleWebSocket(hub *Hub, mc MembershipChecker, pr PATResolver, w http.ResponseWriter, r *http.Request) { - tokenStr := r.URL.Query().Get("token") workspaceID := r.URL.Query().Get("workspace_id") - - if tokenStr == "" || workspaceID == "" { - http.Error(w, `{"error":"token and workspace_id required"}`, http.StatusUnauthorized) - return - } - - var userID string - - if strings.HasPrefix(tokenStr, "mul_") { - // PAT authentication - if pr == nil { - http.Error(w, `{"error":"invalid token"}`, http.StatusUnauthorized) - return - } - uid, ok := pr.ResolveToken(r.Context(), tokenStr) - if !ok { - http.Error(w, `{"error":"invalid token"}`, http.StatusUnauthorized) - return - } - userID = uid - } else { - // JWT authentication - token, err := jwt.Parse(tokenStr, func(token *jwt.Token) (any, error) { - if _, ok := token.Method.(*jwt.SigningMethodHMAC); !ok { - return nil, jwt.ErrSignatureInvalid - } - return auth.JWTSecret(), nil - }) - if err != nil || !token.Valid { - http.Error(w, `{"error":"invalid token"}`, http.StatusUnauthorized) - return - } - - claims, ok := token.Claims.(jwt.MapClaims) - if !ok { - http.Error(w, `{"error":"invalid claims"}`, http.StatusUnauthorized) - return - } - - uid, ok := claims["sub"].(string) - if !ok || strings.TrimSpace(uid) == "" { - http.Error(w, `{"error":"invalid claims"}`, http.StatusUnauthorized) - return - } - userID = uid - } - - // Verify user is a member of the workspace - if !mc.IsMember(r.Context(), userID, workspaceID) { - http.Error(w, `{"error":"not a member of this workspace"}`, http.StatusForbidden) + if workspaceID == "" { + http.Error(w, `{"error":"workspace_id required"}`, http.StatusBadRequest) return } @@ -283,6 +246,75 @@ func HandleWebSocket(hub *Hub, mc MembershipChecker, pr PATResolver, w http.Resp return } + sendErr := func(msg string) { + errResp, _ := json.Marshal(map[string]string{"type": "auth_error", "error": msg}) + conn.WriteMessage(websocket.TextMessage, errResp) + conn.Close() + } + + conn.SetReadDeadline(time.Now().Add(authTimeout)) + + _, raw, err := conn.ReadMessage() + if err != nil { + slog.Debug("ws auth read failed", "error", err) + conn.Close() + return + } + + var msg authMessage + if err := json.Unmarshal(raw, &msg); err != nil || msg.Type != "auth" || msg.Token == "" { + sendErr("invalid auth message") + return + } + + tokenStr := msg.Token + var userID string + + if strings.HasPrefix(tokenStr, "mul_") { + if pr == nil { + sendErr("invalid token") + return + } + uid, ok := pr.ResolveToken(r.Context(), tokenStr) + if !ok { + sendErr("invalid token") + return + } + userID = uid + } else { + token, err := jwt.Parse(tokenStr, func(token *jwt.Token) (any, error) { + if _, ok := token.Method.(*jwt.SigningMethodHMAC); !ok { + return nil, jwt.ErrSignatureInvalid + } + return auth.JWTSecret(), nil + }) + if err != nil || !token.Valid { + sendErr("invalid token") + return + } + + claims, ok := token.Claims.(jwt.MapClaims) + if !ok { + sendErr("invalid claims") + return + } + + uid, ok := claims["sub"].(string) + if !ok || strings.TrimSpace(uid) == "" { + sendErr("invalid claims") + return + } + userID = uid + } + + if !mc.IsMember(r.Context(), userID, workspaceID) { + sendErr("not a member of this workspace") + return + } + + conn.SetReadDeadline(time.Time{}) + conn.WriteMessage(websocket.TextMessage, []byte(`{"type":"auth_ok"}`)) + client := &Client{ hub: hub, conn: conn, diff --git a/server/internal/realtime/hub_test.go b/server/internal/realtime/hub_test.go index 2187f8d58..9c31e39e0 100644 --- a/server/internal/realtime/hub_test.go +++ b/server/internal/realtime/hub_test.go @@ -52,11 +52,26 @@ func newTestHub(t *testing.T) (*Hub, *httptest.Server) { func connectWS(t *testing.T, server *httptest.Server) *websocket.Conn { t.Helper() token := makeTestToken(t) - wsURL := "ws" + strings.TrimPrefix(server.URL, "http") + "/ws?token=" + token + "&workspace_id=" + testWorkspaceID + wsURL := "ws" + strings.TrimPrefix(server.URL, "http") + "/ws?workspace_id=" + testWorkspaceID conn, _, err := websocket.DefaultDialer.Dial(wsURL, nil) if err != nil { t.Fatalf("failed to connect WebSocket: %v", err) } + + if err := conn.WriteMessage(websocket.TextMessage, []byte(`{"type":"auth","token":"`+token+`"}`)); err != nil { + t.Fatalf("failed to send auth message: %v", err) + } + + conn.SetReadDeadline(time.Now().Add(2 * time.Second)) + _, msg, err := conn.ReadMessage() + if err != nil { + t.Fatalf("failed to read auth response: %v", err) + } + if string(msg) != `{"type":"auth_ok"}` { + t.Fatalf("expected auth_ok, got: %s", msg) + } + conn.SetReadDeadline(time.Time{}) + return conn }