Compare commits

...

13 Commits

Author SHA1 Message Date
sagit 129fa0aa4c fix(frontend): constrain service monitor records modal (#557) 2026-09-24 16:12:03 +08:00
sagit cf7246c71b feat: adaptive traffic units and precise MB–PB quotas (#556)
## Summary

- Format traffic amounts adaptively through PB across the user,
dashboard, forwarding, node, monitoring, and panel sharing views. Closes
#548.
- Let administrators choose MB, GB, TB, or PB when setting user, tunnel
permission, and panel sharing traffic limits. Closes #549.
- Persist exact MiB limits for users and tunnel permissions while
retaining the legacy GB field and existing data. Apply the precise limit
in forwarding policy checks and preserve it in backups.

## Verification

- `go test ./...` in `go-backend` (724 passed)
- `pnpm run build` in `vite-frontend`
- ESLint on changed frontend files
- `git diff --check`
2026-09-24 15:40:11 +08:00
sagit b4c2989285 Merge branch 'main' into codex/issues-548-549-traffic-units 2026-09-24 15:38:49 +08:00
sagit f2783713a5 feat: separate light and dark wallpapers (#555)
## Summary

- Add separate light and dark wallpapers with fallback to the existing
background setting, resolving #554.
- Make both wallpaper settings available to unauthenticated pages and
update the background when the theme changes.
- Include the pre-existing installer regression changes in this commit,
as requested.

## Verification

- `go test ./...` in `go-backend` (723 passed)
- `pnpm run build` in `vite-frontend`
- ESLint on changed frontend files
- `bash test-install-scripts-proxy.sh`
- `git diff --check`
2026-09-24 15:38:22 +08:00
sagitchu 5d60c4fbe1 feat: support adaptive traffic units and precise quotas 2026-09-24 15:33:53 +08:00
sagitchu 62c56e0a79 feat: add mode-specific wallpapers and installer regression coverage 2026-09-24 14:59:52 +08:00
sagit 2e845030de fix: build frontend assets on native platform (#553) 2026-09-01 16:27:00 +08:00
sagit 2269f2e2d5 fix: include pnpm config in frontend image build (#552) 2026-09-01 15:48:54 +08:00
sagit da7bef88f1 fix: harden panel updates and nftables deployment (#551)
Fix panel update deployment discovery and rollback safety, add nftables compatibility and atomic replacement, and restore reproducible frontend CI installs.
2026-09-01 15:37:33 +08:00
sagit f26014579b fix: stabilize federated agent runtime updates (#550) 2026-08-25 17:04:04 +08:00
sagit 5041c722c9 fix: harden license lifecycle (#547) 2026-08-13 09:25:05 +08:00
sagit c56798e991 fix: accept existing license machine binding (#546) 2026-08-12 14:12:46 +08:00
sagit 40e96f3592 fix: preserve active per-IP traffic limiters (#545) 2026-08-12 11:31:03 +08:00
58 changed files with 2833 additions and 410 deletions
+13 -1
View File
@@ -7,6 +7,18 @@ on:
branches: ['**']
jobs:
install-scripts:
name: Test Installer Scripts (systemd and Alpine OpenRC)
runs-on: ubuntu-latest
steps:
- uses: actions/checkout@v4
- name: Run installer regression tests
run: bash test-install-scripts-proxy.sh
- name: Test Alpine bootstrap and OpenRC lifecycle
run: docker run --rm -v "$PWD:/workspace:ro" alpine:3.22 sh /workspace/test-install-scripts-alpine.sh
frontend:
name: Build Frontend
runs-on: ubuntu-latest
@@ -22,7 +34,7 @@ jobs:
node-version: '20.19.0'
- name: Install pnpm
run: npm install -g pnpm
run: npm install -g pnpm@10.28.1
- name: Install dependencies
run: pnpm install --frozen-lockfile
+4 -1
View File
@@ -229,6 +229,7 @@ jobs:
docker buildx build \
--platform linux/amd64,linux/arm64 \
--build-arg KEYGEN_ACCOUNT_ID=${{ secrets.KEYGEN_ACCOUNT_ID }} \
--push \
-t ${{ env.REGISTRY }}/${OWNER}/flux-panel-backend:latest \
-t ${{ env.REGISTRY }}/${OWNER}/flux-panel-backend:${VERSION} \
@@ -303,6 +304,9 @@ jobs:
sed -i "s|^PINNED_VERSION=\"\"|PINNED_VERSION=\"${VERSION}\"|" ./artifacts/install.sh
sed -i "s|^PINNED_VERSION=\"\"|PINNED_VERSION=\"${VERSION}\"|" ./artifacts/panel_install.sh
- name: Verify release installer on Alpine OpenRC
run: docker run --rm -v "$PWD:/workspace:ro" alpine:3.22 sh /workspace/test-install-scripts-alpine.sh /workspace/artifacts/install.sh
- name: Create Release
env:
GH_TOKEN: ${{ github.token }}
@@ -431,4 +435,3 @@ jobs:
gh release upload "${VERSION}" ./artifacts/gost-arm64.sha256 --clobber
echo "✅ GOST 二进制文件更新完成"
+2 -1
View File
@@ -7,7 +7,8 @@ RUN go mod download
COPY . .
ARG TARGETOS
ARG TARGETARCH
RUN CGO_ENABLED=0 GOOS=${TARGETOS:-linux} env ${TARGETARCH:+GOARCH=${TARGETARCH}} go build -o /out/paneld ./cmd/paneld
ARG KEYGEN_ACCOUNT_ID
RUN CGO_ENABLED=0 GOOS=${TARGETOS:-linux} env ${TARGETARCH:+GOARCH=${TARGETARCH}} go build -ldflags="-X 'go-backend/internal/license.AccountID=${KEYGEN_ACCOUNT_ID}'" -o /out/paneld ./cmd/paneld
FROM docker:27-cli AS dockercli
@@ -29,6 +29,10 @@ func (h *Handler) getPublicConfigByName(w http.ResponseWriter, r *http.Request)
response.WriteJSON(w, response.Err(403, "禁止访问敏感配置"))
return
}
if value, gated := h.unlicensedPublicBrandValue(configName); gated {
response.WriteJSON(w, response.OK(map[string]string{"name": configName, "value": value}))
return
}
cfg, err := h.repo.GetConfigByName(configName)
if err != nil {
@@ -42,3 +46,22 @@ func (h *Handler) getPublicConfigByName(w http.ResponseWriter, r *http.Request)
response.WriteJSON(w, response.OK(cfg))
}
func (h *Handler) unlicensedPublicBrandValue(configName string) (string, bool) {
switch configName {
case "app_name", "app_logo", "app_favicon", "hide_footer_brand":
default:
return "", false
}
isCommercial, _ := h.repo.GetViteConfigValue("is_commercial")
if isCommercial == "true" {
return "", false
}
if configName == "app_name" {
return "FLVX", true
}
if configName == "hide_footer_brand" {
return "false", true
}
return "", true
}
@@ -19,6 +19,8 @@ func TestPublicConfigGetAllowsBrandKeys(t *testing.T) {
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, "app_bg_image_light", "light-bg-data")
seedConfigValue(t, r, "app_bg_image_dark", "dark-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"}`))
@@ -28,6 +30,50 @@ func TestPublicConfigGetAllowsBrandKeys(t *testing.T) {
router.ServeHTTP(resp, req)
assertHandlerCode(t, resp, 0)
for name, want := range map[string]string{
"app_bg_image_light": "light-bg-data",
"app_bg_image_dark": "dark-bg-data",
} {
req := httptest.NewRequest(http.MethodPost, "/api/v1/public/config/get", bytes.NewBufferString(`{"name":"`+name+`"}`))
req.Header.Set("Content-Type", "application/json")
resp := httptest.NewRecorder()
router.ServeHTTP(resp, req)
assertHandlerConfigValue(t, resp, name, want)
}
}
func TestPublicBrandConfigFallsBackWithoutCommercialLicense(t *testing.T) {
router, r := setupConfigAccessTestRouter(t)
seedConfigValue(t, r, "app_name", "Paid Brand")
seedConfigValue(t, r, "app_logo", "logo-data")
seedConfigValue(t, r, "app_favicon", "favicon-data")
seedConfigValue(t, r, "hide_footer_brand", "true")
seedConfigValue(t, r, "is_commercial", "false")
for name, want := range map[string]string{
"app_name": "FLVX",
"app_logo": "",
"app_favicon": "",
"hide_footer_brand": "false",
} {
req := httptest.NewRequest(http.MethodPost, "/api/v1/public/config/get", bytes.NewBufferString(`{"name":"`+name+`"}`))
req.Header.Set("Content-Type", "application/json")
resp := httptest.NewRecorder()
router.ServeHTTP(resp, req)
assertHandlerConfigValue(t, resp, name, want)
}
}
func TestPublicBrandConfigUsesSavedValuesWithCommercialLicense(t *testing.T) {
router, r := setupConfigAccessTestRouter(t)
seedConfigValue(t, r, "app_name", "Paid Brand")
seedConfigValue(t, r, "is_commercial", "true")
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)
assertHandlerConfigValue(t, resp, "app_name", "Paid Brand")
}
func TestPublicConfigGetRejectsSensitiveKeys(t *testing.T) {
@@ -82,6 +128,56 @@ func TestConfigGetAllowsSensitiveKeysForAdmin(t *testing.T) {
assertHandlerConfigValue(t, resp, "jwt_secret", "jwt-secret")
}
func TestConfigGetNeverReturnsLicenseCredentials(t *testing.T) {
router, r := setupConfigAccessTestRouter(t)
adminToken := mustGenerateConfigAccessToken(t, 1, "admin_user", 0)
seedConfigValue(t, r, "license_key", "license-secret")
seedConfigValue(t, r, "machine_fingerprint", "fingerprint-secret")
for _, name := range []string{"license_key", "license_machine_id", "machine_fingerprint"} {
req := httptest.NewRequest(http.MethodPost, "/api/v1/config/get", bytes.NewBufferString(`{"name":"`+name+`"}`))
req.Header.Set("Content-Type", "application/json")
req.Header.Set("Authorization", adminToken)
resp := httptest.NewRecorder()
router.ServeHTTP(resp, req)
assertHandlerCodeMsg(t, resp, 403, "禁止访问系统授权凭据")
}
}
func TestConfigListNeverReturnsLicenseCredentials(t *testing.T) {
router, r := setupConfigAccessTestRouter(t)
adminToken := mustGenerateConfigAccessToken(t, 1, "admin_user", 0)
seedConfigValue(t, r, "license_key", "license-secret")
seedConfigValue(t, r, "license_machine_id", "machine-id")
seedConfigValue(t, r, "machine_fingerprint", "fingerprint-secret")
seedConfigValue(t, r, "is_commercial", "true")
req := httptest.NewRequest(http.MethodPost, "/api/v1/config/list", nil)
req.Header.Set("Authorization", adminToken)
resp := httptest.NewRecorder()
router.ServeHTTP(resp, req)
var out struct {
Code int `json:"code"`
Data map[string]string `json:"data"`
}
if err := json.NewDecoder(resp.Body).Decode(&out); err != nil {
t.Fatalf("decode response: %v", err)
}
if out.Code != 0 || out.Data["is_commercial"] != "true" {
t.Fatalf("unexpected config response: %+v", out)
}
if _, ok := out.Data["license_key"]; ok {
t.Fatal("license_key must not be returned")
}
if _, ok := out.Data["license_machine_id"]; ok {
t.Fatal("license_machine_id must not be returned")
}
if _, ok := out.Data["machine_fingerprint"]; ok {
t.Fatal("machine_fingerprint must not be returned")
}
}
func TestConfigUpdateAllowsSensitiveKeysForAdmin(t *testing.T) {
router, _ := setupConfigAccessTestRouter(t)
adminToken := mustGenerateConfigAccessToken(t, 1, "admin_user", 0)
@@ -154,7 +250,7 @@ func TestConfigUpdateSingleAllowsCloudflareSecretKeyWrite(t *testing.T) {
}
}
func TestConfigUpdateAllowsLicenseKeyWrite(t *testing.T) {
func TestConfigUpdateRejectsLicenseKeyWrite(t *testing.T) {
router, r := setupConfigAccessTestRouter(t)
adminToken := mustGenerateConfigAccessToken(t, 1, "admin_user", 0)
@@ -165,18 +261,13 @@ func TestConfigUpdateAllowsLicenseKeyWrite(t *testing.T) {
router.ServeHTTP(resp, req)
assertHandlerCode(t, resp, 0)
cfg, err := r.GetConfigByName("license_key")
if err != nil {
t.Fatalf("get config: %v", err)
}
if cfg == nil || cfg.Value != "license-secret" {
t.Fatalf("expected license_key to be updated, got %#v", cfg)
assertHandlerCodeMsg(t, resp, -1, "该配置由系统管理")
if cfg, err := r.GetConfigByName("license_key"); err != nil || cfg != nil {
t.Fatalf("license_key should not be written, got %#v err=%v", cfg, err)
}
}
func TestConfigUpdateSingleAllowsLicenseKeyWrite(t *testing.T) {
func TestConfigUpdateSingleRejectsLicenseKeyWrite(t *testing.T) {
router, r := setupConfigAccessTestRouter(t)
adminToken := mustGenerateConfigAccessToken(t, 1, "admin_user", 0)
@@ -187,14 +278,9 @@ func TestConfigUpdateSingleAllowsLicenseKeyWrite(t *testing.T) {
router.ServeHTTP(resp, req)
assertHandlerCode(t, resp, 0)
cfg, err := r.GetConfigByName("license_key")
if err != nil {
t.Fatalf("get config: %v", err)
}
if cfg == nil || cfg.Value != "license-secret" {
t.Fatalf("expected license_key to be updated, got %#v", cfg)
assertHandlerCodeMsg(t, resp, -1, "该配置由系统管理")
if cfg, err := r.GetConfigByName("license_key"); err != nil || cfg != nil {
t.Fatalf("license_key should not be written, got %#v err=%v", cfg, err)
}
}
@@ -4,6 +4,7 @@ import (
"encoding/json"
"errors"
"log"
"math"
"strconv"
"strings"
"time"
@@ -12,12 +13,27 @@ import (
)
const bytesPerGB int64 = 1024 * 1024 * 1024
const bytesPerMiB int64 = 1024 * 1024
func flowLimitBytes(flowGB, flowMiB int64) int64 {
if flowMiB > 0 {
if flowMiB > math.MaxInt64/bytesPerMiB {
return math.MaxInt64
}
return flowMiB * bytesPerMiB
}
if flowGB > math.MaxInt64/bytesPerGB {
return math.MaxInt64
}
return flowGB * bytesPerGB
}
type userTunnelPolicy struct {
ID int64
UserID int64
TunnelID int64
Flow int64
FlowMiB int64
InFlow int64
OutFlow int64
ExpTime int64
@@ -358,7 +374,7 @@ func (h *Handler) ensureUserTunnelForwardAllowed(userID int64, tunnelID int64, n
return errors.New("账号已过期")
}
flowLimit := user.Flow * bytesPerGB
flowLimit := flowLimitBytes(user.Flow, user.FlowMiB)
current := user.InFlow + user.OutFlow
if flowLimit < current {
return errors.New("流量已超额,禁止开启转发")
@@ -400,7 +416,7 @@ func (h *Handler) ensureUserTunnelForwardAllowed(userID int64, tunnelID int64, n
return errors.New("该隧道已过期")
}
utFlowLimit := policy.Flow * bytesPerGB
utFlowLimit := flowLimitBytes(policy.Flow, policy.FlowMiB)
utCurrent := policy.InFlow + policy.OutFlow
if utCurrent >= utFlowLimit {
return errors.New("该隧道流量已超额,禁止开启转发")
@@ -425,7 +441,7 @@ func (h *Handler) shouldPauseUser(userID int64, now int64) bool {
return false
}
flowLimit := user.Flow * bytesPerGB
flowLimit := flowLimitBytes(user.Flow, user.FlowMiB)
current := user.InFlow + user.OutFlow
if flowLimit < current {
return true
@@ -441,7 +457,7 @@ func shouldPauseUserTunnel(policy *userTunnelPolicy, now int64) bool {
return false
}
flowLimit := policy.Flow * bytesPerGB
flowLimit := flowLimitBytes(policy.Flow, policy.FlowMiB)
current := policy.InFlow + policy.OutFlow
if current >= flowLimit {
return true
@@ -465,7 +481,7 @@ func (h *Handler) getUserTunnelPolicy(userTunnelID int64) (*userTunnelPolicy, er
}
return &userTunnelPolicy{
ID: ut.ID, UserID: ut.UserID, TunnelID: ut.TunnelID,
Flow: ut.Flow, InFlow: ut.InFlow, OutFlow: ut.OutFlow,
Flow: ut.Flow, FlowMiB: ut.FlowMiB, InFlow: ut.InFlow, OutFlow: ut.OutFlow,
ExpTime: ut.ExpTime, Status: ut.Status, Num: ut.Num,
}, nil
}
@@ -659,9 +675,21 @@ func (h *Handler) sendDeleteOrphanedForwardService(nodeID int64, serviceName str
}
func (h *Handler) speedLimiterExists(name string) bool {
name = strings.TrimSpace(name)
if name == "" {
return false
}
const forwardRulePrefix = "rule_traffic_limit_"
if strings.HasPrefix(name, forwardRulePrefix) {
forwardID, err := strconv.ParseInt(strings.TrimPrefix(name, forwardRulePrefix), 10, 64)
if err != nil || forwardID <= 0 {
return false
}
forward, err := h.getForwardRecord(forwardID)
return err == nil && forward != nil && forward.IPSpeedID.Valid && forward.IPSpeedID.Int64 > 0
}
id, err := strconv.ParseInt(name, 10, 64)
if err != nil || id <= 0 {
return false
@@ -0,0 +1,38 @@
package handler
import (
"path/filepath"
"testing"
"go-backend/internal/store/repo"
)
func TestSpeedLimiterExistsPreservesForwardRuleLimiter(t *testing.T) {
r, err := repo.Open(filepath.Join(t.TempDir(), "panel.db"))
if err != nil {
t.Fatalf("open repo: %v", err)
}
defer r.Close()
if err := r.DB().Exec(`
INSERT INTO forward(id, user_id, user_name, name, tunnel_id, remote_addr, strategy, in_flow, out_flow, created_time, updated_time, status, inx, ip_speed_id)
VALUES(8, 1, 'user', 'forward', 1, '127.0.0.1:80', 'fifo', 0, 0, 1, 1, 1, 0, 3),
(9, 1, 'user', 'forward-without-ip-limit', 1, '127.0.0.1:81', 'fifo', 0, 0, 1, 1, 1, 0, NULL)
`).Error; err != nil {
t.Fatalf("insert forward: %v", err)
}
h := &Handler{repo: r}
if !h.speedLimiterExists("rule_traffic_limit_8") {
t.Fatal("expected runtime limiter for existing forward to be preserved")
}
if h.speedLimiterExists("rule_traffic_limit_9") {
t.Fatal("expected runtime limiter for forward without per-IP speed limit to be treated as orphaned")
}
if h.speedLimiterExists("rule_traffic_limit_10") {
t.Fatal("expected runtime limiter for missing forward to be treated as orphaned")
}
if h.speedLimiterExists("rule_traffic_limit_invalid") {
t.Fatal("expected malformed runtime limiter name to be treated as orphaned")
}
}
+48 -20
View File
@@ -5,6 +5,7 @@ import (
"database/sql"
"encoding/base64"
"encoding/json"
"errors"
"fmt"
"io"
"log"
@@ -41,10 +42,12 @@ type Handler struct {
captchaMu sync.Mutex
captchaTokens map[string]int64
jobsMu sync.Mutex
jobsCancel context.CancelFunc
jobsStarted bool
jobsWG sync.WaitGroup
jobsMu sync.Mutex
jobsCancel context.CancelFunc
jobsStarted bool
jobsWG sync.WaitGroup
fingerprintMu sync.Mutex
licenseValidationMu sync.Mutex
upgradeMu sync.Mutex
systemUpgradeMu sync.Mutex
@@ -399,6 +402,10 @@ func (h *Handler) getConfigByName(w http.ResponseWriter, r *http.Request) {
return
}
configName := strings.ToLower(strings.TrimSpace(req.Name))
if configName == "license_key" || configName == "license_machine_id" || configName == "machine_fingerprint" {
response.WriteJSON(w, response.Err(403, "禁止访问系统授权凭据"))
return
}
if repo.IsSensitiveConfigKey(configName) && !isAdminRequest(r) {
response.WriteJSON(w, response.Err(403, "禁止访问敏感配置"))
return
@@ -434,11 +441,13 @@ func (h *Handler) getConfigs(w http.ResponseWriter, r *http.Request) {
return
}
ctxClaims := r.Context().Value(middleware.ClaimsContextKey)
if claims, ok := ctxClaims.(auth.Claims); !ok || claims.RoleID != 0 {
delete(cfgMap, "license_key")
delete(cfgMap, "cloudflare_secret_key")
delete(cfgMap, "jwt_secret")
claims, isAdmin := ctxClaims.(auth.Claims)
if !isAdmin || claims.RoleID != 0 {
cfgMap = repo.FilterSensitiveConfigs(cfgMap)
}
delete(cfgMap, "license_key")
delete(cfgMap, "license_machine_id")
delete(cfgMap, "machine_fingerprint")
response.WriteJSON(w, response.OK(cfgMap))
}
@@ -710,6 +719,7 @@ func (h *Handler) userTunnelList(w http.ResponseWriter, r *http.Request) {
"tunnelName": t.TunnelName,
"status": t.Status,
"flow": t.Flow,
"flowMiB": t.FlowMiB,
"num": t.Num,
"expTime": t.ExpTime,
"flowResetTime": t.FlowResetTime,
@@ -877,10 +887,16 @@ func (h *Handler) flowUpload(w http.ResponseWriter, r *http.Request) {
}
func (h *Handler) getOrCreateMachineFingerprint() (string, error) {
fp, _ := h.repo.GetViteConfigValue("machine_fingerprint")
h.fingerprintMu.Lock()
defer h.fingerprintMu.Unlock()
fp, err := h.repo.GetViteConfigValue("machine_fingerprint")
if fp != "" {
return fp, nil
}
if err != nil && !errors.Is(err, sql.ErrNoRows) {
return "", err
}
newFp := uuid.New().String()
now := time.Now().UnixMilli()
@@ -907,10 +923,13 @@ func (h *Handler) licenseActivate(w http.ResponseWriter, r *http.Request) {
response.WriteJSON(w, response.ErrDefault("授权码不能为空"))
return
}
h.licenseValidationMu.Lock()
defer h.licenseValidationMu.Unlock()
valResp, err := h.validateLicenseForMachine(key)
if err != nil {
response.WriteJSON(w, response.ErrDefault("授权校验失败: "+err.Error()))
log.Printf("license activation failed: %v", err)
response.WriteJSON(w, response.ErrDefault(licenseValidationErrorMessage(err)))
return
}
@@ -920,20 +939,19 @@ func (h *Handler) licenseActivate(w http.ResponseWriter, r *http.Request) {
}
now := time.Now().UnixMilli()
if err := h.repo.UpsertConfig("license_key", key, now); err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
if err := h.repo.UpsertConfig("is_commercial", "true", now); err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
expiry := valResp.Data.Attributes.Expiry
if expiry == "" {
expiry = "never"
}
if err := h.repo.UpsertConfig("license_expiry", expiry, now); err != nil {
licenseState := map[string]string{
"license_key": key,
"is_commercial": "true",
"license_expiry": expiry,
}
if valResp.MachineID != "" {
licenseState["license_machine_id"] = valResp.MachineID
}
if err := h.repo.UpsertConfigs(licenseState, now); err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
@@ -975,6 +993,10 @@ func (h *Handler) updateConfigs(w http.ResponseWriter, r *http.Request) {
response.WriteJSON(w, response.Err(403, "禁止访问敏感配置"))
return
}
if repo.IsSystemManagedConfigKey(key) {
response.WriteJSON(w, response.ErrDefault("该配置由系统管理"))
return
}
if protectedKeys[key] && isCommercial != "true" {
response.WriteJSON(w, response.ErrDefault("需要商业版授权"))
@@ -1017,6 +1039,10 @@ func (h *Handler) updateSingleConfig(w http.ResponseWriter, r *http.Request) {
response.WriteJSON(w, response.Err(403, "禁止访问敏感配置"))
return
}
if repo.IsSystemManagedConfigKey(name) {
response.WriteJSON(w, response.ErrDefault("该配置由系统管理"))
return
}
isCommercial, _ := h.repo.GetViteConfigValue("is_commercial")
if (name == "app_name" || name == "app_logo" || name == "app_favicon" || name == "hide_footer_brand") && isCommercial != "true" {
@@ -1188,6 +1214,7 @@ func (h *Handler) userPackage(w http.ResponseWriter, r *http.Request) {
"tunnelName": t.TunnelName,
"tunnelFlow": t.TunnelFlow,
"flow": t.Flow,
"flowMiB": t.FlowMiB,
"inFlow": t.InFlow,
"outFlow": t.OutFlow,
"num": t.Num,
@@ -1237,6 +1264,7 @@ func (h *Handler) userPackage(w http.ResponseWriter, r *http.Request) {
"user": user.User,
"status": user.Status,
"flow": user.Flow,
"flowMiB": user.FlowMiB,
"inFlow": user.InFlow,
"outFlow": user.OutFlow,
"num": user.Num,
+13 -4
View File
@@ -36,6 +36,7 @@ func (h *Handler) StartBackgroundJobs() {
func (h *Handler) runValidateLicenseJob(ctx context.Context) {
defer h.jobsWG.Done()
h.validateLicenseJob()
ticker := time.NewTicker(12 * time.Hour)
defer ticker.Stop()
@@ -53,11 +54,12 @@ func (h *Handler) validateLicenseJob() {
if h == nil || h.repo == nil {
return
}
h.licenseValidationMu.Lock()
defer h.licenseValidationMu.Unlock()
key, _ := h.repo.GetViteConfigValue("license_key")
isCommercial, _ := h.repo.GetViteConfigValue("is_commercial")
if key == "" || isCommercial != "true" {
if key == "" {
return // Nothing to validate
}
@@ -67,7 +69,7 @@ func (h *Handler) validateLicenseJob() {
// Network and decode failures have no validation response, so retain the
// current state as a grace period. A rejected machine binding still has
// the original invalid response and must not stay commercially enabled.
if valResp != nil && !valResp.Meta.Valid {
if licenseValidationErrorIsDefinitive(valResp, err) {
now := time.Now().UnixMilli()
_ = h.repo.UpsertConfig("is_commercial", "false", now)
}
@@ -84,7 +86,14 @@ func (h *Handler) validateLicenseJob() {
if expiry == "" {
expiry = "never"
}
_ = h.repo.UpsertConfig("license_expiry", expiry, now)
licenseState := map[string]string{
"is_commercial": "true",
"license_expiry": expiry,
}
if valResp.MachineID != "" {
licenseState["license_machine_id"] = valResp.MachineID
}
_ = h.repo.UpsertConfigs(licenseState, now)
}
}
@@ -1,16 +1,24 @@
package handler
import (
"errors"
"fmt"
"net/http"
"os"
"strings"
"go-backend/internal/license"
)
const keygenAccountID = "1bc96cac-09de-4cf4-af34-26afdad63a90"
var newLicenseClient = license.NewKeygenClient
func keygenAccountID() string {
if value := strings.TrimSpace(license.AccountID); value != "" {
return value
}
return strings.TrimSpace(os.Getenv("KEYGEN_ACCOUNT_ID"))
}
func licenseNeedsMachineActivation(code string) bool {
switch strings.ToUpper(strings.TrimSpace(code)) {
case "NO_MACHINES", "NO_MACHINE", "MACHINE_SCOPE_REQUIRED", "FINGERPRINT_SCOPE_MISMATCH":
@@ -20,29 +28,82 @@ func licenseNeedsMachineActivation(code string) bool {
}
}
func licenseValidationErrorIsDefinitive(validation *license.ValidateResponse, err error) bool {
if err == nil {
return validation != nil && !validation.Meta.Valid
}
var apiErr *license.APIError
if !errors.As(err, &apiErr) {
return false
}
if apiErr.StatusCode == http.StatusTooManyRequests || apiErr.StatusCode >= http.StatusInternalServerError {
return false
}
return apiErr.Operation == "activate machine" && validation != nil && !validation.Meta.Valid
}
func licenseValidationErrorMessage(err error) string {
if strings.Contains(err.Error(), "keygen account id is not configured") {
return "授权服务配置错误"
}
var apiErr *license.APIError
if errors.As(err, &apiErr) {
if apiErr.HasCode("MACHINE_LIMIT_EXCEEDED") {
return "授权设备数量已达上限"
}
if apiErr.StatusCode == http.StatusUnauthorized || apiErr.StatusCode == http.StatusForbidden {
return "授权码无效或无权绑定设备"
}
}
return "连接授权服务器失败,请稍后重试"
}
func (h *Handler) validateLicenseForMachine(key string) (*license.ValidateResponse, error) {
fingerprint, err := h.getOrCreateMachineFingerprint()
if err != nil {
return nil, fmt.Errorf("prepare machine fingerprint: %w", err)
}
storedMachineID, _ := h.repo.GetViteConfigValue("license_machine_id")
client := newLicenseClient(keygenAccountID, "")
validation, err := client.ValidateKeyWithFingerprint(key, fingerprint)
accountID := keygenAccountID()
if accountID == "" {
return nil, fmt.Errorf("keygen account id is not configured")
}
client := newLicenseClient(accountID, "")
var validation *license.ValidateResponse
if storedMachineID != "" {
validation, err = client.ValidateKeyWithMachine(key, fingerprint, storedMachineID)
} else {
validation, err = client.ValidateKeyWithFingerprint(key, fingerprint)
}
if err != nil {
return nil, err
}
if validation.Meta.Valid || !licenseNeedsMachineActivation(validation.Meta.Code) {
if validation.Meta.Valid {
validation.MachineID = storedMachineID
}
return validation, nil
}
client.Token = key
if err := client.ActivateMachine(validation.Data.ID, fingerprint); err != nil {
machineID, err := client.ActivateMachine(validation.Data.ID, fingerprint)
if err != nil {
return validation, err
}
if machineID == "" {
machineID, err = client.GetMachineID(fingerprint)
if err != nil {
return validation, fmt.Errorf("retrieve activated machine: %w", err)
}
}
validation, err = client.ValidateKeyWithFingerprint(key, fingerprint)
validation, err = client.ValidateKeyWithMachine(key, fingerprint, machineID)
if err != nil {
return nil, err
}
if validation.Meta.Valid {
validation.MachineID = machineID
}
return validation, nil
}
@@ -2,6 +2,7 @@ package handler
import (
"bytes"
"encoding/json"
"fmt"
"net/http"
"net/http/httptest"
@@ -32,7 +33,9 @@ func TestValidateLicenseJobRepairsMissingMachineBinding(t *testing.T) {
_, _ = fmt.Fprint(w, `{"meta":{"valid":true,"code":"VALID"},"data":{"id":"license-id","attributes":{"expiry":"2030-01-02T00:00:00.000Z"}}}`)
case strings.HasSuffix(req.URL.Path, "/machines"):
w.WriteHeader(http.StatusCreated)
_, _ = fmt.Fprint(w, `{}`)
_, _ = fmt.Fprint(w, `{"data":{"type":"machines","id":"machine-id"}}`)
case strings.Contains(req.URL.Path, "/machines/"):
_, _ = fmt.Fprint(w, `{"data":{"type":"machines","id":"machine-id"}}`)
default:
http.NotFound(w, req)
}
@@ -54,6 +57,43 @@ func TestValidateLicenseJobRepairsMissingMachineBinding(t *testing.T) {
}
}
func TestValidateLicenseJobAcceptsExistingMachineActivation(t *testing.T) {
r := openLicenseTestRepository(t)
now := time.Now().UnixMilli()
seedLicenseConfig(t, r, "license_key", "license-secret", now)
seedLicenseConfig(t, r, "is_commercial", "true", now)
var validations atomic.Int32
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) {
switch {
case strings.HasSuffix(req.URL.Path, "/licenses/actions/validate-key"):
if validations.Add(1) == 1 {
_, _ = fmt.Fprint(w, `{"meta":{"valid":false,"code":"NO_MACHINE"},"data":{"id":"license-id","attributes":{}}}`)
return
}
_, _ = fmt.Fprint(w, `{"meta":{"valid":true,"code":"VALID"},"data":{"id":"license-id","attributes":{"expiry":"never"}}}`)
case strings.HasSuffix(req.URL.Path, "/machines"):
w.WriteHeader(http.StatusUnprocessableEntity)
_, _ = fmt.Fprint(w, `{"errors":[{"code":"FINGERPRINT_TAKEN"},{"code":"MACHINE_LIMIT_EXCEEDED"}]}`)
case strings.Contains(req.URL.Path, "/machines/"):
_, _ = fmt.Fprint(w, `{"data":{"type":"machines","id":"machine-id"}}`)
default:
http.NotFound(w, req)
}
}))
defer server.Close()
restoreLicenseClientFactory(t, server.URL)
h := &Handler{repo: r}
h.validateLicenseJob()
assertLicenseConfig(t, r, "is_commercial", "true")
assertLicenseConfig(t, r, "license_expiry", "never")
if got := validations.Load(); got != 2 {
t.Fatalf("validation calls = %d, want 2", got)
}
}
func TestLicenseActivateRequiresSuccessfulPostActivationValidation(t *testing.T) {
r := openLicenseTestRepository(t)
@@ -68,7 +108,9 @@ func TestLicenseActivateRequiresSuccessfulPostActivationValidation(t *testing.T)
_, _ = fmt.Fprintf(w, `{"meta":{"valid":false,"code":%q},"data":{"id":"license-id","attributes":{}}}`, code)
case strings.HasSuffix(req.URL.Path, "/machines"):
w.WriteHeader(http.StatusCreated)
_, _ = fmt.Fprint(w, `{}`)
_, _ = fmt.Fprint(w, `{"data":{"type":"machines","id":"machine-id"}}`)
case strings.Contains(req.URL.Path, "/machines/"):
_, _ = fmt.Fprint(w, `{"data":{"type":"machines","id":"machine-id"}}`)
default:
http.NotFound(w, req)
}
@@ -84,11 +126,41 @@ func TestLicenseActivateRequiresSuccessfulPostActivationValidation(t *testing.T)
if !strings.Contains(res.Body.String(), "FINGERPRINT_SCOPE_MISMATCH") {
t.Fatalf("expected post-activation validation failure, got %s", res.Body.String())
}
if value, err := r.GetViteConfigValue("is_commercial"); err == nil || value != "" {
t.Fatalf("commercial status should not be persisted, got value=%q err=%v", value, err)
assertLicenseConfig(t, r, "is_commercial", "false")
for _, name := range []string{"license_key", "license_expiry"} {
if value, err := r.GetViteConfigValue(name); err == nil || value != "" {
t.Fatalf("%s should not be persisted, got value=%q err=%v", name, value, err)
}
}
}
func TestLicenseActivatePersistsValidatedState(t *testing.T) {
r := openLicenseTestRepository(t)
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) {
switch {
case strings.HasSuffix(req.URL.Path, "/licenses/actions/validate-key"):
_, _ = fmt.Fprint(w, `{"meta":{"valid":true,"code":"VALID"},"data":{"id":"license-id","attributes":{"expiry":"2030-01-02T00:00:00.000Z"}}}`)
default:
http.NotFound(w, req)
}
}))
defer server.Close()
restoreLicenseClientFactory(t, server.URL)
h := &Handler{repo: r}
req := httptest.NewRequest(http.MethodPost, "/api/v1/license/activate", bytes.NewBufferString(`{"license_key":"license-secret"}`))
res := httptest.NewRecorder()
h.licenseActivate(res, req)
if !strings.Contains(res.Body.String(), `"code":0`) {
t.Fatalf("expected activation success, got %s", res.Body.String())
}
assertLicenseConfig(t, r, "license_key", "license-secret")
assertLicenseConfig(t, r, "is_commercial", "true")
assertLicenseConfig(t, r, "license_expiry", "2030-01-02T00:00:00.000Z")
}
func TestValidateLicenseJobDowngradesWhenMachineBindingIsRejected(t *testing.T) {
r := openLicenseTestRepository(t)
now := time.Now().UnixMilli()
@@ -115,6 +187,136 @@ func TestValidateLicenseJobDowngradesWhenMachineBindingIsRejected(t *testing.T)
assertLicenseConfig(t, r, "is_commercial", "false")
}
func TestValidateLicenseJobRestoresCommercialStateWhenLicenseRecovers(t *testing.T) {
r := openLicenseTestRepository(t)
now := time.Now().UnixMilli()
seedLicenseConfig(t, r, "license_key", "license-secret", now)
seedLicenseConfig(t, r, "is_commercial", "false", now)
seedLicenseConfig(t, r, "machine_fingerprint", "fingerprint", now)
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) {
if strings.HasSuffix(req.URL.Path, "/licenses/actions/validate-key") {
_, _ = fmt.Fprint(w, `{"meta":{"valid":true,"code":"VALID"},"data":{"id":"license-id","attributes":{"expiry":"never"}}}`)
return
}
http.NotFound(w, req)
}))
defer server.Close()
restoreLicenseClientFactory(t, server.URL)
h := &Handler{repo: r}
h.validateLicenseJob()
assertLicenseConfig(t, r, "is_commercial", "true")
assertLicenseConfig(t, r, "license_expiry", "never")
}
func TestValidateLicenseJobUsesStoredMachineScope(t *testing.T) {
r := openLicenseTestRepository(t)
now := time.Now().UnixMilli()
seedLicenseConfig(t, r, "license_key", "license-secret", now)
seedLicenseConfig(t, r, "is_commercial", "true", now)
seedLicenseConfig(t, r, "machine_fingerprint", "fingerprint", now)
seedLicenseConfig(t, r, "license_machine_id", "machine-id", now)
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) {
if !strings.HasSuffix(req.URL.Path, "/licenses/actions/validate-key") {
t.Fatalf("unexpected request %s %s", req.Method, req.URL.Path)
}
var body struct {
Meta struct {
Scope map[string]string `json:"scope"`
} `json:"meta"`
}
if err := json.NewDecoder(req.Body).Decode(&body); err != nil {
t.Fatalf("decode request: %v", err)
}
if body.Meta.Scope["machine"] != "machine-id" || body.Meta.Scope["fingerprint"] != "fingerprint" {
t.Fatalf("unexpected validation scope: %+v", body.Meta.Scope)
}
_, _ = fmt.Fprint(w, `{"meta":{"valid":true,"code":"VALID"},"data":{"id":"license-id","attributes":{"expiry":"never"}}}`)
}))
defer server.Close()
restoreLicenseClientFactory(t, server.URL)
h := &Handler{repo: r}
h.validateLicenseJob()
assertLicenseConfig(t, r, "is_commercial", "true")
assertLicenseConfig(t, r, "license_machine_id", "machine-id")
}
func TestValidateLicenseJobKeepsStateOnMachineLookupServerFailure(t *testing.T) {
r := openLicenseTestRepository(t)
now := time.Now().UnixMilli()
seedLicenseConfig(t, r, "license_key", "license-secret", now)
seedLicenseConfig(t, r, "is_commercial", "true", now)
seedLicenseConfig(t, r, "machine_fingerprint", "fingerprint", now)
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) {
switch {
case strings.HasSuffix(req.URL.Path, "/licenses/actions/validate-key"):
_, _ = fmt.Fprint(w, `{"meta":{"valid":false,"code":"MACHINE_SCOPE_REQUIRED"},"data":{"id":"license-id","attributes":{}}}`)
case strings.HasSuffix(req.URL.Path, "/machines"):
w.WriteHeader(http.StatusUnprocessableEntity)
_, _ = fmt.Fprint(w, `{"errors":[{"code":"FINGERPRINT_TAKEN"}]}`)
case strings.Contains(req.URL.Path, "/machines/"):
w.WriteHeader(http.StatusServiceUnavailable)
_, _ = fmt.Fprint(w, `{"errors":[{"code":"SERVICE_UNAVAILABLE"}]}`)
default:
http.NotFound(w, req)
}
}))
defer server.Close()
restoreLicenseClientFactory(t, server.URL)
h := &Handler{repo: r}
h.validateLicenseJob()
assertLicenseConfig(t, r, "is_commercial", "true")
}
func TestValidateLicenseJobKeepsStateOnMachineLookupNotFound(t *testing.T) {
r := openLicenseTestRepository(t)
now := time.Now().UnixMilli()
seedLicenseConfig(t, r, "license_key", "license-secret", now)
seedLicenseConfig(t, r, "is_commercial", "true", now)
seedLicenseConfig(t, r, "machine_fingerprint", "fingerprint", now)
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) {
switch {
case strings.HasSuffix(req.URL.Path, "/licenses/actions/validate-key"):
_, _ = fmt.Fprint(w, `{"meta":{"valid":false,"code":"MACHINE_SCOPE_REQUIRED"},"data":{"id":"license-id","attributes":{}}}`)
case strings.HasSuffix(req.URL.Path, "/machines"):
w.WriteHeader(http.StatusUnprocessableEntity)
_, _ = fmt.Fprint(w, `{"errors":[{"code":"FINGERPRINT_TAKEN"}]}`)
case strings.Contains(req.URL.Path, "/machines/"):
http.NotFound(w, req)
default:
http.NotFound(w, req)
}
}))
defer server.Close()
restoreLicenseClientFactory(t, server.URL)
h := &Handler{repo: r}
h.validateLicenseJob()
assertLicenseConfig(t, r, "is_commercial", "true")
}
func TestLicenseValidationErrorMessageDoesNotExposeKeygenResponse(t *testing.T) {
err := &license.APIError{
Operation: "activate machine",
StatusCode: http.StatusUnprocessableEntity,
Body: `{"errors":[{"code":"MACHINE_LIMIT_EXCEEDED","detail":"private detail"}]}`,
}
message := licenseValidationErrorMessage(err)
if message != "授权设备数量已达上限" || strings.Contains(message, "private detail") {
t.Fatalf("unexpected public error message %q", message)
}
}
func openLicenseTestRepository(t *testing.T) *repo.Repository {
t.Helper()
r, err := repo.Open(filepath.Join(t.TempDir(), "license.db"))
@@ -145,6 +347,7 @@ func assertLicenseConfig(t *testing.T, r *repo.Repository, name, want string) {
func restoreLicenseClientFactory(t *testing.T, baseURL string) {
t.Helper()
t.Setenv("KEYGEN_ACCOUNT_ID", "account-id")
previous := newLicenseClient
newLicenseClient = func(accountID, token string) *license.KeygenClient {
client := license.NewKeygenClient(accountID, token)
+65 -9
View File
@@ -57,7 +57,11 @@ func (h *Handler) userCreate(w http.ResponseWriter, r *http.Request) {
}
status := asInt(req["status"], 1)
flow := asInt64(req["flow"], 100)
flow, flowMiB, flowErr := parseTrafficLimit(req, 100)
if flowErr != nil {
response.WriteJSON(w, response.ErrDefault(flowErr.Error()))
return
}
num := asInt(req["num"], 10)
expTime := asInt64(req["expTime"], time.Now().Add(365*24*time.Hour).UnixMilli())
flowResetTime := asInt64(req["flowResetTime"], 1)
@@ -76,7 +80,7 @@ func (h *Handler) userCreate(w http.ResponseWriter, r *http.Request) {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
userID, err := h.repo.CreateUser(username, hashedPassword, roleID, expTime, flow, flowResetTime, num, status, maxConn, now)
userID, err := h.repo.CreateUser(username, hashedPassword, roleID, expTime, flow, flowResetTime, num, status, maxConn, now, flowMiB)
if err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
@@ -164,7 +168,16 @@ func (h *Handler) userUpdate(w http.ResponseWriter, r *http.Request) {
return
}
flow := asInt64(req["flow"], 100)
flow, flowMiB, flowErr := parseTrafficLimit(req, 100)
if flowErr != nil {
response.WriteJSON(w, response.ErrDefault(flowErr.Error()))
return
}
if _, supplied := req["flowMiB"]; !supplied {
if current, err := h.repo.GetUserByID(id); err == nil && current != nil && current.Flow == flow {
flowMiB = current.FlowMiB
}
}
num := asInt(req["num"], 10)
expTime := asInt64(req["expTime"], time.Now().Add(365*24*time.Hour).UnixMilli())
flowResetTime := asInt64(req["flowResetTime"], 1)
@@ -176,7 +189,7 @@ func (h *Handler) userUpdate(w http.ResponseWriter, r *http.Request) {
pwd := asString(req["pwd"])
if strings.TrimSpace(pwd) == "" {
if err := h.repo.UpdateUserWithoutPassword(id, username, flow, num, expTime, flowResetTime, status, maxConn, now); err != nil {
if err := h.repo.UpdateUserWithoutPassword(id, username, flow, num, expTime, flowResetTime, status, maxConn, now, flowMiB); err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
@@ -186,13 +199,13 @@ func (h *Handler) userUpdate(w http.ResponseWriter, r *http.Request) {
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 {
if err := h.repo.UpdateUserWithPassword(id, username, hashedPassword, flow, num, expTime, flowResetTime, status, maxConn, now, flowMiB); err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
}
h.repo.PropagateUserFlowToTunnels(id, flow, num, expTime, flowResetTime)
h.repo.PropagateUserFlowToTunnels(id, flow, num, expTime, flowResetTime, flowMiB)
if hasDailyQuota || hasMonthlyQuota {
dailyQuotaGB := asInt64(req["dailyQuotaGB"], 0)
monthlyQuotaGB := asInt64(req["monthlyQuotaGB"], 0)
@@ -1991,14 +2004,32 @@ func (h *Handler) userTunnelUpdate(w http.ResponseWriter, r *http.Request) {
response.WriteJSON(w, response.Err(-2, oldErr.Error()))
return
}
oldTunnel, oldTunnelErr := h.repo.GetUserTunnelByID(id)
if oldTunnelErr != nil {
response.WriteJSON(w, response.Err(-2, oldTunnelErr.Error()))
return
}
if oldTunnel == nil {
response.WriteJSON(w, response.ErrDefault("隧道权限不存在"))
return
}
flow, flowMiB, flowErr := parseTrafficLimit(req, 0)
if flowErr != nil {
response.WriteJSON(w, response.ErrDefault(flowErr.Error()))
return
}
if _, supplied := req["flowMiB"]; !supplied && oldTunnel.Flow == flow {
flowMiB = oldTunnel.FlowMiB
}
if err := h.repo.UpdateUserTunnel(id,
asInt64(req["flow"], 0),
flow,
asInt(req["num"], 0),
asInt64(req["expTime"], time.Now().Add(365*24*time.Hour).UnixMilli()),
asInt64(req["flowResetTime"], 1),
nullableInt(speedID),
asInt(req["status"], 1),
flowMiB,
); err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
@@ -2013,6 +2044,7 @@ func (h *Handler) userTunnelUpdate(w http.ResponseWriter, r *http.Request) {
oldFlowReset,
oldSpeedID,
oldStatus,
oldTunnel.FlowMiB,
)
if rollbackErr != nil {
response.WriteJSON(w, response.Err(-2, fmt.Sprintf("下发失败且回滚失败: %v; 回滚错误: %v", syncErr, rollbackErr)))
@@ -4781,6 +4813,14 @@ func (h *Handler) upsertUserTunnel(req map[string]interface{}) error {
}
reqFlow := asInt64(req["flow"], -1)
var reqFlowMiB int64
if _, hasFlowMiB := req["flowMiB"]; hasFlowMiB {
var flowErr error
reqFlow, reqFlowMiB, flowErr = parseTrafficLimit(req, 0)
if flowErr != nil {
return flowErr
}
}
reqNum := asInt(req["num"], -1)
reqExpTime := asInt64(req["expTime"], -1)
reqFlowReset := asInt64(req["flowResetTime"], -1)
@@ -4792,6 +4832,9 @@ func (h *Handler) upsertUserTunnel(req map[string]interface{}) error {
if uErr == nil {
if reqFlow < 0 {
reqFlow = uFlow
if user, err := h.repo.GetUserByID(userID); err == nil && user != nil {
reqFlowMiB = user.FlowMiB
}
}
if reqNum < 0 {
reqNum = uNum
@@ -4820,7 +4863,7 @@ func (h *Handler) upsertUserTunnel(req map[string]interface{}) error {
reqStatus = 1
}
if err := h.repo.InsertUserTunnel(userID, tunnelID, nullableInt(speedID), reqNum, reqFlow, reqFlowReset, reqExpTime, reqStatus); err != nil {
if err := h.repo.InsertUserTunnel(userID, tunnelID, nullableInt(speedID), reqNum, reqFlow, reqFlowReset, reqExpTime, reqStatus, reqFlowMiB); err != nil {
return err
}
@@ -4844,8 +4887,20 @@ func (h *Handler) upsertUserTunnel(req map[string]interface{}) error {
}
newFlow := currentFlow
oldTunnel, err := h.repo.GetUserTunnelByID(existingID)
if err != nil {
return err
}
if oldTunnel == nil {
return fmt.Errorf("隧道权限不存在")
}
newFlowMiB := oldTunnel.FlowMiB
if reqFlow >= 0 {
newFlow = reqFlow
newFlowMiB = reqFlowMiB
if _, supplied := req["flowMiB"]; !supplied && reqFlow == currentFlow {
newFlowMiB = oldTunnel.FlowMiB
}
}
newNum := int(currentNum)
@@ -4875,7 +4930,7 @@ func (h *Handler) upsertUserTunnel(req map[string]interface{}) error {
newSpeedID = sql.NullInt64{Valid: false}
}
if err := h.repo.UpdateUserTunnelFields(existingID, newSpeedID, newFlow, newNum, newExpTime, newFlowReset, newStatus); err != nil {
if err := h.repo.UpdateUserTunnelFields(existingID, newSpeedID, newFlow, newNum, newExpTime, newFlowReset, newStatus, newFlowMiB); err != nil {
return err
}
@@ -4888,6 +4943,7 @@ func (h *Handler) upsertUserTunnel(req map[string]interface{}) error {
currentExpTime,
currentFlowReset,
currentStatus,
oldTunnel.FlowMiB,
)
if rollbackErr != nil {
return fmt.Errorf("下发失败且回滚失败: %v; 回滚错误: %w", syncErr, rollbackErr)
@@ -0,0 +1,28 @@
package handler
import (
"fmt"
"math"
"strconv"
)
// flowMiB is optional so older clients can keep sending the GB-based flow field.
// A positive value takes precedence and preserves sub-GB limits exactly.
func parseTrafficLimit(req map[string]interface{}, defaultGB int64) (flowGB, flowMiB int64, err error) {
flowGB = asInt64(req["flow"], defaultGB)
if flowGB < 0 {
return 0, 0, fmt.Errorf("流量限制不能小于0")
}
raw, present := req["flowMiB"]
if !present {
return flowGB, 0, nil
}
flowMiB, err = strconv.ParseInt(asString(raw), 10, 64)
if err != nil || flowMiB < 0 || flowMiB > math.MaxInt64/bytesPerMiB {
return 0, 0, fmt.Errorf("流量限制超出范围")
}
if flowMiB > 0 {
flowGB = (flowMiB-1)/1024 + 1
}
return flowGB, flowMiB, nil
}
@@ -0,0 +1,39 @@
package handler
import (
"testing"
"time"
)
func TestTrafficLimitMiBOverridesLegacyGB(t *testing.T) {
flowGB, flowMiB, err := parseTrafficLimit(map[string]interface{}{
"flow": float64(1), "flowMiB": float64(500),
}, 100)
if err != nil || flowGB != 1 || flowMiB != 500 {
t.Fatalf("parseTrafficLimit = (%d, %d, %v), want (1, 500, nil)", flowGB, flowMiB, err)
}
limit := flowLimitBytes(flowGB, flowMiB)
if limit != 500*bytesPerMiB {
t.Fatalf("limit = %d, want %d", limit, 500*bytesPerMiB)
}
policy := &userTunnelPolicy{Flow: flowGB, FlowMiB: flowMiB, InFlow: limit - 1, Status: 1}
if shouldPauseUserTunnel(policy, time.Now().UnixMilli()) {
t.Fatal("policy paused before reaching 500 MiB")
}
policy.InFlow = limit
if !shouldPauseUserTunnel(policy, time.Now().UnixMilli()) {
t.Fatal("policy did not pause at 500 MiB")
}
}
func TestTrafficLimitLegacyAndInvalidValues(t *testing.T) {
flowGB, flowMiB, err := parseTrafficLimit(map[string]interface{}{"flow": float64(2)}, 100)
if err != nil || flowGB != 2 || flowMiB != 0 || flowLimitBytes(flowGB, flowMiB) != 2*bytesPerGB {
t.Fatalf("legacy GB limit changed: (%d, %d, %v)", flowGB, flowMiB, err)
}
for _, value := range []interface{}{"1.5", -1, "999999999999999999999"} {
if _, _, err := parseTrafficLimit(map[string]interface{}{"flowMiB": value}, 100); err == nil {
t.Fatalf("accepted invalid flowMiB %v", value)
}
}
}
@@ -182,6 +182,8 @@ func requiresAdmin(path string) bool {
return true
case "/api/v1/config/update", "/api/v1/config/update-single":
return true
case "/api/v1/license/activate":
return true
case "/api/v1/announcement/update":
return true
default:
@@ -197,6 +197,39 @@ func TestShouldSkipBypassesPublicConfigGet(t *testing.T) {
}
}
func TestLicenseActivateRequiresAdmin(t *testing.T) {
if !requiresAdmin("/api/v1/license/activate") {
t.Fatal("expected license activation to require admin")
}
}
func TestJWTRejectsNonAdminLicenseActivation(t *testing.T) {
secret := "unit-test-secret"
token, err := auth.GenerateToken(2, "regular_user", 1, 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: 1, 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/license/activate", nil)
req.Header.Set("Authorization", token)
res := httptest.NewRecorder()
wrapped.ServeHTTP(res, req)
assertCodeMsg(t, res, 403, "权限不足,仅管理员可操作")
}
func TestJWTExpiresAfterSevenDays(t *testing.T) {
secret := "unit-test-secret"
token, err := auth.GenerateToken(1, "admin_user", 0, secret)
+112 -10
View File
@@ -6,10 +6,13 @@ import (
"fmt"
"io"
"net/http"
"net/url"
"strings"
"time"
)
var AccountID string
type KeygenClient struct {
AccountID string
Token string
@@ -17,6 +20,20 @@ type KeygenClient struct {
HTTPClient *http.Client
}
type APIError struct {
Operation string
StatusCode int
Body string
}
func (e *APIError) Error() string {
return fmt.Sprintf("keygen %s failed: status %d, response: %s", e.Operation, e.StatusCode, e.Body)
}
func (e *APIError) HasCode(code string) bool {
return e != nil && hasKeygenErrorCode([]byte(e.Body), code)
}
const defaultAPIBaseURL = "https://api.keygen.sh/v1"
func NewKeygenClient(accountID, token string) *KeygenClient {
@@ -47,6 +64,7 @@ type ValidateResponse struct {
Expiry string `json:"expiry"`
} `json:"attributes"`
} `json:"data"`
MachineID string `json:"-"`
}
type ActivateMachineRequest struct {
@@ -66,7 +84,36 @@ type ActivateMachineRequest struct {
} `json:"data"`
}
type keygenErrorResponse struct {
Errors []struct {
Code string `json:"code"`
} `json:"errors"`
}
type MachineResponse struct {
Data struct {
ID string `json:"id"`
} `json:"data"`
}
func hasKeygenErrorCode(body []byte, code string) bool {
var resp keygenErrorResponse
if err := json.Unmarshal(body, &resp); err != nil {
return false
}
for _, item := range resp.Errors {
if strings.EqualFold(strings.TrimSpace(item.Code), code) {
return true
}
}
return false
}
func (c *KeygenClient) ValidateKeyWithFingerprint(key string, fingerprint string) (*ValidateResponse, error) {
return c.ValidateKeyWithMachine(key, fingerprint, "")
}
func (c *KeygenClient) ValidateKeyWithMachine(key, fingerprint, machineID string) (*ValidateResponse, error) {
url := c.apiURL("licenses/actions/validate-key")
meta := map[string]interface{}{
@@ -78,6 +125,14 @@ func (c *KeygenClient) ValidateKeyWithFingerprint(key string, fingerprint string
"fingerprint": fingerprint,
}
}
if machineID != "" {
scope, _ := meta["scope"].(map[string]interface{})
if scope == nil {
scope = make(map[string]interface{})
meta["scope"] = scope
}
scope["machine"] = machineID
}
reqBody := map[string]interface{}{
"meta": meta,
@@ -103,7 +158,8 @@ func (c *KeygenClient) ValidateKeyWithFingerprint(key string, fingerprint string
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
return nil, fmt.Errorf("keygen api error: status %d", resp.StatusCode)
body, _ := io.ReadAll(resp.Body)
return nil, &APIError{Operation: "validate license", StatusCode: resp.StatusCode, Body: string(body)}
}
var valResp ValidateResponse
@@ -114,6 +170,42 @@ func (c *KeygenClient) ValidateKeyWithFingerprint(key string, fingerprint string
return &valResp, nil
}
func (c *KeygenClient) GetMachineID(fingerprint string) (string, error) {
machineURL := c.apiURL("machines/" + url.PathEscape(fingerprint))
req, err := http.NewRequest(http.MethodGet, machineURL, nil)
if err != nil {
return "", err
}
req.Header.Set("Accept", "application/vnd.api+json")
if c.Token != "" {
if !strings.HasPrefix(c.Token, "Bearer ") && !strings.HasPrefix(c.Token, "License ") {
req.Header.Set("Authorization", "License "+c.Token)
} else {
req.Header.Set("Authorization", c.Token)
}
}
resp, err := c.HTTPClient.Do(req)
if err != nil {
return "", err
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
body, _ := io.ReadAll(resp.Body)
return "", &APIError{Operation: "retrieve machine", StatusCode: resp.StatusCode, Body: string(body)}
}
var machineResp MachineResponse
if err := json.NewDecoder(resp.Body).Decode(&machineResp); err != nil {
return "", err
}
machineID := strings.TrimSpace(machineResp.Data.ID)
if machineID == "" {
return "", fmt.Errorf("failed to retrieve machine: empty machine id")
}
return machineID, nil
}
func (c *KeygenClient) ValidateKey(key string) (*ValidateResponse, error) {
url := c.apiURL("licenses/actions/validate-key")
@@ -142,7 +234,8 @@ func (c *KeygenClient) ValidateKey(key string) (*ValidateResponse, error) {
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
return nil, fmt.Errorf("keygen api error: status %d", resp.StatusCode)
body, _ := io.ReadAll(resp.Body)
return nil, &APIError{Operation: "validate license", StatusCode: resp.StatusCode, Body: string(body)}
}
var valResp ValidateResponse
@@ -153,7 +246,7 @@ func (c *KeygenClient) ValidateKey(key string) (*ValidateResponse, error) {
return &valResp, nil
}
func (c *KeygenClient) ActivateMachine(licenseID, fingerprint string) error {
func (c *KeygenClient) ActivateMachine(licenseID, fingerprint string) (string, error) {
url := c.apiURL("machines")
var reqBody ActivateMachineRequest
@@ -177,15 +270,24 @@ func (c *KeygenClient) ActivateMachine(licenseID, fingerprint string) error {
resp, err := c.HTTPClient.Do(req)
if err != nil {
return err
return "", err
}
defer resp.Body.Close()
if resp.StatusCode == http.StatusCreated || resp.StatusCode == http.StatusOK {
return nil
}
body, _ := io.ReadAll(resp.Body)
return fmt.Errorf("failed to activate machine: status %d, response: %s", resp.StatusCode, string(body))
if resp.StatusCode == http.StatusCreated || resp.StatusCode == http.StatusOK {
var machineResp MachineResponse
if json.Unmarshal(body, &machineResp) == nil {
return strings.TrimSpace(machineResp.Data.ID), nil
}
return "", nil
}
if resp.StatusCode == http.StatusUnprocessableEntity && hasKeygenErrorCode(body, "FINGERPRINT_TAKEN") {
// Machine activation is idempotent. Keygen scopes fingerprint uniqueness
// to the target license, so this means the same machine is already bound.
return "", nil
}
return "", &APIError{Operation: "activate machine", StatusCode: resp.StatusCode, Body: string(body)}
}
+104
View File
@@ -0,0 +1,104 @@
package license
import (
"encoding/json"
"fmt"
"net/http"
"net/http/httptest"
"strings"
"testing"
)
func TestValidateKeyWithMachineSendsFingerprintAndMachineScopes(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
var body struct {
Meta struct {
Scope map[string]string `json:"scope"`
} `json:"meta"`
}
if err := json.NewDecoder(r.Body).Decode(&body); err != nil {
t.Fatalf("decode request: %v", err)
}
if body.Meta.Scope["fingerprint"] != "fingerprint" {
t.Fatalf("fingerprint scope = %q", body.Meta.Scope["fingerprint"])
}
if body.Meta.Scope["machine"] != "machine-id" {
t.Fatalf("machine scope = %q", body.Meta.Scope["machine"])
}
_, _ = fmt.Fprint(w, `{"meta":{"valid":true,"code":"VALID"},"data":{"id":"license-id","attributes":{}}}`)
}))
defer server.Close()
client := NewKeygenClient("account-id", "")
client.BaseURL = server.URL
validation, err := client.ValidateKeyWithMachine("license-key", "fingerprint", "machine-id")
if err != nil || !validation.Meta.Valid {
t.Fatalf("ValidateKeyWithMachine() validation=%+v err=%v", validation, err)
}
}
func TestGetMachineIDRetrievesMachineByFingerprint(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodGet || !strings.HasSuffix(r.URL.Path, "/machines/fingerprint") {
t.Fatalf("unexpected request %s %s", r.Method, r.URL.Path)
}
if r.Header.Get("Authorization") != "License license-key" {
t.Fatalf("authorization = %q", r.Header.Get("Authorization"))
}
_, _ = fmt.Fprint(w, `{"data":{"type":"machines","id":"machine-id"}}`)
}))
defer server.Close()
client := NewKeygenClient("account-id", "license-key")
client.BaseURL = server.URL
machineID, err := client.GetMachineID("fingerprint")
if err != nil || machineID != "machine-id" {
t.Fatalf("GetMachineID() = %q, %v", machineID, err)
}
}
func TestActivateMachineTreatsFingerprintTakenAsIdempotent(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.WriteHeader(http.StatusUnprocessableEntity)
_, _ = fmt.Fprint(w, `{"errors":[{"code":"FINGERPRINT_TAKEN"},{"code":"MACHINE_LIMIT_EXCEEDED"}]}`)
}))
defer server.Close()
client := NewKeygenClient("account-id", "license-key")
client.BaseURL = server.URL
if machineID, err := client.ActivateMachine("license-id", "fingerprint"); err != nil || machineID != "" {
t.Fatalf("ActivateMachine() error = %v, want idempotent success", err)
}
}
func TestActivateMachineReturnsCreatedMachineID(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.WriteHeader(http.StatusCreated)
_, _ = fmt.Fprint(w, `{"data":{"type":"machines","id":"machine-id"}}`)
}))
defer server.Close()
client := NewKeygenClient("account-id", "license-key")
client.BaseURL = server.URL
machineID, err := client.ActivateMachine("license-id", "fingerprint")
if err != nil || machineID != "machine-id" {
t.Fatalf("ActivateMachine() = %q, %v", machineID, err)
}
}
func TestActivateMachineRejectsMachineLimitWithoutFingerprintTaken(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.WriteHeader(http.StatusUnprocessableEntity)
_, _ = fmt.Fprint(w, `{"errors":[{"code":"MACHINE_LIMIT_EXCEEDED"}]}`)
}))
defer server.Close()
client := NewKeygenClient("account-id", "license-key")
client.BaseURL = server.URL
_, err := client.ActivateMachine("license-id", "fingerprint")
if err == nil || !strings.Contains(err.Error(), "MACHINE_LIMIT_EXCEEDED") {
t.Fatalf("ActivateMachine() error = %v, want machine limit failure", err)
}
}
@@ -46,7 +46,8 @@ func RenderTable(plan NodePlan) string {
}
targetHost := strings.Trim(strings.TrimSpace(rule.TargetHost), "[]")
for _, protocol := range normalizedProtocols(rule.Protocols) {
b.WriteString(fmt.Sprintf(" ct original proto-dst %d %s daddr %s %s dport %d counter comment %q\n",
b.WriteString(fmt.Sprintf(" meta l4proto %s ct original proto-dst %d %s daddr %s %s dport %d counter comment %q\n",
protocol,
rule.InPort,
family,
targetHost,
@@ -54,7 +55,8 @@ func RenderTable(plan NodePlan) string {
rule.TargetPort,
counterComment(rule.ForwardID, CounterDirectionToTarget, protocol),
))
b.WriteString(fmt.Sprintf(" ct original proto-dst %d %s saddr %s %s sport %d counter comment %q\n",
b.WriteString(fmt.Sprintf(" meta l4proto %s ct original proto-dst %d %s saddr %s %s sport %d counter comment %q\n",
protocol,
rule.InPort,
family,
targetHost,
@@ -62,10 +62,10 @@ func TestRenderTableIncludesForwardAccountingCounters(t *testing.T) {
wantLines := []string{
`tcp dport 12345 counter dnat ip to 198.51.100.20:443 comment "flvx forward:42 dnat tcp"`,
`udp dport 12345 counter dnat ip to 198.51.100.20:443 comment "flvx forward:42 dnat udp"`,
`ct original proto-dst 12345 ip daddr 198.51.100.20 tcp dport 443 counter comment "flvx forward:42 to-target tcp"`,
`ct original proto-dst 12345 ip saddr 198.51.100.20 tcp sport 443 counter comment "flvx forward:42 from-target tcp"`,
`ct original proto-dst 12345 ip daddr 198.51.100.20 udp dport 443 counter comment "flvx forward:42 to-target udp"`,
`ct original proto-dst 12345 ip saddr 198.51.100.20 udp sport 443 counter comment "flvx forward:42 from-target udp"`,
`meta l4proto tcp ct original proto-dst 12345 ip daddr 198.51.100.20 tcp dport 443 counter comment "flvx forward:42 to-target tcp"`,
`meta l4proto tcp ct original proto-dst 12345 ip saddr 198.51.100.20 tcp sport 443 counter comment "flvx forward:42 from-target tcp"`,
`meta l4proto udp ct original proto-dst 12345 ip daddr 198.51.100.20 udp dport 443 counter comment "flvx forward:42 to-target udp"`,
`meta l4proto udp ct original proto-dst 12345 ip saddr 198.51.100.20 udp sport 443 counter comment "flvx forward:42 from-target udp"`,
}
for _, want := range wantLines {
if !strings.Contains(got, want) {
@@ -89,8 +89,8 @@ func TestRenderTableIncludesIPv6ForwardAccountingCounters(t *testing.T) {
got := RenderTable(plan)
wantLines := []string{
`tcp dport 12346 counter dnat ip6 to [2001:db8::20]:8443 comment "flvx forward:43 dnat tcp"`,
`ct original proto-dst 12346 ip6 daddr 2001:db8::20 tcp dport 8443 counter comment "flvx forward:43 to-target tcp"`,
`ct original proto-dst 12346 ip6 saddr 2001:db8::20 tcp sport 8443 counter comment "flvx forward:43 from-target tcp"`,
`meta l4proto tcp ct original proto-dst 12346 ip6 daddr 2001:db8::20 tcp dport 8443 counter comment "flvx forward:43 to-target tcp"`,
`meta l4proto tcp ct original proto-dst 12346 ip6 saddr 2001:db8::20 tcp sport 8443 counter comment "flvx forward:43 from-target tcp"`,
}
for _, want := range wantLines {
if !strings.Contains(got, want) {
@@ -110,8 +110,8 @@ func TestRenderTableAccountingCountersIncludeOriginalPort(t *testing.T) {
got := RenderTable(plan)
wantLines := []string{
`ct original proto-dst 12345 ip daddr 198.51.100.20 tcp dport 443 counter comment "flvx forward:42 to-target tcp"`,
`ct original proto-dst 12346 ip daddr 198.51.100.20 tcp dport 443 counter comment "flvx forward:43 to-target tcp"`,
`meta l4proto tcp ct original proto-dst 12345 ip daddr 198.51.100.20 tcp dport 443 counter comment "flvx forward:42 to-target tcp"`,
`meta l4proto tcp ct original proto-dst 12346 ip daddr 198.51.100.20 tcp dport 443 counter comment "flvx forward:43 to-target tcp"`,
}
for _, want := range wantLines {
if !strings.Contains(got, want) {
+43 -8
View File
@@ -26,23 +26,58 @@ func NewSSHRunner() *SSHRunner {
}
func (r *SSHRunner) Test(ctx context.Context, cfg SSHConfig) error {
return r.run(ctx, cfg, "command -v nft >/dev/null 2>&1 && nft --version >/dev/null 2>&1")
nft := nftBinary(cfg)
tableName := fmt.Sprintf("flvx_capability_%d", time.Now().UnixNano())
return r.run(ctx, cfg, buildCapabilityCheckCommand(nft, tableName))
}
func buildCapabilityCheckCommand(nft, tableName string) string {
script := RenderTable(NodePlan{
Rules: []Rule{{
ForwardID: 1,
InPort: 12345,
TargetHost: "192.0.2.1",
TargetPort: 443,
Protocols: []string{"tcp", "udp"},
}},
})
script = strings.Replace(script, "table inet flvx {", "table inet "+tableName+" {", 1)
return "set -eu\n" +
"command -v nft >/dev/null 2>&1\n" +
nft + " --version >/dev/null 2>&1\n" +
"tmp=$(mktemp /tmp/flvx-nft-capability-XXXXXX.nft)\n" +
"trap 'rm -f \"$tmp\"' EXIT\n" +
"cat > \"$tmp\" <<'EOF'\n" + script + "\nEOF\n" +
"if ! " + nft + " -c -f \"$tmp\"; then\n" +
" echo 'nftables cannot validate the generated FLVX rules' >&2\n" +
" exit 1\n" +
"fi"
}
func (r *SSHRunner) ApplyScript(ctx context.Context, cfg SSHConfig, script string) error {
nft := nftBinary(cfg)
command := "tmp=$(mktemp /tmp/flvx-nft-XXXXXX.nft) || exit 1\n" +
command := buildApplyCommand(nftBinary(cfg), script)
return r.run(ctx, cfg, command)
}
func buildApplyCommand(nft, script string) string {
return "set -eu\n" +
"tmp=$(mktemp /tmp/flvx-nft-XXXXXX.nft)\n" +
"batch=$(mktemp /tmp/flvx-nft-batch-XXXXXX.nft) || { rm -f \"$tmp\"; exit 1; }\n" +
"cleanup() {\n" +
" rm -f \"$tmp\"\n" +
" rm -f \"$tmp\" \"$batch\"\n" +
"}\n" +
"trap cleanup EXIT\n" +
"cat > \"$tmp\" <<'EOF'\n" + script + "\nEOF\n" +
nft + " -c -f \"$tmp\"\n" +
"if " + nft + " list table inet flvx >/dev/null 2>&1; then\n" +
" " + nft + " delete table inet flvx\n" +
" { printf '%s\\n' 'delete table inet flvx'; cat \"$tmp\"; } > \"$batch\"\n" +
"else\n" +
" cp \"$tmp\" \"$batch\"\n" +
"fi\n" +
nft + " -f \"$tmp\""
return r.run(ctx, cfg, command)
"if ! " + nft + " -c -f \"$batch\"; then\n" +
" echo 'nftables rule validation failed; active rules were preserved' >&2\n" +
" exit 1\n" +
"fi\n" +
nft + " -f \"$batch\""
}
func (r *SSHRunner) ListTableJSON(ctx context.Context, cfg SSHConfig) ([]byte, error) {
@@ -5,10 +5,113 @@ import (
"crypto/rsa"
"crypto/x509"
"encoding/pem"
"os"
"os/exec"
"path/filepath"
"strings"
"testing"
)
func TestBuildApplyCommandStopsAfterValidationFailure(t *testing.T) {
dir := t.TempDir()
logPath := filepath.Join(dir, "calls.log")
applyMarker := filepath.Join(dir, "applied")
nftPath := filepath.Join(dir, "nft")
fake := `#!/bin/sh
echo "$*" >> "` + logPath + `"
if [ "$1" = "list" ]; then
exit 0
fi
if [ "$1" = "-c" ]; then
exit 1
fi
touch "` + applyMarker + `"
`
if err := os.WriteFile(nftPath, []byte(fake), 0o755); err != nil {
t.Fatalf("write fake nft: %v", err)
}
command := buildApplyCommand(nftPath, "table inet flvx { }")
result := exec.Command("sh", "-c", command)
if err := result.Run(); err == nil {
t.Fatal("expected validation failure")
}
if _, err := os.Stat(applyMarker); !os.IsNotExist(err) {
t.Fatalf("apply ran after validation failure, stat err=%v", err)
}
calls, err := os.ReadFile(logPath)
if err != nil {
t.Fatalf("read fake nft calls: %v", err)
}
if strings.Count(string(calls), "-f ") != 1 {
t.Fatalf("expected validation only, got calls:\n%s", calls)
}
}
func TestBuildCapabilityCheckCommandValidatesRenderedRulesWithoutApplying(t *testing.T) {
dir := t.TempDir()
logPath := filepath.Join(dir, "calls.log")
nftPath := filepath.Join(dir, "nft")
fake := `#!/bin/sh
echo "$*" >> "` + logPath + `"
if [ "$1" = "--version" ] || [ "$1" = "-c" ]; then
exit 0
fi
exit 1
`
if err := os.WriteFile(nftPath, []byte(fake), 0o755); err != nil {
t.Fatalf("write fake nft: %v", err)
}
command := strings.Replace(buildCapabilityCheckCommand(nftPath, "flvx_capability_test"), "command -v nft", "command -v "+nftPath, 1)
result := exec.Command("sh", "-c", command)
if output, err := result.CombinedOutput(); err != nil {
t.Fatalf("capability command failed: %v: %s", err, output)
}
calls, err := os.ReadFile(logPath)
if err != nil {
t.Fatalf("read fake nft calls: %v", err)
}
if strings.Count(string(calls), "-c -f ") != 1 || strings.Contains(string(calls), "\n-f ") {
t.Fatalf("expected one check-only invocation, got calls:\n%s", calls)
}
if !strings.Contains(command, "table inet flvx_capability_test") || !strings.Contains(command, "meta l4proto tcp ct original proto-dst") {
t.Fatalf("capability check does not contain representative rendered rules:\n%s", command)
}
}
func TestBuildApplyCommandUsesAtomicReplacementBatch(t *testing.T) {
dir := t.TempDir()
batchPath := filepath.Join(dir, "batch.nft")
nftPath := filepath.Join(dir, "nft")
fake := `#!/bin/sh
if [ "$1" = "list" ]; then
exit 0
fi
if [ "$1" = "-f" ]; then
cp "$2" "` + batchPath + `"
fi
exit 0
`
if err := os.WriteFile(nftPath, []byte(fake), 0o755); err != nil {
t.Fatalf("write fake nft: %v", err)
}
script := "table inet flvx {\n chain forward { }\n}"
result := exec.Command("sh", "-c", buildApplyCommand(nftPath, script))
if output, err := result.CombinedOutput(); err != nil {
t.Fatalf("apply command failed: %v: %s", err, output)
}
batch, err := os.ReadFile(batchPath)
if err != nil {
t.Fatalf("read applied batch: %v", err)
}
want := "delete table inet flvx\n" + script + "\n"
if string(batch) != want {
t.Fatalf("atomic batch = %q, want %q", batch, want)
}
}
func TestAuthMethodsDefaultToPrivateKey(t *testing.T) {
privateKey := mustGeneratePrivateKey(t)
methods, err := authMethods(SSHConfig{PrivateKey: privateKey})
+5
View File
@@ -16,6 +16,7 @@ type User struct {
RoleID int `gorm:"column:role_id;not null"`
ExpTime int64 `gorm:"column:exp_time;not null"`
Flow int64 `gorm:"not null"`
FlowMiB int64 `gorm:"column:flow_mib;not null;default:0"`
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"`
@@ -230,6 +231,7 @@ type UserTunnel struct {
SpeedID sql.NullInt64 `gorm:"column:speed_id"`
Num int `gorm:"not null"`
Flow int64 `gorm:"not null"`
FlowMiB int64 `gorm:"column:flow_mib;not null;default:0"`
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"`
@@ -417,6 +419,7 @@ type UserBackup struct {
RoleID int `json:"roleId"`
ExpTime int64 `json:"expTime"`
Flow int64 `json:"flow"`
FlowMiB int64 `json:"flowMiB,omitempty"`
InFlow int64 `json:"inFlow"`
OutFlow int64 `json:"outFlow"`
FlowResetTime int64 `json:"flowResetTime"`
@@ -523,6 +526,7 @@ type UserTunnelBackup struct {
SpeedID int64 `json:"speedId,omitempty"`
Num int `json:"num"`
Flow int64 `json:"flow"`
FlowMiB int64 `json:"flowMiB,omitempty"`
InFlow int64 `json:"inFlow"`
OutFlow int64 `json:"outFlow"`
FlowResetTime int64 `json:"flowResetTime"`
@@ -706,6 +710,7 @@ type UserTunnelDetail struct {
Status int
TunnelFlow int
Flow int64
FlowMiB int64 `gorm:"column:flow_mib"`
InFlow int64
OutFlow int64
Num int
@@ -14,15 +14,30 @@ var publicConfigKeys = map[string]struct{}{
"app_logo": {},
"app_favicon": {},
"app_bg_image": {},
"app_bg_image_light": {},
"app_bg_image_dark": {},
"cloudflare_site_key": {},
"is_commercial": {},
"hide_footer_brand": {},
}
var sensitiveConfigKeys = map[string]struct{}{
"jwt_secret": {},
"license_key": {},
"license_expiry": {},
"license_machine_id": {},
"machine_fingerprint": {},
"cloudflare_secret_key": {},
}
var systemManagedConfigKeys = map[string]struct{}{
"license_key": {},
"license_expiry": {},
"license_machine_id": {},
"is_commercial": {},
"machine_fingerprint": {},
}
func PolicyForConfig(name string) ConfigAccessPolicy {
if IsPublicConfigKey(name) {
return ConfigAccessPublic
@@ -43,6 +58,11 @@ func IsSensitiveConfigKey(name string) bool {
return ok
}
func IsSystemManagedConfigKey(name string) bool {
_, ok := systemManagedConfigKeys[normalizeConfigKey(name)]
return ok
}
func FilterSensitiveConfigs(in map[string]string) map[string]string {
if len(in) == 0 {
return map[string]string{}
@@ -57,6 +77,20 @@ func FilterSensitiveConfigs(in map[string]string) map[string]string {
return out
}
func FilterBackupConfigs(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) || IsSystemManagedConfigKey(name) {
continue
}
out[name] = value
}
return out
}
func normalizeConfigKey(name string) string {
return strings.ToLower(strings.TrimSpace(name))
}
@@ -12,7 +12,11 @@ func TestConfigPolicy(t *testing.T) {
{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: "app_bg_image_light is public", key: "app_bg_image_light", want: ConfigAccessPublic},
{name: "app_bg_image_dark is public", key: "app_bg_image_dark", want: ConfigAccessPublic},
{name: "cloudflare_site_key is public", key: "cloudflare_site_key", want: ConfigAccessPublic},
{name: "is_commercial is public", key: "is_commercial", want: ConfigAccessPublic},
{name: "hide_footer_brand is public", key: "hide_footer_brand", 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},
@@ -29,14 +33,14 @@ func TestConfigPolicy(t *testing.T) {
}
func TestConfigPolicyHelpers(t *testing.T) {
publicKeys := []string{"app_name", "app_logo", "app_favicon", "app_bg_image", "cloudflare_site_key"}
publicKeys := []string{"app_name", "app_logo", "app_favicon", "app_bg_image", "app_bg_image_light", "app_bg_image_dark", "cloudflare_site_key", "is_commercial", "hide_footer_brand"}
for _, key := range publicKeys {
if !IsPublicConfigKey(key) {
t.Fatalf("expected %q to be public", key)
}
}
sensitiveKeys := []string{"jwt_secret", "license_key", "cloudflare_secret_key"}
sensitiveKeys := []string{"jwt_secret", "license_key", "license_expiry", "license_machine_id", "machine_fingerprint", "cloudflare_secret_key"}
for _, key := range sensitiveKeys {
if !IsSensitiveConfigKey(key) {
t.Fatalf("expected %q to be sensitive", key)
@@ -67,3 +71,30 @@ func TestConfigPolicyHelpers(t *testing.T) {
t.Fatal("expected cloudflare_secret_key to be filtered out")
}
}
func TestFilterBackupConfigsOmitsLicenseState(t *testing.T) {
filtered := FilterBackupConfigs(map[string]string{
"app_name": "Brand",
"license_key": "secret-license",
"license_expiry": "never",
"license_machine_id": "machine-id",
"is_commercial": "true",
"machine_fingerprint": "fingerprint",
})
if len(filtered) != 1 || filtered["app_name"] != "Brand" {
t.Fatalf("unexpected backup configs: %+v", filtered)
}
}
func TestSystemManagedConfigKeys(t *testing.T) {
for _, key := range []string{"license_key", "license_expiry", "license_machine_id", "is_commercial", "machine_fingerprint"} {
if !IsSystemManagedConfigKey(key) {
t.Fatalf("expected %s to be system managed", key)
}
}
for _, key := range []string{"jwt_secret", "cloudflare_secret_key", "app_name"} {
if IsSystemManagedConfigKey(key) {
t.Fatalf("did not expect %s to be system managed", key)
}
}
}
+36 -9
View File
@@ -473,6 +473,14 @@ func seedData(db *gorm.DB) {
appNameConfig := model.ViteConfig{ID: 1, Name: "app_name", Value: "flux", Time: 1755147963000}
db.Where("id = ?", 1).FirstOrCreate(&appNameConfig)
now := time.Now().UnixMilli()
for name, value := range map[string]string{
"is_commercial": "false",
"hide_footer_brand": "false",
} {
cfg := model.ViteConfig{Name: name, Value: value, Time: now}
db.Where("name = ?", name).FirstOrCreate(&cfg)
}
}
// ─── User Queries ────────────────────────────────────────────────────
@@ -600,6 +608,23 @@ func (r *Repository) UpsertConfig(name, value string, now int64) error {
}).Create(&model.ViteConfig{Name: name, Value: value, Time: now}).Error
}
func (r *Repository) UpsertConfigs(values map[string]string, now int64) error {
if r == nil || r.db == nil {
return errors.New("repository not initialized")
}
return r.db.Transaction(func(tx *gorm.DB) error {
for name, value := range values {
if err := tx.Clauses(clause.OnConflict{
Columns: []clause.Column{{Name: "name"}},
DoUpdates: clause.AssignmentColumns([]string{"value", "time"}),
}).Create(&model.ViteConfig{Name: name, Value: value, Time: now}).Error; err != nil {
return err
}
}
return nil
})
}
// ─── Announcement Queries ────────────────────────────────────────────
func (r *Repository) GetAnnouncement() (*model.Announcement, error) {
@@ -644,7 +669,7 @@ func (r *Repository) GetUserPackageTunnels(userID int64) ([]model.UserTunnelDeta
}
var items []model.UserTunnelDetail
err := r.db.Model(&model.UserTunnel{}).
Select("user_tunnel.id, user_tunnel.user_id, user_tunnel.tunnel_id, tunnel.name AS tunnel_name, user_tunnel.status, tunnel.flow AS tunnel_flow, user_tunnel.flow, user_tunnel.in_flow, user_tunnel.out_flow, user_tunnel.num, user_tunnel.flow_reset_time, user_tunnel.exp_time, user_tunnel.speed_id, speed_limit.name AS speed_limit, speed_limit.speed").
Select("user_tunnel.id, user_tunnel.user_id, user_tunnel.tunnel_id, tunnel.name AS tunnel_name, user_tunnel.status, tunnel.flow AS tunnel_flow, user_tunnel.flow, user_tunnel.flow_mib, user_tunnel.in_flow, user_tunnel.out_flow, user_tunnel.num, user_tunnel.flow_reset_time, user_tunnel.exp_time, user_tunnel.speed_id, speed_limit.name AS speed_limit, speed_limit.speed").
Joins("LEFT JOIN tunnel ON tunnel.id = user_tunnel.tunnel_id").
Joins("LEFT JOIN speed_limit ON speed_limit.id = user_tunnel.speed_id").
Where("user_tunnel.user_id = ?", userID).
@@ -899,7 +924,7 @@ func (r *Repository) ListUsers() ([]map[string]interface{}, error) {
item := map[string]interface{}{
"id": u.ID, "user": u.User, "name": u.User,
"roleId": u.RoleID, "status": u.Status,
"flow": u.Flow, "num": u.Num, "expTime": u.ExpTime,
"flow": u.Flow, "flowMiB": u.FlowMiB, "num": u.Num, "expTime": u.ExpTime,
"flowResetTime": u.FlowResetTime, "createdTime": u.CreatedTime,
"updatedTime": nullableInt64(u.UpdatedTime),
"inFlow": u.InFlow, "outFlow": u.OutFlow,
@@ -1967,7 +1992,7 @@ func (r *Repository) ExportAll() (*model.BackupData, error) {
if err != nil {
return nil, fmt.Errorf("export configs failed: %w", err)
}
backup.Configs = FilterSensitiveConfigs(configs)
backup.Configs = FilterBackupConfigs(configs)
return backup, nil
}
@@ -2047,7 +2072,7 @@ func (r *Repository) ExportPartial(types []string) (*model.BackupData, error) {
if err != nil {
return nil, fmt.Errorf("export configs failed: %w", err)
}
backup.Configs = FilterSensitiveConfigs(v)
backup.Configs = FilterBackupConfigs(v)
}
return backup, nil
}
@@ -2069,7 +2094,7 @@ func (r *Repository) exportUsers() ([]model.UserBackup, error) {
for _, u := range users {
b := model.UserBackup{
ID: u.ID, User: u.User, Pwd: u.Pwd, RoleID: u.RoleID,
ExpTime: u.ExpTime, Flow: u.Flow, InFlow: u.InFlow, OutFlow: u.OutFlow,
ExpTime: u.ExpTime, Flow: u.Flow, FlowMiB: u.FlowMiB, InFlow: u.InFlow, OutFlow: u.OutFlow,
FlowResetTime: u.FlowResetTime, Num: u.Num,
CreatedTime: u.CreatedTime, Status: u.Status,
}
@@ -2245,7 +2270,7 @@ func (r *Repository) exportUserTunnels() ([]model.UserTunnelBackup, error) {
for _, ut := range uts {
b := model.UserTunnelBackup{
ID: ut.ID, UserID: ut.UserID, TunnelID: ut.TunnelID,
Num: ut.Num, Flow: ut.Flow, InFlow: ut.InFlow, OutFlow: ut.OutFlow,
Num: ut.Num, Flow: ut.Flow, FlowMiB: ut.FlowMiB, InFlow: ut.InFlow, OutFlow: ut.OutFlow,
FlowResetTime: ut.FlowResetTime, ExpTime: ut.ExpTime, Status: ut.Status,
}
if ut.SpeedID.Valid {
@@ -2442,6 +2467,7 @@ func importUsers(tx *gorm.DB, users []model.UserBackup, now int64) (int, error)
RoleID: u.RoleID,
ExpTime: u.ExpTime,
Flow: u.Flow,
FlowMiB: u.FlowMiB,
InFlow: u.InFlow,
OutFlow: u.OutFlow,
FlowResetTime: u.FlowResetTime,
@@ -2454,7 +2480,7 @@ func importUsers(tx *gorm.DB, users []model.UserBackup, now int64) (int, error)
err = tx.Clauses(clause.OnConflict{
Columns: []clause.Column{{Name: "id"}},
DoUpdates: clause.AssignmentColumns([]string{
"user", "pwd", "role_id", "exp_time", "flow", "in_flow", "out_flow",
"user", "pwd", "role_id", "exp_time", "flow", "flow_mib", "in_flow", "out_flow",
"flow_reset_time", "num", "updated_time", "status", "password_changed_at",
}),
}).Create(&item).Error
@@ -2692,6 +2718,7 @@ func importUserTunnels(tx *gorm.DB, userTunnels []model.UserTunnelBackup, _ int6
SpeedID: sql.NullInt64{Int64: ut.SpeedID, Valid: ut.SpeedID > 0},
Num: ut.Num,
Flow: ut.Flow,
FlowMiB: ut.FlowMiB,
InFlow: ut.InFlow,
OutFlow: ut.OutFlow,
FlowResetTime: ut.FlowResetTime,
@@ -2701,7 +2728,7 @@ func importUserTunnels(tx *gorm.DB, userTunnels []model.UserTunnelBackup, _ int6
err := tx.Clauses(clause.OnConflict{
Columns: []clause.Column{{Name: "id"}},
DoUpdates: clause.AssignmentColumns([]string{
"user_id", "tunnel_id", "speed_id", "num", "flow", "in_flow", "out_flow",
"user_id", "tunnel_id", "speed_id", "num", "flow", "flow_mib", "in_flow", "out_flow",
"flow_reset_time", "exp_time", "status",
}),
}).Create(&item).Error
@@ -2839,7 +2866,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)
configs = FilterBackupConfigs(configs)
count := 0
for name, value := range configs {
err := tx.Clauses(clause.OnConflict{
@@ -97,6 +97,10 @@ func TestExportAllOmitsSensitiveConfigs(t *testing.T) {
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, "license_expiry", "2030-01-02T00:00:00.000Z")
seedConfig(t, r, "license_machine_id", "machine-id")
seedConfig(t, r, "is_commercial", "true")
seedConfig(t, r, "machine_fingerprint", "machine-fingerprint")
seedConfig(t, r, "cloudflare_secret_key", "cloudflare-secret")
for _, tc := range []struct {
@@ -117,7 +121,7 @@ func TestExportAllOmitsSensitiveConfigs(t *testing.T) {
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"} {
for _, key := range []string{"jwt_secret", "license_key", "license_expiry", "license_machine_id", "is_commercial", "machine_fingerprint", "cloudflare_secret_key"} {
if _, ok := backup.Configs[key]; ok {
t.Fatalf("expected %s to be omitted from export, got %+v", key, backup.Configs)
}
@@ -136,12 +140,20 @@ func TestImportIgnoresSensitiveConfigs(t *testing.T) {
seedConfig(t, r, "app_name", "before")
seedConfig(t, r, "jwt_secret", "jwt-before")
seedConfig(t, r, "license_key", "license-before")
seedConfig(t, r, "license_expiry", "expiry-before")
seedConfig(t, r, "license_machine_id", "machine-before")
seedConfig(t, r, "is_commercial", "true")
seedConfig(t, r, "machine_fingerprint", "fingerprint-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",
"license_expiry": "expiry-after",
"license_machine_id": "machine-after",
"is_commercial": "false",
"machine_fingerprint": "fingerprint-after",
"cloudflare_secret_key": "cloudflare-after",
}}
@@ -156,6 +168,10 @@ func TestImportIgnoresSensitiveConfigs(t *testing.T) {
assertConfigValue(t, r, "app_name", "after")
assertConfigValue(t, r, "jwt_secret", "jwt-before")
assertConfigValue(t, r, "license_key", "license-before")
assertConfigValue(t, r, "license_expiry", "expiry-before")
assertConfigValue(t, r, "license_machine_id", "machine-before")
assertConfigValue(t, r, "is_commercial", "true")
assertConfigValue(t, r, "machine_fingerprint", "fingerprint-before")
assertConfigValue(t, r, "cloudflare_secret_key", "cloudflare-before")
}
@@ -37,7 +37,14 @@ func (r *Repository) UserExistsExcluding(username string, excludeID int64) (bool
return cnt > 0, err
}
func (r *Repository) CreateUser(username, pwdHash string, roleID int, expTime, flow, flowResetTime int64, num, status, maxConn int, now int64) (int64, error) {
func optionalFlowMiB(values []int64) int64 {
if len(values) > 0 {
return values[0]
}
return 0
}
func (r *Repository) CreateUser(username, pwdHash string, roleID int, expTime, flow, flowResetTime int64, num, status, maxConn int, now int64, flowMiB ...int64) (int64, error) {
if r == nil || r.db == nil {
return 0, errors.New("repository not initialized")
}
@@ -47,6 +54,7 @@ func (r *Repository) CreateUser(username, pwdHash string, roleID int, expTime, f
RoleID: roleID,
ExpTime: expTime,
Flow: flow,
FlowMiB: optionalFlowMiB(flowMiB),
InFlow: 0,
OutFlow: 0,
FlowResetTime: flowResetTime,
@@ -75,7 +83,7 @@ func (r *Repository) GetUserRoleID(userID int64) (int, error) {
return user.RoleID, nil
}
func (r *Repository) UpdateUserWithPassword(id int64, username, pwdHash string, flow int64, num int, expTime, flowResetTime int64, status, maxConn int, now int64) error {
func (r *Repository) UpdateUserWithPassword(id int64, username, pwdHash string, flow int64, num int, expTime, flowResetTime int64, status, maxConn int, now int64, flowMiB ...int64) error {
if r == nil || r.db == nil {
return errors.New("repository not initialized")
}
@@ -85,6 +93,7 @@ func (r *Repository) UpdateUserWithPassword(id int64, username, pwdHash string,
"user": username,
"pwd": pwdHash,
"flow": flow,
"flow_mib": optionalFlowMiB(flowMiB),
"num": num,
"exp_time": expTime,
"flow_reset_time": flowResetTime,
@@ -95,7 +104,7 @@ func (r *Repository) UpdateUserWithPassword(id int64, username, pwdHash string,
}).Error
}
func (r *Repository) UpdateUserWithoutPassword(id int64, username string, flow int64, num int, expTime, flowResetTime int64, status, maxConn int, now int64) error {
func (r *Repository) UpdateUserWithoutPassword(id int64, username string, flow int64, num int, expTime, flowResetTime int64, status, maxConn int, now int64, flowMiB ...int64) error {
if r == nil || r.db == nil {
return errors.New("repository not initialized")
}
@@ -104,6 +113,7 @@ func (r *Repository) UpdateUserWithoutPassword(id int64, username string, flow i
Updates(map[string]interface{}{
"user": username,
"flow": flow,
"flow_mib": optionalFlowMiB(flowMiB),
"num": num,
"exp_time": expTime,
"flow_reset_time": flowResetTime,
@@ -126,7 +136,7 @@ func (r *Repository) UpdateUserPassword(userID int64, pwdHash string, now int64)
}).Error
}
func (r *Repository) PropagateUserFlowToTunnels(userID int64, flow int64, num int, expTime, flowResetTime int64) {
func (r *Repository) PropagateUserFlowToTunnels(userID int64, flow int64, num int, expTime, flowResetTime int64, flowMiB ...int64) {
if r == nil || r.db == nil {
return
}
@@ -134,6 +144,7 @@ func (r *Repository) PropagateUserFlowToTunnels(userID int64, flow int64, num in
Where("user_id = ?", userID).
Updates(map[string]interface{}{
"flow": flow,
"flow_mib": optionalFlowMiB(flowMiB),
"num": num,
"exp_time": expTime,
"flow_reset_time": flowResetTime,
@@ -650,7 +661,7 @@ func (r *Repository) DeleteUserTunnel(id int64) error {
return r.db.Where("id = ?", id).Delete(&model.UserTunnel{}).Error
}
func (r *Repository) UpdateUserTunnel(id int64, flow int64, num int, expTime, flowResetTime int64, speedID interface{}, status int) error {
func (r *Repository) UpdateUserTunnel(id int64, flow int64, num int, expTime, flowResetTime int64, speedID interface{}, status int, flowMiB ...int64) error {
if r == nil || r.db == nil {
return errors.New("repository not initialized")
}
@@ -658,6 +669,7 @@ func (r *Repository) UpdateUserTunnel(id int64, flow int64, num int, expTime, fl
Where("id = ?", id).
Updates(map[string]interface{}{
"flow": flow,
"flow_mib": optionalFlowMiB(flowMiB),
"num": num,
"exp_time": expTime,
"flow_reset_time": flowResetTime,
@@ -692,7 +704,7 @@ func (r *Repository) GetExistingUserTunnel(userID, tunnelID int64) (id int64, fl
return ut.ID, ut.Flow, int64(ut.Num), ut.ExpTime, ut.FlowResetTime, ut.SpeedID, ut.Status, nil
}
func (r *Repository) InsertUserTunnel(userID, tunnelID int64, speedID interface{}, num int, flow, flowResetTime, expTime int64, status int) error {
func (r *Repository) InsertUserTunnel(userID, tunnelID int64, speedID interface{}, num int, flow, flowResetTime, expTime int64, status int, flowMiB ...int64) error {
if r == nil || r.db == nil {
return errors.New("repository not initialized")
}
@@ -702,6 +714,7 @@ func (r *Repository) InsertUserTunnel(userID, tunnelID int64, speedID interface{
SpeedID: nullInt64FromInterface(speedID),
Num: num,
Flow: flow,
FlowMiB: optionalFlowMiB(flowMiB),
InFlow: 0,
OutFlow: 0,
FlowResetTime: flowResetTime,
@@ -711,7 +724,7 @@ func (r *Repository) InsertUserTunnel(userID, tunnelID int64, speedID interface{
return r.db.Create(&ut).Error
}
func (r *Repository) UpdateUserTunnelFields(id int64, speedID interface{}, flow int64, num int, expTime, flowResetTime int64, status int) error {
func (r *Repository) UpdateUserTunnelFields(id int64, speedID interface{}, flow int64, num int, expTime, flowResetTime int64, status int, flowMiB ...int64) error {
if r == nil || r.db == nil {
return errors.New("repository not initialized")
}
@@ -720,6 +733,7 @@ func (r *Repository) UpdateUserTunnelFields(id int64, speedID interface{}, flow
Updates(map[string]interface{}{
"speed_id": nullInt64FromInterface(speedID),
"flow": flow,
"flow_mib": optionalFlowMiB(flowMiB),
"num": num,
"exp_time": expTime,
"flow_reset_time": flowResetTime,
@@ -1293,7 +1307,7 @@ func (r *Repository) EnsureUserTunnelGrant(userID, tunnelID int64) (int64, bool,
return 0, false, err
}
var user model.User
if err := r.db.Select("flow, num, exp_time, flow_reset_time").Where("id = ?", userID).First(&user).Error; err != nil {
if err := r.db.Select("flow, flow_mib, num, exp_time, flow_reset_time").Where("id = ?", userID).First(&user).Error; err != nil {
return 0, false, err
}
flow := user.Flow
@@ -1305,6 +1319,7 @@ func (r *Repository) EnsureUserTunnelGrant(userID, tunnelID int64) (int64, bool,
TunnelID: tunnelID,
Num: num,
Flow: flow,
FlowMiB: user.FlowMiB,
InFlow: 0,
OutFlow: 0,
FlowResetTime: flowReset,
@@ -0,0 +1,67 @@
package repo
import (
"path/filepath"
"testing"
"time"
"go-backend/internal/store/model"
)
func TestTrafficLimitMiBSurvivesBackupRestore(t *testing.T) {
source, err := Open(filepath.Join(t.TempDir(), "source.db"))
if err != nil {
t.Fatal(err)
}
defer source.Close()
now := time.Now().UnixMilli()
userID, err := source.CreateUser("mib-user", "hash", 1, now+86400000, 1, 1, 10, 1, 0, now, 500)
if err != nil {
t.Fatal(err)
}
tunnel := model.Tunnel{Name: "mib-tunnel", TrafficRatio: 1, Type: 1, Protocol: "tls", Flow: 1, CreatedTime: now, UpdatedTime: now, Status: 1, Inx: 1}
if err := source.DB().Create(&tunnel).Error; err != nil {
t.Fatal(err)
}
if _, _, err := source.EnsureUserTunnelGrant(userID, tunnel.ID); err != nil {
t.Fatal(err)
}
grants, err := source.GetUserPackageTunnels(userID)
if err != nil || len(grants) != 1 || grants[0].FlowMiB != 500 {
t.Fatalf("inherited tunnel quota = %+v, err = %v", grants, err)
}
backup, err := source.ExportAll()
if err != nil {
t.Fatal(err)
}
found := false
for _, user := range backup.Users {
if user.User == "mib-user" {
found = user.Flow == 1 && user.FlowMiB == 500
}
}
if !found {
t.Fatal("500 MiB user quota missing from backup")
}
if len(backup.UserTunnels) != 1 || backup.UserTunnels[0].FlowMiB != 500 {
t.Fatalf("tunnel quota missing from backup: %+v", backup.UserTunnels)
}
dest, err := Open(filepath.Join(t.TempDir(), "dest.db"))
if err != nil {
t.Fatal(err)
}
defer dest.Close()
if _, err := dest.Import(backup, []string{"users", "tunnels", "userTunnels"}); err != nil {
t.Fatal(err)
}
user, err := dest.GetUserByUsername("mib-user")
if err != nil || user == nil || user.Flow != 1 || user.FlowMiB != 500 {
t.Fatalf("restored quota = %+v, err = %v", user, err)
}
grants, err = dest.GetUserPackageTunnels(user.ID)
if err != nil || len(grants) != 1 || grants[0].FlowMiB != 500 {
t.Fatalf("restored tunnel quota = %+v, err = %v", grants, err)
}
}
+82 -3
View File
@@ -3,6 +3,7 @@ package config
import (
"encoding/json"
"io"
"reflect"
"sync"
"time"
@@ -30,9 +31,87 @@ func Global() *Config {
globalMux.RLock()
defer globalMux.RUnlock()
cfg := &Config{}
*cfg = *global
return cfg
return cloneConfig(global)
}
// cloneConfig returns a detached snapshot of the runtime config. Runtime
// commands mutate slices, pointers and metadata maps, so a shallow struct copy
// is not sufficient once readers and persistence run concurrently.
func cloneConfig(c *Config) *Config {
if c == nil {
return nil
}
v := cloneConfigValue(reflect.ValueOf(c))
if !v.IsValid() || v.IsNil() {
return nil
}
return v.Interface().(*Config)
}
func cloneConfigValue(v reflect.Value) reflect.Value {
if !v.IsValid() {
return reflect.Value{}
}
switch v.Kind() {
case reflect.Interface:
if v.IsNil() {
return reflect.Zero(v.Type())
}
cloned := cloneConfigValue(v.Elem())
out := reflect.New(v.Type()).Elem()
if cloned.IsValid() && cloned.Type().AssignableTo(v.Type()) {
out.Set(cloned)
} else if cloned.IsValid() && cloned.Type().Implements(v.Type()) {
out.Set(cloned)
} else if cloned.IsValid() {
out.Set(cloned)
}
return out
case reflect.Pointer:
if v.IsNil() {
return reflect.Zero(v.Type())
}
out := reflect.New(v.Type().Elem())
out.Elem().Set(cloneConfigValue(v.Elem()))
return out
case reflect.Map:
if v.IsNil() {
return reflect.Zero(v.Type())
}
out := reflect.MakeMapWithSize(v.Type(), v.Len())
iter := v.MapRange()
for iter.Next() {
out.SetMapIndex(cloneConfigValue(iter.Key()), cloneConfigValue(iter.Value()))
}
return out
case reflect.Slice:
if v.IsNil() {
return reflect.Zero(v.Type())
}
out := reflect.MakeSlice(v.Type(), v.Len(), v.Len())
for i := 0; i < v.Len(); i++ {
out.Index(i).Set(cloneConfigValue(v.Index(i)))
}
return out
case reflect.Array:
out := reflect.New(v.Type()).Elem()
for i := 0; i < v.Len(); i++ {
out.Index(i).Set(cloneConfigValue(v.Index(i)))
}
return out
case reflect.Struct:
out := reflect.New(v.Type()).Elem()
out.Set(v)
for i := 0; i < v.NumField(); i++ {
if out.Field(i).CanSet() && v.Field(i).CanInterface() {
out.Field(i).Set(cloneConfigValue(v.Field(i)))
}
}
return out
default:
return v
}
}
func Set(c *Config) {
+121
View File
@@ -0,0 +1,121 @@
package config
import (
"encoding/json"
"path/filepath"
"sync"
"testing"
)
func TestGlobalReturnsDetachedSnapshot(t *testing.T) {
original := Global()
t.Cleanup(func() { Set(original) })
Set(&Config{Services: []*ServiceConfig{{
Name: "snapshot-service",
Metadata: map[string]any{"paused": false},
Handler: &HandlerConfig{
Type: "relay",
Metadata: map[string]any{"retries": 2},
},
}}})
snapshot := Global()
snapshot.Services[0].Name = "changed"
snapshot.Services[0].Metadata["paused"] = true
snapshot.Services[0].Handler.Metadata["retries"] = 9
current := Global()
if current.Services[0].Name != "snapshot-service" {
t.Fatalf("snapshot mutated global service name: %q", current.Services[0].Name)
}
if paused, _ := current.Services[0].Metadata["paused"].(bool); paused {
t.Fatalf("snapshot mutated global service metadata")
}
if retries, _ := current.Services[0].Handler.Metadata["retries"].(int); retries != 2 {
t.Fatalf("snapshot mutated nested handler metadata: %v", current.Services[0].Handler.Metadata["retries"])
}
}
func TestConcurrentGlobalSnapshotAndUpdate(t *testing.T) {
original := Global()
t.Cleanup(func() { Set(original) })
Set(&Config{Services: []*ServiceConfig{{
Name: "concurrent-service",
Metadata: map[string]any{"generation": 0},
}}})
var wg sync.WaitGroup
for worker := 0; worker < 4; worker++ {
wg.Add(1)
go func(worker int) {
defer wg.Done()
for i := 0; i < 500; i++ {
if err := OnUpdate(func(c *Config) error {
c.Services[0] = &ServiceConfig{
Name: "concurrent-service",
Metadata: map[string]any{"generation": worker*500 + i},
}
return nil
}); err != nil {
t.Errorf("OnUpdate: %v", err)
return
}
}
}(worker)
}
for i := 0; i < 2000; i++ {
if _, err := json.Marshal(Global()); err != nil {
t.Fatalf("marshal snapshot: %v", err)
}
}
wg.Wait()
}
func TestConcurrentPersistProducesValidConfig(t *testing.T) {
original := Global()
originalPath := PersistPath()
persistMu.Lock()
originalEnabled := persistEnable
persistMu.Unlock()
t.Cleanup(func() {
Set(original)
SetPersistPath(originalPath)
persistMu.Lock()
persistEnable = originalEnabled
persistMu.Unlock()
})
path := filepath.Join(t.TempDir(), "gost.json")
SetPersistPath(path)
EnablePersist()
Set(&Config{Services: []*ServiceConfig{{Name: "persist-service"}}})
var wg sync.WaitGroup
for worker := 0; worker < 4; worker++ {
wg.Add(1)
go func(worker int) {
defer wg.Done()
for i := 0; i < 50; i++ {
if err := OnUpdate(func(c *Config) error {
c.Services[0].Metadata = map[string]any{"generation": worker*50 + i}
return nil
}); err != nil {
t.Errorf("persist update: %v", err)
return
}
}
}(worker)
}
wg.Wait()
var persisted Config
if err := persisted.ReadFile(path); err != nil {
t.Fatalf("read persisted config: %v", err)
}
if len(persisted.Services) != 1 || persisted.Services[0] == nil || persisted.Services[0].Name != "persist-service" {
t.Fatalf("unexpected persisted services: %#v", persisted.Services)
}
}
+1 -1
View File
@@ -42,9 +42,9 @@ func EnablePersist() {
// persist writes the current global config to the configured file atomically.
func persist() error {
persistMu.Lock()
defer persistMu.Unlock()
path := persistPath
enabled := persistEnable
persistMu.Unlock()
if !enabled || path == "" {
return nil
+117
View File
@@ -0,0 +1,117 @@
package socket
import (
"net"
"sync"
"testing"
"time"
coreservice "github.com/go-gost/core/service"
"github.com/go-gost/x/config"
"github.com/go-gost/x/registry"
)
type blockingCommandService struct {
started chan struct{}
release chan struct{}
startedOnce sync.Once
}
func (s *blockingCommandService) Serve() error { return nil }
func (s *blockingCommandService) Addr() net.Addr { return nil }
func (s *blockingCommandService) Close() error {
s.startedOnce.Do(func() { close(s.started) })
<-s.release
return nil
}
func TestMutationCommandsAreSerialized(t *testing.T) {
mutations := []string{
"AddService", "UpdateService", "DeleteService", "PauseService", "ResumeService",
"AddChains", "UpdateChains", "DeleteChains",
"AddLimiters", "UpdateLimiters", "DeleteLimiters",
"AddCLimiters", "UpdateCLimiters", "DeleteCLimiters",
"SetProtocol", "UpgradeAgent", "RollbackAgent", "reload",
}
for _, command := range mutations {
if !isMutationCommand(command) {
t.Fatalf("expected %s to use the serialized mutation queue", command)
}
}
}
func TestReadOnlyCommandsRemainBoundedAsync(t *testing.T) {
for _, command := range []string{"TcpPing", "UdpPing", "ServiceMonitorCheck"} {
if isMutationCommand(command) {
t.Fatalf("expected %s to remain a read-only command", command)
}
}
}
func TestCommandResponseTypePreservesRequestContract(t *testing.T) {
if got := commandResponseType("UpdateService"); got != "UpdateServiceResponse" {
t.Fatalf("unexpected response type: %s", got)
}
if got := commandResponseType(""); got != "UnknownCommandResponse" {
t.Fatalf("unexpected empty command response type: %s", got)
}
}
func TestMutationQueueExecutesCommandsInArrivalOrder(t *testing.T) {
originalConfig := config.Global()
t.Cleanup(func() { config.Set(originalConfig) })
firstName := "mutation_queue_first_tdd"
secondName := "mutation_queue_second_tdd"
first := &blockingCommandService{started: make(chan struct{}), release: make(chan struct{})}
second := &blockingCommandService{started: make(chan struct{}), release: make(chan struct{})}
for _, name := range []string{firstName, secondName} {
registry.ServiceRegistry().Unregister(name)
}
t.Cleanup(func() {
select {
case <-first.release:
default:
close(first.release)
}
select {
case <-second.release:
default:
close(second.release)
}
registry.ServiceRegistry().Unregister(firstName)
registry.ServiceRegistry().Unregister(secondName)
})
if err := registry.ServiceRegistry().Register(firstName, coreservice.Service(first)); err != nil {
t.Fatalf("register first service: %v", err)
}
if err := registry.ServiceRegistry().Register(secondName, coreservice.Service(second)); err != nil {
t.Fatalf("register second service: %v", err)
}
config.Set(&config.Config{Services: []*config.ServiceConfig{{Name: firstName}, {Name: secondName}}})
reporter := NewWebSocketReporter("", "mutation-queue-test-secret")
go reporter.runMutationCommands()
t.Cleanup(reporter.Stop)
reporter.dispatchCommand(CommandMessage{Type: "DeleteService", Data: map[string]any{"services": []string{firstName}}})
reporter.dispatchCommand(CommandMessage{Type: "DeleteService", Data: map[string]any{"services": []string{secondName}}})
select {
case <-first.started:
case <-time.After(time.Second):
t.Fatal("first mutation did not start")
}
select {
case <-second.started:
t.Fatal("second mutation started before first mutation completed")
case <-time.After(100 * time.Millisecond):
}
close(first.release)
select {
case <-second.started:
case <-time.After(time.Second):
t.Fatal("second mutation did not start after first mutation completed")
}
close(second.release)
}
+82 -18
View File
@@ -15,6 +15,11 @@ import (
xservice "github.com/go-gost/x/service"
)
type serviceReplacement struct {
config config.ServiceConfig
oldConfig *config.ServiceConfig
}
func createServices(req createServicesRequest) error {
if len(req.Data) == 0 {
@@ -95,11 +100,11 @@ func updateServices(req updateServicesRequest) error {
req.Data[i].Name = name
}
// 第二阶段:逐个更新服务(Upsert模式:存在则更新,不存在则创建)
changedServices := make([]struct {
config config.ServiceConfig
service service.Service
}, 0, len(req.Data))
// 第二阶段:逐个更新服务(Upsert模式:存在则更新,不存在则创建)。
// 配置变更命令由 WebSocket reporter 串行调度,但这里仍保留完整回滚,
// 避免新配置解析或监听失败后旧服务永久消失。
originalConfig := config.Global()
changedServices := make([]serviceReplacement, 0, len(req.Data))
for i := range req.Data {
serviceConfig := &req.Data[i]
name := serviceConfig.Name
@@ -107,33 +112,48 @@ func updateServices(req updateServicesRequest) error {
continue
}
// 1. 获取旧服务
old := registry.ServiceRegistry().Get(name)
var oldConfig *config.ServiceConfig
if originalConfig != nil {
for _, current := range originalConfig.Services {
if current != nil && strings.TrimSpace(current.Name) == name {
oldConfig = current
break
}
}
}
// 2. 关闭旧服务 (如果存在)
if old != nil {
// 3. 从注册表移除旧服务;registry 会负责关闭旧服务。
// 1. 关闭并移除旧服务(如果存在)。同名监听必须先释放端口,
// 才能创建新 listener。
if registry.ServiceRegistry().Get(name) != nil {
registry.ServiceRegistry().Unregister(name)
}
// 4. 解析新服务配置
// 2. 解析新服务配置。
svc, err := parser.ParseService(serviceConfig)
if err != nil {
rollbackErr := restoreServiceRuntime(name, oldConfig)
rollbackErr = errors.Join(rollbackErr, rollbackServiceReplacements(changedServices))
if rollbackErr != nil {
return fmt.Errorf("create service %s failed: %v; restore previous service failed: %w", name, err, rollbackErr)
}
return errors.New("create service " + name + " failed: " + err.Error())
}
changedServices = append(changedServices, struct {
config config.ServiceConfig
service service.Service
}{*serviceConfig, svc})
// 5. 注册新服务
// 3. 注册并启动新服务。
if err := registry.ServiceRegistry().Register(name, svc); err != nil {
svc.Close()
rollbackErr := restoreServiceRuntime(name, oldConfig)
rollbackErr = errors.Join(rollbackErr, rollbackServiceReplacements(changedServices))
if rollbackErr != nil {
return fmt.Errorf("service %s already exists; restore previous service failed: %w", name, rollbackErr)
}
return errors.New("service " + name + " already exists")
}
// 6. 启动新服务
go svc.Serve()
changedServices = append(changedServices, serviceReplacement{
config: *serviceConfig,
oldConfig: oldConfig,
})
}
if len(changedServices) == 0 {
return nil
@@ -158,12 +178,56 @@ func updateServices(req updateServicesRequest) error {
}
return nil
}); err != nil {
config.Set(originalConfig)
if rollbackErr := rollbackServiceReplacements(changedServices); rollbackErr != nil {
return fmt.Errorf("%w; restore previous services failed: %v", err, rollbackErr)
}
return err
}
return nil
}
func restoreServiceRuntime(name string, serviceConfig *config.ServiceConfig) error {
name = strings.TrimSpace(name)
if name == "" || serviceConfig == nil {
return nil
}
if registry.ServiceRegistry().Get(name) != nil {
registry.ServiceRegistry().Unregister(name)
}
cfgCopy := *serviceConfig
cfgCopy.Name = name
svc, err := parser.ParseService(&cfgCopy)
if err != nil {
return err
}
if err := registry.ServiceRegistry().Register(name, svc); err != nil {
svc.Close()
return err
}
go svc.Serve()
return nil
}
func rollbackServiceReplacements(replacements []serviceReplacement) error {
var rollbackErr error
for i := len(replacements) - 1; i >= 0; i-- {
name := strings.TrimSpace(replacements[i].config.Name)
if name == "" {
continue
}
if registry.ServiceRegistry().Get(name) != nil {
registry.ServiceRegistry().Unregister(name)
}
if err := restoreServiceRuntime(name, replacements[i].oldConfig); err != nil {
rollbackErr = errors.Join(rollbackErr, fmt.Errorf("restore service %s: %w", name, err))
}
}
return rollbackErr
}
func serviceConfigUnchanged(name string, next config.ServiceConfig) bool {
cfg := config.Global()
if cfg == nil {
+39
View File
@@ -7,6 +7,8 @@ import (
corelogger "github.com/go-gost/core/logger"
"github.com/go-gost/core/service"
"github.com/go-gost/x/config"
_ "github.com/go-gost/x/handler/auto"
_ "github.com/go-gost/x/listener/tcp"
xlogger "github.com/go-gost/x/logger"
"github.com/go-gost/x/registry"
)
@@ -15,6 +17,43 @@ type recordingService struct {
closed int
}
func TestUpdateServicesParseFailureRestoresPreviousRuntime(t *testing.T) {
corelogger.SetDefault(xlogger.Nop())
name := "restore_after_failed_update_tdd"
existing := &recordingService{}
registry.ServiceRegistry().Unregister(name)
t.Cleanup(func() { registry.ServiceRegistry().Unregister(name) })
if err := registry.ServiceRegistry().Register(name, service.Service(existing)); err != nil {
t.Fatalf("register existing service: %v", err)
}
originalConfig := config.Global()
t.Cleanup(func() { config.Set(originalConfig) })
serviceConfig := config.ServiceConfig{Name: name, Addr: "127.0.0.1:0"}
config.Set(&config.Config{Services: []*config.ServiceConfig{&serviceConfig}})
invalid := serviceConfig
invalid.Listener = &config.ListenerConfig{Type: "listener-does-not-exist"}
if err := updateServices(updateServicesRequest{Data: []config.ServiceConfig{invalid}}); err == nil {
t.Fatalf("expected invalid service update to fail")
}
if existing.closed != 1 {
t.Fatalf("expected old runtime to be closed once, got %d", existing.closed)
}
if registry.ServiceRegistry().Get(name) == nil {
t.Fatalf("expected previous runtime to be restored")
}
cfg := config.Global()
if len(cfg.Services) != 1 || cfg.Services[0] == nil || cfg.Services[0].Name != name {
t.Fatalf("expected previous config to remain, got %#v", cfg.Services)
}
if cfg.Services[0].Listener != nil {
t.Fatalf("expected previous listener config to remain unchanged")
}
}
func (s *recordingService) Serve() error { return nil }
func (s *recordingService) Addr() net.Addr { return nil }
func (s *recordingService) Close() error {
+90 -4
View File
@@ -16,6 +16,7 @@ import (
"os"
"os/exec"
"runtime"
"runtime/debug"
"strconv"
"strings"
"sync"
@@ -152,6 +153,8 @@ const (
maxBackoff = 2 * time.Minute // 重连最大退避
defaultMetricReportInterval = 5 * time.Second
maxConcurrentTCPPings = 8
maxConcurrentReadCommands = 16
maxQueuedMutationCommands = 256
)
type WebSocketReporter struct {
@@ -174,6 +177,8 @@ type WebSocketReporter struct {
connMutex sync.Mutex // 连接状态锁
aesCrypto *crypto.AESCrypto // AES加密器
tcpPingSem chan struct{} // 限制诊断探测并发,避免离线目标耗尽连接
readCommandSem chan struct{} // 限制只读命令并发,避免诊断请求耗尽资源
mutationQueue chan CommandMessage
}
var wsDial = func(dialer *websocket.Dialer, rawURL string) (*websocket.Conn, *http.Response, error) {
@@ -204,6 +209,8 @@ func NewWebSocketReporter(serverURL string, secret string) *WebSocketReporter {
connecting: false,
aesCrypto: aesCrypto,
tcpPingSem: make(chan struct{}, maxConcurrentTCPPings),
readCommandSem: make(chan struct{}, maxConcurrentReadCommands),
mutationQueue: make(chan CommandMessage, maxQueuedMutationCommands),
}
}
@@ -231,6 +238,7 @@ func (w *WebSocketReporter) releaseTCPPingSlot() {
// Start 启动WebSocket报告器
func (w *WebSocketReporter) Start() {
go w.runMutationCommands()
go w.run()
}
@@ -777,8 +785,7 @@ func (w *WebSocketReporter) handleReceivedMessage(messageType int, message []byt
}
if cmdMsg.Type != "call" {
// 所有命令统一异步执行,避免阻塞消息接收循环
go w.routeCommand(cmdMsg)
w.dispatchCommand(cmdMsg)
}
} else {
// 处理普通消息
@@ -789,8 +796,7 @@ func (w *WebSocketReporter) handleReceivedMessage(messageType int, message []byt
return
}
if cmdMsg.Type != "call" {
// 所有命令统一异步执行,避免阻塞消息接收循环
go w.routeCommand(cmdMsg)
w.dispatchCommand(cmdMsg)
}
}
@@ -799,6 +805,86 @@ func (w *WebSocketReporter) handleReceivedMessage(messageType int, message []byt
}
}
// dispatchCommand keeps all runtime mutations ordered while allowing bounded
// concurrency for read-only diagnostics. Mutations share process-wide
// registries and configuration, so running them concurrently can corrupt the
// persisted config or interleave service lifecycle operations.
func (w *WebSocketReporter) dispatchCommand(cmd CommandMessage) {
if isMutationCommand(cmd.Type) {
select {
case w.mutationQueue <- cmd:
case <-w.ctx.Done():
w.sendCommandFailure(cmd, "Agent is shutting down")
default:
w.sendCommandFailure(cmd, "运行时配置命令队列已满,请稍后重试")
}
return
}
select {
case w.readCommandSem <- struct{}{}:
go func() {
defer func() { <-w.readCommandSem }()
w.routeCommandSafely(cmd)
}()
case <-w.ctx.Done():
w.sendCommandFailure(cmd, "Agent is shutting down")
default:
w.sendCommandFailure(cmd, "只读命令并发过多,请稍后重试")
}
}
func (w *WebSocketReporter) runMutationCommands() {
for {
select {
case <-w.ctx.Done():
return
case cmd := <-w.mutationQueue:
w.routeCommandSafely(cmd)
}
}
}
func (w *WebSocketReporter) routeCommandSafely(cmd CommandMessage) {
defer func() {
if recovered := recover(); recovered != nil {
fmt.Printf("❌ 命令处理 panic: type=%s panic=%v\n%s", cmd.Type, recovered, debug.Stack())
w.sendCommandFailure(cmd, fmt.Sprintf("命令处理异常: %v", recovered))
}
}()
w.routeCommand(cmd)
}
func (w *WebSocketReporter) sendCommandFailure(cmd CommandMessage, message string) {
w.sendResponse(CommandResponse{
Type: commandResponseType(cmd.Type),
Success: false,
Message: message,
RequestId: cmd.RequestId,
})
}
func commandResponseType(commandType string) string {
commandType = strings.TrimSpace(commandType)
if commandType == "" {
return "UnknownCommandResponse"
}
return commandType + "Response"
}
func isMutationCommand(commandType string) bool {
switch strings.ToLower(strings.TrimSpace(commandType)) {
case "addservice", "updateservice", "deleteservice", "pauseservice", "resumeservice",
"addchains", "updatechains", "deletechains",
"addlimiters", "updatelimiters", "deletelimiters",
"addclimiters", "updateclimiters", "deleteclimiters",
"setprotocol", "upgradeagent", "rollbackagent", "reload":
return true
default:
return false
}
}
// routeCommand 路由命令到对应的处理函数
func (w *WebSocketReporter) routeCommand(cmd CommandMessage) {
jsonBytes, errs := json.Marshal(cmd)
+1
View File
@@ -405,6 +405,7 @@ After=network.target
[Service]
WorkingDirectory=$INSTALL_DIR
ExecStart=$INSTALL_DIR/flux_agent
Environment=GODEBUG=disablethp=1
Restart=on-failure
StandardOutput=null
StandardError=null
+319 -48
View File
@@ -16,6 +16,9 @@ PINNED_VERSION=""
# 镜像加速配置(可由面板传入或交互式询问)
PROXY_ENABLED="${PROXY_ENABLED:-}"
PROXY_URL="${PROXY_URL:-}"
DEFAULT_PANEL_BACKEND_CONTAINER="flux-panel-backend"
DEFAULT_PANEL_POSTGRES_CONTAINER="flux-panel-postgres"
PANEL_SCRIPT_PATH="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)/$(basename "${BASH_SOURCE[0]}")"
# 镜像加速
maybe_proxy_url() {
@@ -304,8 +307,8 @@ get_env_var() {
get_current_db_type() {
local db_type database_url
db_type=$(get_env_var "DB_TYPE")
database_url=$(get_env_var "DATABASE_URL")
db_type=$(get_env_var "DB_TYPE" || true)
database_url=$(get_env_var "DATABASE_URL" || true)
if [[ "$db_type" == "sqlite" ]]; then
echo "sqlite"
@@ -316,13 +319,289 @@ get_current_db_type() {
fi
}
get_container_compose_label() {
local container="$1"
local label="$2"
local value
value=$(docker inspect -f "{{ index .Config.Labels \"$label\" }}" "$container" 2>/dev/null || true)
if [[ "$value" == "<no value>" ]]; then
value=""
fi
printf '%s' "$value"
}
resolve_panel_deployment() {
local backend_container="${PANEL_BACKEND_CONTAINER:-$DEFAULT_PANEL_BACKEND_CONTAINER}"
local requested_dir="${PANEL_DEPLOY_DIR:-}"
local label_dir label_project deploy_dir project_name
label_dir=$(get_container_compose_label "$backend_container" "com.docker.compose.project.working_dir")
label_project=$(get_container_compose_label "$backend_container" "com.docker.compose.project")
if [[ -n "$requested_dir" ]]; then
deploy_dir="$requested_dir"
elif [[ -n "$label_dir" ]]; then
deploy_dir="$label_dir"
elif [[ -f ".env" && -f "docker-compose.yml" ]]; then
deploy_dir=$(pwd -P)
else
echo "❌ 无法识别面板部署目录:未找到容器 Compose 标签,当前目录也没有完整部署配置"
return 1
fi
if [[ ! -d "$deploy_dir" ]]; then
echo "❌ 面板部署目录不存在:$deploy_dir"
return 1
fi
deploy_dir=$(cd "$deploy_dir" && pwd -P)
if [[ -n "$label_dir" && -d "$label_dir" ]]; then
label_dir=$(cd "$label_dir" && pwd -P)
fi
if [[ -n "$requested_dir" && -n "$label_dir" && "$deploy_dir" != "$label_dir" ]]; then
echo "❌ 指定部署目录与运行中容器标签不一致:$deploy_dir != $label_dir"
return 1
fi
project_name="${COMPOSE_PROJECT_NAME:-}"
if [[ -z "$project_name" && -n "$label_project" ]]; then
project_name="$label_project"
fi
if [[ -n "$project_name" && ! "$project_name" =~ ^[a-z0-9][a-z0-9_-]*$ ]]; then
echo "❌ Compose 项目名不合法:$project_name"
return 1
fi
PANEL_DEPLOY_DIR="$deploy_dir"
PANEL_BACKEND_CONTAINER="$backend_container"
PANEL_POSTGRES_CONTAINER="${PANEL_POSTGRES_CONTAINER:-$DEFAULT_PANEL_POSTGRES_CONTAINER}"
PANEL_COMPOSE_PROJECT="$project_name"
export PANEL_DEPLOY_DIR PANEL_BACKEND_CONTAINER PANEL_POSTGRES_CONTAINER PANEL_COMPOSE_PROJECT
if [[ -n "$project_name" ]]; then
COMPOSE_PROJECT_NAME="$project_name"
export COMPOSE_PROJECT_NAME
fi
cd "$PANEL_DEPLOY_DIR"
echo "📁 部署目录:$PANEL_DEPLOY_DIR"
if [[ -n "$PANEL_COMPOSE_PROJECT" ]]; then
echo "📦 Compose 项目:$PANEL_COMPOSE_PROJECT"
fi
}
run_panel_compose() {
$DOCKER_CMD "$@"
}
validate_panel_update_environment() {
local key value backend_port frontend_port configured_db_type db_type database_url postgres_password
if [[ ! -f ".env" ]]; then
echo "❌ 部署目录缺少 .env,更新终止"
return 1
fi
if [[ ! -f "docker-compose.yml" ]]; then
echo "❌ 部署目录缺少 docker-compose.yml,更新终止"
return 1
fi
for key in JWT_SECRET BACKEND_PORT FRONTEND_PORT; do
value=$(get_env_var "$key" || true)
if [[ -z "${value//[[:space:]]/}" ]]; then
echo "❌ .env 中 $key 缺失或为空,更新终止"
return 1
fi
done
backend_port=$(get_env_var "BACKEND_PORT" || true)
frontend_port=$(get_env_var "FRONTEND_PORT" || true)
if [[ ! "$backend_port" =~ ^[0-9]+$ || "$backend_port" -lt 1 || "$backend_port" -gt 65535 ]]; then
echo "❌ BACKEND_PORT 不是有效端口:$backend_port"
return 1
fi
if [[ ! "$frontend_port" =~ ^[0-9]+$ || "$frontend_port" -lt 1 || "$frontend_port" -gt 65535 ]]; then
echo "❌ FRONTEND_PORT 不是有效端口:$frontend_port"
return 1
fi
configured_db_type=$(get_env_var "DB_TYPE" || true)
if [[ -n "$configured_db_type" && "$configured_db_type" != "sqlite" && "$configured_db_type" != "postgres" ]]; then
echo "❌ DB_TYPE 仅支持 sqlite 或 postgres:$configured_db_type"
return 1
fi
db_type=$(get_current_db_type)
if [[ "$db_type" == "postgres" ]]; then
database_url=$(get_env_var "DATABASE_URL" || true)
postgres_password=$(get_env_var "POSTGRES_PASSWORD" || true)
if [[ -z "${database_url//[[:space:]]/}" || -z "${postgres_password//[[:space:]]/}" ]]; then
echo "❌ PostgreSQL 模式要求 DATABASE_URL 和 POSTGRES_PASSWORD 均非空"
return 1
fi
fi
if ! run_panel_compose -f docker-compose.yml config -q; then
echo "❌ 当前 Compose 配置校验失败,更新终止"
return 1
fi
}
download_panel_update_compose() {
local compose_url="$1"
UPDATE_COMPOSE_CANDIDATE=$(mktemp "$PANEL_DEPLOY_DIR/.docker-compose.yml.update.XXXXXX")
if ! curl -fL -o "$UPDATE_COMPOSE_CANDIDATE" "$compose_url"; then
rm -f "$UPDATE_COMPOSE_CANDIDATE"
UPDATE_COMPOSE_CANDIDATE=""
echo "❌ 下载最新 Compose 配置失败,更新终止"
return 1
fi
if [[ ! -s "$UPDATE_COMPOSE_CANDIDATE" ]]; then
rm -f "$UPDATE_COMPOSE_CANDIDATE"
UPDATE_COMPOSE_CANDIDATE=""
echo "❌ 下载的 Compose 配置为空,更新终止"
return 1
fi
}
validate_panel_update_compose() {
if ! FLUX_VERSION="$LATEST_VERSION" run_panel_compose -f "$UPDATE_COMPOSE_CANDIDATE" config -q; then
echo "❌ 新 Compose 配置校验失败,更新终止"
return 1
fi
}
backup_sqlite_for_update() (
local destination="$1"
local running paused="false"
cleanup_paused_backend() {
if [[ "$paused" == "true" ]]; then
docker unpause "$PANEL_BACKEND_CONTAINER" >/dev/null 2>&1 || true
fi
}
trap cleanup_paused_backend EXIT
running=$(docker inspect -f '{{.State.Running}}' "$PANEL_BACKEND_CONTAINER" 2>/dev/null || true)
if [[ "$running" == "true" ]]; then
if ! docker pause "$PANEL_BACKEND_CONTAINER" >/dev/null; then
echo "❌ 暂停后端以创建一致性 SQLite 备份失败"
return 1
fi
paused="true"
fi
mkdir -p "$destination"
if ! docker cp "$PANEL_BACKEND_CONTAINER:/app/data/." "$destination"; then
echo "❌ SQLite 数据备份失败"
return 1
fi
if [[ "$paused" == "true" ]] && ! docker unpause "$PANEL_BACKEND_CONTAINER" >/dev/null; then
echo "❌ SQLite 备份完成,但恢复后端运行失败"
return 1
fi
paused="false"
)
backup_postgres_for_update() {
local destination="$1"
local postgres_db postgres_user
postgres_db=$(get_env_var "POSTGRES_DB" || true)
postgres_user=$(get_env_var "POSTGRES_USER" || true)
postgres_db=${postgres_db:-flux_panel}
postgres_user=${postgres_user:-flux_panel}
if ! docker exec "$PANEL_POSTGRES_CONTAINER" pg_dump -U "$postgres_user" "$postgres_db" > "$destination"; then
rm -f "$destination"
echo "❌ PostgreSQL 数据备份失败"
return 1
fi
}
create_panel_update_backup() {
local db_type="$1"
local timestamp
timestamp=$(date '+%Y%m%d-%H%M%S')
if ! (umask 077 && mkdir -p "$PANEL_DEPLOY_DIR/backups"); then
echo "❌ 创建更新备份目录失败"
return 1
fi
UPDATE_BACKUP_DIR=$(umask 077 && mktemp -d "$PANEL_DEPLOY_DIR/backups/panel-update-$timestamp.XXXXXX") || {
echo "❌ 创建更新备份目录失败"
return 1
}
if ! cp -p .env docker-compose.yml "$UPDATE_BACKUP_DIR/"; then
echo "❌ 备份面板配置失败"
return 1
fi
if [[ "$db_type" == "postgres" ]]; then
backup_postgres_for_update "$UPDATE_BACKUP_DIR/postgres.sql" || return 1
else
backup_sqlite_for_update "$UPDATE_BACKUP_DIR/sqlite" || return 1
fi
echo "💾 更新备份:$UPDATE_BACKUP_DIR"
}
pull_panel_update_images() {
local db_type="$1"
if [[ "$db_type" == "postgres" ]]; then
FLUX_VERSION="$LATEST_VERSION" run_panel_compose -f "$UPDATE_COMPOSE_CANDIDATE" pull backend frontend postgres
else
FLUX_VERSION="$LATEST_VERSION" run_panel_compose -f "$UPDATE_COMPOSE_CANDIDATE" pull backend frontend
fi
}
activate_panel_update_files() {
if ! chmod 0644 "$UPDATE_COMPOSE_CANDIDATE" || ! mv "$UPDATE_COMPOSE_CANDIDATE" docker-compose.yml; then
echo "❌ 替换 Compose 配置失败"
return 1
fi
UPDATE_COMPOSE_CANDIDATE=""
if ! upsert_env_var ".env" "FLUX_VERSION" "$LATEST_VERSION"; then
echo "❌ 更新版本配置失败"
return 1
fi
}
start_panel_after_update() {
local db_type="$1"
if [[ "$db_type" == "postgres" ]]; then
run_panel_compose up -d postgres || return 1
wait_for_postgres_healthy || return 1
fi
run_panel_compose up -d --force-recreate --remove-orphans backend frontend || return 1
wait_for_backend_healthy
}
rollback_panel_update() {
local db_type="$1"
echo "↩️ 正在恢复更新前配置..."
if ! cp -p "$UPDATE_BACKUP_DIR/.env" .env || ! cp -p "$UPDATE_BACKUP_DIR/docker-compose.yml" docker-compose.yml; then
echo "❌ 配置回滚失败,请从 $UPDATE_BACKUP_DIR 手动恢复"
return 1
fi
if [[ "$db_type" == "postgres" ]]; then
run_panel_compose up -d postgres || return 1
wait_for_postgres_healthy || return 1
fi
run_panel_compose up -d --force-recreate --remove-orphans backend frontend || return 1
wait_for_backend_healthy
}
wait_for_postgres_healthy() {
local pg_health
local postgres_container="${PANEL_POSTGRES_CONTAINER:-$DEFAULT_PANEL_POSTGRES_CONTAINER}"
echo "🔍 检查 PostgreSQL 服务状态..."
for i in {1..90}; do
if docker ps --format "{{.Names}}" | grep -q "^flux-panel-postgres$"; then
pg_health=$(docker inspect -f '{{.State.Health.Status}}' flux-panel-postgres 2>/dev/null || echo "unknown")
if docker ps --format "{{.Names}}" | grep -Fxq "$postgres_container"; then
pg_health=$(docker inspect -f '{{.State.Health.Status}}' "$postgres_container" 2>/dev/null || echo "unknown")
if [[ "$pg_health" == "healthy" ]]; then
echo "✅ PostgreSQL 服务健康检查通过"
return 0
@@ -335,7 +614,7 @@ wait_for_postgres_healthy() {
if [ $i -eq 90 ]; then
echo "❌ PostgreSQL 启动超时(90秒)"
echo "🔍 当前状态:$(docker inspect -f '{{.State.Health.Status}}' flux-panel-postgres 2>/dev/null || echo '容器不存在')"
echo "🔍 当前状态:$(docker inspect -f '{{.State.Health.Status}}' "$postgres_container" 2>/dev/null || echo '容器不存在')"
return 1
fi
@@ -348,11 +627,12 @@ wait_for_postgres_healthy() {
wait_for_backend_healthy() {
local backend_health
local backend_container="${PANEL_BACKEND_CONTAINER:-$DEFAULT_PANEL_BACKEND_CONTAINER}"
echo "🔍 检查后端服务状态..."
for i in {1..90}; do
if docker ps --format "{{.Names}}" | grep -q "^flux-panel-backend$"; then
backend_health=$(docker inspect -f '{{.State.Health.Status}}' flux-panel-backend 2>/dev/null || echo "unknown")
if docker ps --format "{{.Names}}" | grep -Fxq "$backend_container"; then
backend_health=$(docker inspect -f '{{.State.Health.Status}}' "$backend_container" 2>/dev/null || echo "unknown")
if [[ "$backend_health" == "healthy" ]]; then
echo "✅ 后端服务健康检查通过"
return 0
@@ -365,7 +645,7 @@ wait_for_backend_healthy() {
if [ $i -eq 90 ]; then
echo "❌ 后端服务启动超时(90秒)"
echo "🔍 当前状态:$(docker inspect -f '{{.State.Health.Status}}' flux-panel-backend 2>/dev/null || echo '容器不存在')"
echo "🔍 当前状态:$(docker inspect -f '{{.State.Health.Status}}' "$backend_container" 2>/dev/null || echo '容器不存在')"
return 1
fi
@@ -380,9 +660,8 @@ wait_for_backend_healthy() {
delete_self() {
echo ""
echo "🗑️ 操作已完成,正在清理脚本文件..."
SCRIPT_PATH="$(readlink -f "$0" 2>/dev/null || realpath "$0" 2>/dev/null || echo "$0")"
sleep 1
rm -f "$SCRIPT_PATH" && echo "✅ 脚本文件已删除" || echo "❌ 删除脚本文件失败"
rm -f "$PANEL_SCRIPT_PATH" && echo "✅ 脚本文件已删除" || echo "❌ 删除脚本文件失败"
}
@@ -486,12 +765,11 @@ EOF
# 更新功能
update_panel() {
echo "🔄 开始更新面板..."
ask_proxy_config
check_docker
resolve_panel_deployment || return 1
ask_proxy_config
validate_panel_update_environment || return 1
if [[ ! -f ".env" ]]; then
echo "⚠️ 未找到 .env,默认按 SQLite 模式更新"
fi
CURRENT_DB_TYPE=$(get_current_db_type)
echo "🗄️ 当前数据库类型:$CURRENT_DB_TYPE"
@@ -502,52 +780,45 @@ update_panel() {
}
echo "🆕 最新版本:$LATEST_VERSION"
set_compose_urls_by_version "$LATEST_VERSION"
upsert_env_var ".env" "FLUX_VERSION" "$LATEST_VERSION"
echo "🔽 下载最新配置文件..."
DOCKER_COMPOSE_URL=$(get_docker_compose_url)
echo "📡 选择配置文件:$(basename "$DOCKER_COMPOSE_URL")"
curl -L -o docker-compose.yml "$DOCKER_COMPOSE_URL"
echo "✅ 下载完成"
# 自动检测并配置 IPv6 支持
if check_ipv6_support; then
echo "🚀 系统支持 IPv6,自动启用 IPv6 配置..."
configure_docker_ipv6
download_panel_update_compose "$DOCKER_COMPOSE_URL" || return 1
if ! validate_panel_update_compose; then
rm -f "$UPDATE_COMPOSE_CANDIDATE"
UPDATE_COMPOSE_CANDIDATE=""
return 1
fi
# 先发送 SIGTERM 信号,让应用优雅关闭
docker stop -t 30 flux-panel-backend 2>/dev/null || true
docker stop -t 10 vite-frontend 2>/dev/null || true
# 等待 WAL 文件同步
echo "⏳ 等待数据同步..."
sleep 5
# 然后再完全停止
$DOCKER_CMD down
echo "💾 备份当前配置和数据库..."
if ! create_panel_update_backup "$CURRENT_DB_TYPE"; then
rm -f "$UPDATE_COMPOSE_CANDIDATE"
UPDATE_COMPOSE_CANDIDATE=""
return 1
fi
echo "⬇️ 拉取最新镜像..."
if [[ "$CURRENT_DB_TYPE" == "postgres" ]]; then
$DOCKER_CMD pull backend frontend postgres
else
$DOCKER_CMD pull backend frontend
if ! pull_panel_update_images "$CURRENT_DB_TYPE"; then
rm -f "$UPDATE_COMPOSE_CANDIDATE"
UPDATE_COMPOSE_CANDIDATE=""
echo "❌ 镜像拉取失败,现有服务保持运行"
return 1
fi
echo "🚀 启动更新后的服务..."
if [[ "$CURRENT_DB_TYPE" == "postgres" ]]; then
$DOCKER_CMD up -d postgres
wait_for_postgres_healthy
$DOCKER_CMD up -d backend frontend
else
$DOCKER_CMD up -d backend frontend
if ! activate_panel_update_files; then
rollback_panel_update "$CURRENT_DB_TYPE" || true
return 1
fi
# 等待服务启动
echo "⏳ 等待服务启动..."
if ! wait_for_backend_healthy; then
echo "🛑 更新终止"
echo "🚀 原地重建更新后的服务..."
if ! start_panel_after_update "$CURRENT_DB_TYPE"; then
echo "🛑 新版本启动失败"
if rollback_panel_update "$CURRENT_DB_TYPE"; then
echo "✅ 已恢复更新前版本"
else
echo "❌ 自动回滚失败,请从 $UPDATE_BACKUP_DIR 手动恢复"
fi
return 1
fi
+86
View File
@@ -0,0 +1,86 @@
#!/bin/sh
set -eu
# Run only in a disposable Alpine container: this exercises real OpenRC services.
# docker run --rm -v "$PWD:/workspace:ro" alpine:3.22 sh /workspace/test-install-scripts-alpine.sh
[ -f /etc/alpine-release ] && [ "$(id -u)" = 0 ] || {
echo "Run this test as root in a disposable Alpine container." >&2
exit 1
}
ROOT_DIR=$(CDPATH= cd -- "$(dirname -- "$0")" && pwd)
INSTALL_SCRIPT=${1:-"$ROOT_DIR/install.sh"}
TEST_DIR=$(mktemp -d)
export TEST_DIR
cleanup() {
rc-service flux_agent stop >/dev/null 2>&1 || true
rc-update del flux_agent default >/dev/null 2>&1 || true
rm -f /etc/init.d/flux_agent
rm -rf "$TEST_DIR"
}
trap cleanup EXIT
# Docker supplies the network; an empty interfaces file lets OpenRC register
# that dependency without changing the container's network configuration.
apk add --no-cache openrc
mkdir -p /run/openrc /etc/network
touch /run/openrc/softlevel /etc/network/interfaces
rc-service networking start
# Fail immediately if any path accidentally calls systemd on Alpine.
mkdir -p "$TEST_DIR/bin"
cat > "$TEST_DIR/bin/systemctl" <<'EOF'
#!/bin/sh
touch "$TEST_DIR/systemctl-called"
exit 1
EOF
chmod +x "$TEST_DIR/bin/systemctl"
export PATH="$TEST_DIR/bin:$PATH"
cat > "$TEST_DIR/agent" <<'EOF'
#!/bin/sh
if [ "${1:-}" = -V ]; then
echo "Alpine installer test agent"
exit 0
fi
pwd > ../working-directory
exec sleep 300
EOF
# Keep the installer bootstrap, service detection and all lifecycle operations.
# Substitute only the network download and optional tcpkill dependency.
sed '/^# 执行主函数$/,$d' "$INSTALL_SCRIPT" > "$TEST_DIR/install.sh"
cat >> "$TEST_DIR/install.sh" <<'EOF'
INSTALL_DIR="$TEST_DIR/flux_agent"
ensure_download_url_initialized() {
ensure_alpine_runtime_dependencies || return 1
DOWNLOAD_URL="file://$TEST_DIR/agent"
}
check_and_install_tcpkill() { :; }
main
EOF
chmod +x "$TEST_DIR/install.sh"
cp "$TEST_DIR/install.sh" "$TEST_DIR/manage.sh"
PROXY_ENABLED=false "$TEST_DIR/install.sh" -a http://127.0.0.1:9 -s alpine-test
rc-service flux_agent status
test -L /etc/runlevels/default/flux_agent
test "$(cat "$TEST_DIR/working-directory")" = "$TEST_DIR/flux_agent"
test -s "$TEST_DIR/flux_agent/config.json"
rc-service flux_agent restart
rc-service flux_agent status
# Update must retain config and restart through OpenRC.
cp "$TEST_DIR/flux_agent/config.json" "$TEST_DIR/config-before-update.json"
cp "$TEST_DIR/manage.sh" "$TEST_DIR/update.sh"
printf '2\n' | PROXY_ENABLED=false "$TEST_DIR/update.sh"
cmp "$TEST_DIR/config-before-update.json" "$TEST_DIR/flux_agent/config.json"
rc-service flux_agent status
printf '3\ny\n' | "$TEST_DIR/manage.sh"
test ! -e /etc/init.d/flux_agent
test ! -e /etc/runlevels/default/flux_agent
test ! -e "$TEST_DIR/flux_agent"
test ! -e "$TEST_DIR/systemctl-called"
echo "Alpine installer lifecycle tests passed"
+110 -6
View File
@@ -277,6 +277,8 @@ EOF
local expected=$'{\n "addr": "panel\\"addr",\n "secret": "sec\\\\ret\\"1"\n}'
assert_equals "$expected" "$actual" "install_flux_agent should JSON-escape config values"
grep -Fqx 'Environment=GODEBUG=disablethp=1' "$FLUX_AGENT_SYSTEMD_SERVICE_FILE" || \
fail "systemd service should disable transparent huge pages for the Go heap"
)
test_install_script_bootstraps_bash_for_alpine() (
@@ -556,6 +558,7 @@ test_update_panel_asks_for_proxy_config() (
load_script_without_main "$ROOT_DIR/panel_install.sh"
local ask_called="0"
local calls=""
ask_proxy_config() {
ask_called="1"
@@ -568,6 +571,15 @@ test_update_panel_asks_for_proxy_config() (
DOCKER_CMD="true"
}
resolve_panel_deployment() {
calls+=" resolve"
PANEL_DEPLOY_DIR=$(mktemp -d)
}
validate_panel_update_environment() {
calls+=" validate-current"
}
get_current_db_type() {
echo "sqlite"
}
@@ -576,13 +588,33 @@ test_update_panel_asks_for_proxy_config() (
echo "v-test"
}
upsert_env_var() { :; }
download_panel_update_compose() {
calls+=" download"
UPDATE_COMPOSE_CANDIDATE="candidate"
}
validate_panel_update_compose() {
calls+=" validate-new"
}
create_panel_update_backup() {
calls+=" backup"
UPDATE_BACKUP_DIR="backup"
}
pull_panel_update_images() {
calls+=" pull"
}
activate_panel_update_files() {
calls+=" activate"
}
start_panel_after_update() {
calls+=" start"
}
check_ipv6_support() { return 1; }
configure_docker_ipv6() { :; }
docker() { return 0; }
wait_for_backend_healthy() { return 0; }
sleep() { :; }
curl() { :; }
update_panel >/dev/null
@@ -591,6 +623,75 @@ test_update_panel_asks_for_proxy_config() (
"https://github.com/${REPO}/releases/download/v-test/docker-compose-v4.yml" \
"$DOCKER_COMPOSEV4_URL" \
"update_panel should honor the prompted proxy choice"
assert_equals \
" resolve validate-current download validate-new backup pull activate start" \
"$calls" \
"update_panel should validate, back up, and pull before activating the update"
)
test_resolve_panel_deployment_uses_container_labels() (
set -euo pipefail
load_script_without_main "$ROOT_DIR/panel_install.sh"
local expected_dir
expected_dir=$(mktemp -d)
expected_dir=$(cd "$expected_dir" && pwd -P)
printf 'JWT_SECRET=test\nBACKEND_PORT=6365\nFRONTEND_PORT=6366\nDB_TYPE=sqlite\n' > "$expected_dir/.env"
printf 'services: {}\n' > "$expected_dir/docker-compose.yml"
unset PANEL_DEPLOY_DIR COMPOSE_PROJECT_NAME
docker() {
if [[ "$1" == "inspect" && "$2" == "-f" ]]; then
case "$3" in
*working_dir*) printf '%s\n' "$expected_dir" ;;
*) printf '%s\n' "flvx-panel" ;;
esac
fi
}
resolve_panel_deployment >/dev/null
assert_equals "$expected_dir" "$PANEL_DEPLOY_DIR" "update should discover the Compose working directory from container labels"
assert_equals "flvx-panel" "$COMPOSE_PROJECT_NAME" "update should preserve the existing Compose project name"
assert_equals "$expected_dir" "$(pwd -P)" "update should run from the discovered deployment directory"
)
test_validate_panel_update_environment_rejects_empty_required_values() (
set -euo pipefail
load_script_without_main "$ROOT_DIR/panel_install.sh"
local deploy_dir rc="0"
deploy_dir=$(mktemp -d)
printf 'JWT_SECRET=\nBACKEND_PORT=6365\nFRONTEND_PORT=6366\nDB_TYPE=sqlite\n' > "$deploy_dir/.env"
printf 'services: {}\n' > "$deploy_dir/docker-compose.yml"
cd "$deploy_dir"
DOCKER_CMD="true"
validate_panel_update_environment >/dev/null || rc="$?"
assert_equals "1" "$rc" "update should reject an empty JWT secret before touching the running service"
)
test_backup_sqlite_for_update_pauses_copies_and_unpauses() (
set -euo pipefail
load_script_without_main "$ROOT_DIR/panel_install.sh"
local destination calls log_path
destination=$(mktemp -d)/sqlite
log_path=$(mktemp)
PANEL_BACKEND_CONTAINER="flux-panel-backend"
docker() {
if [[ "$1" == "inspect" ]]; then
printf 'true\n'
return 0
fi
printf ' %s' "$1" >> "$log_path"
}
backup_sqlite_for_update "$destination" >/dev/null
calls=$(cat "$log_path")
assert_equals " pause cp unpause" "$calls" "SQLite backup should pause writes, copy the database volume, then resume the backend"
)
test_panel_install_script_accepts_proxy_url_env_without_prompt() (
@@ -661,6 +762,9 @@ test_install_script_accepts_proxy_url_env_without_prompt
test_panel_install_script_can_disable_proxy
test_panel_install_script_recomputes_compose_urls_after_prompt
test_update_panel_asks_for_proxy_config
test_resolve_panel_deployment_uses_container_labels
test_validate_panel_update_environment_rejects_empty_required_values
test_backup_sqlite_for_update_pauses_copies_and_unpauses
test_panel_install_script_uses_default_proxy
test_panel_install_script_accepts_proxy_url_env_without_prompt
test_panel_install_script_defaults_proxy_on_eof
+4 -4
View File
@@ -1,10 +1,10 @@
# 多阶段构建 - 构建阶段
FROM node:22-alpine AS builder
FROM --platform=$BUILDPLATFORM node:22-alpine AS builder
WORKDIR /app
COPY package.json pnpm-lock.yaml* ./
RUN corepack prepare pnpm@10 --activate && corepack enable pnpm && pnpm install --frozen-lockfile
COPY package.json pnpm-lock.yaml pnpm-workspace.yaml ./
RUN corepack enable pnpm && pnpm install --frozen-lockfile
COPY . .
RUN pnpm run build
@@ -18,4 +18,4 @@ COPY --from=builder /app/dist /usr/share/nginx/html
EXPOSE 80
CMD ["nginx", "-g", "daemon off;"]
CMD ["nginx", "-g", "daemon off;"]
+1 -7
View File
@@ -2,6 +2,7 @@
"name": "flvx",
"private": true,
"version": "0.0.0",
"packageManager": "pnpm@10.28.1",
"type": "module",
"scripts": {
"dev": "vite",
@@ -81,12 +82,5 @@
"vite": "npm:rolldown-vite@^7.3.1",
"vite-plugin-pwa": "^1.1.0",
"vite-tsconfig-paths": "^6.0.5"
},
"pnpm": {
"overrides": {
"@babel/plugin-transform-modules-systemjs": "7.29.4",
"fast-uri": "3.1.2",
"serialize-javascript": "7.0.5"
}
}
}
+8
View File
@@ -1,2 +1,10 @@
packages:
- '.'
overrides:
'@babel/plugin-transform-modules-systemjs': 7.29.4
fast-uri: 3.1.2
serialize-javascript: 7.0.5
allowBuilds:
'@tailwindcss/oxide': true
+7 -2
View File
@@ -28,6 +28,7 @@ import { isLoggedIn } from "@/utils/auth";
import { siteConfig, updateSiteConfig } from "@/config/site";
import { useH5Mode } from "@/hooks/useH5Mode";
import { SESSION_UPDATED_EVENT } from "@/utils/session";
import { useThemeContext } from "@/themes/context";
const ProtectedRoute = ({
children,
@@ -79,6 +80,7 @@ const LoginRoute = () => {
function App() {
const location = useLocation();
const navigate = useNavigate();
const { effectiveMode } = useThemeContext();
// 全局登录状态监听,当检测到未登录且不在首页时,跳转到首页
useEffect(() => {
@@ -98,7 +100,10 @@ function App() {
// 处理自定义背景图片
useEffect(() => {
const updateBg = () => {
const customBg = siteConfig.app_bg_image;
const customBg =
(effectiveMode === "dark"
? siteConfig.app_bg_image_dark
: siteConfig.app_bg_image_light) || siteConfig.app_bg_image;
if (customBg) {
if (customBg === "theme") {
@@ -149,7 +154,7 @@ function App() {
return () => {
window.removeEventListener("site-config-updated", updateBg);
};
}, []);
}, [effectiveMode]);
// 立即设置页面标题(使用已从缓存读取的配置)
useEffect(() => {
+3
View File
@@ -19,6 +19,7 @@ export interface UserApiItem {
name?: string;
status: number;
flow: number;
flowMiB?: number;
num: number;
expTime?: number;
flowResetTime?: number;
@@ -106,6 +107,7 @@ export interface UserTunnelPermissionApiItem {
tunnelName: string;
status: number;
flow: number;
flowMiB?: number;
num: number;
expTime: number;
flowResetTime: number;
@@ -387,6 +389,7 @@ export interface UserTunnelAssignPayload {
id?: number;
tunnelId?: number;
flow?: number;
flowMiB?: number;
num?: number;
expTime?: number;
flowResetTime?: number;
@@ -0,0 +1,62 @@
import { Input } from "@/shadcn-bridge/heroui/input";
import { Select, SelectItem } from "@/shadcn-bridge/heroui/select";
import {
parseTrafficInput,
TRAFFIC_UNIT_MIB,
type TrafficUnit,
} from "@/utils/traffic";
const UNITS: TrafficUnit[] = ["MB", "GB", "TB", "PB"];
interface TrafficLimitFieldProps {
label: string;
value: string;
unit: TrafficUnit;
onChange: (value: string, unit: TrafficUnit) => void;
description?: string;
isRequired?: boolean;
}
export function TrafficLimitField({
label,
value,
unit,
onChange,
description,
isRequired,
}: TrafficLimitFieldProps) {
return (
<div className="grid grid-cols-[minmax(0,1fr)_7rem] gap-2">
<Input
description={description}
isRequired={isRequired}
label={label}
min="0"
step="any"
type="number"
value={value}
onChange={(event) => onChange(event.target.value, unit)}
/>
<Select
aria-label={`${label}单位`}
label="单位"
selectedKeys={[unit]}
onSelectionChange={(keys) => {
const nextUnit = Array.from(keys)[0] as TrafficUnit | undefined;
if (!nextUnit) return;
const mib = parseTrafficInput(value, unit);
onChange(
mib === null ? value : String(mib / TRAFFIC_UNIT_MIB[nextUnit]),
nextUnit,
);
}}
>
{UNITS.map((option) => (
<SelectItem key={option}>{option}</SelectItem>
))}
</Select>
</div>
);
}
+52 -15
View File
@@ -14,10 +14,16 @@ const PUBLIC_BRAND_CONFIG_KEYS = [
"app_logo",
"app_favicon",
"app_bg_image",
"app_bg_image_light",
"app_bg_image_dark",
"is_commercial",
"hide_footer_brand",
] as const;
const SENSITIVE_CONFIG_KEYS = new Set([
"jwt_secret",
"license_key",
"license_machine_id",
"machine_fingerprint",
"cloudflare_secret_key",
]);
const GITHUB_REPO =
@@ -110,6 +116,8 @@ const getInitialConfig = () => {
app_logo: "",
app_favicon: "",
app_bg_image: "",
app_bg_image_light: "",
app_bg_image_dark: "",
is_commercial: false,
hide_footer_brand: false,
};
@@ -123,6 +131,10 @@ const getInitialConfig = () => {
localStorage.getItem(CACHE_PREFIX + "app_favicon") || "";
const cachedAppBgImage =
localStorage.getItem(CACHE_PREFIX + "app_bg_image") || "";
const cachedAppBgImageLight =
localStorage.getItem(CACHE_PREFIX + "app_bg_image_light") || "";
const cachedAppBgImageDark =
localStorage.getItem(CACHE_PREFIX + "app_bg_image_dark") || "";
const isCommercial =
localStorage.getItem(CACHE_PREFIX + "is_commercial") === "true";
const hideFooterBrand =
@@ -130,15 +142,17 @@ const getInitialConfig = () => {
if (cachedAppName) {
return {
name: cachedAppName,
name: isCommercial ? cachedAppName : "FLVX",
version: VERSION,
app_version: APP_VERSION,
github_repo: GITHUB_REPO,
app_logo: cachedAppLogo,
app_favicon: cachedAppFavicon,
app_logo: isCommercial ? cachedAppLogo : "",
app_favicon: isCommercial ? cachedAppFavicon : "",
app_bg_image: cachedAppBgImage,
app_bg_image_light: cachedAppBgImageLight,
app_bg_image_dark: cachedAppBgImageDark,
is_commercial: isCommercial,
hide_footer_brand: hideFooterBrand,
hide_footer_brand: isCommercial && hideFooterBrand,
};
}
@@ -147,11 +161,13 @@ const getInitialConfig = () => {
version: VERSION,
app_version: APP_VERSION,
github_repo: GITHUB_REPO,
app_logo: cachedAppLogo,
app_favicon: cachedAppFavicon,
app_logo: isCommercial ? cachedAppLogo : "",
app_favicon: isCommercial ? cachedAppFavicon : "",
app_bg_image: cachedAppBgImage,
app_bg_image_light: cachedAppBgImageLight,
app_bg_image_dark: cachedAppBgImageDark,
is_commercial: isCommercial,
hide_footer_brand: hideFooterBrand,
hide_footer_brand: isCommercial && hideFooterBrand,
};
};
@@ -378,7 +394,21 @@ export const updateSiteConfig = async (configMap?: Record<string, string>) => {
resolvedConfigMap,
"app_bg_image",
);
const hasAppBgImageLight = Object.prototype.hasOwnProperty.call(
resolvedConfigMap,
"app_bg_image_light",
);
const hasAppBgImageDark = Object.prototype.hasOwnProperty.call(
resolvedConfigMap,
"app_bg_image_dark",
);
const resolvedCommercial = Object.prototype.hasOwnProperty.call(
resolvedConfigMap,
"is_commercial",
)
? resolvedConfigMap.is_commercial === "true"
: siteConfig.is_commercial;
const appName = hasAppName
? String(resolvedConfigMap.app_name || "").trim()
: siteConfig.name;
@@ -391,24 +421,31 @@ export const updateSiteConfig = async (configMap?: Record<string, string>) => {
const appBgImage = hasAppBgImage
? String(resolvedConfigMap.app_bg_image || "").trim()
: (siteConfig.app_bg_image || "").trim();
const appBgImageLight = hasAppBgImageLight
? String(resolvedConfigMap.app_bg_image_light || "").trim()
: (siteConfig.app_bg_image_light || "").trim();
const appBgImageDark = hasAppBgImageDark
? String(resolvedConfigMap.app_bg_image_dark || "").trim()
: (siteConfig.app_bg_image_dark || "").trim();
if (appName && appName !== siteConfig.name) {
siteConfig.name = appName;
}
siteConfig.app_logo = appLogo;
siteConfig.app_favicon = appFavicon;
siteConfig.name = resolvedCommercial && appName ? appName : "FLVX";
siteConfig.app_logo = resolvedCommercial ? appLogo : "";
siteConfig.app_favicon = resolvedCommercial ? appFavicon : "";
siteConfig.app_bg_image = appBgImage;
siteConfig.app_bg_image_light = appBgImageLight;
siteConfig.app_bg_image_dark = appBgImageDark;
if (
Object.prototype.hasOwnProperty.call(resolvedConfigMap, "is_commercial")
) {
siteConfig.is_commercial = resolvedConfigMap.is_commercial === "true";
siteConfig.is_commercial = resolvedCommercial;
}
if (
Object.prototype.hasOwnProperty.call(resolvedConfigMap, "hide_footer_brand")
) {
siteConfig.hide_footer_brand =
resolvedConfigMap.hide_footer_brand === "true";
resolvedCommercial && resolvedConfigMap.hide_footer_brand === "true";
} else if (!resolvedCommercial) {
siteConfig.hide_footer_brand = false;
}
if (typeof document !== "undefined") {
+80 -37
View File
@@ -90,6 +90,8 @@ interface ConfigItem {
dependsValue?: string; // 依赖的配置项值
}
type BgImageKey = "app_bg_image" | "app_bg_image_light" | "app_bg_image_dark";
const BRAND_PREVIEW_KEYS = ["app_logo", "app_favicon"] as const;
type BrandPreviewKey = (typeof BRAND_PREVIEW_KEYS)[number];
@@ -107,9 +109,9 @@ const toBrandAssetKind = (key: BrandPreviewKey): BrandAssetKind => {
const CONFIG_ITEMS: ConfigItem[] = [
{
key: "app_bg_image",
label: "自定义背景",
label: "背景壁纸",
description:
"上传自定义背景图片(建议使用深色/浅色均可看清的图片,或使用半透明模糊效果)",
"默认背景用于未单独设置壁纸的模式。可分别上传亮色和暗色壁纸,保存后随外观模式自动切换。",
type: "bg_image",
},
{
@@ -263,6 +265,9 @@ const getInitialConfigs = (): Record<string, string> => {
"panel_domain",
"app_logo",
"app_favicon",
"app_bg_image",
"app_bg_image_light",
"app_bg_image_dark",
"github_proxy_enabled",
"github_proxy_url",
"allow_local_remote_addr",
@@ -312,8 +317,12 @@ export default function ConfigPage() {
const logoFileInputRef = useRef<HTMLInputElement>(null);
const faviconFileInputRef = useRef<HTMLInputElement>(null);
const bgImageFileInputRef = useRef<HTMLInputElement>(null);
const [bgImageUploading, setBgImageUploading] = useState(false);
const bgImageFileInputRefs = useRef<
Partial<Record<BgImageKey, HTMLInputElement>>
>({});
const [bgImageUploading, setBgImageUploading] = useState<BgImageKey | null>(
null,
);
const [announcement, setAnnouncement] = useState<AnnouncementData>({
content: "",
@@ -689,7 +698,14 @@ export default function ConfigPage() {
if (
changedKeys.some((key) =>
["app_name", "app_logo", "app_favicon"].includes(key),
[
"app_name",
"app_logo",
"app_favicon",
"app_bg_image",
"app_bg_image_light",
"app_bg_image_dark",
].includes(key),
)
) {
await updateSiteConfig(configs);
@@ -800,6 +816,7 @@ export default function ConfigPage() {
const handleBgImageUpload = async (
e: React.ChangeEvent<HTMLInputElement>,
key: BgImageKey,
) => {
const file = e.target.files?.[0];
@@ -811,7 +828,7 @@ export default function ConfigPage() {
return;
}
setBgImageUploading(true);
setBgImageUploading(key);
try {
const compressedImage = await new Promise<string>((resolve, reject) => {
const reader = new FileReader();
@@ -861,18 +878,19 @@ export default function ConfigPage() {
reader.readAsDataURL(file);
});
handleConfigChange("app_bg_image", compressedImage);
toast.success("自定义背景上传成功");
handleConfigChange(key, compressedImage);
toast.success("壁纸上传成功,保存配置后生效");
} catch (error) {
toast.error(error instanceof Error ? error.message : "图片处理失败");
} finally {
setBgImageUploading(false);
setBgImageUploading(null);
e.target.value = "";
}
};
const renderBgImageUploader = () => {
const bgImage = configs["app_bg_image"] || "";
const renderBgImageUploader = (key: BgImageKey, label: string) => {
const bgImage = configs[key] || "";
const isDefault = key === "app_bg_image";
const isImage =
bgImage.startsWith("http") ||
bgImage.startsWith("data:") ||
@@ -882,47 +900,64 @@ export default function ConfigPage() {
const isSolidColor = bgImage && !isImage && !isTheme;
return (
<div className="flex flex-col gap-4 w-full">
<div className="flex flex-col gap-3 w-full rounded-xl border border-divider p-4">
<div>
<p className="text-sm font-medium text-gray-700 dark:text-gray-300">
{label}
</p>
{!isDefault && !bgImage && (
<p className="text-xs text-gray-500 dark:text-gray-400 mt-1">
未单独设置,使用默认背景
</p>
)}
</div>
<div className="flex flex-wrap items-center gap-4">
<input
ref={bgImageFileInputRef}
ref={(node) => {
if (node) bgImageFileInputRefs.current[key] = node;
}}
accept="image/*"
className="hidden"
type="file"
onChange={handleBgImageUpload}
onChange={(event) => void handleBgImageUpload(event, key)}
/>
<Button
color="primary"
isLoading={bgImageUploading}
isDisabled={bgImageUploading !== null && bgImageUploading !== key}
isLoading={bgImageUploading === key}
variant="flat"
onPress={() => bgImageFileInputRef.current?.click()}
onPress={() => bgImageFileInputRefs.current[key]?.click()}
>
上传图片
</Button>
<Button
color="secondary"
isDisabled={bgImageUploading || isTheme}
variant="flat"
onPress={() => handleConfigChange("app_bg_image", "theme")}
>
自适应纯色 (跟随深色模式)
</Button>
<Button
color="default"
isDisabled={bgImageUploading || bgImage === "#ffffff"}
variant="flat"
onPress={() => handleConfigChange("app_bg_image", "#ffffff")}
>
白色纯色
</Button>
{isDefault && (
<Button
color="secondary"
isDisabled={bgImageUploading !== null || isTheme}
variant="flat"
onPress={() => handleConfigChange(key, "theme")}
>
自适应纯色 (跟随深色模式)
</Button>
)}
{isDefault && (
<Button
color="default"
isDisabled={bgImageUploading !== null || bgImage === "#ffffff"}
variant="flat"
onPress={() => handleConfigChange(key, "#ffffff")}
>
白色纯色
</Button>
)}
{bgImage && (
<Button
color="danger"
isDisabled={bgImageUploading}
isDisabled={bgImageUploading !== null}
variant="flat"
onPress={() => handleConfigChange("app_bg_image", "")}
onPress={() => handleConfigChange(key, "")}
>
恢复默认
{isDefault ? "恢复内置背景" : "清除专属壁纸"}
</Button>
)}
</div>
@@ -930,7 +965,7 @@ export default function ConfigPage() {
{bgImage && isImage && (
<div className="relative rounded-xl overflow-hidden border border-divider">
<img
alt="背景预览"
alt={`${label}预览`}
className="w-full max-h-48 object-cover opacity-80"
src={bgImage}
/>
@@ -1118,7 +1153,15 @@ export default function ConfigPage() {
switch (item.type) {
case "bg_image":
return renderBgImageUploader();
return (
<div className="grid grid-cols-1 lg:grid-cols-2 gap-4">
<div className="lg:col-span-2">
{renderBgImageUploader("app_bg_image", "默认背景")}
</div>
{renderBgImageUploader("app_bg_image_light", "☀️ 亮色模式壁纸")}
{renderBgImageUploader("app_bg_image_dark", "🌙 暗色模式壁纸")}
</div>
);
case "input": {
if (isBrandPreviewKey(item.key)) {
+15 -27
View File
@@ -24,6 +24,11 @@ import { FlowChartCard } from "@/pages/dashboard/components/flow-chart-card";
import { MetricCard } from "@/pages/dashboard/components/metric-card";
import { getSessionName } from "@/utils/session";
import { safeLogout } from "@/utils/logout";
import {
formatTraffic,
formatFlowLimit,
flowLimitBytes,
} from "@/utils/traffic";
import {
formatNodeRenewalTime,
getNodeRenewalCycleLabel,
@@ -70,24 +75,7 @@ export default function DashboardPage() {
const [addressModalTitle, setAddressModalTitle] = useState("");
const [addressList, setAddressList] = useState<AddressItem[]>([]);
const formatFlow = (value: number, unit: string = "bytes"): string => {
// 99999 表示无限制
if (value === 99999) {
return "无限制";
}
if (unit === "gb") {
return value + " GB";
} else {
if (value === 0) return "0 B";
if (value < 1024) return value + " B";
if (value < 1024 * 1024) return (value / 1024).toFixed(2) + " KB";
if (value < 1024 * 1024 * 1024)
return (value / (1024 * 1024)).toFixed(2) + " MB";
return (value / (1024 * 1024 * 1024)).toFixed(2) + " GB";
}
};
const formatFlow = formatTraffic;
const formatNumber = (value: number): string => {
// 99999 表示无限制
@@ -279,10 +267,10 @@ export default function DashboardPage() {
const calculateUsagePercentage = (type: "flow" | "forwards"): number => {
if (type === "flow") {
const totalUsed = calculateUserTotalUsedFlow();
const totalLimit = (userInfo.flow || 0) * 1024 * 1024 * 1024;
const totalLimit = flowLimitBytes(userInfo.flow || 0, userInfo.flowMiB);
// 无限制时返回0%
if (userInfo.flow === 99999) return 0;
if (userInfo.flow === 99999 && !userInfo.flowMiB) return 0;
return totalLimit > 0 ? Math.min((totalUsed / totalLimit) * 100, 100) : 0;
} else if (type === "forwards") {
@@ -351,10 +339,10 @@ export default function DashboardPage() {
const calculateTunnelFlowPercentage = (tunnel: UserTunnel): number => {
const totalUsed = calculateTunnelUsedFlow(tunnel);
const totalLimit = (tunnel.flow || 0) * 1024 * 1024 * 1024;
const totalLimit = flowLimitBytes(tunnel.flow || 0, tunnel.flowMiB);
// 无限制时返回0%
if (tunnel.flow === 99999) return 0;
if (tunnel.flow === 99999 && !tunnel.flowMiB) return 0;
return totalLimit > 0 ? Math.min((totalUsed / totalLimit) * 100, 100) : 0;
};
@@ -741,7 +729,7 @@ export default function DashboardPage() {
}
iconClassName="bg-blue-100 dark:bg-blue-500/20"
title="总流量"
value={formatFlow(userInfo.flow, "gb")}
value={formatFlowLimit(userInfo.flow, userInfo.flowMiB)}
/>
<MetricCard
@@ -750,11 +738,11 @@ export default function DashboardPage() {
{renderProgressBar(
calculateUsagePercentage("flow"),
"sm",
userInfo.flow === 99999,
userInfo.flow === 99999 && !userInfo.flowMiB,
)}
<div className="flex items-center justify-between mt-1">
<p className="text-xs text-default-500 truncate">
{userInfo.flow === 99999
{userInfo.flow === 99999 && !userInfo.flowMiB
? "无限制"
: `${calculateUsagePercentage("flow").toFixed(1)}%`}
</p>
@@ -981,7 +969,7 @@ export default function DashboardPage() {
流量配额
</p>
<p className="font-semibold text-foreground">
{formatFlow(tunnel.flow, "gb")}
{formatFlowLimit(tunnel.flow, tunnel.flowMiB)}
</p>
</div>
<div>
@@ -995,7 +983,7 @@ export default function DashboardPage() {
{renderProgressBar(
calculateTunnelFlowPercentage(tunnel),
"sm",
tunnel.flow === 99999,
tunnel.flow === 99999 && !tunnel.flowMiB,
)}
</div>
</div>
@@ -14,6 +14,7 @@ import { getAdminFlag } from "@/utils/session";
export interface DashboardUserInfo {
flow: number;
flowMiB?: number;
inFlow: number;
outFlow: number;
num: number;
@@ -26,6 +27,7 @@ export interface DashboardUserTunnel {
tunnelId: number;
tunnelName: string;
flow: number;
flowMiB?: number;
inFlow: number;
outFlow: number;
num: number;
+2 -7
View File
@@ -27,6 +27,7 @@ import { useSortable } from "@dnd-kit/sortable";
import { CSS } from "@dnd-kit/utilities";
import { AnimatedPage } from "@/components/animated-page";
import { formatTraffic } from "@/utils/traffic";
import { BatchActionResultModal } from "@/components/batch-action-result-modal";
import { Card, CardBody, CardHeader } from "@/shadcn-bridge/heroui/card";
import { Button } from "@/shadcn-bridge/heroui/button";
@@ -2704,13 +2705,7 @@ export default function ForwardPage() {
// 格式化流量
const formatFlow = (value: number): string => {
if (value === 0) return "0 B";
if (value < 1024) return value + " B";
if (value < 1024 * 1024) return (value / 1024).toFixed(2) + " KB";
if (value < 1024 * 1024 * 1024)
return (value / (1024 * 1024)).toFixed(2) + " MB";
return (value / (1024 * 1024 * 1024)).toFixed(2) + " GB";
return formatTraffic(value);
};
// 显示地址列表弹窗
+2 -9
View File
@@ -20,6 +20,7 @@ import { CSS } from "@dnd-kit/utilities";
import { LayoutGrid, List } from "lucide-react";
import { SearchBar } from "@/components/search-bar";
import { formatTraffic } from "@/utils/traffic";
import { AnimatedPage } from "@/components/animated-page";
import {
Table,
@@ -734,15 +735,7 @@ export default function NodePage() {
// 格式化流量
const formatFlow = (bytes: number): string => {
if (!Number.isFinite(bytes) || bytes <= 0) {
return "0 B";
}
if (bytes < 1024) return `${bytes} B`;
if (bytes < 1024 * 1024) return `${(bytes / 1024).toFixed(2)} KB`;
if (bytes < 1024 * 1024 * 1024)
return `${(bytes / (1024 * 1024)).toFixed(2)} MB`;
return `${(bytes / (1024 * 1024 * 1024)).toFixed(2)} GB`;
return formatTraffic(bytes);
};
const formatChainType = (chainType: number, hopInx: number) => {
+7 -12
View File
@@ -33,6 +33,7 @@ import {
} from "lucide-react";
import toast from "react-hot-toast";
import { formatTraffic } from "@/utils/traffic";
import {
DistroIcon,
parseDistroFromVersion,
@@ -137,15 +138,7 @@ const formatDateTime = (ts: number): string => {
});
};
const formatBytes = (bytes: number): string => {
if (!Number.isFinite(bytes) || bytes <= 0) return "0 B";
const k = 1024;
const sizes = ["B", "KB", "MB", "GB", "TB"];
const i = Math.floor(Math.log(bytes) / Math.log(k));
return `${parseFloat((bytes / Math.pow(k, i)).toFixed(2))} ${sizes[i]}`;
};
const formatBytes = formatTraffic;
const formatBytesPerSecond = (bytesPerSecond: number): string => {
if (!Number.isFinite(bytesPerSecond) || bytesPerSecond <= 0) return "0 B/s";
@@ -1938,13 +1931,15 @@ export function MonitorView({ nodeMap, viewMode = "grid" }: MonitorViewProps) {
<Modal
isOpen={resultsModalOpen}
scrollBehavior="inside"
size="xl"
onClose={() => {
setResultsModalOpen(false);
setResultsMonitorId(null);
}}
>
<ModalContent>
<ModalHeader className="flex flex-row items-center justify-between gap-3">
<ModalHeader className="flex flex-col gap-3 sm:flex-row sm:items-center sm:justify-between">
<div className="min-w-0">
<div className="text-base font-semibold truncate">
监控记录
@@ -1972,7 +1967,7 @@ export function MonitorView({ nodeMap, viewMode = "grid" }: MonitorViewProps) {
) : null}
</div>
{resultsMonitorId != null ? (
<div className="flex items-center gap-2 shrink-0">
<div className="flex shrink-0 items-center gap-2">
<Select
className="w-28"
selectedKeys={[String(resultsLimit)]}
@@ -2018,7 +2013,7 @@ export function MonitorView({ nodeMap, viewMode = "grid" }: MonitorViewProps) {
) : modalResults.length > 0 ? (
<Table
aria-label="监控记录"
className="w-full overflow-x-auto"
className="min-w-[32rem]"
classNames={{
wrapper:
"bg-transparent p-0 shadow-none border-none overflow-auto rounded-2xl",
@@ -36,6 +36,7 @@ import {
} from "lucide-react";
import toast from "react-hot-toast";
import { formatTraffic } from "@/utils/traffic";
import {
getMonitorTunnels,
getTunnelMetrics,
@@ -390,12 +391,7 @@ const TrafficChartCard = React.memo(function TrafficChartCard({
const yFormatter = (value: unknown) => {
const n = Number(value);
if (!Number.isFinite(n) || n <= 0) return "0 B";
const k = 1024;
const sizes = ["B", "KB", "MB", "GB", "TB"];
const i = Math.floor(Math.log(n) / Math.log(k));
return `${parseFloat((n / Math.pow(k, i)).toFixed(2))} ${sizes[i]}`;
return formatTraffic(n);
};
return (
+71 -39
View File
@@ -5,6 +5,15 @@ import { Button } from "@/shadcn-bridge/heroui/button";
import { Card, CardBody, CardHeader } from "@/shadcn-bridge/heroui/card";
import { Tabs, Tab } from "@/shadcn-bridge/heroui/tabs";
import { Input } from "@/shadcn-bridge/heroui/input";
import { TrafficLimitField } from "@/components/traffic-limit-field";
import {
formatTraffic,
MIB,
parseTrafficInput,
preferredTrafficUnit,
TRAFFIC_UNIT_MIB,
type TrafficUnit,
} from "@/utils/traffic";
import {
Modal,
ModalContent,
@@ -82,6 +91,8 @@ interface RemoteUsageNode {
syncError?: string;
}
const MAX_SAFE_BANDWIDTH_MIB = Math.floor(Number.MAX_SAFE_INTEGER / MIB);
export default function PanelSharingPage() {
const [selectedTab, setSelectedTab] = useState("my-shares");
const [shares, setShares] = useState<PeerShare[]>([]);
@@ -108,6 +119,7 @@ export default function PanelSharingPage() {
allowedDomains: "",
allowedIps: "",
});
const [shareUnit, setShareUnit] = useState<TrafficUnit>("GB");
const [importForm, setImportForm] = useState({
remoteUrl: "",
@@ -124,6 +136,9 @@ export default function PanelSharingPage() {
allowedDomains: "",
allowedIps: "",
});
const [editUnit, setEditUnit] = useState<TrafficUnit>("GB");
const [editOriginalMaxBandwidth, setEditOriginalMaxBandwidth] = useState(0);
const [editBandwidthChanged, setEditBandwidthChanged] = useState(false);
const loadShares = useCallback(async () => {
setLoading(true);
@@ -212,8 +227,17 @@ export default function PanelSharingPage() {
return;
}
if (shareForm.maxBandwidth < 0) {
toast.error("流量上限不能为负数");
const limitMiB =
shareForm.maxBandwidth === 0
? 0
: parseTrafficInput(String(shareForm.maxBandwidth), shareUnit);
if (
limitMiB === null ||
limitMiB > MAX_SAFE_BANDWIDTH_MIB ||
shareForm.maxBandwidth < 0
) {
toast.error("请输入有效的流量上限,0 表示不限流量");
return;
}
@@ -223,7 +247,7 @@ export default function PanelSharingPage() {
const res = await createPeerShare({
name: shareForm.name,
nodeId,
maxBandwidth: Math.max(0, shareForm.maxBandwidth) * 1024 * 1024 * 1024,
maxBandwidth: limitMiB * MIB,
expiryTime: shareForm.expiryDays === 0 ? 0 : expiryTime,
portRangeStart: shareForm.portRangeStart,
portRangeEnd: shareForm.portRangeEnd,
@@ -274,13 +298,16 @@ export default function PanelSharingPage() {
};
const openEditShare = (share: PeerShare) => {
const mib = share.maxBandwidth / MIB;
const unit = preferredTrafficUnit(mib);
setEditUnit(unit);
setEditOriginalMaxBandwidth(share.maxBandwidth);
setEditBandwidthChanged(false);
setEditForm({
id: share.id,
name: share.name,
maxBandwidth:
share.maxBandwidth > 0
? Math.round(share.maxBandwidth / (1024 * 1024 * 1024))
: 0,
maxBandwidth: share.maxBandwidth > 0 ? mib / TRAFFIC_UNIT_MIB[unit] : 0,
expiryTime: share.expiryTime,
portRangeStart: share.portRangeStart,
portRangeEnd: share.portRangeEnd,
@@ -296,8 +323,18 @@ export default function PanelSharingPage() {
return;
}
if (editForm.maxBandwidth < 0) {
toast.error("流量上限不能为负数");
const limitMiB =
editForm.maxBandwidth === 0
? 0
: parseTrafficInput(String(editForm.maxBandwidth), editUnit);
if (
editBandwidthChanged &&
(limitMiB === null ||
limitMiB > MAX_SAFE_BANDWIDTH_MIB ||
editForm.maxBandwidth < 0)
) {
toast.error("请输入有效的流量上限,0 表示不限流量");
return;
}
@@ -305,7 +342,9 @@ export default function PanelSharingPage() {
const res = await updatePeerShare({
id: editForm.id,
name: editForm.name,
maxBandwidth: Math.max(0, editForm.maxBandwidth) * 1024 * 1024 * 1024,
maxBandwidth: editBandwidthChanged
? (limitMiB as number) * MIB
: editOriginalMaxBandwidth,
expiryTime: editForm.expiryTime,
portRangeStart: editForm.portRangeStart,
portRangeEnd: editForm.portRangeEnd,
@@ -362,17 +401,7 @@ export default function PanelSharingPage() {
toast.success("Token已复制");
};
const formatFlowGB = (bytes: number) => {
if (!Number.isFinite(bytes) || bytes <= 0) {
return "0 B";
}
if (bytes < 1024) return bytes + " B";
if (bytes < 1024 * 1024) return (bytes / 1024).toFixed(2) + " KB";
if (bytes < 1024 * 1024 * 1024)
return (bytes / (1024 * 1024)).toFixed(2) + " MB";
return (bytes / (1024 * 1024 * 1024)).toFixed(2) + " GB";
};
const formatFlowGB = formatTraffic;
const formatChainType = (chainType: number, hopInx: number) => {
if (chainType === 1) {
@@ -741,17 +770,18 @@ export default function PanelSharingPage() {
})
}
/>
<Input
<TrafficLimitField
description="0 表示不限流量"
label="流量上限 (GB)"
type="number"
label="流量上限"
unit={shareUnit}
value={shareForm.maxBandwidth.toString()}
onChange={(e) =>
setShareForm({
...shareForm,
maxBandwidth: parseInt(e.target.value, 10) || 0,
})
}
onChange={(value, unit) => {
setShareForm((prev) => ({
...prev,
maxBandwidth: Number(value) || 0,
}));
setShareUnit(unit);
}}
/>
<Input
description="限制使用此Token的来源面板域名,多个域名用逗号分隔,留空不限制"
@@ -826,17 +856,19 @@ export default function PanelSharingPage() {
}
/>
</div>
<Input
<TrafficLimitField
description="0 表示不限流量"
label="流量上限 (GB)"
type="number"
label="流量上限"
unit={editUnit}
value={editForm.maxBandwidth.toString()}
onChange={(e) =>
setEditForm({
...editForm,
maxBandwidth: parseInt(e.target.value, 10) || 0,
})
}
onChange={(value, unit) => {
setEditForm((prev) => ({
...prev,
maxBandwidth: Number(value) || 0,
}));
setEditUnit(unit);
setEditBandwidthChanged(true);
}}
/>
<Input
description="留空或清除表示永久有效"
+85 -48
View File
@@ -71,23 +71,22 @@ import {
SearchIcon,
} from "@/components/icons";
import { PageLoadingState } from "@/components/page-state";
import { TrafficLimitField } from "@/components/traffic-limit-field";
import { useLocalStorageState } from "@/hooks/use-local-storage-state";
import { removeItemsById, replaceItemById } from "@/utils/list-state";
import {
formatTraffic,
formatFlowLimit,
flowLimitBytes,
flowLimitMiB,
parseTrafficInput,
preferredTrafficUnit,
TRAFFIC_UNIT_MIB,
type TrafficUnit,
} from "@/utils/traffic";
// 工具函数
const formatFlow = (value: number, unit: string = "bytes"): string => {
if (unit === "gb") {
return `${value} GB`;
} else {
if (value === 0) return "0 B";
if (value < 1024) return `${value} B`;
if (value < 1024 * 1024) return `${(value / 1024).toFixed(2)} KB`;
if (value < 1024 * 1024 * 1024)
return `${(value / (1024 * 1024)).toFixed(2)} MB`;
return `${(value / (1024 * 1024 * 1024)).toFixed(2)} GB`;
}
};
const formatFlow = formatTraffic;
const formatQuotaLimit = (value?: number): string => {
const limit = Number(value ?? 0);
@@ -96,7 +95,14 @@ const formatQuotaLimit = (value?: number): string => {
return "不限";
}
return `${limit} GB`;
return formatTraffic(limit * 1024 ** 3);
};
const trafficInputFor = (flowGB: number, flowMiB?: number) => {
const mib = flowLimitMiB(flowGB, flowMiB);
const unit = preferredTrafficUnit(mib);
return { value: String(mib / TRAFFIC_UNIT_MIB[unit]), unit };
};
const formatDate = (timestamp: number): string => {
@@ -148,6 +154,7 @@ const normalizeUserItem = (item: Partial<User>): User => {
user: String(item.user ?? ""),
status: Number(item.status ?? 0),
flow: Number(item.flow ?? 0),
flowMiB: Number(item.flowMiB ?? 0),
num: Number(item.num ?? 0),
expTime: item.expTime,
flowResetTime: item.flowResetTime ?? 0,
@@ -172,6 +179,7 @@ const normalizeUserTunnelItem = (item: Partial<UserTunnel>): UserTunnel => {
tunnelName: String(item.tunnelName ?? ""),
status: Number(item.status ?? 0),
flow: Number(item.flow ?? 0),
flowMiB: Number(item.flowMiB ?? 0),
num: Number(item.num ?? 0),
expTime: Number(item.expTime ?? 0),
flowResetTime: Number(item.flowResetTime ?? 0),
@@ -223,6 +231,8 @@ export default function UserPage() {
maxConn: 0,
});
const [userFormLoading, setUserFormLoading] = useState(false);
const [userFlowInput, setUserFlowInput] = useState("1000");
const [userFlowUnit, setUserFlowUnit] = useState<TrafficUnit>("GB");
const [quotaResetLoading, setQuotaResetLoading] = useState(false);
const editingUser = useMemo(
@@ -263,6 +273,8 @@ export default function UserPage() {
onClose: onEditTunnelModalClose,
} = useDisclosure();
const [editTunnelForm, setEditTunnelForm] = useState<UserTunnel | null>(null);
const [tunnelFlowInput, setTunnelFlowInput] = useState("");
const [tunnelFlowUnit, setTunnelFlowUnit] = useState<TrafficUnit>("GB");
const [editTunnelLoading, setEditTunnelLoading] = useState(false);
// 删除确认相关状态
@@ -506,6 +518,8 @@ export default function UserPage() {
const handleAdd = () => {
setIsEdit(false);
setUserFlowInput("1000");
setUserFlowUnit("GB");
setUserForm({
user: "",
pwd: "",
@@ -524,6 +538,10 @@ export default function UserPage() {
const handleEdit = async (user: User) => {
setIsEdit(true);
const trafficInput = trafficInputFor(user.flow, user.flowMiB);
setUserFlowInput(trafficInput.value);
setUserFlowUnit(trafficInput.unit);
let currentGroupIds: number[] = [];
try {
@@ -591,10 +609,21 @@ export default function UserPage() {
return;
}
const flowMiB = parseTrafficInput(userFlowInput, userFlowUnit);
if (flowMiB === null) {
toast.error("请输入有效的流量限制,最小单位为 1 MB");
return;
}
setUserFormLoading(true);
try {
const submitData: any = {
...userForm,
flow: Math.ceil(flowMiB / 1024),
flowMiB:
userFlowInput === "99999" && userFlowUnit === "GB" ? 0 : flowMiB,
expTime: userForm.expTime.getTime(),
groupIds: userForm.groupIds ?? [],
};
@@ -756,6 +785,10 @@ export default function UserPage() {
};
const handleEditTunnel = (userTunnel: UserTunnel) => {
const trafficInput = trafficInputFor(userTunnel.flow, userTunnel.flowMiB);
setTunnelFlowInput(trafficInput.value);
setTunnelFlowUnit(trafficInput.unit);
setEditTunnelForm({
...userTunnel,
speedId: normalizeSpeedId(userTunnel.speedId),
@@ -767,12 +800,24 @@ export default function UserPage() {
const handleUpdateTunnel = async () => {
if (!editTunnelForm) return;
const flowMiB = parseTrafficInput(tunnelFlowInput, tunnelFlowUnit);
if (flowMiB === null) {
toast.error("请输入有效的流量限制,最小单位为 1 MB");
return;
}
const flow = Math.ceil(flowMiB / 1024);
const storedFlowMiB =
tunnelFlowInput === "99999" && tunnelFlowUnit === "GB" ? 0 : flowMiB;
setEditTunnelLoading(true);
try {
const speedLimitAutoCleared = isMissingSpeedLimit(editTunnelForm.speedId);
const response = await updateUserTunnel({
id: editTunnelForm.id,
flow: editTunnelForm.flow,
flow,
flowMiB: storedFlowMiB,
num: editTunnelForm.num,
expTime: editTunnelForm.expTime,
flowResetTime: editTunnelForm.flowResetTime,
@@ -792,6 +837,8 @@ export default function UserPage() {
if (currentUser) {
const nextTunnel = normalizeUserTunnelItem({
...editTunnelForm,
flow,
flowMiB: storedFlowMiB,
speedId: normalizeSpeedId(editTunnelForm.speedId),
speedLimitName:
normalizeSpeedId(editTunnelForm.speedId) !== null
@@ -1148,7 +1195,7 @@ export default function UserPage() {
<div className="flex items-center gap-1 text-xs">
<span className="text-default-500">限制:</span>
<span className="text-default-700 font-medium whitespace-nowrap">
{formatFlow(user.flow, "gb")}
{formatFlowLimit(user.flow, user.flowMiB)}
</span>
</div>
</div>
@@ -1258,9 +1305,9 @@ export default function UserPage() {
: null;
const usedFlow = calculateUserTotalUsedFlow(user);
const flowPercent =
user.flow > 0
user.flow > 0 && !(user.flow === 99999 && !user.flowMiB)
? Math.min(
(usedFlow / (user.flow * 1024 * 1024 * 1024)) * 100,
(usedFlow / flowLimitBytes(user.flow, user.flowMiB)) * 100,
100,
)
: 0;
@@ -1316,7 +1363,7 @@ export default function UserPage() {
<div className="flex justify-between text-sm">
<span className="text-default-600">流量限制</span>
<span className="font-medium text-xs">
{formatFlow(user.flow, "gb")}
{formatFlowLimit(user.flow, user.flowMiB)}
</span>
</div>
<div className="flex justify-between text-sm">
@@ -1504,20 +1551,14 @@ export default function UserPage() {
setUserForm((prev) => ({ ...prev, pwd: e.target.value }))
}
/>
<Input
<TrafficLimitField
isRequired
label="流量限制(GB)"
max="99999"
min="1"
type="number"
value={userForm.flow.toString()}
onChange={(e) => {
const value = Math.min(
Math.max(Number(e.target.value) || 0, 1),
99999,
);
setUserForm((prev) => ({ ...prev, flow: value }));
label="流量限制"
unit={userFlowUnit}
value={userFlowInput}
onChange={(value, unit) => {
setUserFlowInput(value);
setUserFlowUnit(unit);
}}
/>
<Input
@@ -1963,7 +2004,10 @@ export default function UserPage() {
<div className="flex justify-between text-small">
<span className="text-gray-600">限制:</span>
<span className="font-medium">
{formatFlow(userTunnel.flow, "gb")}
{formatFlowLimit(
userTunnel.flow,
userTunnel.flowMiB,
)}
</span>
</div>
<div className="flex justify-between text-small">
@@ -2083,21 +2127,14 @@ export default function UserPage() {
{editTunnelForm && (
<>
<div className="grid grid-cols-1 md:grid-cols-2 gap-4">
<Input
label="流量限制(GB)"
max="99999"
min="1"
type="number"
value={editTunnelForm.flow.toString()}
onChange={(e) => {
const value = Math.min(
Math.max(Number(e.target.value) || 0, 1),
99999,
);
setEditTunnelForm((prev) =>
prev ? { ...prev, flow: value } : null,
);
<TrafficLimitField
isRequired
label="流量限制"
unit={tunnelFlowUnit}
value={tunnelFlowInput}
onChange={(value, unit) => {
setTunnelFlowInput(value);
setTunnelFlowUnit(unit);
}}
/>
+2
View File
@@ -12,6 +12,7 @@ export interface User {
pwd?: string;
status: number; // 1-正常, 0-禁用
flow: number; // 流量限制(GB)
flowMiB?: number; // 精确流量限制(MiB),0 表示沿用旧版 GB 字段
num: number; // 转发数量
expTime?: number; // 过期时间戳
flowResetTime?: number; // 流量重置日期(1-31号)
@@ -56,6 +57,7 @@ export interface UserTunnel {
tunnelName: string;
status: number; // 1-正常, 0-禁用
flow: number; // 流量限制(GB)
flowMiB?: number;
num: number; // 转发数量
expTime: number; // 过期时间戳
flowResetTime: number;
+72
View File
@@ -0,0 +1,72 @@
export type TrafficUnit = "MB" | "GB" | "TB" | "PB";
export const MIB = 1024 * 1024;
export const GIB = 1024 * MIB;
export const TRAFFIC_UNIT_MIB: Record<TrafficUnit, number> = {
MB: 1,
GB: 1024,
TB: 1024 ** 2,
PB: 1024 ** 3,
};
const BYTE_UNITS = ["B", "KB", "MB", "GB", "TB", "PB"];
export function formatTraffic(bytes: number): string {
if (!Number.isFinite(bytes) || bytes <= 0) return "0 B";
let value = bytes;
let unit = 0;
while (value >= 1024 && unit < BYTE_UNITS.length - 1) {
value /= 1024;
unit++;
}
return `${unit === 0 ? Math.floor(value) : value.toFixed(2)} ${BYTE_UNITS[unit]}`;
}
export function flowLimitMiB(flowGB: number, flowMiB?: number): number {
return flowMiB && flowMiB > 0 ? flowMiB : flowGB * 1024;
}
export function flowLimitBytes(flowGB: number, flowMiB?: number): number {
return flowLimitMiB(flowGB, flowMiB) * MIB;
}
export function formatFlowLimit(flowGB: number, flowMiB?: number): string {
if (flowGB === 99999 && !flowMiB) return "无限制";
return formatTraffic(flowLimitBytes(flowGB, flowMiB));
}
export function preferredTrafficUnit(mib: number): TrafficUnit {
if (mib <= 0) return "GB";
if (mib > 0 && mib % TRAFFIC_UNIT_MIB.PB === 0) return "PB";
if (mib > 0 && mib % TRAFFIC_UNIT_MIB.TB === 0) return "TB";
if (mib > 0 && mib % TRAFFIC_UNIT_MIB.GB === 0) return "GB";
return "MB";
}
export function parseTrafficInput(
value: string,
unit: TrafficUnit,
): number | null {
const amount = Number(value);
const mib = amount * TRAFFIC_UNIT_MIB[unit];
if (
!value.trim() ||
!Number.isFinite(amount) ||
amount <= 0 ||
!Number.isSafeInteger(mib)
) {
return null;
}
// Backend stores bytes as int64. Keep the converted value within that range.
if (mib > 8_796_093_022_207) return null;
return mib;
}