mirror of
https://github.com/Sagit-chu/flvx.git
synced 2026-09-29 16:06:36 +08:00
fix: tighten websocket auth and client cache handling
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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;
|
||||
|
||||
Reference in New Issue
Block a user