From 106a30bf9d2f5a71790d322a89d21ae93f974aaf Mon Sep 17 00:00:00 2001 From: sagitchu Date: Thu, 14 May 2026 00:24:56 +0800 Subject: [PATCH] fix: tighten websocket auth and client cache handling --- go-backend/internal/ws/server.go | 83 ++++++++++++++++++--------- go-backend/internal/ws/server_test.go | 59 ++++++++++++++++++- vite-frontend/src/config/site.ts | 35 ++++++++++- 3 files changed, 147 insertions(+), 30 deletions(-) diff --git a/go-backend/internal/ws/server.go b/go-backend/internal/ws/server.go index 34fa84b..abb84bb 100644 --- a/go-backend/internal/ws/server.go +++ b/go-backend/internal/ws/server.go @@ -42,6 +42,12 @@ type nodeSession struct { crypto *security.AESCrypto // 缓存的 AES 加密器,避免每条消息重建 } +type adminSession struct { + userID int64 + claims auth.Claims + conn *connWrap +} + type commandResponse struct { Type string `json:"type"` Success bool `json:"success"` @@ -77,7 +83,7 @@ type Server struct { getUserAuthState func(userID int64) (*auth.UserAuthState, error) mu sync.RWMutex - admins map[*connWrap]struct{} + admins map[*adminSession]struct{} nodes map[int64]*nodeSession byConn map[*websocket.Conn]*nodeSession pending map[string]pendingRequest @@ -124,7 +130,7 @@ func NewServer(repo *repo.Repository, jwtSecret string) *Server { upgrader: websocket.Upgrader{ CheckOrigin: func(r *http.Request) bool { return true }, }, - admins: make(map[*connWrap]struct{}), + admins: make(map[*adminSession]struct{}), nodes: make(map[int64]*nodeSession), byConn: make(map[*websocket.Conn]*nodeSession), pending: make(map[string]pendingRequest), @@ -161,26 +167,27 @@ func (s *Server) ServeHTTP(w http.ResponseWriter, r *http.Request) { http.Error(w, "forbidden", http.StatusForbidden) return } - if s.getUserAuthState != nil { - userID, err := strconv.ParseInt(claims.Sub, 10, 64) - if err != nil { - http.Error(w, "forbidden", http.StatusForbidden) - return - } - state, err := s.getUserAuthState(userID) - if err != nil || state == nil || state.Status != 1 || state.RoleID != claims.RoleID || claims.IatMs <= state.PasswordChangedAt { - http.Error(w, "forbidden", http.StatusForbidden) - return - } + userID, err := strconv.ParseInt(claims.Sub, 10, 64) + if err != nil { + http.Error(w, "forbidden", http.StatusForbidden) + return } - s.handleAdmin(w, r) + if claims.RoleID != 0 { + http.Error(w, "forbidden", http.StatusForbidden) + return + } + if !s.validateAdminSession(userID, claims) { + http.Error(w, "forbidden", http.StatusForbidden) + return + } + s.handleAdmin(w, r, userID, claims) return } http.Error(w, "bad request", http.StatusBadRequest) } -func (s *Server) handleAdmin(w http.ResponseWriter, r *http.Request) { +func (s *Server) handleAdmin(w http.ResponseWriter, r *http.Request, userID int64, claims auth.Claims) { conn, err := s.upgrader.Upgrade(w, r, nil) if err != nil { return @@ -191,16 +198,19 @@ func (s *Server) handleAdmin(w http.ResponseWriter, r *http.Request) { return conn.SetReadDeadline(time.Now().Add(wsPongWait)) }) done := make(chan struct{}) - go startKeepalive(cw, done) + session := &adminSession{userID: userID, claims: claims, conn: cw} + go startKeepalive(cw, done, func() bool { + return s.validateAdminSession(session.userID, session.claims) + }) s.mu.Lock() - s.admins[cw] = struct{}{} + s.admins[session] = struct{}{} s.mu.Unlock() defer func() { close(done) s.mu.Lock() - delete(s.admins, cw) + delete(s.admins, session) s.mu.Unlock() _ = conn.Close() }() @@ -223,7 +233,7 @@ func (s *Server) handleNode(w http.ResponseWriter, r *http.Request, nodeID int64 return conn.SetReadDeadline(time.Now().Add(wsPongWait)) }) done := make(chan struct{}) - go startKeepalive(cw, done) + go startKeepalive(cw, done, nil) version := r.URL.Query().Get("version") httpVal := parseIntDefault(r.URL.Query().Get("http"), 0) @@ -577,18 +587,21 @@ func (s *Server) broadcastTyped(nodeID int64, msgType string, data string) { func (s *Server) broadcastToAdmins(message string) { s.mu.RLock() - admins := make([]*connWrap, 0, len(s.admins)) + admins := make([]*adminSession, 0, len(s.admins)) for c := range s.admins { admins = append(admins, c) } s.mu.RUnlock() for _, c := range admins { - c.mu.Lock() - _ = c.conn.SetWriteDeadline(time.Now().Add(wsWriteWait)) - err := c.conn.WriteMessage(websocket.TextMessage, []byte(message)) - _ = c.conn.SetWriteDeadline(time.Time{}) - c.mu.Unlock() + if c == nil || c.conn == nil || c.conn.conn == nil { + continue + } + c.conn.mu.Lock() + _ = c.conn.conn.SetWriteDeadline(time.Now().Add(wsWriteWait)) + err := c.conn.conn.WriteMessage(websocket.TextMessage, []byte(message)) + _ = c.conn.conn.SetWriteDeadline(time.Time{}) + c.conn.mu.Unlock() if err != nil { log.Printf("websocket broadcast failed: %v", err) } @@ -625,7 +638,21 @@ func parseIntDefault(v string, fallback int) int { return x } -func startKeepalive(cw *connWrap, done <-chan struct{}) { +func (s *Server) validateAdminSession(userID int64, claims auth.Claims) bool { + if s == nil { + return false + } + if s.getUserAuthState == nil { + return true + } + state, err := s.getUserAuthState(userID) + if err != nil || state == nil || state.Status != 1 || state.RoleID != claims.RoleID || claims.IatMs <= state.PasswordChangedAt { + return false + } + return true +} + +func startKeepalive(cw *connWrap, done <-chan struct{}, validate func() bool) { if cw == nil || cw.conn == nil { return } @@ -637,6 +664,10 @@ func startKeepalive(cw *connWrap, done <-chan struct{}) { case <-done: return case <-ticker.C: + if validate != nil && !validate() { + _ = cw.conn.Close() + return + } cw.mu.Lock() _ = cw.conn.SetWriteDeadline(time.Now().Add(wsWriteWait)) err := cw.conn.WriteMessage(websocket.PingMessage, nil) diff --git a/go-backend/internal/ws/server_test.go b/go-backend/internal/ws/server_test.go index 50c7b9f..51f902b 100644 --- a/go-backend/internal/ws/server_test.go +++ b/go-backend/internal/ws/server_test.go @@ -3,6 +3,7 @@ package ws import ( "net/http" "net/http/httptest" + "net/url" "testing" "go-backend/internal/auth" @@ -20,7 +21,7 @@ func TestServeHTTPRejectsDisabledAdminToken(t *testing.T) { return &auth.UserAuthState{ID: userID, RoleID: 0, Status: 0, PasswordChangedAt: 0}, nil }) - req := httptest.NewRequest(http.MethodGet, "/system-info?type=0&secret="+token, nil) + req := httptest.NewRequest(http.MethodGet, "/system-info?type=0&secret="+url.QueryEscape(token), nil) rec := httptest.NewRecorder() server.ServeHTTP(rec, req) @@ -29,3 +30,59 @@ func TestServeHTTPRejectsDisabledAdminToken(t *testing.T) { t.Fatalf("expected forbidden for disabled admin token, got %d", rec.Code) } } + +func TestServeHTTPRejectsNonAdminToken(t *testing.T) { + secret := "unit-test-secret" + token, err := auth.GenerateToken(2, "normal_user", 1, secret) + if err != nil { + t.Fatalf("generate token: %v", err) + } + + server := NewServer(nil, secret) + server.SetUserAuthStateLookup(func(userID int64) (*auth.UserAuthState, error) { + return &auth.UserAuthState{ID: userID, RoleID: 1, Status: 1, PasswordChangedAt: 0}, nil + }) + + req := httptest.NewRequest(http.MethodGet, "/system-info?type=0&secret="+url.QueryEscape(token), nil) + rec := httptest.NewRecorder() + + server.ServeHTTP(rec, req) + + if rec.Code != http.StatusForbidden { + t.Fatalf("expected forbidden for non-admin token, got %d", rec.Code) + } +} + +func TestValidateAdminSessionRejectsAuthStateChanges(t *testing.T) { + secret := "unit-test-secret" + token, err := auth.GenerateToken(1, "admin_user", 0, secret) + if err != nil { + t.Fatalf("generate token: %v", err) + } + claims, err := auth.ParseClaims(token, secret) + if err != nil { + t.Fatalf("parse claims: %v", err) + } + + tests := []struct { + name string + state *auth.UserAuthState + }{ + {name: "disabled", state: &auth.UserAuthState{ID: 1, RoleID: 0, Status: 0, PasswordChangedAt: 0}}, + {name: "role changed", state: &auth.UserAuthState{ID: 1, RoleID: 1, Status: 1, PasswordChangedAt: 0}}, + {name: "password changed", state: &auth.UserAuthState{ID: 1, RoleID: 0, Status: 1, PasswordChangedAt: claims.IatMs + 1}}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + server := NewServer(nil, secret) + server.SetUserAuthStateLookup(func(userID int64) (*auth.UserAuthState, error) { + return tt.state, nil + }) + + if ok := server.validateAdminSession(1, claims); ok { + t.Fatalf("expected session validation to fail for %s state", tt.name) + } + }) + } +} diff --git a/vite-frontend/src/config/site.ts b/vite-frontend/src/config/site.ts index c47a0b9..0b6b066 100644 --- a/vite-frontend/src/config/site.ts +++ b/vite-frontend/src/config/site.ts @@ -15,9 +15,24 @@ const PUBLIC_BRAND_CONFIG_KEYS = [ "app_favicon", "app_bg_image", ] as const; +const SENSITIVE_CONFIG_KEYS = new Set([ + "jwt_secret", + "license_key", + "cloudflare_secret_key", +]); const GITHUB_REPO = import.meta.env.VITE_GITHUB_REPO || "https://github.com/Sagit-chu/flux-panel"; +const shouldPersistConfigKey = (key: string) => { + return !SENSITIVE_CONFIG_KEYS.has(key.trim().toLowerCase()); +}; + +const purgeSensitiveConfigCache = () => { + SENSITIVE_CONFIG_KEYS.forEach((key) => { + localStorage.removeItem(CACHE_PREFIX + key); + }); +}; + const readCachedConfigs = (keys: readonly string[]) => { const cachedConfigs: Record = {}; let hasCachedData = false; @@ -76,6 +91,8 @@ const getInitialConfig = () => { }; } + purgeSensitiveConfigCache(); + const cachedAppName = localStorage.getItem(CACHE_PREFIX + "app_name"); const cachedAppLogo = localStorage.getItem(CACHE_PREFIX + "app_logo") || ""; const cachedAppFavicon = @@ -169,7 +186,9 @@ export const getCachedConfig = async (key: string): Promise => { ) { const value = response.data.value; - configCache.set(key, value); + if (shouldPersistConfigKey(key)) { + configCache.set(key, value); + } return value; } @@ -200,9 +219,19 @@ export const getCachedConfigs = async (): Promise> => { if (response.code === 0 && response.data) { const configs = response.data; - // 将所有配置存入缓存 + // 仅将安全配置存入缓存,敏感项会从 localStorage 中移除 Object.entries(configs).forEach(([key, value]) => { - configCache.set(key, value as string); + const normalizedKey = key.trim().toLowerCase(); + + if (SENSITIVE_CONFIG_KEYS.has(normalizedKey)) { + configCache.remove(normalizedKey); + + return; + } + + if (shouldPersistConfigKey(normalizedKey)) { + configCache.set(normalizedKey, value as string); + } }); return configs;