fix: tighten websocket auth and client cache handling

This commit is contained in:
sagitchu
2026-05-14 00:24:56 +08:00
parent 465815cf34
commit 106a30bf9d
3 changed files with 147 additions and 30 deletions
+57 -26
View File
@@ -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)
+58 -1
View File
@@ -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)
}
})
}
}
+32 -3
View File
@@ -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<string, string> = {};
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<string | null> => {
) {
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<Record<string, string>> => {
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;