mirror of
https://github.com/Sagit-chu/flvx.git
synced 2026-09-28 07:36:38 +08:00
Compare commits
15 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 129fa0aa4c | |||
| cf7246c71b | |||
| b4c2989285 | |||
| f2783713a5 | |||
| 5d60c4fbe1 | |||
| 62c56e0a79 | |||
| 2e845030de | |||
| 2269f2e2d5 | |||
| da7bef88f1 | |||
| f26014579b | |||
| 5041c722c9 | |||
| c56798e991 | |||
| 40e96f3592 | |||
| 0e24b53a5b | |||
| a820c49c94 |
@@ -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
|
||||
|
||||
@@ -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 二进制文件更新完成"
|
||||
|
||||
|
||||
+1
-1
@@ -70,7 +70,7 @@ Alpine Linux 最小化安装若未包含 `curl`,可使用系统自带的 `wget
|
||||
wget -O install.sh https://raw.githubusercontent.com/Sagit-chu/flux-panel/main/install.sh && chmod +x install.sh && ./install.sh
|
||||
```
|
||||
|
||||
脚本会在 Alpine 上自动安装 Bash,并使用 OpenRC 注册、启动和管理 `flux_agent` 服务;其他受支持的 Linux 发行版继续使用 systemd。
|
||||
脚本会在 Alpine 上自动安装 Bash、`curl` 和 CA 证书,并使用 OpenRC 注册、启动和管理 `flux_agent` 服务;其他受支持的 Linux 发行版继续使用 systemd。
|
||||
|
||||
**安装过程中会提示输入:**
|
||||
- **服务器地址**: 面板端的通信地址(通常是 `http://<面板IP>:<后端端口>`,例如 `http://1.2.3.4:6365`)。
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -5,6 +5,7 @@ import (
|
||||
"database/sql"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"log"
|
||||
@@ -20,7 +21,6 @@ import (
|
||||
"go-backend/internal/health"
|
||||
"go-backend/internal/http/middleware"
|
||||
"go-backend/internal/http/response"
|
||||
"go-backend/internal/license"
|
||||
"go-backend/internal/metrics"
|
||||
"go-backend/internal/monitoring"
|
||||
runtimenft "go-backend/internal/runtime/nftables"
|
||||
@@ -42,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
|
||||
@@ -400,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
|
||||
@@ -435,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))
|
||||
}
|
||||
|
||||
@@ -711,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,
|
||||
@@ -878,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()
|
||||
@@ -908,56 +923,35 @@ func (h *Handler) licenseActivate(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.ErrDefault("授权码不能为空"))
|
||||
return
|
||||
}
|
||||
h.licenseValidationMu.Lock()
|
||||
defer h.licenseValidationMu.Unlock()
|
||||
|
||||
accountID := "1bc96cac-09de-4cf4-af34-26afdad63a90"
|
||||
|
||||
fingerprint, err := h.getOrCreateMachineFingerprint()
|
||||
valResp, err := h.validateLicenseForMachine(key)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault("生成设备指纹失败"))
|
||||
return
|
||||
}
|
||||
|
||||
client := license.NewKeygenClient(accountID, "")
|
||||
valResp, err := client.ValidateKeyWithFingerprint(key, fingerprint)
|
||||
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
|
||||
}
|
||||
|
||||
if !valResp.Meta.Valid {
|
||||
if valResp.Meta.Code == "NO_MACHINES" || valResp.Meta.Code == "NO_MACHINE" || valResp.Meta.Code == "MACHINE_SCOPE_REQUIRED" || valResp.Meta.Code == "FINGERPRINT_SCOPE_MISMATCH" {
|
||||
// Needs machine activation
|
||||
client.Token = key
|
||||
err = client.ActivateMachine(valResp.Data.ID, fingerprint)
|
||||
if err != nil {
|
||||
// Translate specific error messages or log them
|
||||
response.WriteJSON(w, response.ErrDefault("设备绑定失败: "+err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
// Validation might still fail with scope if we don't query via machine id, but since activate machine succeeded
|
||||
// we can consider the license valid for our simple usecase
|
||||
} else {
|
||||
response.WriteJSON(w, response.ErrDefault("授权码无效或已过期 (Code: "+valResp.Meta.Code+")"))
|
||||
return
|
||||
}
|
||||
response.WriteJSON(w, response.ErrDefault("授权码无效或已过期 (Code: "+valResp.Meta.Code+")"))
|
||||
return
|
||||
}
|
||||
|
||||
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
|
||||
}
|
||||
@@ -999,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("需要商业版授权"))
|
||||
@@ -1041,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" {
|
||||
@@ -1212,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,
|
||||
@@ -1261,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,
|
||||
|
||||
@@ -4,8 +4,6 @@ import (
|
||||
"context"
|
||||
"log"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/license"
|
||||
)
|
||||
|
||||
var nftablesTrafficCollectInterval = 30 * time.Second
|
||||
@@ -38,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()
|
||||
|
||||
@@ -55,22 +54,25 @@ func (h *Handler) validateLicenseJob() {
|
||||
if h == nil || h.repo == nil {
|
||||
return
|
||||
}
|
||||
|
||||
accountID := "1bc96cac-09de-4cf4-af34-26afdad63a90"
|
||||
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
|
||||
}
|
||||
|
||||
fingerprint, _ := h.repo.GetViteConfigValue("machine_fingerprint")
|
||||
client := license.NewKeygenClient(accountID, "")
|
||||
valResp, err := client.ValidateKeyWithFingerprint(key, fingerprint)
|
||||
valResp, err := h.validateLicenseForMachine(key)
|
||||
|
||||
if err != nil {
|
||||
// Network error or timeout. Grace period by not revoking immediately here.
|
||||
// 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 licenseValidationErrorIsDefinitive(valResp, err) {
|
||||
now := time.Now().UnixMilli()
|
||||
_ = h.repo.UpsertConfig("is_commercial", "false", now)
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -0,0 +1,109 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"os"
|
||||
"strings"
|
||||
|
||||
"go-backend/internal/license"
|
||||
)
|
||||
|
||||
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":
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
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")
|
||||
|
||||
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
|
||||
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.ValidateKeyWithMachine(key, fingerprint, machineID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if validation.Meta.Valid {
|
||||
validation.MachineID = machineID
|
||||
}
|
||||
return validation, nil
|
||||
}
|
||||
@@ -0,0 +1,358 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/license"
|
||||
"go-backend/internal/store/repo"
|
||||
)
|
||||
|
||||
func TestValidateLicenseJobRepairsMissingMachineBinding(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":"2030-01-02T00:00:00.000Z"}}}`)
|
||||
case strings.HasSuffix(req.URL.Path, "/machines"):
|
||||
w.WriteHeader(http.StatusCreated)
|
||||
_, _ = 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)
|
||||
}
|
||||
}))
|
||||
defer server.Close()
|
||||
restoreLicenseClientFactory(t, server.URL)
|
||||
|
||||
h := &Handler{repo: r}
|
||||
h.validateLicenseJob()
|
||||
|
||||
assertLicenseConfig(t, r, "is_commercial", "true")
|
||||
assertLicenseConfig(t, r, "license_expiry", "2030-01-02T00:00:00.000Z")
|
||||
fingerprint, err := r.GetViteConfigValue("machine_fingerprint")
|
||||
if err != nil || strings.TrimSpace(fingerprint) == "" {
|
||||
t.Fatalf("expected persisted machine fingerprint, got value=%q err=%v", fingerprint, err)
|
||||
}
|
||||
if got := validations.Load(); got != 2 {
|
||||
t.Fatalf("validation calls = %d, want 2", got)
|
||||
}
|
||||
}
|
||||
|
||||
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)
|
||||
|
||||
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"):
|
||||
code := "NO_MACHINE"
|
||||
if validations.Add(1) > 1 {
|
||||
code = "FINGERPRINT_SCOPE_MISMATCH"
|
||||
}
|
||||
_, _ = 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, `{"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)
|
||||
}
|
||||
}))
|
||||
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(), "FINGERPRINT_SCOPE_MISMATCH") {
|
||||
t.Fatalf("expected post-activation validation failure, got %s", res.Body.String())
|
||||
}
|
||||
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()
|
||||
seedLicenseConfig(t, r, "license_key", "license-secret", now)
|
||||
seedLicenseConfig(t, r, "is_commercial", "true", 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":"NO_MACHINE"},"data":{"id":"license-id","attributes":{}}}`)
|
||||
case strings.HasSuffix(req.URL.Path, "/machines"):
|
||||
w.WriteHeader(http.StatusUnprocessableEntity)
|
||||
_, _ = fmt.Fprint(w, `{"errors":[{"code":"MACHINE_LIMIT_EXCEEDED"}]}`)
|
||||
default:
|
||||
http.NotFound(w, req)
|
||||
}
|
||||
}))
|
||||
defer server.Close()
|
||||
restoreLicenseClientFactory(t, server.URL)
|
||||
|
||||
h := &Handler{repo: r}
|
||||
h.validateLicenseJob()
|
||||
|
||||
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"))
|
||||
if err != nil {
|
||||
t.Fatalf("repo.Open() error = %v", err)
|
||||
}
|
||||
t.Cleanup(func() { _ = r.Close() })
|
||||
return r
|
||||
}
|
||||
|
||||
func seedLicenseConfig(t *testing.T, r *repo.Repository, name, value string, now int64) {
|
||||
t.Helper()
|
||||
if err := r.UpsertConfig(name, value, now); err != nil {
|
||||
t.Fatalf("UpsertConfig(%q) error = %v", name, err)
|
||||
}
|
||||
}
|
||||
|
||||
func assertLicenseConfig(t *testing.T, r *repo.Repository, name, want string) {
|
||||
t.Helper()
|
||||
got, err := r.GetViteConfigValue(name)
|
||||
if err != nil {
|
||||
t.Fatalf("GetViteConfigValue(%q) error = %v", name, err)
|
||||
}
|
||||
if got != want {
|
||||
t.Fatalf("config %q = %q, want %q", name, got, want)
|
||||
}
|
||||
}
|
||||
|
||||
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)
|
||||
client.BaseURL = baseURL
|
||||
return client
|
||||
}
|
||||
t.Cleanup(func() { newLicenseClient = previous })
|
||||
}
|
||||
@@ -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)
|
||||
|
||||
@@ -6,24 +6,53 @@ import (
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
var AccountID string
|
||||
|
||||
type KeygenClient struct {
|
||||
AccountID string
|
||||
Token string
|
||||
BaseURL string
|
||||
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 {
|
||||
return &KeygenClient{
|
||||
AccountID: accountID,
|
||||
Token: token,
|
||||
AccountID: accountID,
|
||||
Token: token,
|
||||
BaseURL: defaultAPIBaseURL,
|
||||
HTTPClient: &http.Client{Timeout: 10 * time.Second},
|
||||
}
|
||||
}
|
||||
|
||||
func (c *KeygenClient) apiURL(path string) string {
|
||||
baseURL := strings.TrimRight(c.BaseURL, "/")
|
||||
if baseURL == "" {
|
||||
baseURL = defaultAPIBaseURL
|
||||
}
|
||||
return fmt.Sprintf("%s/accounts/%s/%s", baseURL, c.AccountID, strings.TrimLeft(path, "/"))
|
||||
}
|
||||
|
||||
type ValidateResponse struct {
|
||||
Meta struct {
|
||||
Valid bool `json:"valid"`
|
||||
@@ -35,6 +64,7 @@ type ValidateResponse struct {
|
||||
Expiry string `json:"expiry"`
|
||||
} `json:"attributes"`
|
||||
} `json:"data"`
|
||||
MachineID string `json:"-"`
|
||||
}
|
||||
|
||||
type ActivateMachineRequest struct {
|
||||
@@ -54,8 +84,37 @@ 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) {
|
||||
url := fmt.Sprintf("https://api.keygen.sh/v1/accounts/%s/licenses/actions/validate-key", c.AccountID)
|
||||
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{}{
|
||||
"key": key,
|
||||
@@ -66,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,
|
||||
@@ -91,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
|
||||
@@ -102,8 +170,44 @@ 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 := fmt.Sprintf("https://api.keygen.sh/v1/accounts/%s/licenses/actions/validate-key", c.AccountID)
|
||||
url := c.apiURL("licenses/actions/validate-key")
|
||||
|
||||
reqBody := map[string]interface{}{
|
||||
"meta": map[string]string{
|
||||
@@ -130,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
|
||||
@@ -141,8 +246,8 @@ func (c *KeygenClient) ValidateKey(key string) (*ValidateResponse, error) {
|
||||
return &valResp, nil
|
||||
}
|
||||
|
||||
func (c *KeygenClient) ActivateMachine(licenseID, fingerprint string) error {
|
||||
url := fmt.Sprintf("https://api.keygen.sh/v1/accounts/%s/machines", c.AccountID)
|
||||
func (c *KeygenClient) ActivateMachine(licenseID, fingerprint string) (string, error) {
|
||||
url := c.apiURL("machines")
|
||||
|
||||
var reqBody ActivateMachineRequest
|
||||
reqBody.Data.Type = "machines"
|
||||
@@ -165,23 +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()
|
||||
body, _ := io.ReadAll(resp.Body)
|
||||
|
||||
if resp.StatusCode == http.StatusCreated || resp.StatusCode == http.StatusOK {
|
||||
return nil
|
||||
}
|
||||
|
||||
body, _ := io.ReadAll(resp.Body)
|
||||
|
||||
if resp.StatusCode == http.StatusConflict || resp.StatusCode == http.StatusUnprocessableEntity {
|
||||
if strings.Contains(string(body), "FINGERPRINT_TAKEN") || strings.Contains(string(body), "MACHINE_LIMIT_EXCEEDED") {
|
||||
// Machine already registered to this license or limit reached because it's already us.
|
||||
// The subsequent ValidateKey check will determine if the existing machine is actually us.
|
||||
return nil
|
||||
var machineResp MachineResponse
|
||||
if json.Unmarshal(body, &machineResp) == nil {
|
||||
return strings.TrimSpace(machineResp.Data.ID), nil
|
||||
}
|
||||
return "", nil
|
||||
}
|
||||
|
||||
return fmt.Errorf("failed to activate machine: status %d, response: %s", resp.StatusCode, string(body))
|
||||
}
|
||||
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)}
|
||||
}
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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})
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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) {
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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 {
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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)
|
||||
|
||||
+46
-5
@@ -64,6 +64,38 @@ SERVICE_MANAGER="${SERVICE_MANAGER:-}"
|
||||
PROXY_ENABLED="${PROXY_ENABLED:-}"
|
||||
PROXY_URL="${PROXY_URL:-}"
|
||||
|
||||
ensure_alpine_runtime_dependencies() {
|
||||
[[ -f /etc/alpine-release ]] || return 0
|
||||
|
||||
local missing_packages=()
|
||||
local privileged_command=""
|
||||
|
||||
command -v curl >/dev/null 2>&1 || missing_packages+=(curl)
|
||||
[[ -f /etc/ssl/certs/ca-certificates.crt ]] || missing_packages+=(ca-certificates)
|
||||
|
||||
if [[ ${#missing_packages[@]} -eq 0 ]]; then
|
||||
return 0
|
||||
fi
|
||||
|
||||
if [[ $EUID -ne 0 ]]; then
|
||||
if command -v sudo >/dev/null 2>&1; then
|
||||
privileged_command="sudo"
|
||||
elif command -v doas >/dev/null 2>&1; then
|
||||
privileged_command="doas"
|
||||
else
|
||||
echo "❌ Alpine 安装需要 root 权限,或已配置 sudo/doas 来安装依赖: ${missing_packages[*]}。" >&2
|
||||
return 1
|
||||
fi
|
||||
fi
|
||||
|
||||
echo "📦 Alpine 缺少运行依赖,正在安装: ${missing_packages[*]}"
|
||||
if [[ -n "$privileged_command" ]]; then
|
||||
"$privileged_command" apk add --no-cache "${missing_packages[@]}"
|
||||
else
|
||||
apk add --no-cache "${missing_packages[@]}"
|
||||
fi
|
||||
}
|
||||
|
||||
# 镜像加速
|
||||
maybe_proxy_url() {
|
||||
local url="$1"
|
||||
@@ -169,6 +201,8 @@ build_download_url() {
|
||||
}
|
||||
|
||||
ensure_download_url_initialized() {
|
||||
ensure_alpine_runtime_dependencies || return 1
|
||||
|
||||
if [[ -n "${DOWNLOAD_URL:-}" ]]; then
|
||||
return 0
|
||||
fi
|
||||
@@ -299,11 +333,16 @@ ensure_service_manager() {
|
||||
esac
|
||||
fi
|
||||
|
||||
if command -v systemctl >/dev/null 2>&1 && [[ -d /run/systemd/system ]]; then
|
||||
# Alpine uses OpenRC even if a systemctl compatibility command happens to be installed.
|
||||
if [[ -f /etc/alpine-release ]]; then
|
||||
if command -v rc-service >/dev/null 2>&1 && command -v rc-update >/dev/null 2>&1; then
|
||||
SERVICE_MANAGER="openrc"
|
||||
return 0
|
||||
fi
|
||||
elif command -v systemctl >/dev/null 2>&1 && [[ -d /run/systemd/system ]]; then
|
||||
SERVICE_MANAGER="systemd"
|
||||
return 0
|
||||
fi
|
||||
if command -v rc-service >/dev/null 2>&1 && command -v rc-update >/dev/null 2>&1; then
|
||||
elif command -v rc-service >/dev/null 2>&1 && command -v rc-update >/dev/null 2>&1; then
|
||||
SERVICE_MANAGER="openrc"
|
||||
return 0
|
||||
fi
|
||||
@@ -366,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
|
||||
@@ -501,7 +541,8 @@ cleanup_legacy_gost_installation() {
|
||||
return 0
|
||||
fi
|
||||
|
||||
if systemctl list-units --full -all 2>/dev/null | grep -Fq "gost.service"; then
|
||||
if [[ "$SERVICE_MANAGER" == "systemd" ]] && \
|
||||
systemctl list-units --full -all 2>/dev/null | grep -Fq "gost.service"; then
|
||||
systemctl stop gost 2>/dev/null || true
|
||||
systemctl disable gost 2>/dev/null || true
|
||||
fi
|
||||
@@ -520,7 +561,7 @@ cleanup_legacy_gost_installation() {
|
||||
rm -f "$LEGACY_GOST_CONFIG_DIR/gost"
|
||||
fi
|
||||
|
||||
if [[ "$removed_service_file" == "1" ]]; then
|
||||
if [[ "$removed_service_file" == "1" && "$SERVICE_MANAGER" == "systemd" ]]; then
|
||||
systemctl daemon-reload 2>/dev/null || true
|
||||
fi
|
||||
}
|
||||
|
||||
+319
-48
@@ -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
|
||||
|
||||
|
||||
@@ -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"
|
||||
@@ -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() (
|
||||
@@ -399,6 +401,7 @@ test_cleanup_legacy_gost_installation_removes_service_and_binary() (
|
||||
LEGACY_GOST_SERVICE_FILE_LIB=$(mktemp -u)
|
||||
LEGACY_GOST_SERVICE_FILE_USR_LIB=$(mktemp -u)
|
||||
LEGACY_GOST_CONFIG_DIR=$(mktemp -d)
|
||||
SERVICE_MANAGER="systemd"
|
||||
cat > "$LEGACY_GOST_SERVICE_FILE_ETC" <<EOF
|
||||
[Unit]
|
||||
Description=Gost Proxy Service
|
||||
@@ -442,6 +445,7 @@ test_cleanup_legacy_gost_installation_preserves_unrelated_gost() (
|
||||
LEGACY_GOST_SERVICE_FILE_LIB=$(mktemp -u)
|
||||
LEGACY_GOST_SERVICE_FILE_USR_LIB=$(mktemp -u)
|
||||
LEGACY_GOST_CONFIG_DIR=$(mktemp -d)
|
||||
SERVICE_MANAGER="systemd"
|
||||
cat > "$LEGACY_GOST_SERVICE_FILE_ETC" <<'EOF'
|
||||
[Unit]
|
||||
Description=Unrelated Gost Service
|
||||
@@ -469,6 +473,35 @@ EOF
|
||||
[[ "$systemctl_calls" != *"disable gost"* ]] || fail "cleanup_legacy_gost_installation should not disable unrelated gost services"
|
||||
)
|
||||
|
||||
test_cleanup_legacy_gost_installation_skips_systemd_on_openrc() (
|
||||
set -euo pipefail
|
||||
load_script_without_main "$ROOT_DIR/install.sh"
|
||||
|
||||
LEGACY_GOST_SERVICE_FILE_ETC=$(mktemp)
|
||||
LEGACY_GOST_SERVICE_FILE_LIB=$(mktemp -u)
|
||||
LEGACY_GOST_SERVICE_FILE_USR_LIB=$(mktemp -u)
|
||||
LEGACY_GOST_CONFIG_DIR=$(mktemp -d)
|
||||
SERVICE_MANAGER="openrc"
|
||||
cat > "$LEGACY_GOST_SERVICE_FILE_ETC" <<EOF
|
||||
[Unit]
|
||||
WorkingDirectory=$LEGACY_GOST_CONFIG_DIR
|
||||
ExecStart=$LEGACY_GOST_CONFIG_DIR/gost
|
||||
EOF
|
||||
: > "$LEGACY_GOST_CONFIG_DIR/config.json"
|
||||
: > "$LEGACY_GOST_CONFIG_DIR/gost.json"
|
||||
|
||||
local systemctl_calls=""
|
||||
systemctl() {
|
||||
systemctl_calls+=$'\n'"$*"
|
||||
return 1
|
||||
}
|
||||
|
||||
cleanup_legacy_gost_installation >/dev/null
|
||||
|
||||
[[ -z "$systemctl_calls" ]] || fail "OpenRC cleanup should not invoke systemctl"
|
||||
[[ ! -e "$LEGACY_GOST_SERVICE_FILE_ETC" ]] || fail "OpenRC cleanup should still remove the legacy service file"
|
||||
)
|
||||
|
||||
test_install_script_accepts_proxy_url_env_without_prompt() (
|
||||
set -euo pipefail
|
||||
load_script_without_main "$ROOT_DIR/install.sh"
|
||||
@@ -525,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"
|
||||
@@ -537,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"
|
||||
}
|
||||
@@ -545,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
|
||||
|
||||
@@ -560,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() (
|
||||
@@ -625,10 +757,14 @@ test_install_flux_agent_uses_openrc
|
||||
test_remove_flux_agent_service_uses_openrc
|
||||
test_cleanup_legacy_gost_installation_removes_service_and_binary
|
||||
test_cleanup_legacy_gost_installation_preserves_unrelated_gost
|
||||
test_cleanup_legacy_gost_installation_skips_systemd_on_openrc
|
||||
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
|
||||
|
||||
@@ -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;"]
|
||||
|
||||
@@ -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"
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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(() => {
|
||||
|
||||
@@ -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>
|
||||
);
|
||||
}
|
||||
@@ -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 =
|
||||
@@ -49,6 +55,30 @@ const readCachedConfigs = (keys: readonly string[]) => {
|
||||
return { cachedConfigs, hasCachedData };
|
||||
};
|
||||
|
||||
const readAllCachedSafeConfigs = () => {
|
||||
const cachedConfigs: Record<string, string> = {};
|
||||
|
||||
Object.keys(localStorage).forEach((storageKey) => {
|
||||
if (!storageKey.startsWith(CACHE_PREFIX)) {
|
||||
return;
|
||||
}
|
||||
|
||||
const key = storageKey.slice(CACHE_PREFIX.length).trim().toLowerCase();
|
||||
|
||||
if (!key || SENSITIVE_CONFIG_KEYS.has(key)) {
|
||||
return;
|
||||
}
|
||||
|
||||
const value = localStorage.getItem(storageKey);
|
||||
|
||||
if (value !== null) {
|
||||
cachedConfigs[key] = value;
|
||||
}
|
||||
});
|
||||
|
||||
return cachedConfigs;
|
||||
};
|
||||
|
||||
const fetchPublicBrandConfigs = async (): Promise<Record<string, string>> => {
|
||||
const publicConfigMap: Record<string, string> = {};
|
||||
|
||||
@@ -86,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,
|
||||
};
|
||||
@@ -99,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 =
|
||||
@@ -106,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,
|
||||
};
|
||||
}
|
||||
|
||||
@@ -123,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,
|
||||
};
|
||||
};
|
||||
|
||||
@@ -206,20 +246,24 @@ export const getCachedConfig = async (key: string): Promise<string | null> => {
|
||||
|
||||
// 获取所有配置(优先从缓存)
|
||||
export const getCachedConfigs = async (): Promise<Record<string, string>> => {
|
||||
const { cachedConfigs, hasCachedData } = readCachedConfigs(
|
||||
PUBLIC_BRAND_CONFIG_KEYS,
|
||||
);
|
||||
const {
|
||||
cachedConfigs: publicCachedConfigs,
|
||||
hasCachedData: hasPublicCachedData,
|
||||
} = readCachedConfigs(PUBLIC_BRAND_CONFIG_KEYS);
|
||||
|
||||
if (!isLoggedIn()) {
|
||||
const publicConfigs = await fetchPublicBrandConfigs();
|
||||
|
||||
if (Object.keys(publicConfigs).length > 0) {
|
||||
return { ...cachedConfigs, ...publicConfigs };
|
||||
return { ...publicCachedConfigs, ...publicConfigs };
|
||||
}
|
||||
|
||||
return cachedConfigs;
|
||||
return publicCachedConfigs;
|
||||
}
|
||||
|
||||
const cachedConfigs = readAllCachedSafeConfigs();
|
||||
const hasCachedData = Object.keys(cachedConfigs).length > 0;
|
||||
|
||||
// 从API获取最新配置
|
||||
try {
|
||||
const response = await getConfigs();
|
||||
@@ -249,14 +293,20 @@ export const getCachedConfigs = async (): Promise<Record<string, string>> => {
|
||||
return cachedConfigs;
|
||||
}
|
||||
|
||||
return await fetchPublicBrandConfigs();
|
||||
const publicConfigs = await fetchPublicBrandConfigs();
|
||||
|
||||
return { ...publicCachedConfigs, ...publicConfigs };
|
||||
} catch {
|
||||
// API失败时返回缓存的数据
|
||||
if (hasCachedData) {
|
||||
return cachedConfigs;
|
||||
}
|
||||
|
||||
return await fetchPublicBrandConfigs();
|
||||
const publicConfigs = await fetchPublicBrandConfigs();
|
||||
|
||||
return hasPublicCachedData
|
||||
? { ...publicCachedConfigs, ...publicConfigs }
|
||||
: publicConfigs;
|
||||
}
|
||||
};
|
||||
|
||||
@@ -344,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;
|
||||
@@ -357,16 +421,32 @@ 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.is_commercial = resolvedConfigMap.is_commercial === "true";
|
||||
siteConfig.hide_footer_brand = resolvedConfigMap.hide_footer_brand === "true";
|
||||
siteConfig.app_bg_image_light = appBgImageLight;
|
||||
siteConfig.app_bg_image_dark = appBgImageDark;
|
||||
if (
|
||||
Object.prototype.hasOwnProperty.call(resolvedConfigMap, "is_commercial")
|
||||
) {
|
||||
siteConfig.is_commercial = resolvedCommercial;
|
||||
}
|
||||
if (
|
||||
Object.prototype.hasOwnProperty.call(resolvedConfigMap, "hide_footer_brand")
|
||||
) {
|
||||
siteConfig.hide_footer_brand =
|
||||
resolvedCommercial && resolvedConfigMap.hide_footer_brand === "true";
|
||||
} else if (!resolvedCommercial) {
|
||||
siteConfig.hide_footer_brand = false;
|
||||
}
|
||||
|
||||
if (typeof document !== "undefined") {
|
||||
document.title = siteConfig.name;
|
||||
|
||||
@@ -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,9 +265,15 @@ 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",
|
||||
"is_commercial",
|
||||
"license_expiry",
|
||||
"hide_footer_brand",
|
||||
];
|
||||
const initialConfigs: Record<string, string> = {};
|
||||
|
||||
@@ -309,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: "",
|
||||
@@ -686,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);
|
||||
@@ -797,6 +816,7 @@ export default function ConfigPage() {
|
||||
|
||||
const handleBgImageUpload = async (
|
||||
e: React.ChangeEvent<HTMLInputElement>,
|
||||
key: BgImageKey,
|
||||
) => {
|
||||
const file = e.target.files?.[0];
|
||||
|
||||
@@ -808,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();
|
||||
@@ -858,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:") ||
|
||||
@@ -879,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>
|
||||
@@ -927,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}
|
||||
/>
|
||||
@@ -1115,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)) {
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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);
|
||||
};
|
||||
|
||||
// 显示地址列表弹窗
|
||||
|
||||
@@ -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) => {
|
||||
|
||||
@@ -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 (
|
||||
|
||||
@@ -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="留空或清除表示永久有效"
|
||||
|
||||
@@ -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);
|
||||
}}
|
||||
/>
|
||||
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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;
|
||||
}
|
||||
Reference in New Issue
Block a user