mirror of
https://github.com/Sagit-chu/flvx.git
synced 2026-10-08 18:56:37 +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 加密器,避免每条消息重建
|
crypto *security.AESCrypto // 缓存的 AES 加密器,避免每条消息重建
|
||||||
}
|
}
|
||||||
|
|
||||||
|
type adminSession struct {
|
||||||
|
userID int64
|
||||||
|
claims auth.Claims
|
||||||
|
conn *connWrap
|
||||||
|
}
|
||||||
|
|
||||||
type commandResponse struct {
|
type commandResponse struct {
|
||||||
Type string `json:"type"`
|
Type string `json:"type"`
|
||||||
Success bool `json:"success"`
|
Success bool `json:"success"`
|
||||||
@@ -77,7 +83,7 @@ type Server struct {
|
|||||||
getUserAuthState func(userID int64) (*auth.UserAuthState, error)
|
getUserAuthState func(userID int64) (*auth.UserAuthState, error)
|
||||||
|
|
||||||
mu sync.RWMutex
|
mu sync.RWMutex
|
||||||
admins map[*connWrap]struct{}
|
admins map[*adminSession]struct{}
|
||||||
nodes map[int64]*nodeSession
|
nodes map[int64]*nodeSession
|
||||||
byConn map[*websocket.Conn]*nodeSession
|
byConn map[*websocket.Conn]*nodeSession
|
||||||
pending map[string]pendingRequest
|
pending map[string]pendingRequest
|
||||||
@@ -124,7 +130,7 @@ func NewServer(repo *repo.Repository, jwtSecret string) *Server {
|
|||||||
upgrader: websocket.Upgrader{
|
upgrader: websocket.Upgrader{
|
||||||
CheckOrigin: func(r *http.Request) bool { return true },
|
CheckOrigin: func(r *http.Request) bool { return true },
|
||||||
},
|
},
|
||||||
admins: make(map[*connWrap]struct{}),
|
admins: make(map[*adminSession]struct{}),
|
||||||
nodes: make(map[int64]*nodeSession),
|
nodes: make(map[int64]*nodeSession),
|
||||||
byConn: make(map[*websocket.Conn]*nodeSession),
|
byConn: make(map[*websocket.Conn]*nodeSession),
|
||||||
pending: make(map[string]pendingRequest),
|
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)
|
http.Error(w, "forbidden", http.StatusForbidden)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
if s.getUserAuthState != nil {
|
userID, err := strconv.ParseInt(claims.Sub, 10, 64)
|
||||||
userID, err := strconv.ParseInt(claims.Sub, 10, 64)
|
if err != nil {
|
||||||
if err != nil {
|
http.Error(w, "forbidden", http.StatusForbidden)
|
||||||
http.Error(w, "forbidden", http.StatusForbidden)
|
return
|
||||||
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
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
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
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
http.Error(w, "bad request", http.StatusBadRequest)
|
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)
|
conn, err := s.upgrader.Upgrade(w, r, nil)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return
|
return
|
||||||
@@ -191,16 +198,19 @@ func (s *Server) handleAdmin(w http.ResponseWriter, r *http.Request) {
|
|||||||
return conn.SetReadDeadline(time.Now().Add(wsPongWait))
|
return conn.SetReadDeadline(time.Now().Add(wsPongWait))
|
||||||
})
|
})
|
||||||
done := make(chan struct{})
|
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.mu.Lock()
|
||||||
s.admins[cw] = struct{}{}
|
s.admins[session] = struct{}{}
|
||||||
s.mu.Unlock()
|
s.mu.Unlock()
|
||||||
|
|
||||||
defer func() {
|
defer func() {
|
||||||
close(done)
|
close(done)
|
||||||
s.mu.Lock()
|
s.mu.Lock()
|
||||||
delete(s.admins, cw)
|
delete(s.admins, session)
|
||||||
s.mu.Unlock()
|
s.mu.Unlock()
|
||||||
_ = conn.Close()
|
_ = 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))
|
return conn.SetReadDeadline(time.Now().Add(wsPongWait))
|
||||||
})
|
})
|
||||||
done := make(chan struct{})
|
done := make(chan struct{})
|
||||||
go startKeepalive(cw, done)
|
go startKeepalive(cw, done, nil)
|
||||||
|
|
||||||
version := r.URL.Query().Get("version")
|
version := r.URL.Query().Get("version")
|
||||||
httpVal := parseIntDefault(r.URL.Query().Get("http"), 0)
|
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) {
|
func (s *Server) broadcastToAdmins(message string) {
|
||||||
s.mu.RLock()
|
s.mu.RLock()
|
||||||
admins := make([]*connWrap, 0, len(s.admins))
|
admins := make([]*adminSession, 0, len(s.admins))
|
||||||
for c := range s.admins {
|
for c := range s.admins {
|
||||||
admins = append(admins, c)
|
admins = append(admins, c)
|
||||||
}
|
}
|
||||||
s.mu.RUnlock()
|
s.mu.RUnlock()
|
||||||
|
|
||||||
for _, c := range admins {
|
for _, c := range admins {
|
||||||
c.mu.Lock()
|
if c == nil || c.conn == nil || c.conn.conn == nil {
|
||||||
_ = c.conn.SetWriteDeadline(time.Now().Add(wsWriteWait))
|
continue
|
||||||
err := c.conn.WriteMessage(websocket.TextMessage, []byte(message))
|
}
|
||||||
_ = c.conn.SetWriteDeadline(time.Time{})
|
c.conn.mu.Lock()
|
||||||
c.mu.Unlock()
|
_ = 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 {
|
if err != nil {
|
||||||
log.Printf("websocket broadcast failed: %v", err)
|
log.Printf("websocket broadcast failed: %v", err)
|
||||||
}
|
}
|
||||||
@@ -625,7 +638,21 @@ func parseIntDefault(v string, fallback int) int {
|
|||||||
return x
|
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 {
|
if cw == nil || cw.conn == nil {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
@@ -637,6 +664,10 @@ func startKeepalive(cw *connWrap, done <-chan struct{}) {
|
|||||||
case <-done:
|
case <-done:
|
||||||
return
|
return
|
||||||
case <-ticker.C:
|
case <-ticker.C:
|
||||||
|
if validate != nil && !validate() {
|
||||||
|
_ = cw.conn.Close()
|
||||||
|
return
|
||||||
|
}
|
||||||
cw.mu.Lock()
|
cw.mu.Lock()
|
||||||
_ = cw.conn.SetWriteDeadline(time.Now().Add(wsWriteWait))
|
_ = cw.conn.SetWriteDeadline(time.Now().Add(wsWriteWait))
|
||||||
err := cw.conn.WriteMessage(websocket.PingMessage, nil)
|
err := cw.conn.WriteMessage(websocket.PingMessage, nil)
|
||||||
|
|||||||
@@ -3,6 +3,7 @@ package ws
|
|||||||
import (
|
import (
|
||||||
"net/http"
|
"net/http"
|
||||||
"net/http/httptest"
|
"net/http/httptest"
|
||||||
|
"net/url"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
"go-backend/internal/auth"
|
"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
|
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()
|
rec := httptest.NewRecorder()
|
||||||
|
|
||||||
server.ServeHTTP(rec, req)
|
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)
|
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_favicon",
|
||||||
"app_bg_image",
|
"app_bg_image",
|
||||||
] as const;
|
] as const;
|
||||||
|
const SENSITIVE_CONFIG_KEYS = new Set([
|
||||||
|
"jwt_secret",
|
||||||
|
"license_key",
|
||||||
|
"cloudflare_secret_key",
|
||||||
|
]);
|
||||||
const GITHUB_REPO =
|
const GITHUB_REPO =
|
||||||
import.meta.env.VITE_GITHUB_REPO || "https://github.com/Sagit-chu/flux-panel";
|
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 readCachedConfigs = (keys: readonly string[]) => {
|
||||||
const cachedConfigs: Record<string, string> = {};
|
const cachedConfigs: Record<string, string> = {};
|
||||||
let hasCachedData = false;
|
let hasCachedData = false;
|
||||||
@@ -76,6 +91,8 @@ const getInitialConfig = () => {
|
|||||||
};
|
};
|
||||||
}
|
}
|
||||||
|
|
||||||
|
purgeSensitiveConfigCache();
|
||||||
|
|
||||||
const cachedAppName = localStorage.getItem(CACHE_PREFIX + "app_name");
|
const cachedAppName = localStorage.getItem(CACHE_PREFIX + "app_name");
|
||||||
const cachedAppLogo = localStorage.getItem(CACHE_PREFIX + "app_logo") || "";
|
const cachedAppLogo = localStorage.getItem(CACHE_PREFIX + "app_logo") || "";
|
||||||
const cachedAppFavicon =
|
const cachedAppFavicon =
|
||||||
@@ -169,7 +186,9 @@ export const getCachedConfig = async (key: string): Promise<string | null> => {
|
|||||||
) {
|
) {
|
||||||
const value = response.data.value;
|
const value = response.data.value;
|
||||||
|
|
||||||
configCache.set(key, value);
|
if (shouldPersistConfigKey(key)) {
|
||||||
|
configCache.set(key, value);
|
||||||
|
}
|
||||||
|
|
||||||
return value;
|
return value;
|
||||||
}
|
}
|
||||||
@@ -200,9 +219,19 @@ export const getCachedConfigs = async (): Promise<Record<string, string>> => {
|
|||||||
if (response.code === 0 && response.data) {
|
if (response.code === 0 && response.data) {
|
||||||
const configs = response.data;
|
const configs = response.data;
|
||||||
|
|
||||||
// 将所有配置存入缓存
|
// 仅将安全配置存入缓存,敏感项会从 localStorage 中移除
|
||||||
Object.entries(configs).forEach(([key, value]) => {
|
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;
|
return configs;
|
||||||
|
|||||||
Reference in New Issue
Block a user