fix: harden auth, config access, and backups

This commit is contained in:
sagitchu
2026-05-13 23:53:06 +08:00
parent ec9fb77eb5
commit 465815cf34
28 changed files with 1376 additions and 97 deletions
+7 -2
View File
@@ -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,
+8
View File
@@ -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)
}
}
+47 -5
View File
@@ -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
}
+12 -2
View File
@@ -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
}
+19 -2
View File
@@ -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)
}
}
+1 -1
View File
@@ -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
+25
View File
@@ -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)
}
}
+15 -14
View File
@@ -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")
}
}
+7 -5
View File
@@ -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
+24 -1
View File
@@ -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
}
+31
View File
@@ -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)
}
}