Compare commits

...

16 Commits

Author SHA1 Message Date
sagit 2e845030de fix: build frontend assets on native platform (#553) 2026-09-01 16:27:00 +08:00
sagit 2269f2e2d5 fix: include pnpm config in frontend image build (#552) 2026-09-01 15:48:54 +08:00
sagit da7bef88f1 fix: harden panel updates and nftables deployment (#551)
Fix panel update deployment discovery and rollback safety, add nftables compatibility and atomic replacement, and restore reproducible frontend CI installs.
2026-09-01 15:37:33 +08:00
sagit f26014579b fix: stabilize federated agent runtime updates (#550) 2026-08-25 17:04:04 +08:00
sagit 5041c722c9 fix: harden license lifecycle (#547) 2026-08-13 09:25:05 +08:00
sagit c56798e991 fix: accept existing license machine binding (#546) 2026-08-12 14:12:46 +08:00
sagit 40e96f3592 fix: preserve active per-IP traffic limiters (#545) 2026-08-12 11:31:03 +08:00
sagit 0e24b53a5b fix: preserve license state across panel upgrades (#544) 2026-08-11 17:04:20 +08:00
sagit a820c49c94 fix(agent): harden Alpine OpenRC installation (#543)
Ensure Alpine installs use OpenRC without invoking systemd cleanup paths, and install missing download dependencies.
2026-08-08 11:45:39 +08:00
sagit 538e64ffc0 Fix modal scroll position jumps (#542)
Preserve page scroll positions while Radix modals acquire focus and forward the dialog overlay ref correctly.
2026-08-07 22:38:02 +08:00
sagit 0b23d6f7d7 fix(agent): retire replaced tunnel sessions (#541) 2026-08-07 14:35:30 +08:00
sagit 9e6f80019d feat(monitor): show backup nodes in topology (#538) 2026-08-03 14:14:01 +08:00
sagit a8fd01d4d8 fix(node): allow IPv6-only addresses (#537) 2026-08-03 11:05:04 +08:00
sagit cbe2fc492e feat(monitor): show backup tunnel latencies (#535)
Closes #508
2026-08-03 10:15:22 +08:00
sagit ae370382d3 feat(agent): support Alpine installation (#534)
Add Alpine bootstrap and OpenRC lifecycle support to the agent installer.

Closes #527
2026-07-31 17:02:52 +08:00
sagit e112d81697 fix: bound and configure tunnel quality probes (#533)
Closes #528 and #532.
2026-07-31 15:37:26 +08:00
78 changed files with 4723 additions and 446 deletions
+1 -1
View File
@@ -22,7 +22,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
+1 -1
View File
@@ -229,6 +229,7 @@ jobs:
docker buildx build \
--platform linux/amd64,linux/arm64 \
--build-arg KEYGEN_ACCOUNT_ID=${{ secrets.KEYGEN_ACCOUNT_ID }} \
--push \
-t ${{ env.REGISTRY }}/${OWNER}/flux-panel-backend:latest \
-t ${{ env.REGISTRY }}/${OWNER}/flux-panel-backend:${VERSION} \
@@ -431,4 +432,3 @@ jobs:
gh release upload "${VERSION}" ./artifacts/gost-arm64.sha256 --clobber
echo "✅ GOST 二进制文件更新完成"
+10 -1
View File
@@ -64,6 +64,14 @@ curl -L https://raw.githubusercontent.com/Sagit-chu/flux-panel/main/panel_instal
curl -L https://raw.githubusercontent.com/Sagit-chu/flux-panel/main/install.sh -o install.sh && chmod +x install.sh && ./install.sh
```
Alpine Linux 最小化安装若未包含 `curl`,可使用系统自带的 `wget` 下载:
```bash
wget -O install.sh https://raw.githubusercontent.com/Sagit-chu/flux-panel/main/install.sh && chmod +x install.sh && ./install.sh
```
脚本会在 Alpine 上自动安装 Bash、`curl` 和 CA 证书,并使用 OpenRC 注册、启动和管理 `flux_agent` 服务;其他受支持的 Linux 发行版继续使用 systemd。
**安装过程中会提示输入:**
- **服务器地址**: 面板端的通信地址(通常是 `http://<面板IP>:<后端端口>`,例如 `http://1.2.3.4:6365`)。
- **密钥**: 刚才在面板中获取的节点密钥。
@@ -77,7 +85,8 @@ curl -L https://raw.githubusercontent.com/Sagit-chu/flux-panel/main/install.sh -
### 3. 验证安装
安装完成后,服务会自动启动。
- 查看状态: `systemctl status flux_agent`
- systemd 查看状态: `systemctl status flux_agent`
- Alpine/OpenRC 查看状态: `rc-service flux_agent status`
- 回到面板 **节点管理** 页面,该节点状态应显示为 **在线**。
---
+2 -1
View File
@@ -7,7 +7,8 @@ RUN go mod download
COPY . .
ARG TARGETOS
ARG TARGETARCH
RUN CGO_ENABLED=0 GOOS=${TARGETOS:-linux} env ${TARGETARCH:+GOARCH=${TARGETARCH}} go build -o /out/paneld ./cmd/paneld
ARG KEYGEN_ACCOUNT_ID
RUN CGO_ENABLED=0 GOOS=${TARGETOS:-linux} env ${TARGETARCH:+GOARCH=${TARGETARCH}} go build -ldflags="-X 'go-backend/internal/license.AccountID=${KEYGEN_ACCOUNT_ID}'" -o /out/paneld ./cmd/paneld
FROM docker:27-cli AS dockercli
@@ -29,6 +29,10 @@ func (h *Handler) getPublicConfigByName(w http.ResponseWriter, r *http.Request)
response.WriteJSON(w, response.Err(403, "禁止访问敏感配置"))
return
}
if value, gated := h.unlicensedPublicBrandValue(configName); gated {
response.WriteJSON(w, response.OK(map[string]string{"name": configName, "value": value}))
return
}
cfg, err := h.repo.GetConfigByName(configName)
if err != nil {
@@ -42,3 +46,22 @@ func (h *Handler) getPublicConfigByName(w http.ResponseWriter, r *http.Request)
response.WriteJSON(w, response.OK(cfg))
}
func (h *Handler) unlicensedPublicBrandValue(configName string) (string, bool) {
switch configName {
case "app_name", "app_logo", "app_favicon", "hide_footer_brand":
default:
return "", false
}
isCommercial, _ := h.repo.GetViteConfigValue("is_commercial")
if isCommercial == "true" {
return "", false
}
if configName == "app_name" {
return "FLVX", true
}
if configName == "hide_footer_brand" {
return "false", true
}
return "", true
}
@@ -30,6 +30,40 @@ func TestPublicConfigGetAllowsBrandKeys(t *testing.T) {
assertHandlerCode(t, resp, 0)
}
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) {
router, _ := setupConfigAccessTestRouter(t)
@@ -82,6 +116,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 +238,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 +249,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 +266,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)
}
}
@@ -59,6 +59,7 @@ type diagnosisWorkItem struct {
type diagnosisExecOptions struct {
commandTimeout time.Duration
pingTimeoutMS int
pingCount int
timeoutMessage string
}
@@ -1596,10 +1597,14 @@ func (h *Handler) tcpPingViaNode(nodeID int64, ip string, port int, options diag
if options.pingTimeoutMS <= 0 {
options.pingTimeoutMS = int(diagnosisCommandTimeout / time.Millisecond)
}
pingCount := options.pingCount
if pingCount <= 0 {
pingCount = 4
}
res, err := h.sendNodeCommandWithTimeout(nodeID, "TcpPing", map[string]interface{}{
"ip": ip,
"port": port,
"count": 4,
"count": pingCount,
"timeout": options.pingTimeoutMS,
}, options.commandTimeout, false, false)
if err != nil {
@@ -1626,12 +1631,16 @@ func (h *Handler) tcpPingViaRemoteNode(node *nodeRecord, ip string, port int, op
if options.pingTimeoutMS <= 0 {
options.pingTimeoutMS = int(diagnosisCommandTimeout / time.Millisecond)
}
pingCount := options.pingCount
if pingCount <= 0 {
pingCount = 4
}
fc := client.NewFederationClientWithTimeout(options.commandTimeout)
return fc.Diagnose(remoteURL, remoteToken, h.federationLocalDomain(), client.RuntimeDiagnoseRequest{
IP: strings.TrimSpace(ip),
Port: port,
Count: 4,
Count: pingCount,
Timeout: options.pingTimeoutMS,
Protocol: "tcp",
})
@@ -659,9 +659,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")
}
}
+62 -47
View File
@@ -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))
}
@@ -878,10 +886,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 +922,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 +992,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("需要商业版授权"))
@@ -1015,6 +1012,7 @@ func (h *Handler) updateConfigs(w http.ResponseWriter, r *http.Request) {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
h.notifyTunnelQualityConfigChanged(key)
}
response.WriteJSON(w, response.OKEmpty())
@@ -1040,6 +1038,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" {
@@ -1062,6 +1064,7 @@ func (h *Handler) updateSingleConfig(w http.ResponseWriter, r *http.Request) {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
h.notifyTunnelQualityConfigChanged(name)
response.WriteJSON(w, response.OKEmpty())
}
@@ -1110,11 +1113,23 @@ func normalizeAndValidateConfigValue(key, value string) (string, error) {
}
case monitoring.ConfigMonitorRetentionDays:
return monitoring.NormalizeMonitoringRetentionDays(value)
case monitoring.ConfigTunnelQualityProbeIntervalSec:
return monitoring.NormalizeTunnelQualityProbeIntervalSeconds(value)
default:
return value, nil
}
}
func (h *Handler) notifyTunnelQualityConfigChanged(key string) {
if h == nil || h.qualityProber == nil {
return
}
switch strings.TrimSpace(key) {
case monitorTunnelQualityEnabledConfigKey, monitoring.ConfigTunnelQualityProbeIntervalSec:
h.qualityProber.NotifyConfigChanged()
}
}
func (h *Handler) isTunnelQualityMonitoringEnabled() bool {
if h == nil || h.repo == nil {
return true
+20 -11
View File
@@ -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 })
}
@@ -3,6 +3,7 @@ package handler
import (
"fmt"
"net"
"net/netip"
"strings"
)
@@ -75,6 +76,12 @@ func IsValidNodeAddress(addr string) error {
if strings.ContainsAny(addr, "/?") {
return fmt.Errorf("address must not contain path or query parameters")
}
// A bare IPv6 literal contains multiple colons, so net.SplitHostPort treats
// it as a malformed host:port pair. Accept IP literals before attempting
// host:port parsing; netip also handles scoped IPv6 addresses.
if _, err := netip.ParseAddr(addr); err == nil {
return nil
}
_, _, err := net.SplitHostPort(addr)
if err != nil {
@@ -0,0 +1,46 @@
package handler
import "testing"
func TestIssue515IsValidNodeAddressAcceptsBareIPv6(t *testing.T) {
for _, addr := range []string{
"2001:db8::1",
"::1",
"fe80::1%eth0",
} {
t.Run(addr, func(t *testing.T) {
if err := IsValidNodeAddress(addr); err != nil {
t.Fatalf("expected bare IPv6 address %q to be accepted: %v", addr, err)
}
})
}
}
func TestIsValidNodeAddressKeepsExistingAddressForms(t *testing.T) {
for _, addr := range []string{
"203.0.113.10",
"node.example.com",
"node.example.com:6365",
"[2001:db8::1]:6365",
} {
t.Run(addr, func(t *testing.T) {
if err := IsValidNodeAddress(addr); err != nil {
t.Fatalf("expected node address %q to be accepted: %v", addr, err)
}
})
}
}
func TestIsValidNodeAddressRejectsURLComponents(t *testing.T) {
for _, addr := range []string{
"https://node.example.com",
"node.example.com/path",
"node.example.com?transport=tcp",
} {
t.Run(addr, func(t *testing.T) {
if err := IsValidNodeAddress(addr); err == nil {
t.Fatalf("expected node address %q to be rejected", addr)
}
})
}
}
@@ -152,7 +152,11 @@ func evaluateBestExitOwner(owner chainNodeRecord, exits []chainNodeRecord, nodes
ownerNode := nodes[owner.NodeID]
for _, exit := range exits {
exitNode := nodes[exit.NodeID]
if exitNode == nil {
if !isTunnelProbeNodeOnline(ownerNode) {
scores = append(scores, failedBestExitCandidate(owner.NodeID, exit, "owner node offline"))
continue
}
if !isTunnelProbeNodeOnline(exitNode) {
scores = append(scores, failedBestExitCandidate(owner.NodeID, exit, "exit node unavailable"))
continue
}
@@ -358,9 +358,9 @@ func TestEvaluateBestExitOwnerScoresAllCandidates(t *testing.T) {
{NodeID: 31, NodeName: "exit-b", Port: 30031},
}
nodes := map[int64]*nodeRecord{
10: {ID: 10, ServerIP: "10.0.0.10", ServerIPv4: "10.0.0.10", TCPListenAddr: "[::]"},
30: {ID: 30, ServerIP: "10.0.0.30", ServerIPv4: "10.0.0.30", TCPListenAddr: "[::]"},
31: {ID: 31, ServerIP: "10.0.0.31", ServerIPv4: "10.0.0.31", TCPListenAddr: "[::]"},
10: {ID: 10, Status: 1, ServerIP: "10.0.0.10", ServerIPv4: "10.0.0.10", TCPListenAddr: "[::]"},
30: {ID: 30, Status: 1, ServerIP: "10.0.0.30", ServerIPv4: "10.0.0.30", TCPListenAddr: "[::]"},
31: {ID: 31, Status: 1, ServerIP: "10.0.0.31", ServerIPv4: "10.0.0.31", TCPListenAddr: "[::]"},
}
pinger := func(nodeID int64, ip string, port int, _ diagnosisExecOptions) (float64, float64, error) {
switch {
@@ -387,12 +387,30 @@ func TestEvaluateBestExitOwnerScoresAllCandidates(t *testing.T) {
}
}
func TestEvaluateBestExitOwnerSkipsOfflineCandidate(t *testing.T) {
owner := chainNodeRecord{NodeID: 10, NodeName: "entry"}
exits := []chainNodeRecord{{NodeID: 30, NodeName: "exit-a", Port: 30030}}
nodes := map[int64]*nodeRecord{
10: {ID: 10, Status: 1, ServerIP: "10.0.0.10", ServerIPv4: "10.0.0.10"},
30: {ID: 30, Status: 0, ServerIP: "10.0.0.30", ServerIPv4: "10.0.0.30"},
}
ping := func(nodeID int64, ip string, port int, options diagnosisExecOptions) (float64, float64, error) {
t.Fatalf("offline best-exit candidate should not be probed: node=%d target=%s:%d", nodeID, ip, port)
return 0, 100, nil
}
scores := evaluateBestExitOwner(owner, exits, nodes, "", diagnosisExecOptions{}, defaultTunnelProbeTarget(), ping)
if len(scores) != 1 || scores[0].Success {
t.Fatalf("expected one failed offline candidate, got %+v", scores)
}
}
func TestEvaluateBestExitOwnerUsesConfiguredPublicProbeTarget(t *testing.T) {
owner := chainNodeRecord{NodeID: 10, NodeName: "entry-a"}
exits := []chainNodeRecord{{NodeID: 30, NodeName: "exit-a", Port: 30001}}
nodes := map[int64]*nodeRecord{
10: {ID: 10, Name: "entry-a", ServerIP: "10.0.0.10", ServerIPv4: "10.0.0.10"},
30: {ID: 30, Name: "exit-a", ServerIP: "10.0.0.30", ServerIPv4: "10.0.0.30"},
10: {ID: 10, Name: "entry-a", Status: 1, ServerIP: "10.0.0.10", ServerIPv4: "10.0.0.10"},
30: {ID: 30, Name: "exit-a", Status: 1, ServerIP: "10.0.0.30", ServerIPv4: "10.0.0.30"},
}
target := tunnelProbeTarget{Host: "speed.example.com", Port: 8443}
var calls []string
@@ -419,8 +437,8 @@ func TestEvaluateBestExitOwnerMarksCandidateFailedWhenOwnerToExitFails(t *testin
owner := chainNodeRecord{NodeID: 10, NodeName: "entry"}
exits := []chainNodeRecord{{NodeID: 30, NodeName: "exit-a", Port: 30030}}
nodes := map[int64]*nodeRecord{
10: {ID: 10, ServerIP: "10.0.0.10", ServerIPv4: "10.0.0.10", TCPListenAddr: "[::]"},
30: {ID: 30, ServerIP: "10.0.0.30", ServerIPv4: "10.0.0.30", TCPListenAddr: "[::]"},
10: {ID: 10, Status: 1, ServerIP: "10.0.0.10", ServerIPv4: "10.0.0.10", TCPListenAddr: "[::]"},
30: {ID: 30, Status: 1, ServerIP: "10.0.0.30", ServerIPv4: "10.0.0.30", TCPListenAddr: "[::]"},
}
pinger := func(nodeID int64, ip string, port int, _ diagnosisExecOptions) (float64, float64, error) {
return 0, 100, errBestExitProbeForTest
@@ -436,8 +454,8 @@ func TestEvaluateBestExitOwnerMarksCandidateFailedWhenTargetResolutionFails(t *t
owner := chainNodeRecord{NodeID: 10, NodeName: "entry"}
exits := []chainNodeRecord{{NodeID: 30, NodeName: "exit-v6", Port: 30030}}
nodes := map[int64]*nodeRecord{
10: {ID: 10, Name: "entry", ServerIP: "10.0.0.10", ServerIPv4: "10.0.0.10", TCPListenAddr: "[::]"},
30: {ID: 30, Name: "exit-v6", ServerIP: "2001:db8::30", ServerIPv6: "2001:db8::30", TCPListenAddr: "[::]"},
10: {ID: 10, Name: "entry", Status: 1, ServerIP: "10.0.0.10", ServerIPv4: "10.0.0.10", TCPListenAddr: "[::]"},
30: {ID: 30, Name: "exit-v6", Status: 1, ServerIP: "2001:db8::30", ServerIPv6: "2001:db8::30", TCPListenAddr: "[::]"},
}
pinger := func(nodeID int64, ip string, port int, _ diagnosisExecOptions) (float64, float64, error) {
t.Fatalf("ping should not be called when target resolution fails: node=%d ip=%s port=%d", nodeID, ip, port)
@@ -3,6 +3,8 @@ package handler
import (
"context"
"encoding/json"
"errors"
"fmt"
"log"
"sync"
"sync/atomic"
@@ -13,7 +15,6 @@ import (
)
const (
tunnelQualityProbeInterval = 1 * time.Second
tunnelQualityProbeTimeout = 8 * time.Second
tunnelQualityPingTimeoutMs = 5000
tunnelQualityPruneInterval = 10 * time.Minute
@@ -31,6 +32,26 @@ type TunnelQualityHop struct {
TargetPort int `json:"targetPort,omitempty"`
}
type TunnelQualityCandidateHop struct {
TunnelQualityHop
FromRole string `json:"fromRole"`
ToRole string `json:"toRole"`
HopIndex int `json:"hopIndex"`
Selected bool `json:"selected"`
ErrorMessage string `json:"errorMessage,omitempty"`
}
type tunnelQualityChainDetails struct {
PrimaryPath []TunnelQualityHop `json:"primaryPath,omitempty"`
CandidateHops []TunnelQualityCandidateHop `json:"candidateHops,omitempty"`
}
type tunnelQualityCandidateGroup struct {
role string
roleIndex int
nodes []chainNodeRecord
}
// tunnelQualitySnapshot is the in-memory latest probe result for a tunnel.
type tunnelQualitySnapshot struct {
TunnelID int64 `json:"tunnelId"`
@@ -56,7 +77,7 @@ type tunnelQualityProber struct {
cache sync.Map // tunnelID (int64) → *tunnelQualitySnapshot
ctx context.Context
cancel context.CancelFunc
interval time.Duration
wake chan struct{}
lastPrune int64
probing int32 // atomic flag: 1 = probeAll running, 0 = idle
probeNode bestExitProbeFunc
@@ -65,8 +86,8 @@ type tunnelQualityProber struct {
// newTunnelQualityProber creates a new prober (not yet running).
func newTunnelQualityProber(h *Handler) *tunnelQualityProber {
return &tunnelQualityProber{
handler: h,
interval: tunnelQualityProbeInterval,
handler: h,
wake: make(chan struct{}, 1),
}
}
@@ -86,6 +107,16 @@ func (p *tunnelQualityProber) Stop() {
p.cancel()
}
func (p *tunnelQualityProber) NotifyConfigChanged() {
if p == nil || p.wake == nil {
return
}
select {
case p.wake <- struct{}{}:
default:
}
}
// GetAll returns all cached quality snapshots (latest per tunnel).
func (p *tunnelQualityProber) GetAll() []tunnelQualitySnapshot {
var items []tunnelQualitySnapshot
@@ -109,20 +140,44 @@ func (p *tunnelQualityProber) loop() {
// Run once immediately
p.probeAll()
ticker := time.NewTicker(p.interval)
defer ticker.Stop()
for {
timer := time.NewTimer(p.probeInterval())
select {
case <-p.ctx.Done():
stopAndDrainTunnelQualityTimer(timer)
return
case <-ticker.C:
case <-p.wake:
stopAndDrainTunnelQualityTimer(timer)
continue
case <-timer.C:
p.probeAll()
p.maybePrune()
}
}
}
func stopAndDrainTunnelQualityTimer(timer *time.Timer) {
if timer == nil || timer.Stop() {
return
}
select {
case <-timer.C:
default:
}
}
func (p *tunnelQualityProber) probeInterval() time.Duration {
if p == nil || p.handler == nil || p.handler.repo == nil {
return time.Duration(monitoring.DefaultTunnelQualityProbeIntervalSec) * time.Second
}
cfg, err := p.handler.repo.GetConfigsByNames([]string{monitoring.ConfigTunnelQualityProbeIntervalSec})
if err != nil {
return time.Duration(monitoring.DefaultTunnelQualityProbeIntervalSec) * time.Second
}
seconds := monitoring.TunnelQualityProbeIntervalSecondsFromConfigMap(cfg)
return time.Duration(seconds) * time.Second
}
func (p *tunnelQualityProber) isEnabled() bool {
if p == nil || p.handler == nil {
return true
@@ -246,15 +301,28 @@ func (p *tunnelQualityProber) probeTunnel(tunnelID int64) {
options := diagnosisExecOptions{
commandTimeout: tunnelQualityProbeTimeout,
pingTimeoutMS: tunnelQualityPingTimeoutMs,
pingCount: 1,
timeoutMessage: "探测超时",
}
p.probeBestExitOwners(tunnelID, inNodes, midNodesGrouped, outNodes, ipPreference, options, probeTarget)
roundPinger := newBestExitRoundPinger(p.pingNode)
p.probeBestExitOwners(tunnelID, inNodes, midNodesGrouped, outNodes, ipPreference, options, probeTarget, roundPinger)
entry, _, entryOnline := p.firstOnlineChainNode(inNodes)
exit, _, exitOnline := p.firstOnlineChainNode(outNodes)
selectedNodeIDs := make(map[string]int64, 2+len(midNodesGrouped))
if entryOnline {
selectedNodeIDs[tunnelQualityGroupKey("entry", 0)] = entry.NodeID
}
if exitOnline {
selectedNodeIDs[tunnelQualityGroupKey("exit", 0)] = exit.NodeID
}
var primaryHops []TunnelQualityHop
switch tunnel.Type {
case 1:
// Port forwarding: entry → public probe target only.
if len(inNodes) > 0 {
lat, loss, err := p.pingNode(inNodes[0].NodeID, probeTarget.Host, probeTarget.Port, options)
if entryOnline {
lat, loss, err := roundPinger(entry.NodeID, probeTarget.Host, probeTarget.Port, options)
if err == nil {
snap.ExitToBingLatency = lat
snap.ExitToBingLoss = loss
@@ -262,24 +330,42 @@ func (p *tunnelQualityProber) probeTunnel(tunnelID int64) {
} else {
snap.ErrorMessage = err.Error()
}
} else {
snap.ErrorMessage = "入口节点均不在线"
}
case 2:
// Tunnel forwarding: entry → exit + exit → Bing
probeOK := true
if len(inNodes) > 0 && len(outNodes) > 0 {
var hops []TunnelQualityHop
if !entryOnline {
probeOK = false
snap.ErrorMessage = "入口节点均不在线"
snap.EntryToExitLatency = -1
snap.EntryToExitLoss = 100
} else if !exitOnline {
probeOK = false
snap.ErrorMessage = "出口节点均不在线"
snap.EntryToExitLatency = -1
snap.EntryToExitLoss = 100
} else {
var totalLat float64
remainingSuccessProb := 1.0
nodesInPath := make([]chainNodeRecord, 0, 2+len(midNodesGrouped))
nodesInPath = append(nodesInPath, inNodes[0])
for _, midGroup := range midNodesGrouped {
if len(midGroup) > 0 {
nodesInPath = append(nodesInPath, midGroup[0])
nodesInPath = append(nodesInPath, entry)
for midIndex, midGroup := range midNodesGrouped {
mid, _, online := p.firstOnlineChainNode(midGroup)
if !online {
probeOK = false
snap.ErrorMessage = "中间节点组均不在线"
break
}
nodesInPath = append(nodesInPath, mid)
selectedNodeIDs[tunnelQualityGroupKey("middle", midIndex)] = mid.NodeID
}
if probeOK {
nodesInPath = append(nodesInPath, exit)
}
nodesInPath = append(nodesInPath, outNodes[0])
for i := 0; i < len(nodesInPath)-1; i++ {
source := nodesInPath[i]
@@ -293,12 +379,12 @@ func (p *tunnelQualityProber) probeTunnel(tunnelID int64) {
}
targetNode, nodeErr := h.getNodeRecord(target.NodeID)
if nodeErr != nil || targetNode == nil {
if nodeErr != nil || !isTunnelProbeNodeOnline(targetNode) {
snap.ErrorMessage = "节点 " + target.NodeName + " 不可用"
probeOK = false
hop.Latency = -1
hop.Loss = 100
hops = append(hops, hop)
primaryHops = append(primaryHops, hop)
break
}
@@ -309,25 +395,25 @@ func (p *tunnelQualityProber) probeTunnel(tunnelID int64) {
probeOK = false
hop.Latency = -1
hop.Loss = 100
hops = append(hops, hop)
primaryHops = append(primaryHops, hop)
break
}
hop.TargetIP = targetIP
hop.TargetPort = targetPort
lat, loss, err := p.pingNode(source.NodeID, targetIP, targetPort, options)
lat, loss, err := roundPinger(source.NodeID, targetIP, targetPort, options)
if err == nil {
hop.Latency = lat
hop.Loss = loss
totalLat += lat
remainingSuccessProb *= (1.0 - loss/100.0)
hops = append(hops, hop)
primaryHops = append(primaryHops, hop)
} else {
probeOK = false
hop.Latency = -1
hop.Loss = 100
hops = append(hops, hop)
primaryHops = append(primaryHops, hop)
if snap.ErrorMessage == "" {
snap.ErrorMessage = err.Error()
}
@@ -342,17 +428,11 @@ func (p *tunnelQualityProber) probeTunnel(tunnelID int64) {
snap.EntryToExitLatency = -1
snap.EntryToExitLoss = 100
}
if len(hops) > 0 {
if b, err := json.Marshal(hops); err == nil {
snap.ChainDetails = string(b)
}
}
}
// Exit → Bing
if len(outNodes) > 0 {
lat, loss, err := p.pingNode(outNodes[0].NodeID, probeTarget.Host, probeTarget.Port, options)
if exitOnline {
lat, loss, err := roundPinger(exit.NodeID, probeTarget.Host, probeTarget.Port, options)
if err == nil {
snap.ExitToBingLatency = lat
snap.ExitToBingLoss = loss
@@ -367,8 +447,8 @@ func (p *tunnelQualityProber) probeTunnel(tunnelID int64) {
snap.Success = probeOK
default:
// Unknown type: entry → public probe target.
if len(inNodes) > 0 {
lat, loss, err := p.pingNode(inNodes[0].NodeID, probeTarget.Host, probeTarget.Port, options)
if entryOnline {
lat, loss, err := roundPinger(entry.NodeID, probeTarget.Host, probeTarget.Port, options)
if err == nil {
snap.ExitToBingLatency = lat
snap.ExitToBingLoss = loss
@@ -376,13 +456,215 @@ func (p *tunnelQualityProber) probeTunnel(tunnelID int64) {
} else {
snap.ErrorMessage = err.Error()
}
} else {
snap.ErrorMessage = "入口节点均不在线"
}
}
candidateHops := p.probeTunnelCandidateHops(
tunnel.Type,
inNodes,
midNodesGrouped,
outNodes,
selectedNodeIDs,
ipPreference,
options,
probeTarget,
roundPinger,
)
if len(primaryHops) > 0 || len(candidateHops) > 0 {
details := tunnelQualityChainDetails{
PrimaryPath: primaryHops,
CandidateHops: candidateHops,
}
if b, err := json.Marshal(details); err == nil {
snap.ChainDetails = string(b)
}
}
p.storeResult(snap)
}
func (p *tunnelQualityProber) probeBestExitOwners(tunnelID int64, inNodes []chainNodeRecord, chainHops [][]chainNodeRecord, outNodes []chainNodeRecord, ipPreference string, options diagnosisExecOptions, probeTarget tunnelProbeTarget) {
func tunnelQualityGroupKey(role string, index int) string {
return fmt.Sprintf("%s:%d", role, index)
}
func (p *tunnelQualityProber) probeTunnelCandidateHops(
tunnelType int,
inNodes []chainNodeRecord,
chainHops [][]chainNodeRecord,
outNodes []chainNodeRecord,
selectedNodeIDs map[string]int64,
ipPreference string,
options diagnosisExecOptions,
probeTarget tunnelProbeTarget,
ping bestExitProbeFunc,
) []TunnelQualityCandidateHop {
if p == nil || p.handler == nil || ping == nil {
return nil
}
if tunnelType != 2 {
return p.probePublicTargetCandidates("entry", 0, inNodes, selectedNodeIDs, options, probeTarget, ping)
}
groups := make([]tunnelQualityCandidateGroup, 0, 2+len(chainHops))
groups = append(groups, tunnelQualityCandidateGroup{role: "entry", roleIndex: 0, nodes: inNodes})
for i, hop := range chainHops {
groups = append(groups, tunnelQualityCandidateGroup{role: "middle", roleIndex: i, nodes: hop})
}
groups = append(groups, tunnelQualityCandidateGroup{role: "exit", roleIndex: 0, nodes: outNodes})
var items []TunnelQualityCandidateHop
for i := 0; i < len(groups)-1; i++ {
items = append(items, p.probeCandidateGroupLinks(
groups[i],
groups[i+1],
i,
selectedNodeIDs,
ipPreference,
options,
ping,
)...)
}
items = append(items, p.probePublicTargetCandidates(
"exit",
0,
outNodes,
selectedNodeIDs,
options,
probeTarget,
ping,
)...)
return items
}
func (p *tunnelQualityProber) probeCandidateGroupLinks(
fromGroup tunnelQualityCandidateGroup,
toGroup tunnelQualityCandidateGroup,
hopIndex int,
selectedNodeIDs map[string]int64,
ipPreference string,
options diagnosisExecOptions,
ping bestExitProbeFunc,
) []TunnelQualityCandidateHop {
items := make([]TunnelQualityCandidateHop, 0, len(fromGroup.nodes)*len(toGroup.nodes))
for _, source := range fromGroup.nodes {
for _, target := range toGroup.nodes {
item := TunnelQualityCandidateHop{
TunnelQualityHop: TunnelQualityHop{
FromNodeID: source.NodeID,
FromNodeName: source.NodeName,
ToNodeID: target.NodeID,
ToNodeName: target.NodeName,
Latency: -1,
Loss: 100,
},
FromRole: fromGroup.role,
ToRole: toGroup.role,
HopIndex: hopIndex,
Selected: selectedNodeIDs[tunnelQualityGroupKey(fromGroup.role, fromGroup.roleIndex)] == source.NodeID &&
selectedNodeIDs[tunnelQualityGroupKey(toGroup.role, toGroup.roleIndex)] == target.NodeID,
}
sourceNode, sourceErr := p.handler.getNodeRecord(source.NodeID)
if sourceErr != nil || !isTunnelProbeNodeOnline(sourceNode) {
item.ErrorMessage = "来源节点不在线"
items = append(items, item)
continue
}
targetNode, targetErr := p.handler.getNodeRecord(target.NodeID)
if targetErr != nil || !isTunnelProbeNodeOnline(targetNode) {
item.ErrorMessage = "目标节点不在线"
items = append(items, item)
continue
}
targetIP, targetPort, resolveErr := resolveChainProbeTarget(sourceNode, targetNode, target.Port, ipPreference, target.ConnectIP)
if resolveErr != nil {
item.ErrorMessage = resolveErr.Error()
items = append(items, item)
continue
}
item.TargetIP = targetIP
item.TargetPort = targetPort
latency, loss, probeErr := ping(source.NodeID, targetIP, targetPort, options)
if probeErr != nil {
item.ErrorMessage = probeErr.Error()
items = append(items, item)
continue
}
item.Latency = latency
item.Loss = loss
items = append(items, item)
}
}
return items
}
func (p *tunnelQualityProber) probePublicTargetCandidates(
fromRole string,
fromIndex int,
nodes []chainNodeRecord,
selectedNodeIDs map[string]int64,
options diagnosisExecOptions,
probeTarget tunnelProbeTarget,
ping bestExitProbeFunc,
) []TunnelQualityCandidateHop {
items := make([]TunnelQualityCandidateHop, 0, len(nodes))
for _, source := range nodes {
item := TunnelQualityCandidateHop{
TunnelQualityHop: TunnelQualityHop{
FromNodeID: source.NodeID,
FromNodeName: source.NodeName,
ToNodeName: formatTunnelProbeTarget(probeTarget),
Latency: -1,
Loss: 100,
TargetIP: probeTarget.Host,
TargetPort: probeTarget.Port,
},
FromRole: fromRole,
ToRole: "target",
HopIndex: fromIndex,
Selected: selectedNodeIDs[tunnelQualityGroupKey(fromRole, fromIndex)] == source.NodeID,
}
sourceNode, sourceErr := p.handler.getNodeRecord(source.NodeID)
if sourceErr != nil || !isTunnelProbeNodeOnline(sourceNode) {
item.ErrorMessage = "来源节点不在线"
items = append(items, item)
continue
}
latency, loss, probeErr := ping(source.NodeID, probeTarget.Host, probeTarget.Port, options)
if probeErr != nil {
item.ErrorMessage = probeErr.Error()
items = append(items, item)
continue
}
item.Latency = latency
item.Loss = loss
items = append(items, item)
}
return items
}
func isTunnelProbeNodeOnline(node *nodeRecord) bool {
return node != nil && (node.IsRemote == 1 || node.Status == 1)
}
func (p *tunnelQualityProber) firstOnlineChainNode(nodes []chainNodeRecord) (chainNodeRecord, *nodeRecord, bool) {
if p == nil || p.handler == nil {
return chainNodeRecord{}, nil, false
}
for _, candidate := range nodes {
node, err := p.handler.getNodeRecord(candidate.NodeID)
if err == nil && isTunnelProbeNodeOnline(node) {
return candidate, node, true
}
}
return chainNodeRecord{}, nil, false
}
func (p *tunnelQualityProber) probeBestExitOwners(tunnelID int64, inNodes []chainNodeRecord, chainHops [][]chainNodeRecord, outNodes []chainNodeRecord, ipPreference string, options diagnosisExecOptions, probeTarget tunnelProbeTarget, roundPinger bestExitProbeFunc) {
if p == nil || p.handler == nil || p.handler.bestExit == nil || len(outNodes) <= 1 {
return
}
@@ -404,9 +686,6 @@ func (p *tunnelQualityProber) probeBestExitOwners(tunnelID int64, inNodes []chai
nodeMap[exit.NodeID] = node
}
}
// This best-exit decision cache is per decision round; the display-oriented
// tunnel quality snapshot may still collect its own first-exit public probe.
roundPinger := newBestExitRoundPinger(p.pingNode)
for _, owner := range owners {
if nodeMap[owner.NodeID] == nil {
continue
@@ -444,6 +723,9 @@ func (p *tunnelQualityProber) tcpPingNode(nodeID int64, ip string, port int, opt
if nodeErr != nil {
return 0, 100, nodeErr
}
if !isTunnelProbeNodeOnline(node) {
return 0, 100, errors.New("节点不在线")
}
var pingData map[string]interface{}
var pingErr error
@@ -1,6 +1,7 @@
package handler
import (
"encoding/json"
"fmt"
"slices"
"testing"
@@ -26,6 +27,9 @@ func TestTunnelQualityProberUsesConfiguredProbeTarget(t *testing.T) {
p := newTunnelQualityProber(h)
var calls []string
p.probeNode = func(nodeID int64, ip string, port int, options diagnosisExecOptions) (float64, float64, error) {
if options.pingCount != 1 {
t.Fatalf("expected real-time quality probe count 1, got %d", options.pingCount)
}
calls = append(calls, fmt.Sprintf("%d|%s|%d", nodeID, ip, port))
return 10, 0, nil
}
@@ -46,6 +50,114 @@ func TestTunnelQualityProberUsesConfiguredProbeTarget(t *testing.T) {
}
}
func TestTunnelQualityProberSkipsAllOfflineExits(t *testing.T) {
h := setupProbeTargetTunnelHandler(t)
seedQualityForwardTunnel(t, h, 81, []int{0, 0, 0})
p := newTunnelQualityProber(h)
probeCalls := 0
p.probeNode = func(nodeID int64, ip string, port int, options diagnosisExecOptions) (float64, float64, error) {
probeCalls++
return 0, 100, fmt.Errorf("unexpected probe node=%d target=%s:%d", nodeID, ip, port)
}
p.probeTunnel(81)
if probeCalls != 0 {
t.Fatalf("expected no TCP probes when all exits are offline, got %d", probeCalls)
}
snaps := p.GetAll()
if len(snaps) != 1 {
t.Fatalf("expected one quality snapshot, got %+v", snaps)
}
if snaps[0].Success || snaps[0].ErrorMessage != "出口节点均不在线" {
t.Fatalf("expected offline exit snapshot, got %+v", snaps[0])
}
if snaps[0].EntryToExitLoss != 100 {
t.Fatalf("expected 100%% entry-to-exit loss, got %+v", snaps[0])
}
}
func TestTunnelQualityProberUsesOnlineBackupExit(t *testing.T) {
h := setupProbeTargetTunnelHandler(t)
seedQualityForwardTunnel(t, h, 82, []int{0, 1})
p := newTunnelQualityProber(h)
var calls []string
p.probeNode = func(nodeID int64, ip string, port int, options diagnosisExecOptions) (float64, float64, error) {
if options.pingCount != 1 {
t.Fatalf("expected real-time quality probe count 1, got %d", options.pingCount)
}
calls = append(calls, fmt.Sprintf("%d|%s|%d", nodeID, ip, port))
return 10, 0, nil
}
p.probeTunnel(82)
if slices.Contains(calls, "10|10.0.0.30|30030") {
t.Fatalf("did not expect probe to offline primary exit, calls=%+v", calls)
}
if !slices.Contains(calls, "10|10.0.0.31|30031") {
t.Fatalf("expected entry probe to online backup exit, calls=%+v", calls)
}
if !slices.Contains(calls, "31|www.bing.com|443") {
t.Fatalf("expected public probe from online backup exit, calls=%+v", calls)
}
snaps := p.GetAll()
if len(snaps) != 1 || !snaps[0].Success {
t.Fatalf("expected successful backup exit snapshot, got %+v", snaps)
}
}
func TestTunnelQualityProberReportsAllExitCandidateLatencies(t *testing.T) {
h := setupProbeTargetTunnelHandler(t)
seedQualityForwardTunnel(t, h, 83, []int{1, 1})
p := newTunnelQualityProber(h)
p.probeNode = func(nodeID int64, ip string, port int, options diagnosisExecOptions) (float64, float64, error) {
switch fmt.Sprintf("%d|%s|%d", nodeID, ip, port) {
case "10|10.0.0.30|30030":
return 20, 0, nil
case "10|10.0.0.31|30031":
return 35, 0, nil
case "30|www.bing.com|443":
return 50, 0, nil
case "31|www.bing.com|443":
return 65, 0, nil
default:
return 0, 100, fmt.Errorf("unexpected probe node=%d target=%s:%d", nodeID, ip, port)
}
}
p.probeTunnel(83)
snaps := p.GetAll()
if len(snaps) != 1 {
t.Fatalf("expected one quality snapshot, got %+v", snaps)
}
if snaps[0].EntryToExitLatency != 20 || snaps[0].ExitToBingLatency != 50 {
t.Fatalf("expected primary path metrics to remain unchanged, got %+v", snaps[0])
}
var details tunnelQualityChainDetails
if err := json.Unmarshal([]byte(snaps[0].ChainDetails), &details); err != nil {
t.Fatalf("decode chain details: %v", err)
}
assertCandidateHop := func(fromID, toID int64, latency float64, selected bool) {
t.Helper()
for _, hop := range details.CandidateHops {
if hop.FromNodeID == fromID && hop.ToNodeID == toID {
if hop.Latency != latency || hop.Selected != selected || hop.ErrorMessage != "" {
t.Fatalf("unexpected candidate hop: %+v", hop)
}
return
}
}
t.Fatalf("candidate hop %d -> %d not found in %+v", fromID, toID, details.CandidateHops)
}
assertCandidateHop(10, 30, 20, true)
assertCandidateHop(10, 31, 35, false)
assertCandidateHop(30, 0, 50, true)
assertCandidateHop(31, 0, 65, false)
}
func TestTunnelQualityProberStoresProbeTargetWhenChainIncomplete(t *testing.T) {
h := setupProbeTargetTunnelHandler(t)
seedProbeTargetTunnel(t, h, 78, "quality-target-incomplete", "speed.example.com", 8443)
@@ -67,3 +179,69 @@ func TestTunnelQualityProberStoresProbeTargetWhenChainIncomplete(t *testing.T) {
t.Fatalf("unexpected snapshot target metadata: %+v", snaps[0])
}
}
func TestTunnelQualityProberUsesConfiguredInterval(t *testing.T) {
h := setupProbeTargetTunnelHandler(t)
if err := h.repo.UpsertConfig("monitor_tunnel_quality_interval_sec", "15", time.Now().UnixMilli()); err != nil {
t.Fatalf("upsert interval config: %v", err)
}
p := newTunnelQualityProber(h)
if got := p.probeInterval(); got != 15*time.Second {
t.Fatalf("probe interval = %s, want 15s", got)
}
}
func TestTunnelQualityProberConfigNotificationIsCoalesced(t *testing.T) {
p := newTunnelQualityProber(nil)
p.NotifyConfigChanged()
p.NotifyConfigChanged()
if got := len(p.wake); got != 1 {
t.Fatalf("wake notifications = %d, want 1", got)
}
}
func TestNormalizeTunnelQualityProbeIntervalConfigValue(t *testing.T) {
got, err := normalizeAndValidateConfigValue("monitor_tunnel_quality_interval_sec", " 15 ")
if err != nil || got != "15" {
t.Fatalf("normalize interval = %q, %v", got, err)
}
if _, err := normalizeAndValidateConfigValue("monitor_tunnel_quality_interval_sec", "0"); err == nil {
t.Fatalf("expected invalid interval to be rejected")
}
}
func seedQualityForwardTunnel(t *testing.T, h *Handler, tunnelID int64, exitStatuses []int) {
t.Helper()
now := time.Now().UnixMilli()
if err := h.repo.DB().Exec(`
INSERT INTO tunnel(id, name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, inx, ip_preference, probe_target_host, probe_target_port)
VALUES(?, ?, 1, 2, 'tls', 1, ?, ?, 1, ?, '', '', 0)
`, tunnelID, fmt.Sprintf("quality-forward-%d", tunnelID), now, now, tunnelID).Error; err != nil {
t.Fatalf("insert forwarding tunnel: %v", err)
}
if err := h.repo.DB().Exec(`
INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol)
VALUES(?, '1', 10, 30001, 'fifo', 1, 'tls')
`, tunnelID).Error; err != nil {
t.Fatalf("insert entry chain: %v", err)
}
for i, status := range exitStatuses {
nodeID := int64(30 + i)
port := 30030 + i
ip := fmt.Sprintf("10.0.0.%d", nodeID)
if err := h.repo.DB().Exec(`
INSERT INTO node(id, name, secret, server_ip, server_ip_v4, server_ip_v6, port, interface_name, version, http, tls, socks, created_time, updated_time, status, tcp_listen_addr, udp_listen_addr, inx)
VALUES(?, ?, ?, ?, ?, '', '30000-30100', '', 'v1', 1, 1, 1, ?, ?, ?, '[::]', '[::]', 0)
`, nodeID, fmt.Sprintf("exit-%d", i+1), fmt.Sprintf("exit-secret-%d", i+1), ip, ip, now, now, status).Error; err != nil {
t.Fatalf("insert exit node %d: %v", nodeID, err)
}
if err := h.repo.DB().Exec(`
INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol)
VALUES(?, '3', ?, ?, 'fifo', ?, 'tls')
`, tunnelID, nodeID, port, i+1).Error; err != nil {
t.Fatalf("insert exit chain %d: %v", nodeID, err)
}
}
}
@@ -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)
+127 -21
View File
@@ -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)}
}
+104
View File
@@ -0,0 +1,104 @@
package license
import (
"encoding/json"
"fmt"
"net/http"
"net/http/httptest"
"strings"
"testing"
)
func TestValidateKeyWithMachineSendsFingerprintAndMachineScopes(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
var body struct {
Meta struct {
Scope map[string]string `json:"scope"`
} `json:"meta"`
}
if err := json.NewDecoder(r.Body).Decode(&body); err != nil {
t.Fatalf("decode request: %v", err)
}
if body.Meta.Scope["fingerprint"] != "fingerprint" {
t.Fatalf("fingerprint scope = %q", body.Meta.Scope["fingerprint"])
}
if body.Meta.Scope["machine"] != "machine-id" {
t.Fatalf("machine scope = %q", body.Meta.Scope["machine"])
}
_, _ = fmt.Fprint(w, `{"meta":{"valid":true,"code":"VALID"},"data":{"id":"license-id","attributes":{}}}`)
}))
defer server.Close()
client := NewKeygenClient("account-id", "")
client.BaseURL = server.URL
validation, err := client.ValidateKeyWithMachine("license-key", "fingerprint", "machine-id")
if err != nil || !validation.Meta.Valid {
t.Fatalf("ValidateKeyWithMachine() validation=%+v err=%v", validation, err)
}
}
func TestGetMachineIDRetrievesMachineByFingerprint(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodGet || !strings.HasSuffix(r.URL.Path, "/machines/fingerprint") {
t.Fatalf("unexpected request %s %s", r.Method, r.URL.Path)
}
if r.Header.Get("Authorization") != "License license-key" {
t.Fatalf("authorization = %q", r.Header.Get("Authorization"))
}
_, _ = fmt.Fprint(w, `{"data":{"type":"machines","id":"machine-id"}}`)
}))
defer server.Close()
client := NewKeygenClient("account-id", "license-key")
client.BaseURL = server.URL
machineID, err := client.GetMachineID("fingerprint")
if err != nil || machineID != "machine-id" {
t.Fatalf("GetMachineID() = %q, %v", machineID, err)
}
}
func TestActivateMachineTreatsFingerprintTakenAsIdempotent(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.WriteHeader(http.StatusUnprocessableEntity)
_, _ = fmt.Fprint(w, `{"errors":[{"code":"FINGERPRINT_TAKEN"},{"code":"MACHINE_LIMIT_EXCEEDED"}]}`)
}))
defer server.Close()
client := NewKeygenClient("account-id", "license-key")
client.BaseURL = server.URL
if machineID, err := client.ActivateMachine("license-id", "fingerprint"); err != nil || machineID != "" {
t.Fatalf("ActivateMachine() error = %v, want idempotent success", err)
}
}
func TestActivateMachineReturnsCreatedMachineID(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.WriteHeader(http.StatusCreated)
_, _ = fmt.Fprint(w, `{"data":{"type":"machines","id":"machine-id"}}`)
}))
defer server.Close()
client := NewKeygenClient("account-id", "license-key")
client.BaseURL = server.URL
machineID, err := client.ActivateMachine("license-id", "fingerprint")
if err != nil || machineID != "machine-id" {
t.Fatalf("ActivateMachine() = %q, %v", machineID, err)
}
}
func TestActivateMachineRejectsMachineLimitWithoutFingerprintTaken(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.WriteHeader(http.StatusUnprocessableEntity)
_, _ = fmt.Fprint(w, `{"errors":[{"code":"MACHINE_LIMIT_EXCEEDED"}]}`)
}))
defer server.Close()
client := NewKeygenClient("account-id", "license-key")
client.BaseURL = server.URL
_, err := client.ActivateMachine("license-id", "fingerprint")
if err == nil || !strings.Contains(err.Error(), "MACHINE_LIMIT_EXCEEDED") {
t.Fatalf("ActivateMachine() error = %v, want machine limit failure", err)
}
}
@@ -0,0 +1,52 @@
package monitoring
import (
"fmt"
"strconv"
"strings"
)
const (
ConfigTunnelQualityProbeIntervalSec = "monitor_tunnel_quality_interval_sec"
DefaultTunnelQualityProbeIntervalSec = 1
MinTunnelQualityProbeIntervalSec = 1
MaxTunnelQualityProbeIntervalSec = 3600
)
func TunnelQualityProbeIntervalSecondsFromConfigMap(cfg map[string]string) int {
if cfg == nil {
return DefaultTunnelQualityProbeIntervalSec
}
seconds, err := parseTunnelQualityProbeIntervalSeconds(cfg[ConfigTunnelQualityProbeIntervalSec])
if err != nil {
return DefaultTunnelQualityProbeIntervalSec
}
return seconds
}
func NormalizeTunnelQualityProbeIntervalSeconds(value string) (string, error) {
seconds, err := parseTunnelQualityProbeIntervalSeconds(value)
if err != nil {
return "", err
}
return strconv.Itoa(seconds), nil
}
func parseTunnelQualityProbeIntervalSeconds(value string) (int, error) {
trimmed := strings.TrimSpace(value)
if trimmed == "" {
return 0, fmt.Errorf("隧道质量探测间隔不能为空")
}
seconds, err := strconv.Atoi(trimmed)
if err != nil {
return 0, fmt.Errorf("隧道质量探测间隔必须是整数")
}
if seconds < MinTunnelQualityProbeIntervalSec || seconds > MaxTunnelQualityProbeIntervalSec {
return 0, fmt.Errorf(
"隧道质量探测间隔必须在 %d 到 %d 秒之间",
MinTunnelQualityProbeIntervalSec,
MaxTunnelQualityProbeIntervalSec,
)
}
return seconds, nil
}
@@ -0,0 +1,37 @@
package monitoring
import "testing"
func TestTunnelQualityProbeIntervalSecondsFromConfigMap(t *testing.T) {
tests := []struct {
name string
cfg map[string]string
want int
}{
{name: "missing config", cfg: nil, want: DefaultTunnelQualityProbeIntervalSec},
{name: "configured", cfg: map[string]string{ConfigTunnelQualityProbeIntervalSec: "15"}, want: 15},
{name: "invalid", cfg: map[string]string{ConfigTunnelQualityProbeIntervalSec: "0"}, want: DefaultTunnelQualityProbeIntervalSec},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if got := TunnelQualityProbeIntervalSecondsFromConfigMap(tt.cfg); got != tt.want {
t.Fatalf("interval = %d, want %d", got, tt.want)
}
})
}
}
func TestNormalizeTunnelQualityProbeIntervalSeconds(t *testing.T) {
for _, value := range []string{"1", "15", "3600"} {
if got, err := NormalizeTunnelQualityProbeIntervalSeconds(value); err != nil || got != value {
t.Fatalf("normalize %q = %q, %v", value, got, err)
}
}
for _, value := range []string{"", "0", "3601", "1.5", "abc"} {
if got, err := NormalizeTunnelQualityProbeIntervalSeconds(value); err == nil {
t.Fatalf("normalize %q unexpectedly succeeded with %q", value, got)
}
}
}
@@ -46,7 +46,8 @@ func RenderTable(plan NodePlan) string {
}
targetHost := strings.Trim(strings.TrimSpace(rule.TargetHost), "[]")
for _, protocol := range normalizedProtocols(rule.Protocols) {
b.WriteString(fmt.Sprintf(" ct original proto-dst %d %s daddr %s %s dport %d counter comment %q\n",
b.WriteString(fmt.Sprintf(" meta l4proto %s ct original proto-dst %d %s daddr %s %s dport %d counter comment %q\n",
protocol,
rule.InPort,
family,
targetHost,
@@ -54,7 +55,8 @@ func RenderTable(plan NodePlan) string {
rule.TargetPort,
counterComment(rule.ForwardID, CounterDirectionToTarget, protocol),
))
b.WriteString(fmt.Sprintf(" ct original proto-dst %d %s saddr %s %s sport %d counter comment %q\n",
b.WriteString(fmt.Sprintf(" meta l4proto %s ct original proto-dst %d %s saddr %s %s sport %d counter comment %q\n",
protocol,
rule.InPort,
family,
targetHost,
@@ -62,10 +62,10 @@ func TestRenderTableIncludesForwardAccountingCounters(t *testing.T) {
wantLines := []string{
`tcp dport 12345 counter dnat ip to 198.51.100.20:443 comment "flvx forward:42 dnat tcp"`,
`udp dport 12345 counter dnat ip to 198.51.100.20:443 comment "flvx forward:42 dnat udp"`,
`ct original proto-dst 12345 ip daddr 198.51.100.20 tcp dport 443 counter comment "flvx forward:42 to-target tcp"`,
`ct original proto-dst 12345 ip saddr 198.51.100.20 tcp sport 443 counter comment "flvx forward:42 from-target tcp"`,
`ct original proto-dst 12345 ip daddr 198.51.100.20 udp dport 443 counter comment "flvx forward:42 to-target udp"`,
`ct original proto-dst 12345 ip saddr 198.51.100.20 udp sport 443 counter comment "flvx forward:42 from-target udp"`,
`meta l4proto tcp ct original proto-dst 12345 ip daddr 198.51.100.20 tcp dport 443 counter comment "flvx forward:42 to-target tcp"`,
`meta l4proto tcp ct original proto-dst 12345 ip saddr 198.51.100.20 tcp sport 443 counter comment "flvx forward:42 from-target tcp"`,
`meta l4proto udp ct original proto-dst 12345 ip daddr 198.51.100.20 udp dport 443 counter comment "flvx forward:42 to-target udp"`,
`meta l4proto udp ct original proto-dst 12345 ip saddr 198.51.100.20 udp sport 443 counter comment "flvx forward:42 from-target udp"`,
}
for _, want := range wantLines {
if !strings.Contains(got, want) {
@@ -89,8 +89,8 @@ func TestRenderTableIncludesIPv6ForwardAccountingCounters(t *testing.T) {
got := RenderTable(plan)
wantLines := []string{
`tcp dport 12346 counter dnat ip6 to [2001:db8::20]:8443 comment "flvx forward:43 dnat tcp"`,
`ct original proto-dst 12346 ip6 daddr 2001:db8::20 tcp dport 8443 counter comment "flvx forward:43 to-target tcp"`,
`ct original proto-dst 12346 ip6 saddr 2001:db8::20 tcp sport 8443 counter comment "flvx forward:43 from-target tcp"`,
`meta l4proto tcp ct original proto-dst 12346 ip6 daddr 2001:db8::20 tcp dport 8443 counter comment "flvx forward:43 to-target tcp"`,
`meta l4proto tcp ct original proto-dst 12346 ip6 saddr 2001:db8::20 tcp sport 8443 counter comment "flvx forward:43 from-target tcp"`,
}
for _, want := range wantLines {
if !strings.Contains(got, want) {
@@ -110,8 +110,8 @@ func TestRenderTableAccountingCountersIncludeOriginalPort(t *testing.T) {
got := RenderTable(plan)
wantLines := []string{
`ct original proto-dst 12345 ip daddr 198.51.100.20 tcp dport 443 counter comment "flvx forward:42 to-target tcp"`,
`ct original proto-dst 12346 ip daddr 198.51.100.20 tcp dport 443 counter comment "flvx forward:43 to-target tcp"`,
`meta l4proto tcp ct original proto-dst 12345 ip daddr 198.51.100.20 tcp dport 443 counter comment "flvx forward:42 to-target tcp"`,
`meta l4proto tcp ct original proto-dst 12346 ip daddr 198.51.100.20 tcp dport 443 counter comment "flvx forward:43 to-target tcp"`,
}
for _, want := range wantLines {
if !strings.Contains(got, want) {
+43 -8
View File
@@ -26,23 +26,58 @@ func NewSSHRunner() *SSHRunner {
}
func (r *SSHRunner) Test(ctx context.Context, cfg SSHConfig) error {
return r.run(ctx, cfg, "command -v nft >/dev/null 2>&1 && nft --version >/dev/null 2>&1")
nft := nftBinary(cfg)
tableName := fmt.Sprintf("flvx_capability_%d", time.Now().UnixNano())
return r.run(ctx, cfg, buildCapabilityCheckCommand(nft, tableName))
}
func buildCapabilityCheckCommand(nft, tableName string) string {
script := RenderTable(NodePlan{
Rules: []Rule{{
ForwardID: 1,
InPort: 12345,
TargetHost: "192.0.2.1",
TargetPort: 443,
Protocols: []string{"tcp", "udp"},
}},
})
script = strings.Replace(script, "table inet flvx {", "table inet "+tableName+" {", 1)
return "set -eu\n" +
"command -v nft >/dev/null 2>&1\n" +
nft + " --version >/dev/null 2>&1\n" +
"tmp=$(mktemp /tmp/flvx-nft-capability-XXXXXX.nft)\n" +
"trap 'rm -f \"$tmp\"' EXIT\n" +
"cat > \"$tmp\" <<'EOF'\n" + script + "\nEOF\n" +
"if ! " + nft + " -c -f \"$tmp\"; then\n" +
" echo 'nftables cannot validate the generated FLVX rules' >&2\n" +
" exit 1\n" +
"fi"
}
func (r *SSHRunner) ApplyScript(ctx context.Context, cfg SSHConfig, script string) error {
nft := nftBinary(cfg)
command := "tmp=$(mktemp /tmp/flvx-nft-XXXXXX.nft) || exit 1\n" +
command := buildApplyCommand(nftBinary(cfg), script)
return r.run(ctx, cfg, command)
}
func buildApplyCommand(nft, script string) string {
return "set -eu\n" +
"tmp=$(mktemp /tmp/flvx-nft-XXXXXX.nft)\n" +
"batch=$(mktemp /tmp/flvx-nft-batch-XXXXXX.nft) || { rm -f \"$tmp\"; exit 1; }\n" +
"cleanup() {\n" +
" rm -f \"$tmp\"\n" +
" rm -f \"$tmp\" \"$batch\"\n" +
"}\n" +
"trap cleanup EXIT\n" +
"cat > \"$tmp\" <<'EOF'\n" + script + "\nEOF\n" +
nft + " -c -f \"$tmp\"\n" +
"if " + nft + " list table inet flvx >/dev/null 2>&1; then\n" +
" " + nft + " delete table inet flvx\n" +
" { printf '%s\\n' 'delete table inet flvx'; cat \"$tmp\"; } > \"$batch\"\n" +
"else\n" +
" cp \"$tmp\" \"$batch\"\n" +
"fi\n" +
nft + " -f \"$tmp\""
return r.run(ctx, cfg, command)
"if ! " + nft + " -c -f \"$batch\"; then\n" +
" echo 'nftables rule validation failed; active rules were preserved' >&2\n" +
" exit 1\n" +
"fi\n" +
nft + " -f \"$batch\""
}
func (r *SSHRunner) ListTableJSON(ctx context.Context, cfg SSHConfig) ([]byte, error) {
@@ -5,10 +5,113 @@ import (
"crypto/rsa"
"crypto/x509"
"encoding/pem"
"os"
"os/exec"
"path/filepath"
"strings"
"testing"
)
func TestBuildApplyCommandStopsAfterValidationFailure(t *testing.T) {
dir := t.TempDir()
logPath := filepath.Join(dir, "calls.log")
applyMarker := filepath.Join(dir, "applied")
nftPath := filepath.Join(dir, "nft")
fake := `#!/bin/sh
echo "$*" >> "` + logPath + `"
if [ "$1" = "list" ]; then
exit 0
fi
if [ "$1" = "-c" ]; then
exit 1
fi
touch "` + applyMarker + `"
`
if err := os.WriteFile(nftPath, []byte(fake), 0o755); err != nil {
t.Fatalf("write fake nft: %v", err)
}
command := buildApplyCommand(nftPath, "table inet flvx { }")
result := exec.Command("sh", "-c", command)
if err := result.Run(); err == nil {
t.Fatal("expected validation failure")
}
if _, err := os.Stat(applyMarker); !os.IsNotExist(err) {
t.Fatalf("apply ran after validation failure, stat err=%v", err)
}
calls, err := os.ReadFile(logPath)
if err != nil {
t.Fatalf("read fake nft calls: %v", err)
}
if strings.Count(string(calls), "-f ") != 1 {
t.Fatalf("expected validation only, got calls:\n%s", calls)
}
}
func TestBuildCapabilityCheckCommandValidatesRenderedRulesWithoutApplying(t *testing.T) {
dir := t.TempDir()
logPath := filepath.Join(dir, "calls.log")
nftPath := filepath.Join(dir, "nft")
fake := `#!/bin/sh
echo "$*" >> "` + logPath + `"
if [ "$1" = "--version" ] || [ "$1" = "-c" ]; then
exit 0
fi
exit 1
`
if err := os.WriteFile(nftPath, []byte(fake), 0o755); err != nil {
t.Fatalf("write fake nft: %v", err)
}
command := strings.Replace(buildCapabilityCheckCommand(nftPath, "flvx_capability_test"), "command -v nft", "command -v "+nftPath, 1)
result := exec.Command("sh", "-c", command)
if output, err := result.CombinedOutput(); err != nil {
t.Fatalf("capability command failed: %v: %s", err, output)
}
calls, err := os.ReadFile(logPath)
if err != nil {
t.Fatalf("read fake nft calls: %v", err)
}
if strings.Count(string(calls), "-c -f ") != 1 || strings.Contains(string(calls), "\n-f ") {
t.Fatalf("expected one check-only invocation, got calls:\n%s", calls)
}
if !strings.Contains(command, "table inet flvx_capability_test") || !strings.Contains(command, "meta l4proto tcp ct original proto-dst") {
t.Fatalf("capability check does not contain representative rendered rules:\n%s", command)
}
}
func TestBuildApplyCommandUsesAtomicReplacementBatch(t *testing.T) {
dir := t.TempDir()
batchPath := filepath.Join(dir, "batch.nft")
nftPath := filepath.Join(dir, "nft")
fake := `#!/bin/sh
if [ "$1" = "list" ]; then
exit 0
fi
if [ "$1" = "-f" ]; then
cp "$2" "` + batchPath + `"
fi
exit 0
`
if err := os.WriteFile(nftPath, []byte(fake), 0o755); err != nil {
t.Fatalf("write fake nft: %v", err)
}
script := "table inet flvx {\n chain forward { }\n}"
result := exec.Command("sh", "-c", buildApplyCommand(nftPath, script))
if output, err := result.CombinedOutput(); err != nil {
t.Fatalf("apply command failed: %v: %s", err, output)
}
batch, err := os.ReadFile(batchPath)
if err != nil {
t.Fatalf("read applied batch: %v", err)
}
want := "delete table inet flvx\n" + script + "\n"
if string(batch) != want {
t.Fatalf("atomic batch = %q, want %q", batch, want)
}
}
func TestAuthMethodsDefaultToPrivateKey(t *testing.T) {
privateKey := mustGeneratePrivateKey(t)
methods, err := authMethods(SSHConfig{PrivateKey: privateKey})
@@ -15,14 +15,27 @@ var publicConfigKeys = map[string]struct{}{
"app_favicon": {},
"app_bg_image": {},
"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 +56,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 +75,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))
}
@@ -13,6 +13,8 @@ func TestConfigPolicy(t *testing.T) {
{name: "app_favicon is public", key: "app_favicon", want: ConfigAccessPublic},
{name: "app_bg_image is public", key: "app_bg_image", want: ConfigAccessPublic},
{name: "cloudflare_site_key is public", key: "cloudflare_site_key", want: ConfigAccessPublic},
{name: "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 +31,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", "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 +69,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)
}
}
}
+28 -3
View File
@@ -473,6 +473,14 @@ func seedData(db *gorm.DB) {
appNameConfig := model.ViteConfig{ID: 1, Name: "app_name", Value: "flux", Time: 1755147963000}
db.Where("id = ?", 1).FirstOrCreate(&appNameConfig)
now := time.Now().UnixMilli()
for name, value := range map[string]string{
"is_commercial": "false",
"hide_footer_brand": "false",
} {
cfg := model.ViteConfig{Name: name, Value: value, Time: now}
db.Where("name = ?", name).FirstOrCreate(&cfg)
}
}
// ─── User Queries ────────────────────────────────────────────────────
@@ -600,6 +608,23 @@ func (r *Repository) UpsertConfig(name, value string, now int64) error {
}).Create(&model.ViteConfig{Name: name, Value: value, Time: now}).Error
}
func (r *Repository) UpsertConfigs(values map[string]string, now int64) error {
if r == nil || r.db == nil {
return errors.New("repository not initialized")
}
return r.db.Transaction(func(tx *gorm.DB) error {
for name, value := range values {
if err := tx.Clauses(clause.OnConflict{
Columns: []clause.Column{{Name: "name"}},
DoUpdates: clause.AssignmentColumns([]string{"value", "time"}),
}).Create(&model.ViteConfig{Name: name, Value: value, Time: now}).Error; err != nil {
return err
}
}
return nil
})
}
// ─── Announcement Queries ────────────────────────────────────────────
func (r *Repository) GetAnnouncement() (*model.Announcement, error) {
@@ -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
}
@@ -2839,7 +2864,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")
}
@@ -132,3 +132,31 @@ func TestUpsertTunnelMetricBucketsIsSafeUnderConcurrency(t *testing.T) {
t.Fatalf("expected bytesOut %d, got %d", wantOut, rows[0].BytesOut)
}
}
func TestGetLatestTunnelQualitiesIncludesChainDetails(t *testing.T) {
r, err := Open(":memory:")
if err != nil {
t.Fatalf("open repo: %v", err)
}
defer r.Close()
if err := r.InsertTunnelQuality(&model.TunnelQuality{
TunnelID: 7,
Timestamp: time.Now().UnixMilli(),
Success: 1,
ChainDetails: `{"primaryPath":[],"candidateHops":[{"fromNodeId":10,"toNodeId":31}]}`,
}); err != nil {
t.Fatalf("insert tunnel quality: %v", err)
}
items, err := r.GetLatestTunnelQualities()
if err != nil {
t.Fatalf("get latest tunnel qualities: %v", err)
}
if len(items) != 1 {
t.Fatalf("expected one latest tunnel quality, got %+v", items)
}
if items[0].ChainDetails == "" {
t.Fatalf("expected chain details in latest quality row, got %+v", items[0])
}
}
@@ -44,7 +44,8 @@ func (r *Repository) GetLatestTunnelQualities() ([]model.TunnelQuality, error) {
// Use window function (works on modern SQLite 3.25+ and PostgreSQL).
q := `
SELECT id, tunnel_id, entry_to_exit_latency, exit_to_bing_latency,
entry_to_exit_loss, exit_to_bing_loss, success, error_message, timestamp
entry_to_exit_loss, exit_to_bing_loss, success, error_message, timestamp,
chain_details
FROM (
SELECT *, ROW_NUMBER() OVER (PARTITION BY tunnel_id ORDER BY timestamp DESC, id DESC) AS rn
FROM tunnel_quality
+46 -6
View File
@@ -2,6 +2,8 @@ package chain
import (
"context"
"errors"
"io"
"github.com/go-gost/core/chain"
"github.com/go-gost/core/hop"
@@ -38,11 +40,12 @@ type chainNamer interface {
}
type Chain struct {
name string
hops []hop.Hop
marker selector.Marker
metadata metadata.Metadata
logger logger.Logger
name string
hops []hop.Hop
ownedHops []hop.Hop
marker selector.Marker
metadata metadata.Metadata
logger logger.Logger
}
func NewChain(name string, opts ...ChainOption) *Chain {
@@ -61,8 +64,15 @@ func NewChain(name string, opts ...ChainOption) *Chain {
}
}
func (c *Chain) AddHop(hop hop.Hop) {
func (c *Chain) AddHop(hop hop.Hop, owned ...bool) {
c.hops = append(c.hops, hop)
isOwned := true
if len(owned) > 0 {
isOwned = owned[0]
}
if isOwned {
c.ownedHops = append(c.ownedHops, hop)
}
}
// Metadata implements metadata.Metadatable interface.
@@ -112,6 +122,36 @@ func (c *Chain) Route(ctx context.Context, network, address string, opts ...chai
return rt
}
// Retire gracefully drains resources owned by a chain that has been replaced.
func (c *Chain) Retire() {
if c == nil {
return
}
for _, h := range c.ownedHops {
if retirer, ok := h.(interface{ Retire() }); ok {
retirer.Retire()
continue
}
if closer, ok := h.(io.Closer); ok {
_ = closer.Close()
}
}
}
// Close immediately releases all resources owned by the chain.
func (c *Chain) Close() error {
if c == nil {
return nil
}
var errs []error
for _, h := range c.ownedHops {
if closer, ok := h.(io.Closer); ok {
errs = append(errs, closer.Close())
}
}
return errors.Join(errs...)
}
type chainGroup struct {
chains []chain.Chainer
selector selector.Selector[chain.Chainer]
+64
View File
@@ -0,0 +1,64 @@
package chain
import (
"context"
"testing"
corechain "github.com/go-gost/core/chain"
corehop "github.com/go-gost/core/hop"
)
type lifecycleTestHop struct {
selected int
retired int
closed int
}
func (h *lifecycleTestHop) Select(context.Context, ...corehop.SelectOption) *corechain.Node {
h.selected++
return nil
}
func (h *lifecycleTestHop) Retire() {
h.retired++
}
func (h *lifecycleTestHop) Close() error {
h.closed++
return nil
}
func TestChainRoutesThroughSharedHopWithoutOwningLifecycle(t *testing.T) {
hop := &lifecycleTestHop{}
chain := NewChain("shared-hop")
chain.AddHop(hop, false)
if route := chain.Route(context.Background(), "tcp", "example.com:443"); route == nil {
t.Fatal("route is nil")
}
if hop.selected != 1 {
t.Fatalf("shared hop selected %d times, want 1", hop.selected)
}
chain.Retire()
if err := chain.Close(); err != nil {
t.Fatalf("close chain: %v", err)
}
if hop.retired != 0 || hop.closed != 0 {
t.Fatalf("shared hop lifecycle changed: retired=%d closed=%d", hop.retired, hop.closed)
}
}
func TestChainRetiresAndClosesOwnedHop(t *testing.T) {
hop := &lifecycleTestHop{}
chain := NewChain("owned-hop")
chain.AddHop(hop)
chain.Retire()
if err := chain.Close(); err != nil {
t.Fatalf("close chain: %v", err)
}
if hop.retired != 1 || hop.closed != 1 {
t.Fatalf("owned hop lifecycle: retired=%d closed=%d, want 1/1", hop.retired, hop.closed)
}
}
+32 -1
View File
@@ -2,6 +2,8 @@ package chain
import (
"context"
"errors"
"io"
"net"
"github.com/go-gost/core/chain"
@@ -102,5 +104,34 @@ func (tr *Transport) Options() *chain.TransportOptions {
func (tr *Transport) Copy() chain.Transporter {
tr2 := &Transport{}
*tr2 = *tr
return tr
return tr2
}
// Retire prevents long-lived dialer sessions owned by an obsolete chain from
// accepting new streams while allowing existing streams to drain.
func (tr *Transport) Retire() {
if tr == nil {
return
}
if retirer, ok := tr.dialer.(interface{ Retire() }); ok {
retirer.Retire()
}
if retirer, ok := tr.connector.(interface{ Retire() }); ok {
retirer.Retire()
}
}
// Close immediately releases transport-owned dialer and connector resources.
func (tr *Transport) Close() error {
if tr == nil {
return nil
}
var errs []error
if closer, ok := tr.dialer.(io.Closer); ok {
errs = append(errs, closer.Close())
}
if closer, ok := tr.connector.(io.Closer); ok {
errs = append(errs, closer.Close())
}
return errors.Join(errs...)
}
+45
View File
@@ -0,0 +1,45 @@
package chain
import (
"context"
"net"
"testing"
corechain "github.com/go-gost/core/chain"
)
type copyTestRoute struct{}
func (copyTestRoute) Dial(context.Context, string, string, ...corechain.DialOption) (net.Conn, error) {
return nil, nil
}
func (copyTestRoute) Bind(context.Context, string, string, ...corechain.BindOption) (net.Listener, error) {
return nil, nil
}
func (copyTestRoute) Nodes() []*corechain.Node {
return nil
}
func TestTransportCopyReturnsIndependentTransport(t *testing.T) {
originalRoute := copyTestRoute{}
replacementRoute := &copyTestRoute{}
original := NewTransport(nil, nil, corechain.RouteTransportOption(originalRoute))
copied, ok := original.Copy().(*Transport)
if !ok {
t.Fatalf("copy type = %T, want *Transport", original.Copy())
}
if copied == original {
t.Fatal("Copy returned the original transport")
}
copied.Options().Route = replacementRoute
if original.Options().Route != originalRoute {
t.Fatal("mutating copied transport changed original route")
}
if copied.Options().Route != replacementRoute {
t.Fatal("copied transport did not retain its independent route")
}
}
+82 -3
View File
@@ -3,6 +3,7 @@ package config
import (
"encoding/json"
"io"
"reflect"
"sync"
"time"
@@ -30,9 +31,87 @@ func Global() *Config {
globalMux.RLock()
defer globalMux.RUnlock()
cfg := &Config{}
*cfg = *global
return cfg
return cloneConfig(global)
}
// cloneConfig returns a detached snapshot of the runtime config. Runtime
// commands mutate slices, pointers and metadata maps, so a shallow struct copy
// is not sufficient once readers and persistence run concurrently.
func cloneConfig(c *Config) *Config {
if c == nil {
return nil
}
v := cloneConfigValue(reflect.ValueOf(c))
if !v.IsValid() || v.IsNil() {
return nil
}
return v.Interface().(*Config)
}
func cloneConfigValue(v reflect.Value) reflect.Value {
if !v.IsValid() {
return reflect.Value{}
}
switch v.Kind() {
case reflect.Interface:
if v.IsNil() {
return reflect.Zero(v.Type())
}
cloned := cloneConfigValue(v.Elem())
out := reflect.New(v.Type()).Elem()
if cloned.IsValid() && cloned.Type().AssignableTo(v.Type()) {
out.Set(cloned)
} else if cloned.IsValid() && cloned.Type().Implements(v.Type()) {
out.Set(cloned)
} else if cloned.IsValid() {
out.Set(cloned)
}
return out
case reflect.Pointer:
if v.IsNil() {
return reflect.Zero(v.Type())
}
out := reflect.New(v.Type().Elem())
out.Elem().Set(cloneConfigValue(v.Elem()))
return out
case reflect.Map:
if v.IsNil() {
return reflect.Zero(v.Type())
}
out := reflect.MakeMapWithSize(v.Type(), v.Len())
iter := v.MapRange()
for iter.Next() {
out.SetMapIndex(cloneConfigValue(iter.Key()), cloneConfigValue(iter.Value()))
}
return out
case reflect.Slice:
if v.IsNil() {
return reflect.Zero(v.Type())
}
out := reflect.MakeSlice(v.Type(), v.Len(), v.Len())
for i := 0; i < v.Len(); i++ {
out.Index(i).Set(cloneConfigValue(v.Index(i)))
}
return out
case reflect.Array:
out := reflect.New(v.Type()).Elem()
for i := 0; i < v.Len(); i++ {
out.Index(i).Set(cloneConfigValue(v.Index(i)))
}
return out
case reflect.Struct:
out := reflect.New(v.Type()).Elem()
out.Set(v)
for i := 0; i < v.NumField(); i++ {
if out.Field(i).CanSet() && v.Field(i).CanInterface() {
out.Field(i).Set(cloneConfigValue(v.Field(i)))
}
}
return out
default:
return v
}
}
func Set(c *Config) {
+121
View File
@@ -0,0 +1,121 @@
package config
import (
"encoding/json"
"path/filepath"
"sync"
"testing"
)
func TestGlobalReturnsDetachedSnapshot(t *testing.T) {
original := Global()
t.Cleanup(func() { Set(original) })
Set(&Config{Services: []*ServiceConfig{{
Name: "snapshot-service",
Metadata: map[string]any{"paused": false},
Handler: &HandlerConfig{
Type: "relay",
Metadata: map[string]any{"retries": 2},
},
}}})
snapshot := Global()
snapshot.Services[0].Name = "changed"
snapshot.Services[0].Metadata["paused"] = true
snapshot.Services[0].Handler.Metadata["retries"] = 9
current := Global()
if current.Services[0].Name != "snapshot-service" {
t.Fatalf("snapshot mutated global service name: %q", current.Services[0].Name)
}
if paused, _ := current.Services[0].Metadata["paused"].(bool); paused {
t.Fatalf("snapshot mutated global service metadata")
}
if retries, _ := current.Services[0].Handler.Metadata["retries"].(int); retries != 2 {
t.Fatalf("snapshot mutated nested handler metadata: %v", current.Services[0].Handler.Metadata["retries"])
}
}
func TestConcurrentGlobalSnapshotAndUpdate(t *testing.T) {
original := Global()
t.Cleanup(func() { Set(original) })
Set(&Config{Services: []*ServiceConfig{{
Name: "concurrent-service",
Metadata: map[string]any{"generation": 0},
}}})
var wg sync.WaitGroup
for worker := 0; worker < 4; worker++ {
wg.Add(1)
go func(worker int) {
defer wg.Done()
for i := 0; i < 500; i++ {
if err := OnUpdate(func(c *Config) error {
c.Services[0] = &ServiceConfig{
Name: "concurrent-service",
Metadata: map[string]any{"generation": worker*500 + i},
}
return nil
}); err != nil {
t.Errorf("OnUpdate: %v", err)
return
}
}
}(worker)
}
for i := 0; i < 2000; i++ {
if _, err := json.Marshal(Global()); err != nil {
t.Fatalf("marshal snapshot: %v", err)
}
}
wg.Wait()
}
func TestConcurrentPersistProducesValidConfig(t *testing.T) {
original := Global()
originalPath := PersistPath()
persistMu.Lock()
originalEnabled := persistEnable
persistMu.Unlock()
t.Cleanup(func() {
Set(original)
SetPersistPath(originalPath)
persistMu.Lock()
persistEnable = originalEnabled
persistMu.Unlock()
})
path := filepath.Join(t.TempDir(), "gost.json")
SetPersistPath(path)
EnablePersist()
Set(&Config{Services: []*ServiceConfig{{Name: "persist-service"}}})
var wg sync.WaitGroup
for worker := 0; worker < 4; worker++ {
wg.Add(1)
go func(worker int) {
defer wg.Done()
for i := 0; i < 50; i++ {
if err := OnUpdate(func(c *Config) error {
c.Services[0].Metadata = map[string]any{"generation": worker*50 + i}
return nil
}); err != nil {
t.Errorf("persist update: %v", err)
return
}
}
}(worker)
}
wg.Wait()
var persisted Config
if err := persisted.ReadFile(path); err != nil {
t.Fatalf("read persisted config: %v", err)
}
if len(persisted.Services) != 1 || persisted.Services[0] == nil || persisted.Services[0].Name != "persist-service" {
t.Fatalf("unexpected persisted services: %#v", persisted.Services)
}
}
+3 -1
View File
@@ -35,16 +35,18 @@ func ParseChain(cfg *config.ChainConfig, log logger.Logger) (chain.Chainer, erro
for _, ch := range cfg.Hops {
var hop hop.Hop
var err error
owned := false
if ch.Nodes != nil || ch.Plugin != nil {
if hop, err = hop_parser.ParseHop(ch, log); err != nil {
return nil, err
}
owned = true
} else {
hop = registry.HopRegistry().Get(ch.Name)
}
if hop != nil {
c.AddHop(hop)
c.AddHop(hop, owned)
}
}
+1 -1
View File
@@ -42,9 +42,9 @@ func EnablePersist() {
// persist writes the current global config to the configured file atomically.
func persist() error {
persistMu.Lock()
defer persistMu.Unlock()
path := persistPath
enabled := persistEnable
persistMu.Unlock()
if !enabled || path == "" {
return nil
+5 -2
View File
@@ -19,20 +19,23 @@ func (session *muxSession) Accept() (net.Conn, error) {
}
func (session *muxSession) Close() error {
if session.session == nil {
if session == nil || session.session == nil {
return nil
}
return session.session.Close()
}
func (session *muxSession) IsClosed() bool {
if session.session == nil {
if session == nil || session.session == nil {
return true
}
return session.session.IsClosed()
}
func (session *muxSession) NumStreams() int {
if session == nil || session.session == nil {
return 0
}
return session.session.NumStreams()
}
+38
View File
@@ -11,6 +11,7 @@ import (
"github.com/go-gost/core/logger"
md "github.com/go-gost/core/metadata"
kcp_util "github.com/go-gost/x/internal/util/kcp"
"github.com/go-gost/x/internal/util/sessionretire"
mdutil "github.com/go-gost/x/metadata/util"
"github.com/go-gost/x/registry"
"github.com/xtaci/kcp-go/v5"
@@ -25,6 +26,7 @@ func init() {
type kcpDialer struct {
sessions map[string]*muxSession
sessionMutex sync.Mutex
retired bool
logger logger.Logger
md metadata
options dialer.Options
@@ -64,6 +66,9 @@ func (d *kcpDialer) Dial(ctx context.Context, addr string, opts ...dialer.DialOp
d.sessionMutex.Lock()
defer d.sessionMutex.Unlock()
if d.retired {
return nil, net.ErrClosed
}
session, ok := d.sessions[addr]
if session != nil && session.IsClosed() {
@@ -171,3 +176,36 @@ func (d *kcpDialer) initSession(ctx context.Context, addr net.Addr, conn net.Pac
func (d *kcpDialer) Multiplex() bool {
return true
}
// Retire drains existing streams and closes their backing sessions once idle.
func (d *kcpDialer) Retire() {
for _, session := range d.detachSessions() {
sessionretire.Gracefully(session)
}
}
// Close immediately releases all cached multiplex sessions.
func (d *kcpDialer) Close() error {
var errs []error
for _, session := range d.detachSessions() {
errs = append(errs, session.Close())
}
return errors.Join(errs...)
}
func (d *kcpDialer) detachSessions() []*muxSession {
if d == nil {
return nil
}
d.sessionMutex.Lock()
d.retired = true
sessions := make([]*muxSession, 0, len(d.sessions))
for _, session := range d.sessions {
if session != nil {
sessions = append(sessions, session)
}
}
d.sessions = make(map[string]*muxSession)
d.sessionMutex.Unlock()
return sessions
}
+17
View File
@@ -0,0 +1,17 @@
package kcp
import (
"context"
"errors"
"net"
"testing"
)
func TestRetiredDialerRejectsNewConnections(t *testing.T) {
dialer := NewDialer().(*kcpDialer)
dialer.Retire()
if _, err := dialer.Dial(context.Background(), "127.0.0.1:1"); !errors.Is(err, net.ErrClosed) {
t.Fatalf("Dial error = %v, want net.ErrClosed", err)
}
}
+14
View File
@@ -20,13 +20,24 @@ func (session *muxSession) Accept() (net.Conn, error) {
}
func (session *muxSession) Close() error {
if session == nil {
return nil
}
if session.session == nil {
if session.conn != nil {
conn := session.conn
session.conn = nil
return conn.Close()
}
return nil
}
return session.session.Close()
}
func (session *muxSession) IsClosed() bool {
if session == nil {
return true
}
if session.session == nil {
return true
}
@@ -34,5 +45,8 @@ func (session *muxSession) IsClosed() bool {
}
func (session *muxSession) NumStreams() int {
if session == nil || session.session == nil {
return 0
}
return session.session.NumStreams()
}
+40
View File
@@ -11,6 +11,7 @@ import (
"github.com/go-gost/core/logger"
md "github.com/go-gost/core/metadata"
"github.com/go-gost/x/internal/util/mux"
"github.com/go-gost/x/internal/util/sessionretire"
"github.com/go-gost/x/registry"
)
@@ -21,6 +22,7 @@ func init() {
type mtcpDialer struct {
sessions map[string]*muxSession
sessionMutex sync.Mutex
retired bool
logger logger.Logger
md metadata
options dialer.Options
@@ -55,6 +57,9 @@ func (d *mtcpDialer) Multiplex() bool {
func (d *mtcpDialer) Dial(ctx context.Context, addr string, opts ...dialer.DialOption) (conn net.Conn, err error) {
d.sessionMutex.Lock()
defer d.sessionMutex.Unlock()
if d.retired {
return nil, net.ErrClosed
}
session, ok := d.sessions[addr]
if session != nil && session.IsClosed() {
@@ -88,6 +93,10 @@ func (d *mtcpDialer) Handshake(ctx context.Context, conn net.Conn, options ...di
d.sessionMutex.Lock()
defer d.sessionMutex.Unlock()
if d.retired {
conn.Close()
return nil, net.ErrClosed
}
if d.md.handshakeTimeout > 0 {
conn.SetDeadline(time.Now().Add(d.md.handshakeTimeout))
@@ -129,3 +138,34 @@ func (d *mtcpDialer) initSession(ctx context.Context, conn net.Conn) (*muxSessio
}
return &muxSession{conn: conn, session: session}, nil
}
func (d *mtcpDialer) Retire() {
for _, session := range d.detachSessions() {
sessionretire.Gracefully(session)
}
}
func (d *mtcpDialer) Close() error {
var errs []error
for _, session := range d.detachSessions() {
errs = append(errs, session.Close())
}
return errors.Join(errs...)
}
func (d *mtcpDialer) detachSessions() []*muxSession {
if d == nil {
return nil
}
d.sessionMutex.Lock()
d.retired = true
sessions := make([]*muxSession, 0, len(d.sessions))
for _, session := range d.sessions {
if session != nil {
sessions = append(sessions, session)
}
}
d.sessions = make(map[string]*muxSession)
d.sessionMutex.Unlock()
return sessions
}
+30
View File
@@ -0,0 +1,30 @@
package mtcp
import (
"context"
"errors"
"net"
"testing"
)
func TestRetiredDialerRejectsNewConnections(t *testing.T) {
dialer := NewDialer().(*mtcpDialer)
dialer.Retire()
if _, err := dialer.Dial(context.Background(), "127.0.0.1:1"); !errors.Is(err, net.ErrClosed) {
t.Fatalf("Dial error = %v, want net.ErrClosed", err)
}
}
func TestSessionCloseReleasesPreHandshakeConnection(t *testing.T) {
conn, peer := net.Pipe()
defer peer.Close()
session := &muxSession{conn: conn}
if err := session.Close(); err != nil {
t.Fatalf("Close: %v", err)
}
if !session.IsClosed() {
t.Fatal("pre-handshake session still reports open after Close")
}
}
+14
View File
@@ -20,13 +20,24 @@ func (session *muxSession) Accept() (net.Conn, error) {
}
func (session *muxSession) Close() error {
if session == nil {
return nil
}
if session.session == nil {
if session.conn != nil {
conn := session.conn
session.conn = nil
return conn.Close()
}
return nil
}
return session.session.Close()
}
func (session *muxSession) IsClosed() bool {
if session == nil {
return true
}
if session.session == nil {
return true
}
@@ -34,5 +45,8 @@ func (session *muxSession) IsClosed() bool {
}
func (session *muxSession) NumStreams() int {
if session == nil || session.session == nil {
return 0
}
return session.session.NumStreams()
}
+40
View File
@@ -12,6 +12,7 @@ import (
"github.com/go-gost/core/logger"
md "github.com/go-gost/core/metadata"
"github.com/go-gost/x/internal/util/mux"
"github.com/go-gost/x/internal/util/sessionretire"
"github.com/go-gost/x/registry"
)
@@ -22,6 +23,7 @@ func init() {
type mtlsDialer struct {
sessions map[string]*muxSession
sessionMutex sync.Mutex
retired bool
logger logger.Logger
md metadata
options dialer.Options
@@ -56,6 +58,9 @@ func (d *mtlsDialer) Multiplex() bool {
func (d *mtlsDialer) Dial(ctx context.Context, addr string, opts ...dialer.DialOption) (conn net.Conn, err error) {
d.sessionMutex.Lock()
defer d.sessionMutex.Unlock()
if d.retired {
return nil, net.ErrClosed
}
session, ok := d.sessions[addr]
if session != nil && session.IsClosed() {
@@ -89,6 +94,10 @@ func (d *mtlsDialer) Handshake(ctx context.Context, conn net.Conn, options ...di
d.sessionMutex.Lock()
defer d.sessionMutex.Unlock()
if d.retired {
conn.Close()
return nil, net.ErrClosed
}
if d.md.handshakeTimeout > 0 {
conn.SetDeadline(time.Now().Add(d.md.handshakeTimeout))
@@ -136,3 +145,34 @@ func (d *mtlsDialer) initSession(ctx context.Context, conn net.Conn) (*muxSessio
}
return &muxSession{conn: conn, session: session}, nil
}
func (d *mtlsDialer) Retire() {
for _, session := range d.detachSessions() {
sessionretire.Gracefully(session)
}
}
func (d *mtlsDialer) Close() error {
var errs []error
for _, session := range d.detachSessions() {
errs = append(errs, session.Close())
}
return errors.Join(errs...)
}
func (d *mtlsDialer) detachSessions() []*muxSession {
if d == nil {
return nil
}
d.sessionMutex.Lock()
d.retired = true
sessions := make([]*muxSession, 0, len(d.sessions))
for _, session := range d.sessions {
if session != nil {
sessions = append(sessions, session)
}
}
d.sessions = make(map[string]*muxSession)
d.sessionMutex.Unlock()
return sessions
}
+30
View File
@@ -0,0 +1,30 @@
package mtls
import (
"context"
"errors"
"net"
"testing"
)
func TestRetiredDialerRejectsNewConnections(t *testing.T) {
dialer := NewDialer().(*mtlsDialer)
dialer.Retire()
if _, err := dialer.Dial(context.Background(), "127.0.0.1:1"); !errors.Is(err, net.ErrClosed) {
t.Fatalf("Dial error = %v, want net.ErrClosed", err)
}
}
func TestSessionCloseReleasesPreHandshakeConnection(t *testing.T) {
conn, peer := net.Pipe()
defer peer.Close()
session := &muxSession{conn: conn}
if err := session.Close(); err != nil {
t.Fatalf("Close: %v", err)
}
if !session.IsClosed() {
t.Fatal("pre-handshake session still reports open after Close")
}
}
+14
View File
@@ -20,13 +20,24 @@ func (session *muxSession) Accept() (net.Conn, error) {
}
func (session *muxSession) Close() error {
if session == nil {
return nil
}
if session.session == nil {
if session.conn != nil {
conn := session.conn
session.conn = nil
return conn.Close()
}
return nil
}
return session.session.Close()
}
func (session *muxSession) IsClosed() bool {
if session == nil {
return true
}
if session.session == nil {
return true
}
@@ -34,5 +45,8 @@ func (session *muxSession) IsClosed() bool {
}
func (session *muxSession) NumStreams() int {
if session == nil || session.session == nil {
return 0
}
return session.session.NumStreams()
}
+40
View File
@@ -12,6 +12,7 @@ import (
"github.com/go-gost/core/logger"
md "github.com/go-gost/core/metadata"
"github.com/go-gost/x/internal/util/mux"
"github.com/go-gost/x/internal/util/sessionretire"
ws_util "github.com/go-gost/x/internal/util/ws"
"github.com/go-gost/x/registry"
"github.com/gorilla/websocket"
@@ -25,6 +26,7 @@ func init() {
type mwsDialer struct {
sessions map[string]*muxSession
sessionMutex sync.Mutex
retired bool
tlsEnabled bool
md metadata
options dialer.Options
@@ -70,6 +72,9 @@ func (d *mwsDialer) Multiplex() bool {
func (d *mwsDialer) Dial(ctx context.Context, addr string, opts ...dialer.DialOption) (conn net.Conn, err error) {
d.sessionMutex.Lock()
defer d.sessionMutex.Unlock()
if d.retired {
return nil, net.ErrClosed
}
session, ok := d.sessions[addr]
if session != nil && session.IsClosed() {
@@ -108,6 +113,10 @@ func (d *mwsDialer) Handshake(ctx context.Context, conn net.Conn, options ...dia
d.sessionMutex.Lock()
defer d.sessionMutex.Unlock()
if d.retired {
conn.Close()
return nil, net.ErrClosed
}
session, ok := d.sessions[opts.Addr]
if session != nil && session.conn != conn {
@@ -208,3 +217,34 @@ func (d *mwsDialer) keepAlive(conn ws_util.WebsocketConn) {
conn.SetWriteDeadline(time.Time{})
}
}
func (d *mwsDialer) Retire() {
for _, session := range d.detachSessions() {
sessionretire.Gracefully(session)
}
}
func (d *mwsDialer) Close() error {
var errs []error
for _, session := range d.detachSessions() {
errs = append(errs, session.Close())
}
return errors.Join(errs...)
}
func (d *mwsDialer) detachSessions() []*muxSession {
if d == nil {
return nil
}
d.sessionMutex.Lock()
d.retired = true
sessions := make([]*muxSession, 0, len(d.sessions))
for _, session := range d.sessions {
if session != nil {
sessions = append(sessions, session)
}
}
d.sessions = make(map[string]*muxSession)
d.sessionMutex.Unlock()
return sessions
}
+30
View File
@@ -0,0 +1,30 @@
package mws
import (
"context"
"errors"
"net"
"testing"
)
func TestRetiredDialerRejectsNewConnections(t *testing.T) {
dialer := NewDialer().(*mwsDialer)
dialer.Retire()
if _, err := dialer.Dial(context.Background(), "127.0.0.1:1"); !errors.Is(err, net.ErrClosed) {
t.Fatalf("Dial error = %v, want net.ErrClosed", err)
}
}
func TestSessionCloseReleasesPreHandshakeConnection(t *testing.T) {
conn, peer := net.Pipe()
defer peer.Close()
session := &muxSession{conn: conn}
if err := session.Close(); err != nil {
t.Fatalf("Close: %v", err)
}
if !session.IsClosed() {
t.Fatal("pre-handshake session still reports open after Close")
}
}
+49 -8
View File
@@ -3,6 +3,7 @@ package hop
import (
"context"
"encoding/json"
"errors"
"io"
"net"
"sort"
@@ -92,6 +93,7 @@ type chainHop struct {
nodes []*chain.Node
mu sync.RWMutex
cancelFunc context.CancelFunc
stopOnce sync.Once
options options
}
@@ -383,13 +385,52 @@ func (p *chainHop) parseNode(r io.Reader) ([]*chain.Node, error) {
return nodes, nil
}
func (p *chainHop) Close() error {
p.cancelFunc()
if p.options.fileLoader != nil {
p.options.fileLoader.Close()
func (p *chainHop) stopReload() {
if p == nil {
return
}
if p.options.redisLoader != nil {
p.options.redisLoader.Close()
}
return nil
p.stopOnce.Do(func() {
p.cancelFunc()
if p.options.fileLoader != nil {
p.options.fileLoader.Close()
}
if p.options.redisLoader != nil {
p.options.redisLoader.Close()
}
if p.options.httpLoader != nil {
p.options.httpLoader.Close()
}
})
}
func (p *chainHop) Retire() {
if p == nil {
return
}
p.stopReload()
for _, node := range p.Nodes() {
if node == nil || node.Options().Transport == nil {
continue
}
if retirer, ok := node.Options().Transport.(interface{ Retire() }); ok {
retirer.Retire()
}
}
}
func (p *chainHop) Close() error {
if p == nil {
return nil
}
p.stopReload()
var errs []error
for _, node := range p.Nodes() {
if node == nil || node.Options().Transport == nil {
continue
}
if closer, ok := node.Options().Transport.(io.Closer); ok {
errs = append(errs, closer.Close())
}
}
return errors.Join(errs...)
}
@@ -0,0 +1,58 @@
package sessionretire
import "time"
const (
defaultIdleGrace = time.Second
defaultPollPeriod = 100 * time.Millisecond
)
// Session is the lifecycle surface shared by the multiplexed dialers.
type Session interface {
Close() error
IsClosed() bool
NumStreams() int
}
// Gracefully closes a retired session after all existing streams have drained.
// A short idle grace covers the Dial/Handshake hand-off used by several dialers.
func Gracefully(session Session) {
if session == nil {
return
}
go waitUntilIdle(session, defaultIdleGrace, defaultPollPeriod)
}
func waitUntilIdle(session Session, idleGrace, pollPeriod time.Duration) {
if session == nil {
return
}
if idleGrace <= 0 {
idleGrace = defaultIdleGrace
}
if pollPeriod <= 0 {
pollPeriod = defaultPollPeriod
}
ticker := time.NewTicker(pollPeriod)
defer ticker.Stop()
var idleSince time.Time
for {
if session.IsClosed() {
_ = session.Close()
return
}
if session.NumStreams() == 0 {
if idleSince.IsZero() {
idleSince = time.Now()
} else if time.Since(idleSince) >= idleGrace {
_ = session.Close()
return
}
} else {
idleSince = time.Time{}
}
<-ticker.C
}
}
@@ -0,0 +1,59 @@
package sessionretire
import (
"sync"
"testing"
"time"
)
type testSession struct {
mu sync.Mutex
streams int
closed bool
}
func (s *testSession) Close() error {
s.mu.Lock()
s.closed = true
s.mu.Unlock()
return nil
}
func (s *testSession) IsClosed() bool {
s.mu.Lock()
defer s.mu.Unlock()
return s.closed
}
func (s *testSession) NumStreams() int {
s.mu.Lock()
defer s.mu.Unlock()
return s.streams
}
func TestWaitUntilIdlePreservesActiveStreams(t *testing.T) {
session := &testSession{streams: 1}
done := make(chan struct{})
go func() {
waitUntilIdle(session, 20*time.Millisecond, time.Millisecond)
close(done)
}()
time.Sleep(30 * time.Millisecond)
if session.IsClosed() {
t.Fatal("active session was closed")
}
session.mu.Lock()
session.streams = 0
session.mu.Unlock()
select {
case <-done:
case <-time.After(250 * time.Millisecond):
t.Fatal("idle session was not closed")
}
if !session.IsClosed() {
t.Fatal("retired session did not close after becoming idle")
}
}
+7 -1
View File
@@ -28,7 +28,13 @@ func (r *chainRegistry) Register(name string, v chain.Chainer) error {
}
func (r *chainRegistry) replace(name string, v chain.Chainer) {
r.m.Store(name, v)
old, loaded := r.m.Swap(name, v)
if !loaded {
return
}
if retirer, ok := old.(interface{ Retire() }); ok {
retirer.Retire()
}
}
func (r *chainRegistry) Get(name string) chain.Chainer {
+26
View File
@@ -16,6 +16,15 @@ func (c testChainer) Route(context.Context, string, string, ...chain.RouteOption
return c.route
}
type retiringTestChainer struct {
testChainer
retired bool
}
func (c *retiringTestChainer) Retire() {
c.retired = true
}
type testRoute struct {
nodes []*chain.Node
}
@@ -49,3 +58,20 @@ func TestReplaceChainOverwritesExistingRegistration(t *testing.T) {
t.Fatalf("expected replacement chain route, got %#v", route)
}
}
func TestReplaceChainRetiresPreviousRegistration(t *testing.T) {
name := "replace_chain_retire_tdd"
ChainRegistry().Unregister(name)
defer ChainRegistry().Unregister(name)
old := &retiringTestChainer{}
if err := ChainRegistry().Register(name, old); err != nil {
t.Fatalf("register old chain: %v", err)
}
if err := ReplaceChain(name, testChainer{}); err != nil {
t.Fatalf("replace chain: %v", err)
}
if !old.retired {
t.Fatal("previous chain was not retired")
}
}
+117
View File
@@ -0,0 +1,117 @@
package socket
import (
"net"
"sync"
"testing"
"time"
coreservice "github.com/go-gost/core/service"
"github.com/go-gost/x/config"
"github.com/go-gost/x/registry"
)
type blockingCommandService struct {
started chan struct{}
release chan struct{}
startedOnce sync.Once
}
func (s *blockingCommandService) Serve() error { return nil }
func (s *blockingCommandService) Addr() net.Addr { return nil }
func (s *blockingCommandService) Close() error {
s.startedOnce.Do(func() { close(s.started) })
<-s.release
return nil
}
func TestMutationCommandsAreSerialized(t *testing.T) {
mutations := []string{
"AddService", "UpdateService", "DeleteService", "PauseService", "ResumeService",
"AddChains", "UpdateChains", "DeleteChains",
"AddLimiters", "UpdateLimiters", "DeleteLimiters",
"AddCLimiters", "UpdateCLimiters", "DeleteCLimiters",
"SetProtocol", "UpgradeAgent", "RollbackAgent", "reload",
}
for _, command := range mutations {
if !isMutationCommand(command) {
t.Fatalf("expected %s to use the serialized mutation queue", command)
}
}
}
func TestReadOnlyCommandsRemainBoundedAsync(t *testing.T) {
for _, command := range []string{"TcpPing", "UdpPing", "ServiceMonitorCheck"} {
if isMutationCommand(command) {
t.Fatalf("expected %s to remain a read-only command", command)
}
}
}
func TestCommandResponseTypePreservesRequestContract(t *testing.T) {
if got := commandResponseType("UpdateService"); got != "UpdateServiceResponse" {
t.Fatalf("unexpected response type: %s", got)
}
if got := commandResponseType(""); got != "UnknownCommandResponse" {
t.Fatalf("unexpected empty command response type: %s", got)
}
}
func TestMutationQueueExecutesCommandsInArrivalOrder(t *testing.T) {
originalConfig := config.Global()
t.Cleanup(func() { config.Set(originalConfig) })
firstName := "mutation_queue_first_tdd"
secondName := "mutation_queue_second_tdd"
first := &blockingCommandService{started: make(chan struct{}), release: make(chan struct{})}
second := &blockingCommandService{started: make(chan struct{}), release: make(chan struct{})}
for _, name := range []string{firstName, secondName} {
registry.ServiceRegistry().Unregister(name)
}
t.Cleanup(func() {
select {
case <-first.release:
default:
close(first.release)
}
select {
case <-second.release:
default:
close(second.release)
}
registry.ServiceRegistry().Unregister(firstName)
registry.ServiceRegistry().Unregister(secondName)
})
if err := registry.ServiceRegistry().Register(firstName, coreservice.Service(first)); err != nil {
t.Fatalf("register first service: %v", err)
}
if err := registry.ServiceRegistry().Register(secondName, coreservice.Service(second)); err != nil {
t.Fatalf("register second service: %v", err)
}
config.Set(&config.Config{Services: []*config.ServiceConfig{{Name: firstName}, {Name: secondName}}})
reporter := NewWebSocketReporter("", "mutation-queue-test-secret")
go reporter.runMutationCommands()
t.Cleanup(reporter.Stop)
reporter.dispatchCommand(CommandMessage{Type: "DeleteService", Data: map[string]any{"services": []string{firstName}}})
reporter.dispatchCommand(CommandMessage{Type: "DeleteService", Data: map[string]any{"services": []string{secondName}}})
select {
case <-first.started:
case <-time.After(time.Second):
t.Fatal("first mutation did not start")
}
select {
case <-second.started:
t.Fatal("second mutation started before first mutation completed")
case <-time.After(100 * time.Millisecond):
}
close(first.release)
select {
case <-second.started:
case <-time.After(time.Second):
t.Fatal("second mutation did not start after first mutation completed")
}
close(second.release)
}
+82 -18
View File
@@ -15,6 +15,11 @@ import (
xservice "github.com/go-gost/x/service"
)
type serviceReplacement struct {
config config.ServiceConfig
oldConfig *config.ServiceConfig
}
func createServices(req createServicesRequest) error {
if len(req.Data) == 0 {
@@ -95,11 +100,11 @@ func updateServices(req updateServicesRequest) error {
req.Data[i].Name = name
}
// 第二阶段:逐个更新服务(Upsert模式:存在则更新,不存在则创建)
changedServices := make([]struct {
config config.ServiceConfig
service service.Service
}, 0, len(req.Data))
// 第二阶段:逐个更新服务(Upsert模式:存在则更新,不存在则创建)。
// 配置变更命令由 WebSocket reporter 串行调度,但这里仍保留完整回滚,
// 避免新配置解析或监听失败后旧服务永久消失。
originalConfig := config.Global()
changedServices := make([]serviceReplacement, 0, len(req.Data))
for i := range req.Data {
serviceConfig := &req.Data[i]
name := serviceConfig.Name
@@ -107,33 +112,48 @@ func updateServices(req updateServicesRequest) error {
continue
}
// 1. 获取旧服务
old := registry.ServiceRegistry().Get(name)
var oldConfig *config.ServiceConfig
if originalConfig != nil {
for _, current := range originalConfig.Services {
if current != nil && strings.TrimSpace(current.Name) == name {
oldConfig = current
break
}
}
}
// 2. 关闭旧服务 (如果存在)
if old != nil {
// 3. 从注册表移除旧服务;registry 会负责关闭旧服务。
// 1. 关闭并移除旧服务(如果存在)。同名监听必须先释放端口,
// 才能创建新 listener。
if registry.ServiceRegistry().Get(name) != nil {
registry.ServiceRegistry().Unregister(name)
}
// 4. 解析新服务配置
// 2. 解析新服务配置。
svc, err := parser.ParseService(serviceConfig)
if err != nil {
rollbackErr := restoreServiceRuntime(name, oldConfig)
rollbackErr = errors.Join(rollbackErr, rollbackServiceReplacements(changedServices))
if rollbackErr != nil {
return fmt.Errorf("create service %s failed: %v; restore previous service failed: %w", name, err, rollbackErr)
}
return errors.New("create service " + name + " failed: " + err.Error())
}
changedServices = append(changedServices, struct {
config config.ServiceConfig
service service.Service
}{*serviceConfig, svc})
// 5. 注册新服务
// 3. 注册并启动新服务。
if err := registry.ServiceRegistry().Register(name, svc); err != nil {
svc.Close()
rollbackErr := restoreServiceRuntime(name, oldConfig)
rollbackErr = errors.Join(rollbackErr, rollbackServiceReplacements(changedServices))
if rollbackErr != nil {
return fmt.Errorf("service %s already exists; restore previous service failed: %w", name, rollbackErr)
}
return errors.New("service " + name + " already exists")
}
// 6. 启动新服务
go svc.Serve()
changedServices = append(changedServices, serviceReplacement{
config: *serviceConfig,
oldConfig: oldConfig,
})
}
if len(changedServices) == 0 {
return nil
@@ -158,12 +178,56 @@ func updateServices(req updateServicesRequest) error {
}
return nil
}); err != nil {
config.Set(originalConfig)
if rollbackErr := rollbackServiceReplacements(changedServices); rollbackErr != nil {
return fmt.Errorf("%w; restore previous services failed: %v", err, rollbackErr)
}
return err
}
return nil
}
func restoreServiceRuntime(name string, serviceConfig *config.ServiceConfig) error {
name = strings.TrimSpace(name)
if name == "" || serviceConfig == nil {
return nil
}
if registry.ServiceRegistry().Get(name) != nil {
registry.ServiceRegistry().Unregister(name)
}
cfgCopy := *serviceConfig
cfgCopy.Name = name
svc, err := parser.ParseService(&cfgCopy)
if err != nil {
return err
}
if err := registry.ServiceRegistry().Register(name, svc); err != nil {
svc.Close()
return err
}
go svc.Serve()
return nil
}
func rollbackServiceReplacements(replacements []serviceReplacement) error {
var rollbackErr error
for i := len(replacements) - 1; i >= 0; i-- {
name := strings.TrimSpace(replacements[i].config.Name)
if name == "" {
continue
}
if registry.ServiceRegistry().Get(name) != nil {
registry.ServiceRegistry().Unregister(name)
}
if err := restoreServiceRuntime(name, replacements[i].oldConfig); err != nil {
rollbackErr = errors.Join(rollbackErr, fmt.Errorf("restore service %s: %w", name, err))
}
}
return rollbackErr
}
func serviceConfigUnchanged(name string, next config.ServiceConfig) bool {
cfg := config.Global()
if cfg == nil {
+39
View File
@@ -7,6 +7,8 @@ import (
corelogger "github.com/go-gost/core/logger"
"github.com/go-gost/core/service"
"github.com/go-gost/x/config"
_ "github.com/go-gost/x/handler/auto"
_ "github.com/go-gost/x/listener/tcp"
xlogger "github.com/go-gost/x/logger"
"github.com/go-gost/x/registry"
)
@@ -15,6 +17,43 @@ type recordingService struct {
closed int
}
func TestUpdateServicesParseFailureRestoresPreviousRuntime(t *testing.T) {
corelogger.SetDefault(xlogger.Nop())
name := "restore_after_failed_update_tdd"
existing := &recordingService{}
registry.ServiceRegistry().Unregister(name)
t.Cleanup(func() { registry.ServiceRegistry().Unregister(name) })
if err := registry.ServiceRegistry().Register(name, service.Service(existing)); err != nil {
t.Fatalf("register existing service: %v", err)
}
originalConfig := config.Global()
t.Cleanup(func() { config.Set(originalConfig) })
serviceConfig := config.ServiceConfig{Name: name, Addr: "127.0.0.1:0"}
config.Set(&config.Config{Services: []*config.ServiceConfig{&serviceConfig}})
invalid := serviceConfig
invalid.Listener = &config.ListenerConfig{Type: "listener-does-not-exist"}
if err := updateServices(updateServicesRequest{Data: []config.ServiceConfig{invalid}}); err == nil {
t.Fatalf("expected invalid service update to fail")
}
if existing.closed != 1 {
t.Fatalf("expected old runtime to be closed once, got %d", existing.closed)
}
if registry.ServiceRegistry().Get(name) == nil {
t.Fatalf("expected previous runtime to be restored")
}
cfg := config.Global()
if len(cfg.Services) != 1 || cfg.Services[0] == nil || cfg.Services[0].Name != name {
t.Fatalf("expected previous config to remain, got %#v", cfg.Services)
}
if cfg.Services[0].Listener != nil {
t.Fatalf("expected previous listener config to remain unchanged")
}
}
func (s *recordingService) Serve() error { return nil }
func (s *recordingService) Addr() net.Addr { return nil }
func (s *recordingService) Close() error {
+121 -5
View File
@@ -16,6 +16,7 @@ import (
"os"
"os/exec"
"runtime"
"runtime/debug"
"strconv"
"strings"
"sync"
@@ -151,6 +152,9 @@ const (
initialBackoff = 2 * time.Second // 重连初始退避
maxBackoff = 2 * time.Minute // 重连最大退避
defaultMetricReportInterval = 5 * time.Second
maxConcurrentTCPPings = 8
maxConcurrentReadCommands = 16
maxQueuedMutationCommands = 256
)
type WebSocketReporter struct {
@@ -172,6 +176,9 @@ type WebSocketReporter struct {
connecting bool // 正在连接状态
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) {
@@ -201,11 +208,37 @@ func NewWebSocketReporter(serverURL string, secret string) *WebSocketReporter {
connected: false,
connecting: false,
aesCrypto: aesCrypto,
tcpPingSem: make(chan struct{}, maxConcurrentTCPPings),
readCommandSem: make(chan struct{}, maxConcurrentReadCommands),
mutationQueue: make(chan CommandMessage, maxQueuedMutationCommands),
}
}
func (w *WebSocketReporter) tryAcquireTCPPingSlot() bool {
if w == nil || w.tcpPingSem == nil {
return false
}
select {
case w.tcpPingSem <- struct{}{}:
return true
default:
return false
}
}
func (w *WebSocketReporter) releaseTCPPingSlot() {
if w == nil || w.tcpPingSem == nil {
return
}
select {
case <-w.tcpPingSem:
default:
}
}
// Start 启动WebSocket报告器
func (w *WebSocketReporter) Start() {
go w.runMutationCommands()
go w.run()
}
@@ -752,8 +785,7 @@ func (w *WebSocketReporter) handleReceivedMessage(messageType int, message []byt
}
if cmdMsg.Type != "call" {
// 所有命令统一异步执行,避免阻塞消息接收循环
go w.routeCommand(cmdMsg)
w.dispatchCommand(cmdMsg)
}
} else {
// 处理普通消息
@@ -764,8 +796,7 @@ func (w *WebSocketReporter) handleReceivedMessage(messageType int, message []byt
return
}
if cmdMsg.Type != "call" {
// 所有命令统一异步执行,避免阻塞消息接收循环
go w.routeCommand(cmdMsg)
w.dispatchCommand(cmdMsg)
}
}
@@ -774,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)
@@ -840,9 +951,14 @@ func (w *WebSocketReporter) routeCommand(cmd CommandMessage) {
// TCP Ping 诊断命令(只读,不需要保存配置)
case "TcpPing":
response.Type = "TcpPingResponse"
if !w.tryAcquireTCPPingSlot() {
err = fmt.Errorf("TCP探测任务过多,请稍后重试")
break
}
defer w.releaseTCPPingSlot()
var tcpPingResult TcpPingResponse
tcpPingResult, err = w.handleTcpPing(cmd.Data)
response.Type = "TcpPingResponse"
response.Data = tcpPingResult
// needSaveConfig = false (默认值)
@@ -148,6 +148,25 @@ func TestNewWebSocketReporterUsesReducedMetricInterval(t *testing.T) {
}
}
func TestWebSocketReporterLimitsConcurrentTCPPings(t *testing.T) {
reporter := &WebSocketReporter{tcpPingSem: make(chan struct{}, maxConcurrentTCPPings)}
for i := 0; i < maxConcurrentTCPPings; i++ {
if !reporter.tryAcquireTCPPingSlot() {
t.Fatalf("expected TCP ping slot %d to be available", i)
}
}
if reporter.tryAcquireTCPPingSlot() {
t.Fatalf("expected TCP ping concurrency limit at %d", maxConcurrentTCPPings)
}
for i := 0; i < maxConcurrentTCPPings; i++ {
reporter.releaseTCPPingSlot()
}
if !reporter.tryAcquireTCPPingSlot() {
t.Fatalf("expected released TCP ping slot to be reusable")
}
reporter.releaseTCPPingSlot()
}
func TestFormatWebSocketDialErrorIncludesHTTPStatus(t *testing.T) {
err := errors.New("websocket: bad handshake")
resp := &http.Response{
+290 -42
View File
@@ -1,4 +1,31 @@
#!/bin/bash
#!/bin/sh
# shellcheck shell=bash
# Alpine 默认不带 Bash。先用系统自带的 /bin/sh 安装/切换到 Bash,
# 后续主体继续使用 Bash 语法,避免要求用户手动准备运行环境。
if [ -z "${BASH_VERSION:-}" ]; then
if command -v bash >/dev/null 2>&1; then
exec bash "$0" "$@"
fi
if [ -f /etc/alpine-release ] && command -v apk >/dev/null 2>&1; then
if [ "$(id -u)" -eq 0 ]; then
apk add --no-cache bash
elif command -v sudo >/dev/null 2>&1; then
sudo apk add --no-cache bash
elif command -v doas >/dev/null 2>&1; then
doas apk add --no-cache bash
else
echo "❌ Alpine 安装需要 root 权限,或已配置 sudo/doas。" >&2
exit 1
fi
exec bash "$0" "$@"
fi
echo "❌ 此安装脚本需要 Bash。" >&2
exit 1
fi
# GitHub repo used for release downloads
REPO="Sagit-chu/flux-panel"
@@ -24,16 +51,51 @@ get_architecture() {
# 安装目录
INSTALL_DIR="/etc/flux_agent"
FLUX_AGENT_SYSTEMD_SERVICE_FILE="/etc/systemd/system/flux_agent.service"
FLUX_AGENT_OPENRC_SERVICE_FILE="/etc/init.d/flux_agent"
LEGACY_GOST_BINARY="/usr/local/bin/gost"
LEGACY_GOST_CONFIG_DIR="/etc/gost"
LEGACY_GOST_SERVICE_FILE_ETC="/etc/systemd/system/gost.service"
LEGACY_GOST_SERVICE_FILE_LIB="/lib/systemd/system/gost.service"
LEGACY_GOST_SERVICE_FILE_USR_LIB="/usr/lib/systemd/system/gost.service"
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"
@@ -139,6 +201,8 @@ build_download_url() {
}
ensure_download_url_initialized() {
ensure_alpine_runtime_dependencies || return 1
if [[ -n "${DOWNLOAD_URL:-}" ]]; then
return 0
fi
@@ -256,6 +320,205 @@ write_flux_agent_config() {
"$(json_escape "$SECRET")" > "$path"
}
ensure_service_manager() {
if [[ -n "$SERVICE_MANAGER" ]]; then
case "$SERVICE_MANAGER" in
systemd|openrc)
return 0
;;
*)
echo "❌ 不支持的服务管理器: $SERVICE_MANAGER" >&2
return 1
;;
esac
fi
# 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
elif command -v rc-service >/dev/null 2>&1 && command -v rc-update >/dev/null 2>&1; then
SERVICE_MANAGER="openrc"
return 0
fi
echo "❌ 未检测到受支持的服务管理器(systemd 或 OpenRC)。" >&2
return 1
}
flux_agent_service_exists() {
ensure_service_manager || return 1
case "$SERVICE_MANAGER" in
systemd)
[[ -f "$FLUX_AGENT_SYSTEMD_SERVICE_FILE" ]] || \
systemctl list-units --full -all 2>/dev/null | grep -Fq "flux_agent.service"
;;
openrc)
[[ -f "$FLUX_AGENT_OPENRC_SERVICE_FILE" ]]
;;
esac
}
stop_flux_agent_service() {
ensure_service_manager || return 1
case "$SERVICE_MANAGER" in
systemd)
systemctl stop flux_agent 2>/dev/null || true
;;
openrc)
rc-service flux_agent stop 2>/dev/null || true
;;
esac
}
disable_flux_agent_service() {
ensure_service_manager || return 1
case "$SERVICE_MANAGER" in
systemd)
systemctl disable flux_agent 2>/dev/null || true
;;
openrc)
rc-update del flux_agent default 2>/dev/null || true
;;
esac
}
write_flux_agent_service() {
ensure_service_manager || return 1
case "$SERVICE_MANAGER" in
systemd)
mkdir -p "$(dirname "$FLUX_AGENT_SYSTEMD_SERVICE_FILE")"
cat > "$FLUX_AGENT_SYSTEMD_SERVICE_FILE" <<EOF
[Unit]
Description=Flux_agent Proxy Service
After=network.target
[Service]
WorkingDirectory=$INSTALL_DIR
ExecStart=$INSTALL_DIR/flux_agent
Environment=GODEBUG=disablethp=1
Restart=on-failure
StandardOutput=null
StandardError=null
[Install]
WantedBy=multi-user.target
EOF
;;
openrc)
mkdir -p "$(dirname "$FLUX_AGENT_OPENRC_SERVICE_FILE")"
cat > "$FLUX_AGENT_OPENRC_SERVICE_FILE" <<EOF
#!/sbin/openrc-run
name="flux_agent"
description="Flux_agent Proxy Service"
command="$INSTALL_DIR/flux_agent"
directory="$INSTALL_DIR"
command_background="yes"
pidfile="/run/\${RC_SVCNAME}.pid"
output_log="/dev/null"
error_log="/dev/null"
depend() {
need net
}
EOF
chmod +x "$FLUX_AGENT_OPENRC_SERVICE_FILE"
;;
esac
}
enable_and_start_flux_agent_service() {
ensure_service_manager || return 1
case "$SERVICE_MANAGER" in
systemd)
systemctl daemon-reload
systemctl enable flux_agent
systemctl start flux_agent
;;
openrc)
rc-update add flux_agent default
rc-service flux_agent start
;;
esac
}
start_flux_agent_service() {
ensure_service_manager || return 1
case "$SERVICE_MANAGER" in
systemd)
systemctl start flux_agent
;;
openrc)
rc-service flux_agent start
;;
esac
}
flux_agent_service_is_active() {
ensure_service_manager || return 1
case "$SERVICE_MANAGER" in
systemd)
systemctl is-active --quiet flux_agent
;;
openrc)
rc-service flux_agent status >/dev/null 2>&1
;;
esac
}
flux_agent_service_status() {
ensure_service_manager || return 1
case "$SERVICE_MANAGER" in
systemd)
systemctl is-active flux_agent 2>/dev/null || true
;;
openrc)
rc-service flux_agent status 2>/dev/null || true
;;
esac
}
remove_flux_agent_service() {
ensure_service_manager || return 1
case "$SERVICE_MANAGER" in
systemd)
rm -f "$FLUX_AGENT_SYSTEMD_SERVICE_FILE"
systemctl daemon-reload 2>/dev/null || true
;;
openrc)
rm -f "$FLUX_AGENT_OPENRC_SERVICE_FILE"
;;
esac
}
flux_agent_service_status_hint() {
ensure_service_manager || return 1
case "$SERVICE_MANAGER" in
systemd)
echo "systemctl status flux_agent --no-pager"
;;
openrc)
echo "rc-service flux_agent status"
;;
esac
}
cleanup_legacy_gost_installation() {
local matched_service_files=()
local service_file=""
@@ -278,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
@@ -297,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
}
@@ -341,19 +605,21 @@ install_flux_agent() {
get_config_params
# 检查并安装 tcpkill
# 检查并安装 tcpkill
check_and_install_tcpkill
ensure_service_manager || exit 1
mkdir -p "$INSTALL_DIR"
local tmp_binary="$INSTALL_DIR/flux_agent.new"
# 停止并禁用已有服务
if systemctl list-units --full -all | grep -Fq "flux_agent.service"; then
if flux_agent_service_exists; then
echo "🔍 检测到已存在的flux_agent服务"
systemctl stop flux_agent 2>/dev/null && echo "🛑 停止服务"
systemctl disable flux_agent 2>/dev/null && echo "🚫 禁用自启"
stop_flux_agent_service
echo "🛑 停止服务"
disable_flux_agent_service
echo "🚫 禁用自启"
fi
# 下载 flux_agent
@@ -392,38 +658,21 @@ EOF
# 加强权限
chmod 600 "$INSTALL_DIR"/*.json
# 创建 systemd 服务
SERVICE_FILE="/etc/systemd/system/flux_agent.service"
cat > "$SERVICE_FILE" <<EOF
[Unit]
Description=Flux_agent Proxy Service
After=network.target
[Service]
WorkingDirectory=$INSTALL_DIR
ExecStart=$INSTALL_DIR/flux_agent
Restart=on-failure
StandardOutput=null
StandardError=null
[Install]
WantedBy=multi-user.target
EOF
# 创建 systemd 或 OpenRC 服务
write_flux_agent_service
# 启动服务
systemctl daemon-reload
systemctl enable flux_agent
systemctl start flux_agent
enable_and_start_flux_agent_service
# 检查状态
echo "🔄 检查服务状态..."
if systemctl is-active --quiet flux_agent; then
if flux_agent_service_is_active; then
echo "✅ 安装完成,flux_agent服务已启动并设置为开机启动。"
echo "📁 配置目录: $INSTALL_DIR"
echo "🔧 服务状态: $(systemctl is-active flux_agent)"
echo "🔧 服务状态: $(flux_agent_service_status)"
else
echo "❌ flux_agent服务启动失败,请执行以下命令查看状态:"
echo "systemctl status flux_agent --no-pager"
flux_agent_service_status_hint
fi
}
@@ -443,6 +692,7 @@ update_flux_agent() {
# 检查并安装 tcpkill
check_and_install_tcpkill
ensure_service_manager || return 1
# 先下载新版本
echo "⬇️ 下载最新版本..."
@@ -455,9 +705,9 @@ update_flux_agent() {
cleanup_legacy_gost_installation
# 停止服务
if systemctl list-units --full -all | grep -Fq "flux_agent.service"; then
if flux_agent_service_exists; then
echo "🛑 停止 flux_agent 服务..."
systemctl stop flux_agent
stop_flux_agent_service
fi
# 替换文件
@@ -469,7 +719,7 @@ update_flux_agent() {
# 重启服务
echo "🔄 重启服务..."
systemctl start flux_agent
start_flux_agent_service
echo "✅ 更新完成,服务已重新启动。"
}
@@ -477,6 +727,7 @@ update_flux_agent() {
# 卸载功能
uninstall_flux_agent() {
echo "🗑️ 开始卸载 flux_agent..."
ensure_service_manager || return 1
read -p "确认卸载 flux_agent 吗?此操作将删除所有相关文件 (y/N): " confirm
if [[ "$confirm" != "y" && "$confirm" != "Y" ]]; then
@@ -485,15 +736,15 @@ uninstall_flux_agent() {
fi
# 停止并禁用服务
if systemctl list-units --full -all | grep -Fq "flux_agent.service"; then
if flux_agent_service_exists; then
echo "🛑 停止并禁用服务..."
systemctl stop flux_agent 2>/dev/null
systemctl disable flux_agent 2>/dev/null
stop_flux_agent_service
disable_flux_agent_service
fi
# 删除服务文件
if [[ -f "/etc/systemd/system/flux_agent.service" ]]; then
rm -f "/etc/systemd/system/flux_agent.service"
if [[ -f "$FLUX_AGENT_SYSTEMD_SERVICE_FILE" || -f "$FLUX_AGENT_OPENRC_SERVICE_FILE" ]]; then
remove_flux_agent_service
echo "🧹 删除服务文件"
fi
@@ -503,9 +754,6 @@ uninstall_flux_agent() {
echo "🧹 删除安装目录: $INSTALL_DIR"
fi
# 重载 systemd
systemctl daemon-reload
echo "✅ 卸载完成"
}
+319 -48
View File
@@ -16,6 +16,9 @@ PINNED_VERSION=""
# 镜像加速配置(可由面板传入或交互式询问)
PROXY_ENABLED="${PROXY_ENABLED:-}"
PROXY_URL="${PROXY_URL:-}"
DEFAULT_PANEL_BACKEND_CONTAINER="flux-panel-backend"
DEFAULT_PANEL_POSTGRES_CONTAINER="flux-panel-postgres"
PANEL_SCRIPT_PATH="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)/$(basename "${BASH_SOURCE[0]}")"
# 镜像加速
maybe_proxy_url() {
@@ -304,8 +307,8 @@ get_env_var() {
get_current_db_type() {
local db_type database_url
db_type=$(get_env_var "DB_TYPE")
database_url=$(get_env_var "DATABASE_URL")
db_type=$(get_env_var "DB_TYPE" || true)
database_url=$(get_env_var "DATABASE_URL" || true)
if [[ "$db_type" == "sqlite" ]]; then
echo "sqlite"
@@ -316,13 +319,289 @@ get_current_db_type() {
fi
}
get_container_compose_label() {
local container="$1"
local label="$2"
local value
value=$(docker inspect -f "{{ index .Config.Labels \"$label\" }}" "$container" 2>/dev/null || true)
if [[ "$value" == "<no value>" ]]; then
value=""
fi
printf '%s' "$value"
}
resolve_panel_deployment() {
local backend_container="${PANEL_BACKEND_CONTAINER:-$DEFAULT_PANEL_BACKEND_CONTAINER}"
local requested_dir="${PANEL_DEPLOY_DIR:-}"
local label_dir label_project deploy_dir project_name
label_dir=$(get_container_compose_label "$backend_container" "com.docker.compose.project.working_dir")
label_project=$(get_container_compose_label "$backend_container" "com.docker.compose.project")
if [[ -n "$requested_dir" ]]; then
deploy_dir="$requested_dir"
elif [[ -n "$label_dir" ]]; then
deploy_dir="$label_dir"
elif [[ -f ".env" && -f "docker-compose.yml" ]]; then
deploy_dir=$(pwd -P)
else
echo "❌ 无法识别面板部署目录:未找到容器 Compose 标签,当前目录也没有完整部署配置"
return 1
fi
if [[ ! -d "$deploy_dir" ]]; then
echo "❌ 面板部署目录不存在:$deploy_dir"
return 1
fi
deploy_dir=$(cd "$deploy_dir" && pwd -P)
if [[ -n "$label_dir" && -d "$label_dir" ]]; then
label_dir=$(cd "$label_dir" && pwd -P)
fi
if [[ -n "$requested_dir" && -n "$label_dir" && "$deploy_dir" != "$label_dir" ]]; then
echo "❌ 指定部署目录与运行中容器标签不一致:$deploy_dir != $label_dir"
return 1
fi
project_name="${COMPOSE_PROJECT_NAME:-}"
if [[ -z "$project_name" && -n "$label_project" ]]; then
project_name="$label_project"
fi
if [[ -n "$project_name" && ! "$project_name" =~ ^[a-z0-9][a-z0-9_-]*$ ]]; then
echo "❌ Compose 项目名不合法:$project_name"
return 1
fi
PANEL_DEPLOY_DIR="$deploy_dir"
PANEL_BACKEND_CONTAINER="$backend_container"
PANEL_POSTGRES_CONTAINER="${PANEL_POSTGRES_CONTAINER:-$DEFAULT_PANEL_POSTGRES_CONTAINER}"
PANEL_COMPOSE_PROJECT="$project_name"
export PANEL_DEPLOY_DIR PANEL_BACKEND_CONTAINER PANEL_POSTGRES_CONTAINER PANEL_COMPOSE_PROJECT
if [[ -n "$project_name" ]]; then
COMPOSE_PROJECT_NAME="$project_name"
export COMPOSE_PROJECT_NAME
fi
cd "$PANEL_DEPLOY_DIR"
echo "📁 部署目录:$PANEL_DEPLOY_DIR"
if [[ -n "$PANEL_COMPOSE_PROJECT" ]]; then
echo "📦 Compose 项目:$PANEL_COMPOSE_PROJECT"
fi
}
run_panel_compose() {
$DOCKER_CMD "$@"
}
validate_panel_update_environment() {
local key value backend_port frontend_port configured_db_type db_type database_url postgres_password
if [[ ! -f ".env" ]]; then
echo "❌ 部署目录缺少 .env,更新终止"
return 1
fi
if [[ ! -f "docker-compose.yml" ]]; then
echo "❌ 部署目录缺少 docker-compose.yml,更新终止"
return 1
fi
for key in JWT_SECRET BACKEND_PORT FRONTEND_PORT; do
value=$(get_env_var "$key" || true)
if [[ -z "${value//[[:space:]]/}" ]]; then
echo "❌ .env 中 $key 缺失或为空,更新终止"
return 1
fi
done
backend_port=$(get_env_var "BACKEND_PORT" || true)
frontend_port=$(get_env_var "FRONTEND_PORT" || true)
if [[ ! "$backend_port" =~ ^[0-9]+$ || "$backend_port" -lt 1 || "$backend_port" -gt 65535 ]]; then
echo "❌ BACKEND_PORT 不是有效端口:$backend_port"
return 1
fi
if [[ ! "$frontend_port" =~ ^[0-9]+$ || "$frontend_port" -lt 1 || "$frontend_port" -gt 65535 ]]; then
echo "❌ FRONTEND_PORT 不是有效端口:$frontend_port"
return 1
fi
configured_db_type=$(get_env_var "DB_TYPE" || true)
if [[ -n "$configured_db_type" && "$configured_db_type" != "sqlite" && "$configured_db_type" != "postgres" ]]; then
echo "❌ DB_TYPE 仅支持 sqlite 或 postgres:$configured_db_type"
return 1
fi
db_type=$(get_current_db_type)
if [[ "$db_type" == "postgres" ]]; then
database_url=$(get_env_var "DATABASE_URL" || true)
postgres_password=$(get_env_var "POSTGRES_PASSWORD" || true)
if [[ -z "${database_url//[[:space:]]/}" || -z "${postgres_password//[[:space:]]/}" ]]; then
echo "❌ PostgreSQL 模式要求 DATABASE_URL 和 POSTGRES_PASSWORD 均非空"
return 1
fi
fi
if ! run_panel_compose -f docker-compose.yml config -q; then
echo "❌ 当前 Compose 配置校验失败,更新终止"
return 1
fi
}
download_panel_update_compose() {
local compose_url="$1"
UPDATE_COMPOSE_CANDIDATE=$(mktemp "$PANEL_DEPLOY_DIR/.docker-compose.yml.update.XXXXXX")
if ! curl -fL -o "$UPDATE_COMPOSE_CANDIDATE" "$compose_url"; then
rm -f "$UPDATE_COMPOSE_CANDIDATE"
UPDATE_COMPOSE_CANDIDATE=""
echo "❌ 下载最新 Compose 配置失败,更新终止"
return 1
fi
if [[ ! -s "$UPDATE_COMPOSE_CANDIDATE" ]]; then
rm -f "$UPDATE_COMPOSE_CANDIDATE"
UPDATE_COMPOSE_CANDIDATE=""
echo "❌ 下载的 Compose 配置为空,更新终止"
return 1
fi
}
validate_panel_update_compose() {
if ! FLUX_VERSION="$LATEST_VERSION" run_panel_compose -f "$UPDATE_COMPOSE_CANDIDATE" config -q; then
echo "❌ 新 Compose 配置校验失败,更新终止"
return 1
fi
}
backup_sqlite_for_update() (
local destination="$1"
local running paused="false"
cleanup_paused_backend() {
if [[ "$paused" == "true" ]]; then
docker unpause "$PANEL_BACKEND_CONTAINER" >/dev/null 2>&1 || true
fi
}
trap cleanup_paused_backend EXIT
running=$(docker inspect -f '{{.State.Running}}' "$PANEL_BACKEND_CONTAINER" 2>/dev/null || true)
if [[ "$running" == "true" ]]; then
if ! docker pause "$PANEL_BACKEND_CONTAINER" >/dev/null; then
echo "❌ 暂停后端以创建一致性 SQLite 备份失败"
return 1
fi
paused="true"
fi
mkdir -p "$destination"
if ! docker cp "$PANEL_BACKEND_CONTAINER:/app/data/." "$destination"; then
echo "❌ SQLite 数据备份失败"
return 1
fi
if [[ "$paused" == "true" ]] && ! docker unpause "$PANEL_BACKEND_CONTAINER" >/dev/null; then
echo "❌ SQLite 备份完成,但恢复后端运行失败"
return 1
fi
paused="false"
)
backup_postgres_for_update() {
local destination="$1"
local postgres_db postgres_user
postgres_db=$(get_env_var "POSTGRES_DB" || true)
postgres_user=$(get_env_var "POSTGRES_USER" || true)
postgres_db=${postgres_db:-flux_panel}
postgres_user=${postgres_user:-flux_panel}
if ! docker exec "$PANEL_POSTGRES_CONTAINER" pg_dump -U "$postgres_user" "$postgres_db" > "$destination"; then
rm -f "$destination"
echo "❌ PostgreSQL 数据备份失败"
return 1
fi
}
create_panel_update_backup() {
local db_type="$1"
local timestamp
timestamp=$(date '+%Y%m%d-%H%M%S')
if ! (umask 077 && mkdir -p "$PANEL_DEPLOY_DIR/backups"); then
echo "❌ 创建更新备份目录失败"
return 1
fi
UPDATE_BACKUP_DIR=$(umask 077 && mktemp -d "$PANEL_DEPLOY_DIR/backups/panel-update-$timestamp.XXXXXX") || {
echo "❌ 创建更新备份目录失败"
return 1
}
if ! cp -p .env docker-compose.yml "$UPDATE_BACKUP_DIR/"; then
echo "❌ 备份面板配置失败"
return 1
fi
if [[ "$db_type" == "postgres" ]]; then
backup_postgres_for_update "$UPDATE_BACKUP_DIR/postgres.sql" || return 1
else
backup_sqlite_for_update "$UPDATE_BACKUP_DIR/sqlite" || return 1
fi
echo "💾 更新备份:$UPDATE_BACKUP_DIR"
}
pull_panel_update_images() {
local db_type="$1"
if [[ "$db_type" == "postgres" ]]; then
FLUX_VERSION="$LATEST_VERSION" run_panel_compose -f "$UPDATE_COMPOSE_CANDIDATE" pull backend frontend postgres
else
FLUX_VERSION="$LATEST_VERSION" run_panel_compose -f "$UPDATE_COMPOSE_CANDIDATE" pull backend frontend
fi
}
activate_panel_update_files() {
if ! chmod 0644 "$UPDATE_COMPOSE_CANDIDATE" || ! mv "$UPDATE_COMPOSE_CANDIDATE" docker-compose.yml; then
echo "❌ 替换 Compose 配置失败"
return 1
fi
UPDATE_COMPOSE_CANDIDATE=""
if ! upsert_env_var ".env" "FLUX_VERSION" "$LATEST_VERSION"; then
echo "❌ 更新版本配置失败"
return 1
fi
}
start_panel_after_update() {
local db_type="$1"
if [[ "$db_type" == "postgres" ]]; then
run_panel_compose up -d postgres || return 1
wait_for_postgres_healthy || return 1
fi
run_panel_compose up -d --force-recreate --remove-orphans backend frontend || return 1
wait_for_backend_healthy
}
rollback_panel_update() {
local db_type="$1"
echo "↩️ 正在恢复更新前配置..."
if ! cp -p "$UPDATE_BACKUP_DIR/.env" .env || ! cp -p "$UPDATE_BACKUP_DIR/docker-compose.yml" docker-compose.yml; then
echo "❌ 配置回滚失败,请从 $UPDATE_BACKUP_DIR 手动恢复"
return 1
fi
if [[ "$db_type" == "postgres" ]]; then
run_panel_compose up -d postgres || return 1
wait_for_postgres_healthy || return 1
fi
run_panel_compose up -d --force-recreate --remove-orphans backend frontend || return 1
wait_for_backend_healthy
}
wait_for_postgres_healthy() {
local pg_health
local postgres_container="${PANEL_POSTGRES_CONTAINER:-$DEFAULT_PANEL_POSTGRES_CONTAINER}"
echo "🔍 检查 PostgreSQL 服务状态..."
for i in {1..90}; do
if docker ps --format "{{.Names}}" | grep -q "^flux-panel-postgres$"; then
pg_health=$(docker inspect -f '{{.State.Health.Status}}' flux-panel-postgres 2>/dev/null || echo "unknown")
if docker ps --format "{{.Names}}" | grep -Fxq "$postgres_container"; then
pg_health=$(docker inspect -f '{{.State.Health.Status}}' "$postgres_container" 2>/dev/null || echo "unknown")
if [[ "$pg_health" == "healthy" ]]; then
echo "✅ PostgreSQL 服务健康检查通过"
return 0
@@ -335,7 +614,7 @@ wait_for_postgres_healthy() {
if [ $i -eq 90 ]; then
echo "❌ PostgreSQL 启动超时(90秒)"
echo "🔍 当前状态:$(docker inspect -f '{{.State.Health.Status}}' flux-panel-postgres 2>/dev/null || echo '容器不存在')"
echo "🔍 当前状态:$(docker inspect -f '{{.State.Health.Status}}' "$postgres_container" 2>/dev/null || echo '容器不存在')"
return 1
fi
@@ -348,11 +627,12 @@ wait_for_postgres_healthy() {
wait_for_backend_healthy() {
local backend_health
local backend_container="${PANEL_BACKEND_CONTAINER:-$DEFAULT_PANEL_BACKEND_CONTAINER}"
echo "🔍 检查后端服务状态..."
for i in {1..90}; do
if docker ps --format "{{.Names}}" | grep -q "^flux-panel-backend$"; then
backend_health=$(docker inspect -f '{{.State.Health.Status}}' flux-panel-backend 2>/dev/null || echo "unknown")
if docker ps --format "{{.Names}}" | grep -Fxq "$backend_container"; then
backend_health=$(docker inspect -f '{{.State.Health.Status}}' "$backend_container" 2>/dev/null || echo "unknown")
if [[ "$backend_health" == "healthy" ]]; then
echo "✅ 后端服务健康检查通过"
return 0
@@ -365,7 +645,7 @@ wait_for_backend_healthy() {
if [ $i -eq 90 ]; then
echo "❌ 后端服务启动超时(90秒)"
echo "🔍 当前状态:$(docker inspect -f '{{.State.Health.Status}}' flux-panel-backend 2>/dev/null || echo '容器不存在')"
echo "🔍 当前状态:$(docker inspect -f '{{.State.Health.Status}}' "$backend_container" 2>/dev/null || echo '容器不存在')"
return 1
fi
@@ -380,9 +660,8 @@ wait_for_backend_healthy() {
delete_self() {
echo ""
echo "🗑️ 操作已完成,正在清理脚本文件..."
SCRIPT_PATH="$(readlink -f "$0" 2>/dev/null || realpath "$0" 2>/dev/null || echo "$0")"
sleep 1
rm -f "$SCRIPT_PATH" && echo "✅ 脚本文件已删除" || echo "❌ 删除脚本文件失败"
rm -f "$PANEL_SCRIPT_PATH" && echo "✅ 脚本文件已删除" || echo "❌ 删除脚本文件失败"
}
@@ -486,12 +765,11 @@ EOF
# 更新功能
update_panel() {
echo "🔄 开始更新面板..."
ask_proxy_config
check_docker
resolve_panel_deployment || return 1
ask_proxy_config
validate_panel_update_environment || return 1
if [[ ! -f ".env" ]]; then
echo "⚠️ 未找到 .env,默认按 SQLite 模式更新"
fi
CURRENT_DB_TYPE=$(get_current_db_type)
echo "🗄️ 当前数据库类型:$CURRENT_DB_TYPE"
@@ -502,52 +780,45 @@ update_panel() {
}
echo "🆕 最新版本:$LATEST_VERSION"
set_compose_urls_by_version "$LATEST_VERSION"
upsert_env_var ".env" "FLUX_VERSION" "$LATEST_VERSION"
echo "🔽 下载最新配置文件..."
DOCKER_COMPOSE_URL=$(get_docker_compose_url)
echo "📡 选择配置文件:$(basename "$DOCKER_COMPOSE_URL")"
curl -L -o docker-compose.yml "$DOCKER_COMPOSE_URL"
echo "✅ 下载完成"
# 自动检测并配置 IPv6 支持
if check_ipv6_support; then
echo "🚀 系统支持 IPv6,自动启用 IPv6 配置..."
configure_docker_ipv6
download_panel_update_compose "$DOCKER_COMPOSE_URL" || return 1
if ! validate_panel_update_compose; then
rm -f "$UPDATE_COMPOSE_CANDIDATE"
UPDATE_COMPOSE_CANDIDATE=""
return 1
fi
# 先发送 SIGTERM 信号,让应用优雅关闭
docker stop -t 30 flux-panel-backend 2>/dev/null || true
docker stop -t 10 vite-frontend 2>/dev/null || true
# 等待 WAL 文件同步
echo "⏳ 等待数据同步..."
sleep 5
# 然后再完全停止
$DOCKER_CMD down
echo "💾 备份当前配置和数据库..."
if ! create_panel_update_backup "$CURRENT_DB_TYPE"; then
rm -f "$UPDATE_COMPOSE_CANDIDATE"
UPDATE_COMPOSE_CANDIDATE=""
return 1
fi
echo "⬇️ 拉取最新镜像..."
if [[ "$CURRENT_DB_TYPE" == "postgres" ]]; then
$DOCKER_CMD pull backend frontend postgres
else
$DOCKER_CMD pull backend frontend
if ! pull_panel_update_images "$CURRENT_DB_TYPE"; then
rm -f "$UPDATE_COMPOSE_CANDIDATE"
UPDATE_COMPOSE_CANDIDATE=""
echo "❌ 镜像拉取失败,现有服务保持运行"
return 1
fi
echo "🚀 启动更新后的服务..."
if [[ "$CURRENT_DB_TYPE" == "postgres" ]]; then
$DOCKER_CMD up -d postgres
wait_for_postgres_healthy
$DOCKER_CMD up -d backend frontend
else
$DOCKER_CMD up -d backend frontend
if ! activate_panel_update_files; then
rollback_panel_update "$CURRENT_DB_TYPE" || true
return 1
fi
# 等待服务启动
echo "⏳ 等待服务启动..."
if ! wait_for_backend_healthy; then
echo "🛑 更新终止"
echo "🚀 原地重建更新后的服务..."
if ! start_panel_after_update "$CURRENT_DB_TYPE"; then
echo "🛑 新版本启动失败"
if rollback_panel_update "$CURRENT_DB_TYPE"; then
echo "✅ 已恢复更新前版本"
else
echo "❌ 自动回滚失败,请从 $UPDATE_BACKUP_DIR 手动恢复"
fi
return 1
fi
+262 -7
View File
@@ -90,6 +90,7 @@ test_update_flux_agent_asks_for_proxy_config() (
set -euo pipefail
load_script_without_main "$ROOT_DIR/install.sh"
SERVICE_MANAGER="systemd"
INSTALL_DIR=$(mktemp -d)
cat > "$INSTALL_DIR/flux_agent" <<'EOF'
#!/bin/bash
@@ -165,6 +166,7 @@ test_install_flux_agent_preserves_legacy_gost_when_download_fails() (
set -euo pipefail
load_script_without_main "$ROOT_DIR/install.sh"
SERVICE_MANAGER="systemd"
INSTALL_DIR=$(mktemp -d)
cat > "$INSTALL_DIR/flux_agent" <<'EOF'
#!/bin/bash
@@ -199,6 +201,7 @@ test_update_flux_agent_preserves_legacy_gost_when_download_fails() (
set -euo pipefail
load_script_without_main "$ROOT_DIR/install.sh"
SERVICE_MANAGER="systemd"
INSTALL_DIR=$(mktemp -d)
cat > "$INSTALL_DIR/flux_agent" <<'EOF'
#!/bin/bash
@@ -236,7 +239,9 @@ test_install_flux_agent_writes_json_safe_config() (
set -euo pipefail
load_script_without_main "$ROOT_DIR/install.sh"
SERVICE_MANAGER="systemd"
INSTALL_DIR=$(mktemp -d)
FLUX_AGENT_SYSTEMD_SERVICE_FILE="$INSTALL_DIR/flux_agent.service"
SERVER_ADDR='panel"addr'
SECRET='sec\ret"1'
DOWNLOAD_URL="https://example.com/gost"
@@ -272,6 +277,119 @@ 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() (
set -euo pipefail
local shebang
shebang=$(head -n 1 "$ROOT_DIR/install.sh")
assert_equals "#!/bin/sh" "$shebang" "install.sh should start with Alpine's default shell"
grep -Fq 'apk add --no-cache bash' "$ROOT_DIR/install.sh" || \
fail "install.sh should bootstrap Bash through apk on Alpine"
)
test_install_flux_agent_uses_openrc() (
set -euo pipefail
load_script_without_main "$ROOT_DIR/install.sh"
local temp_root
temp_root=$(mktemp -d)
INSTALL_DIR="$temp_root/flux_agent"
FLUX_AGENT_OPENRC_SERVICE_FILE="$temp_root/init.d/flux_agent"
SERVICE_MANAGER="openrc"
SERVER_ADDR="panel.example.com:443"
SECRET="secret"
DOWNLOAD_URL="https://example.com/gost"
local rc_service_calls=""
local rc_update_calls=""
ask_proxy_config() { :; }
ensure_download_url_initialized() { :; }
get_config_params() { :; }
check_and_install_tcpkill() { :; }
cleanup_legacy_gost_installation() { :; }
curl() {
local output=""
while [[ $# -gt 0 ]]; do
if [[ "$1" == "-o" ]]; then
output="$2"
shift 2
continue
fi
shift
done
cat > "$output" <<'EOF'
#!/bin/sh
echo "new version"
EOF
chmod +x "$output"
}
rc-service() {
rc_service_calls+=$'\n'"$*"
if [[ "$2" == "status" ]]; then
echo "status: started"
fi
return 0
}
rc-update() {
rc_update_calls+=$'\n'"$*"
return 0
}
install_flux_agent >/dev/null
[[ -x "$FLUX_AGENT_OPENRC_SERVICE_FILE" ]] || fail "OpenRC service file should be executable"
grep -Fq '#!/sbin/openrc-run' "$FLUX_AGENT_OPENRC_SERVICE_FILE" || \
fail "OpenRC service should use openrc-run"
grep -Fq "command=\"$INSTALL_DIR/flux_agent\"" "$FLUX_AGENT_OPENRC_SERVICE_FILE" || \
fail "OpenRC service should launch the installed flux_agent binary"
grep -Fq 'command_background="yes"' "$FLUX_AGENT_OPENRC_SERVICE_FILE" || \
fail "OpenRC service should run flux_agent in the background"
if command -v openrc-run >/dev/null 2>&1; then
"$FLUX_AGENT_OPENRC_SERVICE_FILE" describe >/dev/null 2>&1
fi
[[ "$rc_update_calls" == *"add flux_agent default"* ]] || \
fail "OpenRC install should enable flux_agent in the default runlevel"
[[ "$rc_service_calls" == *"start"* ]] || fail "OpenRC install should start flux_agent"
[[ "$rc_service_calls" == *"status"* ]] || fail "OpenRC install should verify flux_agent status"
)
test_remove_flux_agent_service_uses_openrc() (
set -euo pipefail
load_script_without_main "$ROOT_DIR/install.sh"
local temp_root
temp_root=$(mktemp -d)
SERVICE_MANAGER="openrc"
FLUX_AGENT_OPENRC_SERVICE_FILE="$temp_root/init.d/flux_agent"
mkdir -p "$(dirname "$FLUX_AGENT_OPENRC_SERVICE_FILE")"
: > "$FLUX_AGENT_OPENRC_SERVICE_FILE"
local rc_service_calls=""
local rc_update_calls=""
rc-service() {
rc_service_calls+=$'\n'"$*"
return 0
}
rc-update() {
rc_update_calls+=$'\n'"$*"
return 0
}
stop_flux_agent_service
disable_flux_agent_service
remove_flux_agent_service
[[ "$rc_service_calls" == *"stop"* ]] || fail "OpenRC uninstall should stop flux_agent"
[[ "$rc_update_calls" == *"del flux_agent default"* ]] || \
fail "OpenRC uninstall should remove flux_agent from the default runlevel"
[[ ! -e "$FLUX_AGENT_OPENRC_SERVICE_FILE" ]] || fail "OpenRC uninstall should remove its service file"
)
test_cleanup_legacy_gost_installation_removes_service_and_binary() (
@@ -283,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
@@ -326,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
@@ -353,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"
@@ -409,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"
@@ -421,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"
}
@@ -429,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
@@ -444,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() (
@@ -504,14 +752,21 @@ test_update_flux_agent_skips_proxy_prompt_when_not_installed
test_install_flux_agent_preserves_legacy_gost_when_download_fails
test_update_flux_agent_preserves_legacy_gost_when_download_fails
test_install_flux_agent_writes_json_safe_config
test_install_script_bootstraps_bash_for_alpine
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
echo "install script proxy tests passed"
echo "install script tests passed"
+4 -4
View File
@@ -1,10 +1,10 @@
# 多阶段构建 - 构建阶段
FROM node:22-alpine AS builder
FROM --platform=$BUILDPLATFORM node:22-alpine AS builder
WORKDIR /app
COPY package.json pnpm-lock.yaml* ./
RUN corepack prepare pnpm@10 --activate && corepack enable pnpm && pnpm install --frozen-lockfile
COPY package.json pnpm-lock.yaml pnpm-workspace.yaml ./
RUN corepack enable pnpm && pnpm install --frozen-lockfile
COPY . .
RUN pnpm run build
@@ -18,4 +18,4 @@ COPY --from=builder /app/dist /usr/share/nginx/html
EXPOSE 80
CMD ["nginx", "-g", "daemon off;"]
CMD ["nginx", "-g", "daemon off;"]
+1 -7
View File
@@ -2,6 +2,7 @@
"name": "flvx",
"private": true,
"version": "0.0.0",
"packageManager": "pnpm@10.28.1",
"type": "module",
"scripts": {
"dev": "vite",
@@ -81,12 +82,5 @@
"vite": "npm:rolldown-vite@^7.3.1",
"vite-plugin-pwa": "^1.1.0",
"vite-tsconfig-paths": "^6.0.5"
},
"pnpm": {
"overrides": {
"@babel/plugin-transform-modules-systemjs": "7.29.4",
"fast-uri": "3.1.2",
"serialize-javascript": "7.0.5"
}
}
}
+8
View File
@@ -1,2 +1,10 @@
packages:
- '.'
overrides:
'@babel/plugin-transform-modules-systemjs': 7.29.4
fast-uri: 3.1.2
serialize-javascript: 7.0.5
allowBuilds:
'@tailwindcss/oxide': true
+14
View File
@@ -595,6 +595,20 @@ export interface TunnelQualityHopApiItem {
targetPort?: number;
}
export interface TunnelQualityCandidateHopApiItem
extends TunnelQualityHopApiItem {
fromRole: "entry" | "middle" | "exit";
toRole: "middle" | "exit" | "target";
hopIndex: number;
selected: boolean;
errorMessage?: string;
}
export interface TunnelQualityChainDetailsApiItem {
primaryPath?: TunnelQualityHopApiItem[];
candidateHops?: TunnelQualityCandidateHopApiItem[];
}
export interface TunnelQualityApiItem {
tunnelId: number;
entryToExitLatency: number;
+8 -5
View File
@@ -28,12 +28,13 @@ function DialogClose({
return <DialogPrimitive.Close data-slot="dialog-close" {...props} />;
}
function DialogOverlay({
className,
...props
}: React.ComponentProps<typeof DialogPrimitive.Overlay>) {
const DialogOverlay = React.forwardRef<
React.ElementRef<typeof DialogPrimitive.Overlay>,
React.ComponentPropsWithoutRef<typeof DialogPrimitive.Overlay>
>(({ className, ...props }, ref) => {
return (
<DialogPrimitive.Overlay
ref={ref}
className={cn(
"fixed inset-0 z-50 bg-black/30 backdrop-blur-md data-[state=open]:animate-in data-[state=closed]:animate-out data-[state=closed]:fade-out-0 data-[state=open]:fade-in-0",
className,
@@ -42,7 +43,9 @@ function DialogOverlay({
{...props}
/>
);
}
});
DialogOverlay.displayName = DialogPrimitive.Overlay.displayName;
function DialogContent({
className,
+74 -22
View File
@@ -14,10 +14,14 @@ const PUBLIC_BRAND_CONFIG_KEYS = [
"app_logo",
"app_favicon",
"app_bg_image",
"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 +53,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> = {};
@@ -106,15 +134,15 @@ 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,
is_commercial: isCommercial,
hide_footer_brand: hideFooterBrand,
hide_footer_brand: isCommercial && hideFooterBrand,
};
}
@@ -123,11 +151,11 @@ 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,
is_commercial: isCommercial,
hide_footer_brand: hideFooterBrand,
hide_footer_brand: isCommercial && hideFooterBrand,
};
};
@@ -206,20 +234,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 +281,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;
}
};
@@ -345,6 +383,12 @@ export const updateSiteConfig = async (configMap?: Record<string, string>) => {
"app_bg_image",
);
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;
@@ -358,15 +402,23 @@ export const updateSiteConfig = async (configMap?: Record<string, string>) => {
? String(resolvedConfigMap.app_bg_image || "").trim()
: (siteConfig.app_bg_image || "").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";
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;
@@ -0,0 +1,39 @@
export const TUNNEL_QUALITY_INTERVAL_CONFIG_KEY =
"monitor_tunnel_quality_interval_sec";
export const DEFAULT_TUNNEL_QUALITY_INTERVAL_SEC = 1;
export const MIN_TUNNEL_QUALITY_INTERVAL_SEC = 1;
export const MAX_TUNNEL_QUALITY_INTERVAL_SEC = 3600;
export const parseTunnelQualityIntervalSeconds = (value: unknown): number => {
const seconds = Number(value);
return Number.isInteger(seconds) &&
seconds >= MIN_TUNNEL_QUALITY_INTERVAL_SEC &&
seconds <= MAX_TUNNEL_QUALITY_INTERVAL_SEC
? seconds
: DEFAULT_TUNNEL_QUALITY_INTERVAL_SEC;
};
export const validateTunnelQualityInterval = (value: string): string | null => {
const normalized = value.trim();
if (!normalized) {
return "请输入探测间隔";
}
const seconds = Number(normalized);
if (!Number.isInteger(seconds)) {
return "探测间隔必须是整数";
}
if (
seconds < MIN_TUNNEL_QUALITY_INTERVAL_SEC ||
seconds > MAX_TUNNEL_QUALITY_INTERVAL_SEC
) {
return `探测间隔必须在 ${MIN_TUNNEL_QUALITY_INTERVAL_SEC} 到 ${MAX_TUNNEL_QUALITY_INTERVAL_SEC} 秒之间`;
}
return null;
};
export const tunnelQualityIntervalLabel = (seconds: number): string =>
seconds === 1 ? "每秒" : `每 ${seconds} 秒`;
+75 -4
View File
@@ -43,6 +43,14 @@ import { BackIcon, SettingsIcon } from "@/components/icons";
import { ThemeSettings } from "@/components/theme-settings";
import { isAdmin } from "@/utils/auth";
import { getCachedConfigs, configCache, updateSiteConfig } from "@/config/site";
import {
DEFAULT_TUNNEL_QUALITY_INTERVAL_SEC,
MAX_TUNNEL_QUALITY_INTERVAL_SEC,
MIN_TUNNEL_QUALITY_INTERVAL_SEC,
parseTunnelQualityIntervalSeconds,
TUNNEL_QUALITY_INTERVAL_CONFIG_KEY,
validateTunnelQualityInterval,
} from "@/config/tunnel-quality";
import {
type UpdateReleaseChannel,
getUpdateReleaseChannel,
@@ -157,6 +165,16 @@ const CONFIG_ITEMS: ConfigItem[] = [
"关闭后,前端停止自动刷新,后端停止实时隧道质量探测(全局配置)",
type: "switch",
},
{
key: TUNNEL_QUALITY_INTERVAL_CONFIG_KEY,
label: "隧道质量探测间隔",
placeholder: String(DEFAULT_TUNNEL_QUALITY_INTERVAL_SEC),
description:
"设置实时隧道质量检测的执行频率,单位为秒;允许 1–3600 秒,默认 1 秒。",
type: "input",
dependsOn: "monitor_tunnel_quality_enabled",
dependsValue: "true",
},
{
key: "monitor_retention_days",
label: "监控数据保留天数",
@@ -239,6 +257,7 @@ const getInitialConfigs = (): Record<string, string> => {
"cloudflare_secret_key",
"forward_compact_mode",
"monitor_tunnel_quality_enabled",
TUNNEL_QUALITY_INTERVAL_CONFIG_KEY,
"monitor_retention_days",
"ip",
"panel_domain",
@@ -247,6 +266,9 @@ const getInitialConfigs = (): Record<string, string> => {
"github_proxy_enabled",
"github_proxy_url",
"allow_local_remote_addr",
"is_commercial",
"license_expiry",
"hide_footer_brand",
];
const initialConfigs: Record<string, string> = {};
@@ -622,6 +644,19 @@ export default function ConfigPage() {
// 保存配置
const handleSave = async () => {
const intervalValue = configs[TUNNEL_QUALITY_INTERVAL_CONFIG_KEY];
const intervalChanged =
intervalValue !== originalConfigs[TUNNEL_QUALITY_INTERVAL_CONFIG_KEY];
const intervalError = intervalChanged
? validateTunnelQualityInterval(intervalValue || "")
: null;
if (intervalError) {
toast.error(intervalError);
return;
}
setSaving(true);
try {
const changedKeys = Object.keys(configs).filter(
@@ -667,12 +702,22 @@ export default function ConfigPage() {
}),
);
// 如果隧道质量检测开关变更,通知 tunnel-monitor-view
if (changedKeys.includes("monitor_tunnel_quality_enabled")) {
// 如果隧道质量检测配置变更,通知 tunnel-monitor-view
if (
changedKeys.some((key) =>
[
"monitor_tunnel_quality_enabled",
TUNNEL_QUALITY_INTERVAL_CONFIG_KEY,
].includes(key),
)
) {
window.dispatchEvent(
new CustomEvent("monitorTunnelQualityEnabledChanged", {
detail: {
enabled: configs["monitor_tunnel_quality_enabled"] === "true",
intervalSec: parseTunnelQualityIntervalSeconds(
configs[TUNNEL_QUALITY_INTERVAL_CONFIG_KEY],
),
},
}),
);
@@ -1075,11 +1120,21 @@ export default function ConfigPage() {
case "bg_image":
return renderBgImageUploader();
case "input":
case "input": {
if (isBrandPreviewKey(item.key)) {
return renderBrandAssetUploader(item.key, isChanged);
}
const isTunnelQualityInterval =
item.key === TUNNEL_QUALITY_INTERVAL_CONFIG_KEY;
const intervalValue = isTunnelQualityInterval
? (configs[item.key] ?? String(DEFAULT_TUNNEL_QUALITY_INTERVAL_SEC))
: (configs[item.key] ?? "");
const intervalError =
isTunnelQualityInterval && configs[item.key] !== undefined
? validateTunnelQualityInterval(intervalValue)
: null;
return (
<Input
classNames={{
@@ -1091,14 +1146,30 @@ export default function ConfigPage() {
description={
isCommercialDisabled ? "需商业版授权才能修改此项" : undefined
}
endContent={isTunnelQualityInterval ? "秒" : undefined}
errorMessage={intervalError || undefined}
isDisabled={isCommercialDisabled}
isInvalid={Boolean(intervalError)}
max={
isTunnelQualityInterval
? MAX_TUNNEL_QUALITY_INTERVAL_SEC
: undefined
}
min={
isTunnelQualityInterval
? MIN_TUNNEL_QUALITY_INTERVAL_SEC
: undefined
}
placeholder={item.placeholder}
size="md"
value={configs[item.key] || ""}
step={isTunnelQualityInterval ? 1 : undefined}
type={isTunnelQualityInterval ? "number" : "text"}
value={intervalValue}
variant="bordered"
onChange={(e) => handleConfigChange(item.key, e.target.value)}
/>
);
}
case "switch":
return (
@@ -2,6 +2,8 @@ import type {
MonitorTunnelApiItem,
TunnelMetricApiItem,
TunnelQualityApiItem,
TunnelQualityCandidateHopApiItem,
TunnelQualityChainDetailsApiItem,
TunnelQualityHopApiItem,
} from "@/api/types";
@@ -53,12 +55,17 @@ import {
TableRow,
TableCell,
} from "@/shadcn-bridge/heroui/table";
import {
DEFAULT_TUNNEL_QUALITY_INTERVAL_SEC,
parseTunnelQualityIntervalSeconds,
TUNNEL_QUALITY_INTERVAL_CONFIG_KEY,
tunnelQualityIntervalLabel,
} from "@/config/tunnel-quality";
interface TunnelMonitorViewProps {
viewMode?: "list" | "grid";
}
const QUALITY_POLL_INTERVAL = 1_000; // 1 second
const MONITOR_TUNNEL_QUALITY_ENABLED_CONFIG_KEY =
"monitor_tunnel_quality_enabled";
const MONITOR_TUNNEL_QUALITY_ENABLED_EVENT =
@@ -463,18 +470,249 @@ const TrafficChartCard = React.memo(function TrafficChartCard({
);
});
function ForwardingChainTopology({ hopsStr }: { hopsStr?: string }) {
if (!hopsStr) return null;
let hops: TunnelQualityHopApiItem[] = [];
const parseTunnelQualityChainDetails = (
raw?: string,
): TunnelQualityChainDetailsApiItem => {
if (!raw) return {};
try {
hops = JSON.parse(hopsStr);
const parsed: unknown = JSON.parse(raw);
// Backward-compatible with historical rows that stored the primary path
// directly as a JSON array.
if (Array.isArray(parsed)) {
return { primaryPath: parsed as TunnelQualityHopApiItem[] };
}
if (parsed && typeof parsed === "object") {
return parsed as TunnelQualityChainDetailsApiItem;
}
} catch {
return null;
return {};
}
if (!Array.isArray(hops) || hops.length === 0) return null;
return {};
};
type TunnelTopologyHop = TunnelQualityHopApiItem & {
errorMessage?: string;
};
interface TunnelTopologyPath {
key: string;
hops: TunnelTopologyHop[];
alternativeNodeIndex?: number;
}
const tunnelTopologyHopKey = (fromNodeId: number, toNodeId: number) =>
`${fromNodeId}:${toNodeId}`;
const buildTunnelTopologyPaths = (
details: TunnelQualityChainDetailsApiItem,
): TunnelTopologyPath[] => {
const primaryHops = details.primaryPath ?? [];
const candidates = details.candidateHops ?? [];
if (primaryHops.length === 0) {
const publicCandidates = candidates.filter(
(candidate) => candidate.toRole === "target",
);
const selected = publicCandidates.find((candidate) => candidate.selected);
const paths: TunnelTopologyPath[] = [];
if (selected) {
paths.push({ key: "primary-public", hops: [selected] });
}
for (const candidate of publicCandidates) {
if (candidate.selected) continue;
paths.push({
key: `alternative-public-${candidate.fromNodeId}`,
hops: [candidate],
alternativeNodeIndex: 0,
});
}
return paths;
}
const primaryNodeIds = [
primaryHops[0].fromNodeId,
...primaryHops.map((hop) => hop.toNodeId),
];
const internalCandidates = candidates.filter(
(candidate) => candidate.toRole !== "target",
);
const candidateHopMap = new Map<string, TunnelQualityCandidateHopApiItem>();
const alternativeNodes = new Map<
string,
{ column: number; nodeId: number }
>();
for (const candidate of internalCandidates) {
candidateHopMap.set(
tunnelTopologyHopKey(candidate.fromNodeId, candidate.toNodeId),
candidate,
);
const sourceColumn = candidate.hopIndex;
const targetColumn = candidate.hopIndex + 1;
if (
sourceColumn >= 0 &&
sourceColumn < primaryNodeIds.length &&
candidate.fromNodeId !== primaryNodeIds[sourceColumn]
) {
alternativeNodes.set(`${sourceColumn}:${candidate.fromNodeId}`, {
column: sourceColumn,
nodeId: candidate.fromNodeId,
});
}
if (
targetColumn >= 0 &&
targetColumn < primaryNodeIds.length &&
candidate.toNodeId !== primaryNodeIds[targetColumn]
) {
alternativeNodes.set(`${targetColumn}:${candidate.toNodeId}`, {
column: targetColumn,
nodeId: candidate.toNodeId,
});
}
}
const paths: TunnelTopologyPath[] = [{ key: "primary", hops: primaryHops }];
for (const alternative of alternativeNodes.values()) {
const nodeIds = [...primaryNodeIds];
nodeIds[alternative.column] = alternative.nodeId;
const hops: TunnelTopologyHop[] = [];
for (let index = 0; index < nodeIds.length - 1; index += 1) {
const fromNodeId = nodeIds[index];
const toNodeId = nodeIds[index + 1];
const usesPrimaryEdge =
fromNodeId === primaryNodeIds[index] &&
toNodeId === primaryNodeIds[index + 1];
const hop = usesPrimaryEdge
? primaryHops[index]
: candidateHopMap.get(tunnelTopologyHopKey(fromNodeId, toNodeId));
if (!hop) break;
hops.push(hop);
}
if (hops.length === primaryHops.length) {
paths.push({
key: `alternative-${alternative.column}-${alternative.nodeId}`,
hops,
alternativeNodeIndex: alternative.column,
});
}
}
return paths;
};
function TunnelTopologyPathRow({ path }: { path: TunnelTopologyPath }) {
return (
<div className="flex min-w-max items-center py-2">
{path.hops.map((hop, index) => {
const hasError =
Boolean(hop.errorMessage) || hop.latency < 0 || hop.loss > 0;
const colorClass =
hop.latency < 0 || hop.errorMessage
? "text-danger"
: hop.loss > 0
? "text-warning"
: "text-success";
const borderColor = hasError ? "border-danger" : "";
return (
<React.Fragment
key={`${path.key}-${hop.fromNodeId}-${hop.toNodeId}-${index}`}
>
{index === 0 ? (
<TopologyNodeChip
isAlternative={path.alternativeNodeIndex === 0}
name={hop.fromNodeName}
/>
) : null}
<div
className="relative mx-1 flex min-w-[70px] shrink-0 flex-col items-center justify-center"
title={hop.errorMessage}
>
<span
className={`mb-1 text-[10px] font-mono leading-none ${colorClass}`}
>
{hop.latency >= 0 && !hop.errorMessage
? `${hop.latency.toFixed(0)}ms`
: "超时"}
</span>
<div
className={`relative flex h-[2px] w-full items-center justify-end bg-default-200 ${hop.latency < 0 || hop.errorMessage ? "!bg-danger" : ""}`}
>
<ArrowRight
className={`absolute -right-2 z-10 h-3.5 w-3.5 rounded-full bg-background p-[1px] ${colorClass}`}
/>
</div>
<span
className={`mt-1.5 text-[10px] font-mono leading-none ${hop.errorMessage || hop.loss > 0 ? "text-warning" : "text-default-400"}`}
>
{hop.errorMessage ? "探测失败" : `${hop.loss.toFixed(0)}% 丢包`}
</span>
</div>
<TopologyNodeChip
borderColor={borderColor}
isAlternative={path.alternativeNodeIndex === index + 1}
name={hop.toNodeName}
/>
</React.Fragment>
);
})}
</div>
);
}
function TopologyNodeChip({
name,
isAlternative = false,
borderColor = "",
}: {
name: string;
isAlternative?: boolean;
borderColor?: string;
}) {
return (
<Chip
className={`shrink-0 font-mono shadow-sm ${borderColor}`}
size="sm"
variant="flat"
>
<span className="flex items-center gap-1.5">
<span>{name}</span>
{isAlternative ? (
<span className="rounded-full bg-warning/20 px-1.5 py-0.5 text-[9px] font-semibold leading-none text-warning">
备选
</span>
) : null}
</span>
</Chip>
);
}
const ForwardingChainTopology = React.memo(function ForwardingChainTopology({
hopsStr,
}: {
hopsStr?: string;
}) {
const details = useMemo(
() => parseTunnelQualityChainDetails(hopsStr),
[hopsStr],
);
const topologyPaths = useMemo(
() => buildTunnelTopologyPaths(details),
[details],
);
if (topologyPaths.length === 0) return null;
return (
<Card className="border border-divider/60 shadow-sm transition-shadow bg-gradient-to-br from-background to-default-50/50 mt-4">
@@ -485,62 +723,24 @@ function ForwardingChainTopology({ hopsStr }: { hopsStr?: string }) {
</h3>
</CardHeader>
<CardBody className="py-2 px-4 pb-4">
<div className="flex items-center overflow-x-auto pb-2 py-2">
{hops.map((hop, index) => {
const hasError = hop.latency < 0 || hop.loss > 0;
const colorClass =
hop.latency < 0
? "text-danger"
: hop.loss > 0
? "text-warning"
: "text-success";
const borderColor = hasError ? "border-danger" : "";
return (
<React.Fragment key={index}>
{index === 0 && (
<Chip
className="shrink-0 font-mono shadow-sm"
size="sm"
variant="flat"
>
{hop.fromNodeName}
</Chip>
)}
<div className="flex flex-col items-center justify-center min-w-[70px] mx-1 shrink-0 relative">
<span
className={`text-[10px] font-mono leading-none mb-1 ${colorClass}`}
>
{hop.latency >= 0 ? `${hop.latency.toFixed(0)}ms` : "超时"}
</span>
<div
className={`h-[2px] w-full relative flex items-center justify-end bg-default-200 ${hop.latency < 0 ? "!bg-danger" : ""}`}
>
<ArrowRight
className={`w-3.5 h-3.5 absolute -right-2 ${colorClass} bg-background rounded-full p-[1px] z-10`}
/>
</div>
<span
className={`text-[10px] font-mono leading-none mt-1.5 ${hop.loss > 0 ? "text-warning" : "text-default-400"}`}
>
{hop.loss.toFixed(0)}% 丢包
</span>
</div>
<Chip
className={`shrink-0 font-mono shadow-sm ${borderColor}`}
size="sm"
variant="flat"
>
{hop.toNodeName}
</Chip>
</React.Fragment>
);
})}
<div className="max-h-80 space-y-1 overflow-auto pb-1">
{topologyPaths.map((path, index) => (
<div
key={path.key}
className={
index === 0
? "overflow-x-auto"
: "overflow-x-auto border-t border-dashed border-divider/60"
}
>
<TunnelTopologyPathRow path={path} />
</div>
))}
</div>
</CardBody>
</Card>
);
}
});
export function TunnelMonitorView({
viewMode = "grid",
@@ -562,6 +762,9 @@ export function TunnelMonitorView({
const qualityTimerRef = useRef<number | null>(null);
const [monitorTunnelQualityEnabled, setMonitorTunnelQualityEnabled] =
useState(true);
const [tunnelQualityIntervalSec, setTunnelQualityIntervalSec] = useState(
DEFAULT_TUNNEL_QUALITY_INTERVAL_SEC,
);
// Detail view state
const [detailTunnelId, setDetailTunnelId] = useState<number | null>(null);
@@ -618,26 +821,28 @@ export function TunnelMonitorView({
}
}, []);
const loadMonitorTunnelQualityEnabled = useCallback(async () => {
try {
const response = await getConfigByName(
MONITOR_TUNNEL_QUALITY_ENABLED_CONFIG_KEY,
);
const loadTunnelQualityConfig = useCallback(async () => {
const [enabledResponse, intervalResponse] = await Promise.all([
getConfigByName(MONITOR_TUNNEL_QUALITY_ENABLED_CONFIG_KEY).catch(
() => null,
),
getConfigByName(TUNNEL_QUALITY_INTERVAL_CONFIG_KEY).catch(() => null),
]);
setMonitorTunnelQualityEnabled(
typeof response.data?.value === "string"
? response.data.value === "true"
: true,
);
} catch {
setMonitorTunnelQualityEnabled(true);
}
setMonitorTunnelQualityEnabled(
typeof enabledResponse?.data?.value === "string"
? enabledResponse.data.value === "true"
: true,
);
setTunnelQualityIntervalSec(
parseTunnelQualityIntervalSeconds(intervalResponse?.data?.value),
);
}, []);
useEffect(() => {
void loadTunnels();
void loadMonitorTunnelQualityEnabled();
}, [loadMonitorTunnelQualityEnabled, loadTunnels]);
void loadTunnelQualityConfig();
}, [loadTunnelQualityConfig, loadTunnels]);
useEffect(() => {
const timer = window.setInterval(() => {
@@ -649,8 +854,11 @@ export function TunnelMonitorView({
useEffect(() => {
const handleMonitorTunnelQualityEnabledChanged = (event: Event) => {
const enabled = (event as CustomEvent<{ enabled?: boolean }>).detail
?.enabled;
const detail = (
event as CustomEvent<{ enabled?: boolean; intervalSec?: number }>
).detail;
const enabled = detail?.enabled;
const intervalSec = detail?.intervalSec;
if (typeof enabled === "boolean") {
setMonitorTunnelQualityEnabled(enabled);
@@ -658,7 +866,12 @@ export function TunnelMonitorView({
setQualityLoading(false);
}
} else {
void loadMonitorTunnelQualityEnabled();
void loadTunnelQualityConfig();
}
if (typeof intervalSec === "number") {
setTunnelQualityIntervalSec(
parseTunnelQualityIntervalSeconds(String(intervalSec)),
);
}
};
@@ -673,7 +886,7 @@ export function TunnelMonitorView({
handleMonitorTunnelQualityEnabledChanged as EventListener,
);
};
}, [loadMonitorTunnelQualityEnabled]);
}, [loadTunnelQualityConfig]);
useEffect(() => {
if (tunnels.length > 0 && !initialHistoryFetched.current) {
@@ -726,7 +939,7 @@ export function TunnelMonitorView({
}
}, [tunnels]);
// --- Load quality snapshots (auto-polling every 10s) ---
// --- Load quality snapshots using the configured probe interval ---
const loadQuality = useCallback(async (options?: { silent?: boolean }) => {
const silent = options?.silent ?? false;
@@ -788,7 +1001,7 @@ export function TunnelMonitorView({
qualityTimerRef.current = window.setInterval(() => {
void loadQuality({ silent: true });
}, QUALITY_POLL_INTERVAL);
}, tunnelQualityIntervalSec * 1000);
return () => {
if (qualityTimerRef.current) {
@@ -796,7 +1009,7 @@ export function TunnelMonitorView({
qualityTimerRef.current = null;
}
};
}, [loadQuality, monitorTunnelQualityEnabled]);
}, [loadQuality, monitorTunnelQualityEnabled, tunnelQualityIntervalSec]);
// --- Load quality history for detail chart ---
const loadQualityHistory = useCallback(
@@ -1065,7 +1278,7 @@ export function TunnelMonitorView({
{monitorTunnelQualityEnabled ? (
<>
<LiveDot />
<span>自动探测中(每秒测试,30秒上报)</span>
<span>{`自动探测中(${tunnelQualityIntervalLabel(tunnelQualityIntervalSec)}测试)`}</span>
</>
) : (
<>
@@ -1129,7 +1342,7 @@ export function TunnelMonitorView({
{monitorTunnelQualityEnabled ? (
<>
<LiveDot />
<span>每秒探测 · 更新于 {lastQualityUpdate}</span>
<span>{`${tunnelQualityIntervalLabel(tunnelQualityIntervalSec)}探测 · 更新于 ${lastQualityUpdate}`}</span>
</>
) : (
<>
@@ -49,6 +49,41 @@ function useModalContext() {
return React.useContext(ModalContext);
}
interface ScrollPosition {
element: HTMLElement | null;
left: number;
top: number;
}
function captureScrollPositions(): ScrollPosition[] {
const positions: ScrollPosition[] = [
{ element: null, left: window.scrollX, top: window.scrollY },
];
for (const element of Array.from(
document.querySelectorAll<HTMLElement>("main, [data-scroll-container]"),
)) {
positions.push({
element,
left: element.scrollLeft,
top: element.scrollTop,
});
}
return positions;
}
function restoreScrollPositions(positions: ScrollPosition[]) {
for (const position of positions) {
if (position.element) {
position.element.scrollLeft = position.left;
position.element.scrollTop = position.top;
} else {
window.scrollTo(position.left, position.top);
}
}
}
type ModalSize = "sm" | "md" | "lg" | "xl" | "2xl" | "4xl" | "full";
function mapSize(size: ModalSize | undefined) {
@@ -97,6 +132,46 @@ export function Modal({
scrollBehavior,
size,
}: ModalProps) {
const previousScrollPositionsRef = React.useRef<ScrollPosition[] | null>(
null,
);
// Radix focus management and scroll locking can move an ancestor scroll
// container when a modal is opened from a card/grid item. Capture the
// current positions before the open render and restore them after focus
// settles so opening a modal never changes the page position.
React.useLayoutEffect(() => {
return () => {
if (!isOpen) {
previousScrollPositionsRef.current = captureScrollPositions();
}
};
}, [isOpen]);
React.useLayoutEffect(() => {
const positions = previousScrollPositionsRef.current;
if (!isOpen || !positions) {
return;
}
restoreScrollPositions(positions);
let nestedFrame = 0;
const frame = window.requestAnimationFrame(() => {
restoreScrollPositions(positions);
nestedFrame = window.requestAnimationFrame(() =>
restoreScrollPositions(positions),
);
});
previousScrollPositionsRef.current = null;
return () => {
window.cancelAnimationFrame(frame);
window.cancelAnimationFrame(nestedFrame);
};
}, [isOpen]);
const handleOpenChange = (open: boolean) => {
onOpenChange?.(open);
if (!open) {