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) <noreply@anthropic.com>
This commit is contained in:
yushen
2026-04-13 11:30:07 +08:00
parent 04238bea22
commit 50512ca415
6 changed files with 152 additions and 66 deletions

View File

@@ -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) {

View File

@@ -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)

View File

@@ -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"},

View File

@@ -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)
})
}

View File

@@ -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,

View File

@@ -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
}