mirror of
https://github.com/Sagit-chu/flvx.git
synced 2026-09-28 07:36:38 +08:00
fix: harden auth, config access, and backups
This commit is contained in:
@@ -12,12 +12,13 @@ import (
|
||||
|
||||
const (
|
||||
algorithm = "HmacSHA256"
|
||||
expireTime = 90 * 24 * time.Hour
|
||||
expireTime = 7 * 24 * time.Hour
|
||||
)
|
||||
|
||||
type Claims struct {
|
||||
Sub string `json:"sub"`
|
||||
Iat int64 `json:"iat"`
|
||||
IatMs int64 `json:"iat_ms"`
|
||||
Exp int64 `json:"exp"`
|
||||
User string `json:"user"`
|
||||
Name string `json:"name"`
|
||||
@@ -30,11 +31,15 @@ type tokenHeader struct {
|
||||
}
|
||||
|
||||
func GenerateToken(userID int64, username string, roleID int, secret string) (string, error) {
|
||||
now := time.Now()
|
||||
return GenerateTokenAt(userID, username, roleID, secret, time.Now())
|
||||
}
|
||||
|
||||
func GenerateTokenAt(userID int64, username string, roleID int, secret string, now time.Time) (string, error) {
|
||||
header := tokenHeader{Alg: algorithm, Typ: "JWT"}
|
||||
claims := Claims{
|
||||
Sub: strconv.FormatInt(userID, 10),
|
||||
Iat: now.Unix(),
|
||||
IatMs: now.UnixMilli(),
|
||||
Exp: now.Add(expireTime).Unix(),
|
||||
User: username,
|
||||
Name: username,
|
||||
|
||||
@@ -0,0 +1,8 @@
|
||||
package auth
|
||||
|
||||
type UserAuthState struct {
|
||||
ID int64
|
||||
RoleID int
|
||||
Status int
|
||||
PasswordChangedAt int64
|
||||
}
|
||||
@@ -0,0 +1,44 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"strings"
|
||||
|
||||
"go-backend/internal/http/response"
|
||||
"go-backend/internal/store/repo"
|
||||
)
|
||||
|
||||
func (h *Handler) getPublicConfigByName(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodPost {
|
||||
response.WriteJSON(w, response.ErrDefault("请求失败"))
|
||||
return
|
||||
}
|
||||
|
||||
var req nameRequest
|
||||
if err := decodeJSON(r.Body, &req); err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault("配置名称不能为空"))
|
||||
return
|
||||
}
|
||||
|
||||
configName := strings.ToLower(strings.TrimSpace(req.Name))
|
||||
if configName == "" {
|
||||
response.WriteJSON(w, response.ErrDefault("配置名称不能为空"))
|
||||
return
|
||||
}
|
||||
if !repo.IsPublicConfigKey(configName) {
|
||||
response.WriteJSON(w, response.Err(403, "禁止访问敏感配置"))
|
||||
return
|
||||
}
|
||||
|
||||
cfg, err := h.repo.GetConfigByName(configName)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if cfg == nil {
|
||||
response.WriteJSON(w, response.ErrDefault("配置不存在"))
|
||||
return
|
||||
}
|
||||
|
||||
response.WriteJSON(w, response.OK(cfg))
|
||||
}
|
||||
@@ -0,0 +1,103 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
"go-backend/internal/http/middleware"
|
||||
"go-backend/internal/http/response"
|
||||
"go-backend/internal/store/repo"
|
||||
)
|
||||
|
||||
func TestPublicConfigGetAllowsBrandKeys(t *testing.T) {
|
||||
router, r := setupConfigAccessTestRouter(t)
|
||||
seedConfigValue(t, r, "app_name", "FLVX Brand")
|
||||
seedConfigValue(t, r, "app_logo", "logo-data")
|
||||
seedConfigValue(t, r, "app_favicon", "favicon-data")
|
||||
seedConfigValue(t, r, "app_bg_image", "bg-data")
|
||||
seedConfigValue(t, r, "cloudflare_site_key", "site-key")
|
||||
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/public/config/get", bytes.NewBufferString(`{"name":"app_name"}`))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
resp := httptest.NewRecorder()
|
||||
|
||||
router.ServeHTTP(resp, req)
|
||||
|
||||
assertHandlerCode(t, resp, 0)
|
||||
}
|
||||
|
||||
func TestPublicConfigGetRejectsSensitiveKeys(t *testing.T) {
|
||||
router, _ := setupConfigAccessTestRouter(t)
|
||||
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/public/config/get", bytes.NewBufferString(`{"name":"jwt_secret"}`))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
resp := httptest.NewRecorder()
|
||||
|
||||
router.ServeHTTP(resp, req)
|
||||
|
||||
assertHandlerCodeMsg(t, resp, 403, "禁止访问敏感配置")
|
||||
}
|
||||
|
||||
func TestConfigGetNowRequiresAuth(t *testing.T) {
|
||||
router, _ := setupConfigAccessTestRouter(t)
|
||||
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/config/get", bytes.NewBufferString(`{"name":"app_name"}`))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
resp := httptest.NewRecorder()
|
||||
|
||||
router.ServeHTTP(resp, req)
|
||||
|
||||
assertHandlerCodeMsg(t, resp, 401, "未登录或token已过期")
|
||||
}
|
||||
|
||||
func setupConfigAccessTestRouter(t *testing.T) (http.Handler, *repo.Repository) {
|
||||
t.Helper()
|
||||
r, err := repo.Open(t.TempDir() + "/config-access.db")
|
||||
if err != nil {
|
||||
t.Fatalf("open sqlite: %v", err)
|
||||
}
|
||||
t.Cleanup(func() {
|
||||
_ = r.Close()
|
||||
})
|
||||
|
||||
h := New(r, "unit-test-secret")
|
||||
mux := http.NewServeMux()
|
||||
h.Register(mux)
|
||||
wrapped := middleware.Recover(mux)
|
||||
wrapped = middleware.JWT(middleware.AuthOptions{JWTSecret: "unit-test-secret", GetUserAuthState: h.GetUserAuthState})(wrapped)
|
||||
wrapped = middleware.RequestLog(wrapped)
|
||||
wrapped = middleware.CORS(wrapped)
|
||||
return wrapped, r
|
||||
}
|
||||
|
||||
func seedConfigValue(t *testing.T, r *repo.Repository, name, value string) {
|
||||
t.Helper()
|
||||
if err := r.DB().Exec(`INSERT INTO vite_config(name, value, time) VALUES(?, ?, 0) ON CONFLICT(name) DO UPDATE SET value = excluded.value, time = excluded.time`, name, value).Error; err != nil {
|
||||
t.Fatalf("seed config %s: %v", name, err)
|
||||
}
|
||||
}
|
||||
|
||||
func assertHandlerCode(t *testing.T, rec *httptest.ResponseRecorder, expected int) {
|
||||
t.Helper()
|
||||
var out response.R
|
||||
if err := json.NewDecoder(rec.Body).Decode(&out); err != nil {
|
||||
t.Fatalf("decode response: %v", err)
|
||||
}
|
||||
if out.Code != expected {
|
||||
t.Fatalf("expected code %d, got %d", expected, out.Code)
|
||||
}
|
||||
}
|
||||
|
||||
func assertHandlerCodeMsg(t *testing.T, rec *httptest.ResponseRecorder, expectedCode int, expectedMsg string) {
|
||||
t.Helper()
|
||||
var out response.R
|
||||
if err := json.NewDecoder(rec.Body).Decode(&out); err != nil {
|
||||
t.Fatalf("decode response: %v", err)
|
||||
}
|
||||
if out.Code != expectedCode || out.Msg != expectedMsg {
|
||||
t.Fatalf("expected (%d,%q), got (%d,%q)", expectedCode, expectedMsg, out.Code, out.Msg)
|
||||
}
|
||||
}
|
||||
@@ -136,6 +136,7 @@ func New(repo *repo.Repository, jwtSecret string) *Handler {
|
||||
}
|
||||
h.metrics.RecordNodeMetric(nodeID, metricInfo)
|
||||
})
|
||||
h.wsServer.SetUserAuthStateLookup(h.GetUserAuthState)
|
||||
return h
|
||||
}
|
||||
|
||||
@@ -143,6 +144,10 @@ func (h *Handler) WebSocketHandler() http.Handler {
|
||||
return h.wsServer
|
||||
}
|
||||
|
||||
func (h *Handler) GetUserAuthState(userID int64) (*auth.UserAuthState, error) {
|
||||
return h.repo.GetUserAuthState(userID)
|
||||
}
|
||||
|
||||
func (h *Handler) Register(mux *http.ServeMux) {
|
||||
mux.HandleFunc("/api/v1/user/login", h.login)
|
||||
mux.HandleFunc("/api/v1/user/list", h.userList)
|
||||
@@ -152,6 +157,7 @@ func (h *Handler) Register(mux *http.ServeMux) {
|
||||
mux.HandleFunc("/api/v1/user/reset", h.userResetFlow)
|
||||
mux.HandleFunc("/api/v1/user/quota/reset", h.userQuotaReset)
|
||||
mux.HandleFunc("/api/v1/user/groups", h.userGroups)
|
||||
mux.HandleFunc("/api/v1/public/config/get", h.getPublicConfigByName)
|
||||
mux.HandleFunc("/api/v1/config/get", h.getConfigByName)
|
||||
mux.HandleFunc("/api/v1/config/list", h.getConfigs)
|
||||
mux.HandleFunc("/api/v1/config/update", h.updateConfigs)
|
||||
@@ -334,7 +340,8 @@ func (h *Handler) login(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.ErrDefault("账号或密码错误"))
|
||||
return
|
||||
}
|
||||
if user.Pwd != security.MD5(req.Password) {
|
||||
passwordMatched, passwordWasLegacy := security.VerifyPassword(user.Pwd, req.Password)
|
||||
if !passwordMatched {
|
||||
response.WriteJSON(w, response.ErrDefault("账号或密码错误"))
|
||||
return
|
||||
}
|
||||
@@ -342,8 +349,20 @@ func (h *Handler) login(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.ErrDefault("账号被停用"))
|
||||
return
|
||||
}
|
||||
issueAt := time.Now()
|
||||
if passwordWasLegacy {
|
||||
updatedAt := time.Now().UnixMilli()
|
||||
hashedPassword, err := security.HashPassword(req.Password)
|
||||
if err != nil {
|
||||
log.Printf("legacy password rehash skipped user_id=%d path=login err=%v", user.ID, err)
|
||||
} else if err := h.repo.UpdateUserPassword(user.ID, hashedPassword, updatedAt); err != nil {
|
||||
log.Printf("legacy password rehash update skipped user_id=%d path=login err=%v", user.ID, err)
|
||||
} else {
|
||||
issueAt = time.UnixMilli(updatedAt + 1)
|
||||
}
|
||||
}
|
||||
|
||||
token, err := auth.GenerateToken(user.ID, user.User, user.RoleID, h.jwtSecret)
|
||||
token, err := auth.GenerateTokenAt(user.ID, user.User, user.RoleID, h.jwtSecret, issueAt)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
@@ -557,10 +576,27 @@ func (h *Handler) openAPISubStore(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if user == nil || user.Pwd != security.MD5(password) {
|
||||
if user == nil {
|
||||
response.WriteJSON(w, response.ErrDefault("鉴权失败"))
|
||||
return
|
||||
}
|
||||
passwordMatched, passwordWasLegacy := security.VerifyPassword(user.Pwd, password)
|
||||
if !passwordMatched {
|
||||
response.WriteJSON(w, response.ErrDefault("鉴权失败"))
|
||||
return
|
||||
}
|
||||
if user.Status == 0 {
|
||||
response.WriteJSON(w, response.ErrDefault("账号被停用"))
|
||||
return
|
||||
}
|
||||
if passwordWasLegacy {
|
||||
hashedPassword, err := security.HashPassword(password)
|
||||
if err != nil {
|
||||
log.Printf("legacy password rehash skipped user_id=%d path=sub_store err=%v", user.ID, err)
|
||||
} else if err := h.repo.UpdateUserPassword(user.ID, hashedPassword, time.Now().UnixMilli()); err != nil {
|
||||
log.Printf("legacy password rehash update skipped user_id=%d path=sub_store err=%v", user.ID, err)
|
||||
}
|
||||
}
|
||||
|
||||
const giga = int64(1024 * 1024 * 1024)
|
||||
headerValue := ""
|
||||
@@ -1255,7 +1291,8 @@ func (h *Handler) updatePassword(w http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
|
||||
if user.Pwd != security.MD5(req.CurrentPassword) {
|
||||
passwordMatched, _ := security.VerifyPassword(user.Pwd, req.CurrentPassword)
|
||||
if !passwordMatched {
|
||||
response.WriteJSON(w, response.ErrDefault("当前密码错误"))
|
||||
return
|
||||
}
|
||||
@@ -1270,7 +1307,12 @@ func (h *Handler) updatePassword(w http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
|
||||
if err := h.repo.UpdateUserNameAndPassword(userID, req.NewUsername, security.MD5(req.NewPassword), time.Now().UnixMilli()); err != nil {
|
||||
hashedPassword, err := security.HashPassword(req.NewPassword)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if err := h.repo.UpdateUserNameAndPassword(userID, req.NewUsername, hashedPassword, time.Now().UnixMilli()); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
@@ -71,7 +71,12 @@ func (h *Handler) userCreate(w http.ResponseWriter, r *http.Request) {
|
||||
now := time.Now().UnixMilli()
|
||||
maxConn := asInt(req["maxConn"], 0)
|
||||
|
||||
userID, err := h.repo.CreateUser(username, security.MD5(pwd), roleID, expTime, flow, flowResetTime, num, status, maxConn, now)
|
||||
hashedPassword, err := security.HashPassword(pwd)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
userID, err := h.repo.CreateUser(username, hashedPassword, roleID, expTime, flow, flowResetTime, num, status, maxConn, now)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
@@ -176,7 +181,12 @@ func (h *Handler) userUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
} else {
|
||||
if err := h.repo.UpdateUserWithPassword(id, username, security.MD5(pwd), flow, num, expTime, flowResetTime, status, maxConn, now); err != nil {
|
||||
hashedPassword, err := security.HashPassword(pwd)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if err := h.repo.UpdateUserWithPassword(id, username, hashedPassword, flow, num, expTime, flowResetTime, status, maxConn, now); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
@@ -3,6 +3,7 @@ package middleware
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
"go-backend/internal/auth"
|
||||
@@ -14,7 +15,8 @@ type contextKey string
|
||||
const ClaimsContextKey contextKey = "claims"
|
||||
|
||||
type AuthOptions struct {
|
||||
JWTSecret string
|
||||
JWTSecret string
|
||||
GetUserAuthState func(userID int64) (*auth.UserAuthState, error)
|
||||
}
|
||||
|
||||
func JWT(opts AuthOptions) func(http.Handler) http.Handler {
|
||||
@@ -42,6 +44,19 @@ func JWT(opts AuthOptions) func(http.Handler) http.Handler {
|
||||
return
|
||||
}
|
||||
|
||||
if opts.GetUserAuthState != nil {
|
||||
userID, err := strconv.ParseInt(claims.Sub, 10, 64)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(401, "无效的token或token已过期"))
|
||||
return
|
||||
}
|
||||
state, err := opts.GetUserAuthState(userID)
|
||||
if err != nil || state == nil || state.Status != 1 || state.RoleID != claims.RoleID || claims.IatMs <= state.PasswordChangedAt {
|
||||
response.WriteJSON(w, response.Err(401, "无效的token或token已过期"))
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
if requiresAdmin(r.URL.Path) && claims.RoleID != 0 {
|
||||
response.WriteJSON(w, response.Err(403, "权限不足,仅管理员可操作"))
|
||||
return
|
||||
@@ -78,9 +93,11 @@ func shouldSkip(path string) bool {
|
||||
case strings.HasPrefix(path, "/api/v1/captcha/"):
|
||||
return true
|
||||
case path == "/api/v1/config/get":
|
||||
return true
|
||||
return false
|
||||
case path == "/api/v1/user/login":
|
||||
return true
|
||||
case path == "/api/v1/public/config/get":
|
||||
return true
|
||||
case path == "/api/v1/federation/connect":
|
||||
return true
|
||||
case path == "/api/v1/federation/tunnel/create":
|
||||
|
||||
@@ -0,0 +1,243 @@
|
||||
package middleware
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/auth"
|
||||
"go-backend/internal/http/response"
|
||||
)
|
||||
|
||||
func TestJWTRejectsPasswordChangedToken(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)
|
||||
}
|
||||
|
||||
wrapped := JWT(AuthOptions{
|
||||
JWTSecret: secret,
|
||||
GetUserAuthState: func(userID int64) (*auth.UserAuthState, error) {
|
||||
return &auth.UserAuthState{ID: userID, RoleID: 0, Status: 1, PasswordChangedAt: claims.IatMs + 1}, nil
|
||||
},
|
||||
})(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.OK("pass"))
|
||||
}))
|
||||
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/list", nil)
|
||||
req.Header.Set("Authorization", token)
|
||||
res := httptest.NewRecorder()
|
||||
wrapped.ServeHTTP(res, req)
|
||||
assertAuthDenied(t, res)
|
||||
}
|
||||
|
||||
func TestJWTAcceptsCurrentUserState(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)
|
||||
}
|
||||
|
||||
wrapped := JWT(AuthOptions{
|
||||
JWTSecret: secret,
|
||||
GetUserAuthState: func(userID int64) (*auth.UserAuthState, error) {
|
||||
return &auth.UserAuthState{ID: userID, RoleID: 0, Status: 1, PasswordChangedAt: claims.IatMs - 1}, nil
|
||||
},
|
||||
})(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.OK("pass"))
|
||||
}))
|
||||
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/list", nil)
|
||||
req.Header.Set("Authorization", token)
|
||||
res := httptest.NewRecorder()
|
||||
wrapped.ServeHTTP(res, req)
|
||||
assertCode(t, res, 0)
|
||||
}
|
||||
|
||||
func TestJWTRejectsDisabledUserToken(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)
|
||||
}
|
||||
|
||||
wrapped := JWT(AuthOptions{
|
||||
JWTSecret: secret,
|
||||
GetUserAuthState: func(userID int64) (*auth.UserAuthState, error) {
|
||||
return &auth.UserAuthState{ID: userID, RoleID: 0, Status: 0, PasswordChangedAt: time.Now().Unix()}, nil
|
||||
},
|
||||
})(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.OK("pass"))
|
||||
}))
|
||||
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/list", nil)
|
||||
req.Header.Set("Authorization", token)
|
||||
res := httptest.NewRecorder()
|
||||
wrapped.ServeHTTP(res, req)
|
||||
assertAuthDenied(t, res)
|
||||
}
|
||||
|
||||
func TestJWTRejectsPasswordChangedAtSameMillisecond(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)
|
||||
}
|
||||
|
||||
wrapped := JWT(AuthOptions{
|
||||
JWTSecret: secret,
|
||||
GetUserAuthState: func(userID int64) (*auth.UserAuthState, error) {
|
||||
return &auth.UserAuthState{ID: userID, RoleID: 0, Status: 1, PasswordChangedAt: claims.IatMs}, nil
|
||||
},
|
||||
})(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.OK("pass"))
|
||||
}))
|
||||
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/list", nil)
|
||||
req.Header.Set("Authorization", token)
|
||||
res := httptest.NewRecorder()
|
||||
wrapped.ServeHTTP(res, req)
|
||||
assertAuthDenied(t, res)
|
||||
}
|
||||
|
||||
func TestJWTRejectsRoleMismatch(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)
|
||||
}
|
||||
|
||||
wrapped := JWT(AuthOptions{
|
||||
JWTSecret: secret,
|
||||
GetUserAuthState: func(userID int64) (*auth.UserAuthState, error) {
|
||||
return &auth.UserAuthState{ID: userID, RoleID: 1, Status: 1, PasswordChangedAt: 0}, nil
|
||||
},
|
||||
})(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.OK("pass"))
|
||||
}))
|
||||
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/list", nil)
|
||||
req.Header.Set("Authorization", token)
|
||||
res := httptest.NewRecorder()
|
||||
wrapped.ServeHTTP(res, req)
|
||||
assertAuthDenied(t, res)
|
||||
}
|
||||
|
||||
func TestJWTRejectsMissingUserState(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)
|
||||
}
|
||||
|
||||
wrapped := JWT(AuthOptions{
|
||||
JWTSecret: secret,
|
||||
GetUserAuthState: func(userID int64) (*auth.UserAuthState, error) {
|
||||
return nil, nil
|
||||
},
|
||||
})(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.OK("pass"))
|
||||
}))
|
||||
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/list", nil)
|
||||
req.Header.Set("Authorization", token)
|
||||
res := httptest.NewRecorder()
|
||||
wrapped.ServeHTTP(res, req)
|
||||
assertAuthDenied(t, res)
|
||||
}
|
||||
|
||||
func TestJWTRejectsAuthStateLookupError(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)
|
||||
}
|
||||
|
||||
wrapped := JWT(AuthOptions{
|
||||
JWTSecret: secret,
|
||||
GetUserAuthState: func(userID int64) (*auth.UserAuthState, error) {
|
||||
return nil, errors.New("boom")
|
||||
},
|
||||
})(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.OK("pass"))
|
||||
}))
|
||||
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/list", nil)
|
||||
req.Header.Set("Authorization", token)
|
||||
res := httptest.NewRecorder()
|
||||
wrapped.ServeHTTP(res, req)
|
||||
assertAuthDenied(t, res)
|
||||
}
|
||||
|
||||
func TestShouldSkipDoesNotBypassConfigGet(t *testing.T) {
|
||||
if shouldSkip("/api/v1/config/get") {
|
||||
t.Fatal("expected /api/v1/config/get to require auth")
|
||||
}
|
||||
}
|
||||
|
||||
func TestShouldSkipBypassesPublicConfigGet(t *testing.T) {
|
||||
if !shouldSkip("/api/v1/public/config/get") {
|
||||
t.Fatal("expected /api/v1/public/config/get to remain public")
|
||||
}
|
||||
}
|
||||
|
||||
func TestJWTExpiresAfterSevenDays(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)
|
||||
}
|
||||
if got := claims.Exp - claims.Iat; got != int64(7*24*time.Hour/time.Second) {
|
||||
t.Fatalf("expected 7 day token lifetime, got %d seconds", got)
|
||||
}
|
||||
if claims.IatMs <= 0 {
|
||||
t.Fatalf("expected millisecond issuance time to be populated, got %d", claims.IatMs)
|
||||
}
|
||||
}
|
||||
|
||||
func assertCode(t *testing.T, rec *httptest.ResponseRecorder, expected int) {
|
||||
t.Helper()
|
||||
var out response.R
|
||||
if err := json.NewDecoder(rec.Body).Decode(&out); err != nil {
|
||||
t.Fatalf("decode response: %v", err)
|
||||
}
|
||||
if out.Code != expected {
|
||||
t.Fatalf("expected code %d, got %d", expected, out.Code)
|
||||
}
|
||||
}
|
||||
|
||||
func assertAuthDenied(t *testing.T, rec *httptest.ResponseRecorder) {
|
||||
t.Helper()
|
||||
assertCodeMsg(t, rec, 401, "无效的token或token已过期")
|
||||
}
|
||||
|
||||
func assertCodeMsg(t *testing.T, rec *httptest.ResponseRecorder, expectedCode int, expectedMsg string) {
|
||||
t.Helper()
|
||||
var out response.R
|
||||
if err := json.NewDecoder(rec.Body).Decode(&out); err != nil {
|
||||
t.Fatalf("decode response: %v", err)
|
||||
}
|
||||
if out.Code != expectedCode || out.Msg != expectedMsg {
|
||||
t.Fatalf("expected (%d,%q), got (%d,%q)", expectedCode, expectedMsg, out.Code, out.Msg)
|
||||
}
|
||||
}
|
||||
@@ -13,7 +13,7 @@ func NewRouter(h *handler.Handler, jwtSecret string) http.Handler {
|
||||
mux.Handle("/system-info", h.WebSocketHandler())
|
||||
|
||||
wrapped := middleware.Recover(mux)
|
||||
wrapped = middleware.JWT(middleware.AuthOptions{JWTSecret: jwtSecret})(wrapped)
|
||||
wrapped = middleware.JWT(middleware.AuthOptions{JWTSecret: jwtSecret, GetUserAuthState: h.GetUserAuthState})(wrapped)
|
||||
wrapped = middleware.RequestLog(wrapped)
|
||||
wrapped = middleware.CORS(wrapped)
|
||||
return wrapped
|
||||
|
||||
@@ -0,0 +1,25 @@
|
||||
package security
|
||||
|
||||
import (
|
||||
"strings"
|
||||
|
||||
"golang.org/x/crypto/bcrypt"
|
||||
)
|
||||
|
||||
func HashPassword(plain string) (string, error) {
|
||||
hash, err := bcrypt.GenerateFromPassword([]byte(plain), bcrypt.DefaultCost)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return string(hash), nil
|
||||
}
|
||||
|
||||
func VerifyPassword(storedHash, plain string) (bool, bool) {
|
||||
if strings.HasPrefix(storedHash, "$2") {
|
||||
return bcrypt.CompareHashAndPassword([]byte(storedHash), []byte(plain)) == nil, false
|
||||
}
|
||||
if MD5(plain) == storedHash {
|
||||
return true, true
|
||||
}
|
||||
return false, false
|
||||
}
|
||||
@@ -0,0 +1,25 @@
|
||||
package security
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestHashPasswordProducesBcrypt(t *testing.T) {
|
||||
hash, err := HashPassword("admin_user")
|
||||
if err != nil {
|
||||
t.Fatalf("HashPassword() error = %v", err)
|
||||
}
|
||||
if len(hash) < 50 || !strings.HasPrefix(hash, "$2") {
|
||||
t.Fatalf("expected bcrypt hash, got %q", hash)
|
||||
}
|
||||
if ok, legacy := VerifyPassword(hash, "admin_user"); !ok || legacy {
|
||||
t.Fatalf("VerifyPassword() = (%v,%v), want (true,false)", ok, legacy)
|
||||
}
|
||||
}
|
||||
|
||||
func TestVerifyPasswordAcceptsLegacyMD5(t *testing.T) {
|
||||
if ok, legacy := VerifyPassword("3c85cdebade1c51cf64ca9f3c09d182d", "admin_user"); !ok || !legacy {
|
||||
t.Fatalf("VerifyPassword() = (%v,%v), want (true,true)", ok, legacy)
|
||||
}
|
||||
}
|
||||
@@ -10,20 +10,21 @@ import "database/sql"
|
||||
// User maps to the "user" table. PostgreSQL treats "user" as a reserved
|
||||
// word, so TableName() is required for correct quoting.
|
||||
type User struct {
|
||||
ID int64 `gorm:"primaryKey;autoIncrement"`
|
||||
User string `gorm:"column:user;type:varchar(100);not null"`
|
||||
Pwd string `gorm:"type:varchar(100);not null"`
|
||||
RoleID int `gorm:"column:role_id;not null"`
|
||||
ExpTime int64 `gorm:"column:exp_time;not null"`
|
||||
Flow int64 `gorm:"not null"`
|
||||
InFlow int64 `gorm:"column:in_flow;not null;default:0"`
|
||||
OutFlow int64 `gorm:"column:out_flow;not null;default:0"`
|
||||
FlowResetTime int64 `gorm:"column:flow_reset_time;not null"`
|
||||
Num int `gorm:"not null"`
|
||||
CreatedTime int64 `gorm:"column:created_time;not null"`
|
||||
UpdatedTime sql.NullInt64 `gorm:"column:updated_time"`
|
||||
Status int `gorm:"not null"`
|
||||
MaxConn int `gorm:"column:max_conn;not null;default:0"`
|
||||
ID int64 `gorm:"primaryKey;autoIncrement"`
|
||||
User string `gorm:"column:user;type:varchar(100);not null"`
|
||||
Pwd string `gorm:"type:varchar(100);not null"`
|
||||
RoleID int `gorm:"column:role_id;not null"`
|
||||
ExpTime int64 `gorm:"column:exp_time;not null"`
|
||||
Flow int64 `gorm:"not null"`
|
||||
InFlow int64 `gorm:"column:in_flow;not null;default:0"`
|
||||
OutFlow int64 `gorm:"column:out_flow;not null;default:0"`
|
||||
FlowResetTime int64 `gorm:"column:flow_reset_time;not null"`
|
||||
Num int `gorm:"not null"`
|
||||
CreatedTime int64 `gorm:"column:created_time;not null"`
|
||||
UpdatedTime sql.NullInt64 `gorm:"column:updated_time"`
|
||||
Status int `gorm:"not null"`
|
||||
PasswordChangedAt int64 `gorm:"column:password_changed_at;not null;default:0"`
|
||||
MaxConn int `gorm:"column:max_conn;not null;default:0"`
|
||||
}
|
||||
|
||||
func (User) TableName() string { return "user" }
|
||||
|
||||
@@ -0,0 +1,62 @@
|
||||
package repo
|
||||
|
||||
import "strings"
|
||||
|
||||
type ConfigAccessPolicy string
|
||||
|
||||
const (
|
||||
ConfigAccessPublic ConfigAccessPolicy = "public"
|
||||
ConfigAccessSensitive ConfigAccessPolicy = "sensitive"
|
||||
)
|
||||
|
||||
var publicConfigKeys = map[string]struct{}{
|
||||
"app_name": {},
|
||||
"app_logo": {},
|
||||
"app_favicon": {},
|
||||
"app_bg_image": {},
|
||||
"cloudflare_site_key": {},
|
||||
}
|
||||
|
||||
var sensitiveConfigKeys = map[string]struct{}{
|
||||
"jwt_secret": {},
|
||||
"license_key": {},
|
||||
"cloudflare_secret_key": {},
|
||||
}
|
||||
|
||||
func PolicyForConfig(name string) ConfigAccessPolicy {
|
||||
if IsPublicConfigKey(name) {
|
||||
return ConfigAccessPublic
|
||||
}
|
||||
if IsSensitiveConfigKey(name) {
|
||||
return ConfigAccessSensitive
|
||||
}
|
||||
return ConfigAccessSensitive
|
||||
}
|
||||
|
||||
func IsPublicConfigKey(name string) bool {
|
||||
_, ok := publicConfigKeys[normalizeConfigKey(name)]
|
||||
return ok
|
||||
}
|
||||
|
||||
func IsSensitiveConfigKey(name string) bool {
|
||||
_, ok := sensitiveConfigKeys[normalizeConfigKey(name)]
|
||||
return ok
|
||||
}
|
||||
|
||||
func FilterSensitiveConfigs(in map[string]string) map[string]string {
|
||||
if len(in) == 0 {
|
||||
return map[string]string{}
|
||||
}
|
||||
out := make(map[string]string, len(in))
|
||||
for name, value := range in {
|
||||
if IsSensitiveConfigKey(name) {
|
||||
continue
|
||||
}
|
||||
out[name] = value
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func normalizeConfigKey(name string) string {
|
||||
return strings.ToLower(strings.TrimSpace(name))
|
||||
}
|
||||
@@ -0,0 +1,69 @@
|
||||
package repo
|
||||
|
||||
import "testing"
|
||||
|
||||
func TestConfigPolicy(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
key string
|
||||
want ConfigAccessPolicy
|
||||
}{
|
||||
{name: "app_name is public", key: "app_name", want: ConfigAccessPublic},
|
||||
{name: "app_logo is public", key: "app_logo", want: ConfigAccessPublic},
|
||||
{name: "app_favicon is public", key: "app_favicon", want: ConfigAccessPublic},
|
||||
{name: "app_bg_image is public", key: "app_bg_image", want: ConfigAccessPublic},
|
||||
{name: "cloudflare_site_key is public", key: "cloudflare_site_key", want: ConfigAccessPublic},
|
||||
{name: "jwt_secret is sensitive", key: "jwt_secret", want: ConfigAccessSensitive},
|
||||
{name: "license_key is sensitive", key: "license_key", want: ConfigAccessSensitive},
|
||||
{name: "cloudflare_secret_key is sensitive", key: "cloudflare_secret_key", want: ConfigAccessSensitive},
|
||||
{name: "trimmed public key is public", key: " APP_NAME ", want: ConfigAccessPublic},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
if got := PolicyForConfig(tt.key); got != tt.want {
|
||||
t.Fatalf("PolicyForConfig(%q) = %v, want %v", tt.key, got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestConfigPolicyHelpers(t *testing.T) {
|
||||
publicKeys := []string{"app_name", "app_logo", "app_favicon", "app_bg_image", "cloudflare_site_key"}
|
||||
for _, key := range publicKeys {
|
||||
if !IsPublicConfigKey(key) {
|
||||
t.Fatalf("expected %q to be public", key)
|
||||
}
|
||||
}
|
||||
|
||||
sensitiveKeys := []string{"jwt_secret", "license_key", "cloudflare_secret_key"}
|
||||
for _, key := range sensitiveKeys {
|
||||
if !IsSensitiveConfigKey(key) {
|
||||
t.Fatalf("expected %q to be sensitive", key)
|
||||
}
|
||||
}
|
||||
|
||||
input := map[string]string{
|
||||
"app_name": "FLVX",
|
||||
"license_key": "secret-license",
|
||||
"cloudflare_secret_key": "secret-cloudflare",
|
||||
"jwt_secret": "secret-jwt",
|
||||
"cloudflare_site_key": "site-key",
|
||||
}
|
||||
filtered := FilterSensitiveConfigs(input)
|
||||
if len(filtered) != 2 {
|
||||
t.Fatalf("expected 2 public configs, got %d", len(filtered))
|
||||
}
|
||||
if filtered["app_name"] != "FLVX" || filtered["cloudflare_site_key"] != "site-key" {
|
||||
t.Fatalf("unexpected filtered configs: %+v", filtered)
|
||||
}
|
||||
if _, ok := filtered["jwt_secret"]; ok {
|
||||
t.Fatal("expected jwt_secret to be filtered out")
|
||||
}
|
||||
if _, ok := filtered["license_key"]; ok {
|
||||
t.Fatal("expected license_key to be filtered out")
|
||||
}
|
||||
if _, ok := filtered["cloudflare_secret_key"]; ok {
|
||||
t.Fatal("expected cloudflare_secret_key to be filtered out")
|
||||
}
|
||||
}
|
||||
@@ -486,9 +486,10 @@ func (r *Repository) UpdateUserNameAndPassword(userID int64, username, passwordM
|
||||
return errors.New("repository not initialized")
|
||||
}
|
||||
return r.db.Model(&model.User{}).Where("id = ?", userID).Updates(map[string]interface{}{
|
||||
"user": username,
|
||||
"pwd": passwordMD5,
|
||||
"updated_time": now,
|
||||
"user": username,
|
||||
"pwd": passwordMD5,
|
||||
"password_changed_at": now,
|
||||
"updated_time": now,
|
||||
}).Error
|
||||
}
|
||||
|
||||
@@ -1890,7 +1891,7 @@ func (r *Repository) ExportAll() (*model.BackupData, error) {
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("export configs failed: %w", err)
|
||||
}
|
||||
backup.Configs = configs
|
||||
backup.Configs = FilterSensitiveConfigs(configs)
|
||||
|
||||
return backup, nil
|
||||
}
|
||||
@@ -1970,7 +1971,7 @@ func (r *Repository) ExportPartial(types []string) (*model.BackupData, error) {
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("export configs failed: %w", err)
|
||||
}
|
||||
backup.Configs = v
|
||||
backup.Configs = FilterSensitiveConfigs(v)
|
||||
}
|
||||
return backup, nil
|
||||
}
|
||||
@@ -2726,6 +2727,7 @@ func importPermissions(tx *gorm.DB, permissions []model.PermissionBackup, _ int6
|
||||
}
|
||||
|
||||
func importConfigs(tx *gorm.DB, configs map[string]string, now int64) (int, error) {
|
||||
configs = FilterSensitiveConfigs(configs)
|
||||
count := 0
|
||||
for name, value := range configs {
|
||||
err := tx.Clauses(clause.OnConflict{
|
||||
|
||||
@@ -0,0 +1,29 @@
|
||||
package repo
|
||||
|
||||
import (
|
||||
"errors"
|
||||
|
||||
"go-backend/internal/auth"
|
||||
"go-backend/internal/store/model"
|
||||
)
|
||||
|
||||
func (r *Repository) GetUserAuthState(userID int64) (*auth.UserAuthState, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return nil, errors.New("repository not initialized")
|
||||
}
|
||||
var user struct {
|
||||
ID int64 `gorm:"column:id"`
|
||||
RoleID int `gorm:"column:role_id"`
|
||||
Status int `gorm:"column:status"`
|
||||
PasswordChangedAt int64 `gorm:"column:password_changed_at"`
|
||||
}
|
||||
if err := r.db.Model(&model.User{}).Select("id", "role_id", "status", "password_changed_at").Where("id = ?", userID).First(&user).Error; err != nil {
|
||||
return nil, normalizeNotFoundErr(err)
|
||||
}
|
||||
return &auth.UserAuthState{
|
||||
ID: user.ID,
|
||||
RoleID: user.RoleID,
|
||||
Status: user.Status,
|
||||
PasswordChangedAt: user.PasswordChangedAt,
|
||||
}, nil
|
||||
}
|
||||
@@ -0,0 +1,29 @@
|
||||
package repo
|
||||
|
||||
import (
|
||||
"path/filepath"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestGetUserAuthStateReturnsPasswordChangedAt(t *testing.T) {
|
||||
r, err := Open(filepath.Join(t.TempDir(), "auth.db"))
|
||||
if err != nil {
|
||||
t.Fatalf("Open() error = %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
userID, err := r.CreateUser("admin_user", "pwd", 0, 2727251700000, 99999, 1, 99999, 1, 0, now)
|
||||
if err != nil {
|
||||
t.Fatalf("CreateUser() error = %v", err)
|
||||
}
|
||||
|
||||
state, err := r.GetUserAuthState(userID)
|
||||
if err != nil {
|
||||
t.Fatalf("GetUserAuthState() error = %v", err)
|
||||
}
|
||||
if state == nil || state.PasswordChangedAt != now || state.Status != 1 || state.RoleID != 0 {
|
||||
t.Fatalf("unexpected auth state: %+v", state)
|
||||
}
|
||||
}
|
||||
@@ -4,6 +4,8 @@ import (
|
||||
"path/filepath"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/store/model"
|
||||
)
|
||||
|
||||
func TestBackupRoundTripsTunnelProbeTarget(t *testing.T) {
|
||||
@@ -57,3 +59,101 @@ func TestBackupRoundTripsTunnelProbeTarget(t *testing.T) {
|
||||
t.Fatalf("unexpected imported probe target: %+v", items[0])
|
||||
}
|
||||
}
|
||||
|
||||
func TestExportAllOmitsSensitiveConfigs(t *testing.T) {
|
||||
r, err := Open(filepath.Join(t.TempDir(), "export.db"))
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
seedConfig(t, r, "app_name", "FLVX")
|
||||
seedConfig(t, r, "app_logo", "logo")
|
||||
seedConfig(t, r, "app_favicon", "favicon")
|
||||
seedConfig(t, r, "app_bg_image", "bg")
|
||||
seedConfig(t, r, "cloudflare_site_key", "site-key")
|
||||
seedConfig(t, r, "jwt_secret", "jwt-secret")
|
||||
seedConfig(t, r, "license_key", "license-secret")
|
||||
seedConfig(t, r, "cloudflare_secret_key", "cloudflare-secret")
|
||||
|
||||
for _, tc := range []struct {
|
||||
name string
|
||||
export func() (*model.BackupData, error)
|
||||
}{
|
||||
{name: "ExportAll", export: r.ExportAll},
|
||||
{name: "ExportPartial", export: func() (*model.BackupData, error) { return r.ExportPartial([]string{"configs"}) }},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
backup, err := tc.export()
|
||||
if err != nil {
|
||||
t.Fatalf("export backup: %v", err)
|
||||
}
|
||||
if backup.Configs["app_name"] != "FLVX" {
|
||||
t.Fatalf("expected public config in export, got %+v", backup.Configs)
|
||||
}
|
||||
if backup.Configs["cloudflare_site_key"] != "site-key" {
|
||||
t.Fatalf("expected public config in export, got %+v", backup.Configs)
|
||||
}
|
||||
for _, key := range []string{"jwt_secret", "license_key", "cloudflare_secret_key"} {
|
||||
if _, ok := backup.Configs[key]; ok {
|
||||
t.Fatalf("expected %s to be omitted from export, got %+v", key, backup.Configs)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestImportIgnoresSensitiveConfigs(t *testing.T) {
|
||||
r, err := Open(filepath.Join(t.TempDir(), "import.db"))
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
seedConfig(t, r, "app_name", "before")
|
||||
seedConfig(t, r, "jwt_secret", "jwt-before")
|
||||
seedConfig(t, r, "license_key", "license-before")
|
||||
seedConfig(t, r, "cloudflare_secret_key", "cloudflare-before")
|
||||
|
||||
backup := &model.BackupData{Configs: map[string]string{
|
||||
"app_name": "after",
|
||||
"jwt_secret": "jwt-after",
|
||||
"license_key": "license-after",
|
||||
"cloudflare_secret_key": "cloudflare-after",
|
||||
}}
|
||||
|
||||
result, err := r.Import(backup, []string{"configs"})
|
||||
if err != nil {
|
||||
t.Fatalf("import backup: %v", err)
|
||||
}
|
||||
if result.ConfigsImported != 1 {
|
||||
t.Fatalf("expected one imported config, got %d", result.ConfigsImported)
|
||||
}
|
||||
|
||||
assertConfigValue(t, r, "app_name", "after")
|
||||
assertConfigValue(t, r, "jwt_secret", "jwt-before")
|
||||
assertConfigValue(t, r, "license_key", "license-before")
|
||||
assertConfigValue(t, r, "cloudflare_secret_key", "cloudflare-before")
|
||||
}
|
||||
|
||||
func seedConfig(t *testing.T, r *Repository, name, value string) {
|
||||
t.Helper()
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO vite_config(name, value, time)
|
||||
VALUES(?, ?, ?)
|
||||
ON CONFLICT(name) DO UPDATE SET value = excluded.value, time = excluded.time
|
||||
`, name, value, time.Now().UnixMilli()).Error; err != nil {
|
||||
t.Fatalf("seed config %s: %v", name, err)
|
||||
}
|
||||
}
|
||||
|
||||
func assertConfigValue(t *testing.T, r *Repository, name, want string) {
|
||||
t.Helper()
|
||||
cfg, err := r.GetConfigByName(name)
|
||||
if err != nil {
|
||||
t.Fatalf("get config %s: %v", name, err)
|
||||
}
|
||||
if cfg == nil || cfg.Value != want {
|
||||
t.Fatalf("expected config %s=%q, got %+v", name, want, cfg)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -42,19 +42,20 @@ func (r *Repository) CreateUser(username, pwdHash string, roleID int, expTime, f
|
||||
return 0, errors.New("repository not initialized")
|
||||
}
|
||||
user := model.User{
|
||||
User: username,
|
||||
Pwd: pwdHash,
|
||||
RoleID: roleID,
|
||||
ExpTime: expTime,
|
||||
Flow: flow,
|
||||
InFlow: 0,
|
||||
OutFlow: 0,
|
||||
FlowResetTime: flowResetTime,
|
||||
Num: num,
|
||||
MaxConn: maxConn,
|
||||
CreatedTime: now,
|
||||
UpdatedTime: sql.NullInt64{Int64: now, Valid: true},
|
||||
Status: status,
|
||||
User: username,
|
||||
Pwd: pwdHash,
|
||||
RoleID: roleID,
|
||||
ExpTime: expTime,
|
||||
Flow: flow,
|
||||
InFlow: 0,
|
||||
OutFlow: 0,
|
||||
FlowResetTime: flowResetTime,
|
||||
Num: num,
|
||||
MaxConn: maxConn,
|
||||
CreatedTime: now,
|
||||
UpdatedTime: sql.NullInt64{Int64: now, Valid: true},
|
||||
Status: status,
|
||||
PasswordChangedAt: now,
|
||||
}
|
||||
if err := r.db.Create(&user).Error; err != nil {
|
||||
return 0, err
|
||||
@@ -81,15 +82,16 @@ func (r *Repository) UpdateUserWithPassword(id int64, username, pwdHash string,
|
||||
return r.db.Model(&model.User{}).
|
||||
Where("id = ?", id).
|
||||
Updates(map[string]interface{}{
|
||||
"user": username,
|
||||
"pwd": pwdHash,
|
||||
"flow": flow,
|
||||
"num": num,
|
||||
"exp_time": expTime,
|
||||
"flow_reset_time": flowResetTime,
|
||||
"status": status,
|
||||
"max_conn": maxConn,
|
||||
"updated_time": sql.NullInt64{Int64: now, Valid: true},
|
||||
"user": username,
|
||||
"pwd": pwdHash,
|
||||
"flow": flow,
|
||||
"num": num,
|
||||
"exp_time": expTime,
|
||||
"flow_reset_time": flowResetTime,
|
||||
"status": status,
|
||||
"max_conn": maxConn,
|
||||
"password_changed_at": now,
|
||||
"updated_time": sql.NullInt64{Int64: now, Valid: true},
|
||||
}).Error
|
||||
}
|
||||
|
||||
@@ -111,6 +113,19 @@ func (r *Repository) UpdateUserWithoutPassword(id int64, username string, flow i
|
||||
}).Error
|
||||
}
|
||||
|
||||
func (r *Repository) UpdateUserPassword(userID int64, pwdHash string, now int64) error {
|
||||
if r == nil || r.db == nil {
|
||||
return errors.New("repository not initialized")
|
||||
}
|
||||
return r.db.Model(&model.User{}).
|
||||
Where("id = ?", userID).
|
||||
Updates(map[string]interface{}{
|
||||
"pwd": pwdHash,
|
||||
"password_changed_at": now,
|
||||
"updated_time": sql.NullInt64{Int64: now, Valid: true},
|
||||
}).Error
|
||||
}
|
||||
|
||||
func (r *Repository) PropagateUserFlowToTunnels(userID int64, flow int64, num int, expTime, flowResetTime int64) {
|
||||
if r == nil || r.db == nil {
|
||||
return
|
||||
|
||||
@@ -74,6 +74,7 @@ type Server struct {
|
||||
upgrader websocket.Upgrader
|
||||
onNodeOnline func(nodeID int64)
|
||||
onNodeMetric func(nodeID int64, info SystemInfo)
|
||||
getUserAuthState func(userID int64) (*auth.UserAuthState, error)
|
||||
|
||||
mu sync.RWMutex
|
||||
admins map[*connWrap]struct{}
|
||||
@@ -130,6 +131,15 @@ func NewServer(repo *repo.Repository, jwtSecret string) *Server {
|
||||
}
|
||||
}
|
||||
|
||||
func (s *Server) SetUserAuthStateLookup(fn func(userID int64) (*auth.UserAuthState, error)) {
|
||||
if s == nil {
|
||||
return
|
||||
}
|
||||
s.mu.Lock()
|
||||
s.getUserAuthState = fn
|
||||
s.mu.Unlock()
|
||||
}
|
||||
|
||||
func (s *Server) ServeHTTP(w http.ResponseWriter, r *http.Request) {
|
||||
query := r.URL.Query()
|
||||
typeVal := query.Get("type")
|
||||
@@ -146,10 +156,23 @@ func (s *Server) ServeHTTP(w http.ResponseWriter, r *http.Request) {
|
||||
}
|
||||
|
||||
if typeVal == "0" {
|
||||
if _, ok := auth.ValidateToken(secret, s.jwtSecret); !ok {
|
||||
claims, ok := auth.ValidateToken(secret, s.jwtSecret)
|
||||
if !ok {
|
||||
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
|
||||
}
|
||||
}
|
||||
s.handleAdmin(w, r)
|
||||
return
|
||||
}
|
||||
|
||||
@@ -0,0 +1,31 @@
|
||||
package ws
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
"go-backend/internal/auth"
|
||||
)
|
||||
|
||||
func TestServeHTTPRejectsDisabledAdminToken(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)
|
||||
}
|
||||
|
||||
server := NewServer(nil, secret)
|
||||
server.SetUserAuthStateLookup(func(userID int64) (*auth.UserAuthState, error) {
|
||||
return &auth.UserAuthState{ID: userID, RoleID: 0, Status: 0, PasswordChangedAt: 0}, nil
|
||||
})
|
||||
|
||||
req := httptest.NewRequest(http.MethodGet, "/system-info?type=0&secret="+token, nil)
|
||||
rec := httptest.NewRecorder()
|
||||
|
||||
server.ServeHTTP(rec, req)
|
||||
|
||||
if rec.Code != http.StatusForbidden {
|
||||
t.Fatalf("expected forbidden for disabled admin token, got %d", rec.Code)
|
||||
}
|
||||
}
|
||||
@@ -1,6 +1,7 @@
|
||||
package contract_test
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
@@ -9,6 +10,7 @@ import (
|
||||
"go-backend/internal/auth"
|
||||
"go-backend/internal/http/middleware"
|
||||
"go-backend/internal/http/response"
|
||||
"go-backend/internal/security"
|
||||
)
|
||||
|
||||
func TestJWTMiddlewareContracts(t *testing.T) {
|
||||
@@ -18,7 +20,9 @@ func TestJWTMiddlewareContracts(t *testing.T) {
|
||||
response.WriteJSON(w, response.OK("pass"))
|
||||
})
|
||||
|
||||
wrapped := middleware.JWT(middleware.AuthOptions{JWTSecret: secret})(next)
|
||||
wrapped := middleware.JWT(middleware.AuthOptions{JWTSecret: secret, GetUserAuthState: func(userID int64) (*auth.UserAuthState, error) {
|
||||
return &auth.UserAuthState{ID: userID, RoleID: 0, Status: 1, PasswordChangedAt: 0}, nil
|
||||
}})(next)
|
||||
|
||||
t.Run("login path is excluded", func(t *testing.T) {
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/user/login", nil)
|
||||
@@ -59,6 +63,9 @@ func TestJWTMiddlewareContracts(t *testing.T) {
|
||||
if err != nil {
|
||||
t.Fatalf("generate token: %v", err)
|
||||
}
|
||||
wrapped := middleware.JWT(middleware.AuthOptions{JWTSecret: secret, GetUserAuthState: func(userID int64) (*auth.UserAuthState, error) {
|
||||
return &auth.UserAuthState{ID: userID, RoleID: 1, Status: 1, PasswordChangedAt: 0}, nil
|
||||
}})(next)
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/config/update", nil)
|
||||
req.Header.Set("Authorization", token)
|
||||
res := httptest.NewRecorder()
|
||||
@@ -67,6 +74,83 @@ func TestJWTMiddlewareContracts(t *testing.T) {
|
||||
})
|
||||
}
|
||||
|
||||
func TestLoginTokenValidatesThroughRouter(t *testing.T) {
|
||||
router, r := setupContractRouter(t, "contract-jwt-secret")
|
||||
seedLegacyUser(t, r, 9110, "router-login-user", "router-login-pass")
|
||||
|
||||
body := bytes.NewBufferString(`{"username":"router-login-user","password":"router-login-pass","captchaId":""}`)
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/user/login", body)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
resp := httptest.NewRecorder()
|
||||
router.ServeHTTP(resp, req)
|
||||
|
||||
var out response.R
|
||||
if err := json.NewDecoder(resp.Body).Decode(&out); err != nil {
|
||||
t.Fatalf("decode login response: %v", err)
|
||||
}
|
||||
if out.Code != 0 {
|
||||
t.Fatalf("expected login code 0, got %d (%s)", out.Code, out.Msg)
|
||||
}
|
||||
data, ok := out.Data.(map[string]interface{})
|
||||
if !ok {
|
||||
t.Fatalf("expected login data map, got %T", out.Data)
|
||||
}
|
||||
token, _ := data["token"].(string)
|
||||
if token == "" {
|
||||
t.Fatal("expected login token")
|
||||
}
|
||||
|
||||
checkReq := httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/user/tunnel", nil)
|
||||
checkReq.Header.Set("Authorization", token)
|
||||
checkResp := httptest.NewRecorder()
|
||||
router.ServeHTTP(checkResp, checkReq)
|
||||
assertCode(t, checkResp, 0)
|
||||
}
|
||||
|
||||
func TestLegacyPasswordMigratesOnLogin(t *testing.T) {
|
||||
router, r := setupContractRouter(t, "contract-jwt-secret")
|
||||
legacyChangedAt := seedLegacyUser(t, r, 9101, "legacy-login-user", "legacy-login-pass")
|
||||
|
||||
body := bytes.NewBufferString(`{"username":"legacy-login-user","password":"legacy-login-pass","captchaId":""}`)
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/user/login", body)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
resp := httptest.NewRecorder()
|
||||
|
||||
router.ServeHTTP(resp, req)
|
||||
assertCode(t, resp, 0)
|
||||
assertUserPasswordIsBcrypt(t, r, "legacy-login-user", "legacy-login-pass")
|
||||
if changedAt := mustQueryPasswordChangedAtByUsername(t, r, "legacy-login-user"); changedAt <= legacyChangedAt {
|
||||
t.Fatalf("expected password_changed_at to advance on login migration, got %d <= %d", changedAt, legacyChangedAt)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDisabledLegacyPasswordIsRejectedWithoutMigration(t *testing.T) {
|
||||
router, r := setupContractRouter(t, "contract-jwt-secret")
|
||||
legacyChangedAt := seedLegacyUserWithStatus(t, r, 9105, "disabled-legacy-user", "disabled-legacy-pass", 0)
|
||||
|
||||
body := bytes.NewBufferString(`{"username":"disabled-legacy-user","password":"disabled-legacy-pass","captchaId":""}`)
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/user/login", body)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
resp := httptest.NewRecorder()
|
||||
|
||||
router.ServeHTTP(resp, req)
|
||||
assertCodeMsg(t, resp, -1, "账号被停用")
|
||||
|
||||
user, err := r.GetUserByUsername("disabled-legacy-user")
|
||||
if err != nil {
|
||||
t.Fatalf("get user: %v", err)
|
||||
}
|
||||
if user == nil {
|
||||
t.Fatal("expected disabled user to exist")
|
||||
}
|
||||
if ok, migrated := security.VerifyPassword(user.Pwd, "disabled-legacy-pass"); !ok || !migrated {
|
||||
t.Fatalf("expected disabled user to remain legacy MD5, got (%v,%v) with hash %q", ok, migrated, user.Pwd)
|
||||
}
|
||||
if changedAt := mustQueryPasswordChangedAtByUsername(t, r, "disabled-legacy-user"); changedAt != legacyChangedAt {
|
||||
t.Fatalf("expected password_changed_at to remain unchanged, got %d want %d", changedAt, legacyChangedAt)
|
||||
}
|
||||
}
|
||||
|
||||
func assertCode(t *testing.T, rec *httptest.ResponseRecorder, expected int) {
|
||||
t.Helper()
|
||||
var out response.R
|
||||
|
||||
@@ -17,6 +17,7 @@ import (
|
||||
httpserver "go-backend/internal/http"
|
||||
"go-backend/internal/http/handler"
|
||||
"go-backend/internal/http/response"
|
||||
"go-backend/internal/security"
|
||||
"go-backend/internal/store/repo"
|
||||
|
||||
"gorm.io/gorm"
|
||||
@@ -128,6 +129,46 @@ func TestCaptchaVerifyLoginContract(t *testing.T) {
|
||||
})
|
||||
}
|
||||
|
||||
func TestPublicConfigGetAndAuthConfigContract(t *testing.T) {
|
||||
router, r := setupContractRouter(t, "contract-jwt-secret")
|
||||
|
||||
for name, value := range map[string]string{
|
||||
"app_name": "FLVX Public",
|
||||
"app_logo": "logo",
|
||||
"app_favicon": "favicon",
|
||||
"app_bg_image": "bg",
|
||||
"cloudflare_site_key": "site-key",
|
||||
"cloudflare_secret_key": "secret-key",
|
||||
"jwt_secret": "jwt-secret",
|
||||
} {
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO vite_config(name, value, time)
|
||||
VALUES(?, ?, ?)
|
||||
ON CONFLICT(name) DO UPDATE SET value = excluded.value, time = excluded.time
|
||||
`, name, value, time.Now().UnixMilli()).Error; err != nil {
|
||||
t.Fatalf("seed config %s: %v", name, err)
|
||||
}
|
||||
}
|
||||
|
||||
publicReq := httptest.NewRequest(http.MethodPost, "/api/v1/public/config/get", bytes.NewBufferString(`{"name":"app_name"}`))
|
||||
publicReq.Header.Set("Content-Type", "application/json")
|
||||
publicResp := httptest.NewRecorder()
|
||||
router.ServeHTTP(publicResp, publicReq)
|
||||
assertCode(t, publicResp, 0)
|
||||
|
||||
secretReq := httptest.NewRequest(http.MethodPost, "/api/v1/public/config/get", bytes.NewBufferString(`{"name":"jwt_secret"}`))
|
||||
secretReq.Header.Set("Content-Type", "application/json")
|
||||
secretResp := httptest.NewRecorder()
|
||||
router.ServeHTTP(secretResp, secretReq)
|
||||
assertCodeMsg(t, secretResp, 403, "禁止访问敏感配置")
|
||||
|
||||
configReq := httptest.NewRequest(http.MethodPost, "/api/v1/config/get", bytes.NewBufferString(`{"name":"app_name"}`))
|
||||
configReq.Header.Set("Content-Type", "application/json")
|
||||
configResp := httptest.NewRecorder()
|
||||
router.ServeHTTP(configResp, configReq)
|
||||
assertCodeMsg(t, configResp, 401, "未登录或token已过期")
|
||||
}
|
||||
|
||||
func TestOpenAPISubStoreContracts(t *testing.T) {
|
||||
router, r := setupContractRouter(t, "contract-jwt-secret")
|
||||
|
||||
@@ -209,6 +250,122 @@ func TestOpenAPISubStoreContracts(t *testing.T) {
|
||||
})
|
||||
}
|
||||
|
||||
func TestLegacyPasswordMigratesOnSubStore(t *testing.T) {
|
||||
router, r := setupContractRouter(t, "contract-jwt-secret")
|
||||
legacyChangedAt := seedLegacyUser(t, r, 9102, "legacy-substore-user", "legacy-substore-pass")
|
||||
|
||||
req := httptest.NewRequest(http.MethodGet, "/api/v1/open_api/sub_store?user=legacy-substore-user&pwd=legacy-substore-pass", nil)
|
||||
resp := httptest.NewRecorder()
|
||||
|
||||
router.ServeHTTP(resp, req)
|
||||
|
||||
body, err := io.ReadAll(resp.Body)
|
||||
if err != nil {
|
||||
t.Fatalf("read body: %v", err)
|
||||
}
|
||||
expected := "upload=0; download=0; total=107373108658176; expire=2727251700"
|
||||
if string(body) != expected {
|
||||
t.Fatalf("expected body %q, got %q", expected, string(body))
|
||||
}
|
||||
assertUserPasswordIsBcrypt(t, r, "legacy-substore-user", "legacy-substore-pass")
|
||||
if changedAt := mustQueryPasswordChangedAtByUsername(t, r, "legacy-substore-user"); changedAt <= legacyChangedAt {
|
||||
t.Fatalf("expected password_changed_at to advance on sub-store migration, got %d <= %d", changedAt, legacyChangedAt)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDisabledLegacyPasswordIsRejectedOnSubStoreWithoutMigration(t *testing.T) {
|
||||
router, r := setupContractRouter(t, "contract-jwt-secret")
|
||||
legacyChangedAt := seedLegacyUserWithStatus(t, r, 9106, "disabled-substore-user", "disabled-substore-pass", 0)
|
||||
|
||||
req := httptest.NewRequest(http.MethodGet, "/api/v1/open_api/sub_store?user=disabled-substore-user&pwd=disabled-substore-pass", nil)
|
||||
resp := httptest.NewRecorder()
|
||||
|
||||
router.ServeHTTP(resp, req)
|
||||
assertCodeMsg(t, resp, -1, "账号被停用")
|
||||
|
||||
user, err := r.GetUserByUsername("disabled-substore-user")
|
||||
if err != nil {
|
||||
t.Fatalf("get user: %v", err)
|
||||
}
|
||||
if user == nil {
|
||||
t.Fatal("expected disabled user to exist")
|
||||
}
|
||||
if ok, migrated := security.VerifyPassword(user.Pwd, "disabled-substore-pass"); !ok || !migrated {
|
||||
t.Fatalf("expected disabled user to remain legacy MD5, got (%v,%v) with hash %q", ok, migrated, user.Pwd)
|
||||
}
|
||||
if changedAt := mustQueryPasswordChangedAtByUsername(t, r, "disabled-substore-user"); changedAt != legacyChangedAt {
|
||||
t.Fatalf("expected password_changed_at to remain unchanged, got %d want %d", changedAt, legacyChangedAt)
|
||||
}
|
||||
}
|
||||
|
||||
func TestUserCreateStoresStrongHash(t *testing.T) {
|
||||
router, r := setupContractRouter(t, "contract-jwt-secret")
|
||||
startedAt := time.Now().UnixMilli()
|
||||
adminToken, err := auth.GenerateToken(1, "admin_user", 0, "contract-jwt-secret")
|
||||
if err != nil {
|
||||
t.Fatalf("generate token: %v", err)
|
||||
}
|
||||
|
||||
body := bytes.NewBufferString(`{"user":"created-user","pwd":"created-pass"}`)
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/user/create", body)
|
||||
req.Header.Set("Authorization", adminToken)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
resp := httptest.NewRecorder()
|
||||
|
||||
router.ServeHTTP(resp, req)
|
||||
assertCode(t, resp, 0)
|
||||
assertUserPasswordIsBcrypt(t, r, "created-user", "created-pass")
|
||||
if changedAt := mustQueryPasswordChangedAtByUsername(t, r, "created-user"); changedAt < startedAt {
|
||||
t.Fatalf("expected password_changed_at to be set on create, got %d < %d", changedAt, startedAt)
|
||||
}
|
||||
}
|
||||
|
||||
func TestUserUpdateStoresStrongHash(t *testing.T) {
|
||||
router, r := setupContractRouter(t, "contract-jwt-secret")
|
||||
seedLegacyUser(t, r, 9103, "legacy-update-user", "legacy-update-pass")
|
||||
startedAt := time.Now().UnixMilli()
|
||||
adminToken, err := auth.GenerateToken(1, "admin_user", 0, "contract-jwt-secret")
|
||||
if err != nil {
|
||||
t.Fatalf("generate token: %v", err)
|
||||
}
|
||||
|
||||
body := bytes.NewBufferString(`{"id":9103,"user":"updated-user","pwd":"updated-pass","flow":99999,"num":99999,"expTime":2727251700000,"flowResetTime":1,"status":1,"maxConn":0}`)
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/user/update", body)
|
||||
req.Header.Set("Authorization", adminToken)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
resp := httptest.NewRecorder()
|
||||
|
||||
router.ServeHTTP(resp, req)
|
||||
assertCode(t, resp, 0)
|
||||
assertUserPasswordIsBcrypt(t, r, "updated-user", "updated-pass")
|
||||
if changedAt := mustQueryPasswordChangedAtByUsername(t, r, "updated-user"); changedAt < startedAt {
|
||||
t.Fatalf("expected password_changed_at to be updated on user update, got %d < %d", changedAt, startedAt)
|
||||
}
|
||||
}
|
||||
|
||||
func TestUpdatePasswordStoresStrongHash(t *testing.T) {
|
||||
router, r := setupContractRouter(t, "contract-jwt-secret")
|
||||
seedLegacyUser(t, r, 9104, "legacy-self-user", "legacy-self-pass")
|
||||
startedAt := time.Now().UnixMilli()
|
||||
token, err := auth.GenerateToken(9104, "legacy-self-user", 1, "contract-jwt-secret")
|
||||
if err != nil {
|
||||
t.Fatalf("generate token: %v", err)
|
||||
}
|
||||
|
||||
body := bytes.NewBufferString(`{"newUsername":"self-updated-user","currentPassword":"legacy-self-pass","newPassword":"self-updated-pass","confirmPassword":"self-updated-pass"}`)
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/user/updatePassword", body)
|
||||
req.Header.Set("Authorization", token)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
resp := httptest.NewRecorder()
|
||||
|
||||
router.ServeHTTP(resp, req)
|
||||
assertCode(t, resp, 0)
|
||||
assertUserPasswordIsBcrypt(t, r, "self-updated-user", "self-updated-pass")
|
||||
if changedAt := mustQueryPasswordChangedAtByUsername(t, r, "self-updated-user"); changedAt < startedAt {
|
||||
t.Fatalf("expected password_changed_at to be updated on password change, got %d < %d", changedAt, startedAt)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSpeedLimitTunnelsRouteRemoved(t *testing.T) {
|
||||
secret := "contract-jwt-secret"
|
||||
router, _ := setupContractRouter(t, secret)
|
||||
@@ -232,6 +389,7 @@ func TestSpeedLimitTunnelsRouteRemoved(t *testing.T) {
|
||||
func TestBackupExportImportRestoreContracts(t *testing.T) {
|
||||
secret := "contract-jwt-secret"
|
||||
router, r := setupContractRouter(t, secret)
|
||||
seedContractUser(t, r, 2, "normal_user", 1, 1)
|
||||
|
||||
adminToken, err := auth.GenerateToken(1, "admin_user", 0, secret)
|
||||
if err != nil {
|
||||
@@ -559,6 +717,71 @@ func TestBackupExportImportRestoreContracts(t *testing.T) {
|
||||
})
|
||||
}
|
||||
|
||||
func TestBackupConfigFilteringContract(t *testing.T) {
|
||||
secret := "contract-jwt-secret"
|
||||
router, r := setupContractRouter(t, secret)
|
||||
|
||||
adminToken, err := auth.GenerateToken(1, "admin_user", 0, secret)
|
||||
if err != nil {
|
||||
t.Fatalf("generate admin token: %v", err)
|
||||
}
|
||||
|
||||
configs := map[string]string{
|
||||
"app_name": "contract-before",
|
||||
"cloudflare_site_key": "site-key-before",
|
||||
"jwt_secret": "jwt-before",
|
||||
"license_key": "license-before",
|
||||
"cloudflare_secret_key": "cloudflare-before",
|
||||
}
|
||||
for name, value := range configs {
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO vite_config(name, value, time)
|
||||
VALUES(?, ?, ?)
|
||||
ON CONFLICT(name) DO UPDATE SET value = excluded.value, time = excluded.time
|
||||
`, name, value, time.Now().UnixMilli()).Error; err != nil {
|
||||
t.Fatalf("seed config %s: %v", name, err)
|
||||
}
|
||||
}
|
||||
|
||||
payload := exportBackupPayload(t, router, "/api/v1/backup/export", adminToken)
|
||||
for _, key := range []string{"jwt_secret", "license_key", "cloudflare_secret_key"} {
|
||||
if _, ok := payload.Configs[key]; ok {
|
||||
t.Fatalf("expected %s to be omitted from exported configs: %+v", key, payload.Configs)
|
||||
}
|
||||
}
|
||||
if payload.Configs["app_name"] != "contract-before" {
|
||||
t.Fatalf("expected public config to be exported, got %+v", payload.Configs)
|
||||
}
|
||||
|
||||
payload.Configs["app_name"] = "contract-after"
|
||||
payload.Configs["jwt_secret"] = "jwt-after"
|
||||
payload.Configs["license_key"] = "license-after"
|
||||
payload.Configs["cloudflare_secret_key"] = "cloudflare-after"
|
||||
raw, err := json.Marshal(backupImportPayload{Types: []string{"configs"}, backupExportPayload: payload})
|
||||
if err != nil {
|
||||
t.Fatalf("marshal import payload: %v", err)
|
||||
}
|
||||
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/backup/import", bytes.NewReader(raw))
|
||||
req.Header.Set("Authorization", adminToken)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
resp := httptest.NewRecorder()
|
||||
|
||||
router.ServeHTTP(resp, req)
|
||||
var out response.R
|
||||
if err := json.NewDecoder(resp.Body).Decode(&out); err != nil {
|
||||
t.Fatalf("decode import response: %v", err)
|
||||
}
|
||||
if out.Code != 0 {
|
||||
t.Fatalf("expected import code 0, got %d (%s)", out.Code, out.Msg)
|
||||
}
|
||||
|
||||
assertConfigValue(t, r, "app_name", "contract-after")
|
||||
assertConfigValue(t, r, "jwt_secret", "jwt-before")
|
||||
assertConfigValue(t, r, "license_key", "license-before")
|
||||
assertConfigValue(t, r, "cloudflare_secret_key", "cloudflare-before")
|
||||
}
|
||||
|
||||
type backupExportPayload struct {
|
||||
Version string `json:"version"`
|
||||
ExportedAt int64 `json:"exportedAt"`
|
||||
@@ -619,6 +842,66 @@ func setupContractRouter(t *testing.T, jwtSecret string) (http.Handler, *repo.Re
|
||||
return httpserver.NewRouter(h, jwtSecret), r
|
||||
}
|
||||
|
||||
func assertConfigValue(t *testing.T, r *repo.Repository, name, want string) {
|
||||
t.Helper()
|
||||
cfg, err := r.GetConfigByName(name)
|
||||
if err != nil {
|
||||
t.Fatalf("get config %s: %v", name, err)
|
||||
}
|
||||
if cfg == nil || cfg.Value != want {
|
||||
t.Fatalf("expected config %s=%q, got %+v", name, want, cfg)
|
||||
}
|
||||
}
|
||||
|
||||
func seedLegacyUser(t *testing.T, r *repo.Repository, id int64, username, password string) int64 {
|
||||
return seedLegacyUserWithStatus(t, r, id, username, password, 1)
|
||||
}
|
||||
|
||||
func seedContractUser(t *testing.T, r *repo.Repository, id int64, username string, roleID, status int) int64 {
|
||||
t.Helper()
|
||||
now := time.Now().UnixMilli()
|
||||
passwordChangedAt := now - 10_000
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO user(id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, max_conn, created_time, updated_time, status, password_changed_at)
|
||||
VALUES(?, ?, ?, ?, 2727251700000, 99999, 0, 0, 1, 99999, 0, ?, ?, ?, ?)
|
||||
`, id, username, security.MD5("contract-pass"), roleID, now, now, status, passwordChangedAt).Error; err != nil {
|
||||
t.Fatalf("seed contract user %s: %v", username, err)
|
||||
}
|
||||
return passwordChangedAt
|
||||
}
|
||||
|
||||
func seedLegacyUserWithStatus(t *testing.T, r *repo.Repository, id int64, username, password string, status int) int64 {
|
||||
t.Helper()
|
||||
now := time.Now().UnixMilli()
|
||||
legacyChangedAt := now - 10_000
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO user(id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, max_conn, created_time, updated_time, status, password_changed_at)
|
||||
VALUES(?, ?, ?, 1, 2727251700000, 99999, 0, 0, 1, 99999, 0, ?, ?, ?, ?)
|
||||
`, id, username, security.MD5(password), now, now, status, legacyChangedAt).Error; err != nil {
|
||||
t.Fatalf("seed legacy user %s: %v", username, err)
|
||||
}
|
||||
return legacyChangedAt
|
||||
}
|
||||
|
||||
func mustQueryPasswordChangedAtByUsername(t *testing.T, r *repo.Repository, username string) int64 {
|
||||
t.Helper()
|
||||
return mustQueryInt64(t, r, `SELECT password_changed_at FROM user WHERE user = ?`, username)
|
||||
}
|
||||
|
||||
func assertUserPasswordIsBcrypt(t *testing.T, r *repo.Repository, username, password string) {
|
||||
t.Helper()
|
||||
user, err := r.GetUserByUsername(username)
|
||||
if err != nil {
|
||||
t.Fatalf("get user %s: %v", username, err)
|
||||
}
|
||||
if user == nil {
|
||||
t.Fatalf("expected user %s to exist", username)
|
||||
}
|
||||
if ok, migrated := security.VerifyPassword(user.Pwd, password); !ok || migrated {
|
||||
t.Fatalf("expected bcrypt password for %s, got %q", username, user.Pwd)
|
||||
}
|
||||
}
|
||||
|
||||
func TestOpenMigratesLegacyNodeDualStackColumns(t *testing.T) {
|
||||
dbPath := filepath.Join(t.TempDir(), "legacy-2.0.7-beta.db")
|
||||
legacyDB, err := sql.Open("sqlite", dbPath)
|
||||
|
||||
@@ -16,6 +16,7 @@ import (
|
||||
func TestNodeMetricsEndpoints(t *testing.T) {
|
||||
secret := "monitoring-jwt-secret"
|
||||
router, repo := setupContractRouter(t, secret)
|
||||
seedContractUser(t, repo, 2, "normal_user", 1, 1)
|
||||
|
||||
adminToken, err := auth.GenerateToken(1, "admin_user", 0, secret)
|
||||
if err != nil {
|
||||
@@ -1302,6 +1303,7 @@ func TestMonitoringAuthRequired(t *testing.T) {
|
||||
func TestMonitorAccessEndpoint(t *testing.T) {
|
||||
secret := "monitoring-jwt-secret"
|
||||
router, repo := setupContractRouter(t, secret)
|
||||
seedContractUser(t, repo, 2, "normal_user", 1, 1)
|
||||
|
||||
adminToken, err := auth.GenerateToken(1, "admin_user", 0, secret)
|
||||
if err != nil {
|
||||
@@ -1357,6 +1359,7 @@ func TestMonitorAccessEndpoint(t *testing.T) {
|
||||
func TestMonitoringPermissionRequired(t *testing.T) {
|
||||
secret := "monitoring-jwt-secret"
|
||||
router, repo := setupContractRouter(t, secret)
|
||||
seedContractUser(t, repo, 2, "normal_user", 1, 1)
|
||||
|
||||
userToken, err := auth.GenerateToken(2, "normal_user", 1, secret)
|
||||
if err != nil {
|
||||
|
||||
@@ -12,7 +12,8 @@ import (
|
||||
|
||||
func TestStorageSummaryRequiresAdminAndReturnsSize(t *testing.T) {
|
||||
secret := "storage-contract-secret"
|
||||
router, _ := setupContractRouter(t, secret)
|
||||
router, r := setupContractRouter(t, secret)
|
||||
seedContractUser(t, r, 2, "normal_user", 1, 1)
|
||||
|
||||
adminToken, err := auth.GenerateToken(1, "admin_user", 0, secret)
|
||||
if err != nil {
|
||||
|
||||
@@ -253,6 +253,8 @@ export const getConfigs = () =>
|
||||
Network.post<Record<string, string>>("/config/list");
|
||||
export const getConfigByName = (name: string) =>
|
||||
Network.post<{ name: string; value: string }>("/config/get", { name });
|
||||
export const getPublicConfigByName = (name: string) =>
|
||||
Network.post<{ name: string; value: string }>("/public/config/get", { name });
|
||||
export const updateConfigs = (configMap: Record<string, string>) =>
|
||||
Network.post("/config/update", configMap);
|
||||
export const updateConfig = (name: string, value: string) =>
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
import { getConfigByName, getConfigs } from "@/api";
|
||||
import { getConfigByName, getConfigs, getPublicConfigByName } from "@/api";
|
||||
import { isLoggedIn } from "@/utils/auth";
|
||||
|
||||
export type SiteConfig = typeof siteConfig;
|
||||
|
||||
@@ -8,9 +9,58 @@ const VERSION = import.meta.env.VITE_APP_VERSION || "dev";
|
||||
const APP_VERSION = "1.0.3";
|
||||
const DEFAULT_FAVICON = "/favicon.ico";
|
||||
const FAVICON_LINK_ID = "app-favicon";
|
||||
const PUBLIC_BRAND_CONFIG_KEYS = [
|
||||
"app_name",
|
||||
"app_logo",
|
||||
"app_favicon",
|
||||
"app_bg_image",
|
||||
] as const;
|
||||
const GITHUB_REPO =
|
||||
import.meta.env.VITE_GITHUB_REPO || "https://github.com/Sagit-chu/flux-panel";
|
||||
|
||||
const readCachedConfigs = (keys: readonly string[]) => {
|
||||
const cachedConfigs: Record<string, string> = {};
|
||||
let hasCachedData = false;
|
||||
|
||||
keys.forEach((key) => {
|
||||
const cachedValue = configCache.get(key);
|
||||
|
||||
if (cachedValue !== null) {
|
||||
cachedConfigs[key] = cachedValue;
|
||||
hasCachedData = true;
|
||||
}
|
||||
});
|
||||
|
||||
return { cachedConfigs, hasCachedData };
|
||||
};
|
||||
|
||||
const fetchPublicBrandConfigs = async (): Promise<Record<string, string>> => {
|
||||
const publicConfigMap: Record<string, string> = {};
|
||||
|
||||
await Promise.all(
|
||||
PUBLIC_BRAND_CONFIG_KEYS.map(async (key) => {
|
||||
try {
|
||||
const response = await getPublicConfigByName(key);
|
||||
|
||||
if (
|
||||
response.code === 0 &&
|
||||
response.data &&
|
||||
typeof response.data.value === "string"
|
||||
) {
|
||||
const value = response.data.value;
|
||||
|
||||
publicConfigMap[key] = value;
|
||||
configCache.set(key, value);
|
||||
}
|
||||
} catch {
|
||||
// ignore single key fetch error
|
||||
}
|
||||
}),
|
||||
);
|
||||
|
||||
return publicConfigMap;
|
||||
};
|
||||
|
||||
const getInitialConfig = () => {
|
||||
if (typeof window === "undefined") {
|
||||
return {
|
||||
@@ -129,46 +179,19 @@ export const getCachedConfig = async (key: string): Promise<string | null> => {
|
||||
|
||||
// 获取所有配置(优先从缓存)
|
||||
export const getCachedConfigs = async (): Promise<Record<string, string>> => {
|
||||
// 尝试从缓存获取所有配置
|
||||
const configKeys = ["app_name", "app_logo", "app_favicon", "app_bg_image"];
|
||||
const cachedConfigs: Record<string, string> = {};
|
||||
let hasCachedData = false;
|
||||
const { cachedConfigs, hasCachedData } = readCachedConfigs(
|
||||
PUBLIC_BRAND_CONFIG_KEYS,
|
||||
);
|
||||
|
||||
configKeys.forEach((key) => {
|
||||
const cachedValue = configCache.get(key);
|
||||
if (!isLoggedIn()) {
|
||||
const publicConfigs = await fetchPublicBrandConfigs();
|
||||
|
||||
if (cachedValue !== null) {
|
||||
cachedConfigs[key] = cachedValue;
|
||||
hasCachedData = true;
|
||||
if (Object.keys(publicConfigs).length > 0) {
|
||||
return { ...cachedConfigs, ...publicConfigs };
|
||||
}
|
||||
});
|
||||
|
||||
const fetchPublicConfigs = async (): Promise<Record<string, string>> => {
|
||||
const publicConfigMap: Record<string, string> = {};
|
||||
|
||||
await Promise.all(
|
||||
configKeys.map(async (key) => {
|
||||
try {
|
||||
const response = await getConfigByName(key);
|
||||
|
||||
if (
|
||||
response.code === 0 &&
|
||||
response.data &&
|
||||
typeof response.data.value === "string"
|
||||
) {
|
||||
const value = response.data.value;
|
||||
|
||||
publicConfigMap[key] = value;
|
||||
configCache.set(key, value);
|
||||
}
|
||||
} catch {
|
||||
// ignore single key fetch error
|
||||
}
|
||||
}),
|
||||
);
|
||||
|
||||
return publicConfigMap;
|
||||
};
|
||||
return cachedConfigs;
|
||||
}
|
||||
|
||||
// 从API获取最新配置
|
||||
try {
|
||||
@@ -189,14 +212,14 @@ export const getCachedConfigs = async (): Promise<Record<string, string>> => {
|
||||
return cachedConfigs;
|
||||
}
|
||||
|
||||
return await fetchPublicConfigs();
|
||||
return await fetchPublicBrandConfigs();
|
||||
} catch {
|
||||
// API失败时返回缓存的数据
|
||||
if (hasCachedData) {
|
||||
return cachedConfigs;
|
||||
}
|
||||
|
||||
return await fetchPublicConfigs();
|
||||
return await fetchPublicBrandConfigs();
|
||||
}
|
||||
};
|
||||
|
||||
|
||||
@@ -10,7 +10,7 @@ import { Button } from "@/shadcn-bridge/heroui/button";
|
||||
import { siteConfig } from "@/config/site";
|
||||
import { VersionFooter } from "@/components/version-footer";
|
||||
import { BrandLogo } from "@/components/brand-logo";
|
||||
import { login, LoginData, checkCaptcha, getConfigByName } from "@/api";
|
||||
import { login, LoginData, checkCaptcha, getPublicConfigByName } from "@/api";
|
||||
import { writeLoginSession } from "@/utils/session";
|
||||
import { useWebViewMode } from "@/hooks/useWebViewMode";
|
||||
|
||||
@@ -128,7 +128,7 @@ export default function IndexPage() {
|
||||
if (checkResponse.data === 0) {
|
||||
await performLogin();
|
||||
} else {
|
||||
const configResp = await getConfigByName("cloudflare_site_key");
|
||||
const configResp = await getPublicConfigByName("cloudflare_site_key");
|
||||
|
||||
if (configResp.code === 0 && configResp.data && configResp.data.value) {
|
||||
setSiteKey(configResp.data.value);
|
||||
|
||||
Reference in New Issue
Block a user