fix: harden proxy protocol rollout safety (#494)

This commit is contained in:
sagit
2026-05-07 16:37:51 +08:00
committed by GitHub
parent 1f53a39784
commit 4ebd6703fe
6 changed files with 500 additions and 17 deletions
+19 -5
View File
@@ -158,6 +158,9 @@ type WebSocketReporter struct {
addr string // 保存服务器地址
secret string // 保存密钥
version string // 保存版本号
http int
tls int
socks int
preferredWSScheme string
conn *websocket.Conn
curBackoff time.Duration // 当前重连退避间隔
@@ -296,9 +299,9 @@ func (w *WebSocketReporter) connect() error {
Socks int `json:"socks"`
}
var cfg LocalConfig
cfg := LocalConfig{Http: w.http, Tls: w.tls, Socks: w.socks}
if b, err := os.ReadFile("config.json"); err == nil {
json.Unmarshal(b, &cfg)
_ = json.Unmarshal(b, &cfg)
}
candidates := buildWebSocketCandidates(w.addr, w.secret, w.version, cfg.Http, cfg.Tls, cfg.Socks, w.preferredWSScheme)
@@ -1369,7 +1372,7 @@ func (w *WebSocketReporter) handleUpgradeAgent(data interface{}) error {
// 执行重启脚本
// 使用 systemd-run 在独立的 transient unit 中运行重启脚本,
// 避免 systemctl stop 杀死 flux_agent cgroup 内所有进程(包括此脚本自身)导致 mv 未执行。
script := fmt.Sprintf("sleep 1 && systemctl stop flux_agent && mv %s %s && systemctl start flux_agent", tmpPath, binaryPath)
script := buildAgentRestartScript(tmpPath, binaryPath)
cmd := exec.Command("systemd-run", "--quiet", "/bin/sh", "-c", script)
if err := cmd.Start(); err != nil {
os.Remove(tmpPath)
@@ -1403,6 +1406,14 @@ func (w *WebSocketReporter) handleRollbackAgent(data interface{}) error {
return nil
}
func buildAgentRestartScript(tmpPath, binaryPath string) string {
return fmt.Sprintf(
"sleep 1 && systemctl stop flux_agent && legacy_service='' && for service_file in /etc/systemd/system/gost.service /lib/systemd/system/gost.service /usr/lib/systemd/system/gost.service; do if [ -f \"$service_file\" ] && grep -Fq \"WorkingDirectory=/etc/gost\" \"$service_file\" && (grep -Fq \"ExecStart=/etc/gost/gost\" \"$service_file\" || (grep -Fq \"ExecStart=/usr/local/bin/gost\" \"$service_file\" && [ -f /etc/gost/config.json ] && [ -f /etc/gost/gost.json ])); then legacy_service=\"$service_file\"; break; fi; done && if [ -n \"$legacy_service\" ]; then (systemctl stop gost 2>/dev/null || true) && (systemctl disable gost 2>/dev/null || true) && rm -f /usr/local/bin/gost /etc/gost/gost \"$legacy_service\" && (systemctl daemon-reload 2>/dev/null || true); fi && mv %s %s && systemctl start flux_agent",
tmpPath,
binaryPath,
)
}
// updateLocalConfigJSON 将 http/tls/socks 写入工作目录下的 config.json
func updateLocalConfigJSON(httpVal int, tlsVal int, socksVal int) error {
path := "config.json"
@@ -1650,13 +1661,16 @@ func StartWebSocketReporterWithConfig(addr string, secret string, http int, tls
candidates := buildWebSocketCandidates(addr, secret, version, http, tls, socks, "")
fullURL := candidates[0]
fmt.Printf("🔗 WebSocket连接URL: %s\n", fullURL)
fmt.Printf("🔗 WebSocket连接URL: %s\n", sanitizeWebSocketURL(fullURL))
reporter := NewWebSocketReporter(fullURL, secret)
// 保存 addr, secret, version 供重连时使用
// 保存 addr, secret, version 和协议能力供重连时使用
reporter.addr = addr
reporter.secret = secret
reporter.version = version
reporter.http = http
reporter.tls = tls
reporter.socks = socks
reporter.Start()
return reporter
}
+130
View File
@@ -1,15 +1,45 @@
package socket
import (
"bytes"
"errors"
"io"
"net/http"
"os"
"runtime"
"strings"
"testing"
"time"
"github.com/gorilla/websocket"
)
func captureStdout(t *testing.T, fn func()) string {
t.Helper()
orig := os.Stdout
r, w, err := os.Pipe()
if err != nil {
t.Fatalf("create stdout pipe: %v", err)
}
os.Stdout = w
defer func() {
os.Stdout = orig
_ = w.Close()
_ = r.Close()
}()
fn()
_ = w.Close()
var buf bytes.Buffer
if _, err := io.Copy(&buf, r); err != nil {
t.Fatalf("read stdout: %v", err)
}
return buf.String()
}
func TestBuildWebSocketCandidatesSecureFirst(t *testing.T) {
candidates := buildWebSocketCandidates("panel.example.com:443", "abc", "2.0.2", 1, 0, 1, "")
@@ -133,3 +163,103 @@ func TestFormatWebSocketDialErrorIncludesHTTPStatus(t *testing.T) {
t.Fatalf("expected response body in message, got %s", msg)
}
}
func TestAgentUpgradeRestartScriptStopsLegacyGostService(t *testing.T) {
script := buildAgentRestartScript("/tmp/flux_agent.new", "/etc/flux_agent/flux_agent")
if !strings.Contains(script, "systemctl stop flux_agent") {
t.Fatalf("expected script to stop flux_agent, got %s", script)
}
if !strings.Contains(script, "mv /tmp/flux_agent.new /etc/flux_agent/flux_agent") {
t.Fatalf("expected script to replace the flux_agent binary, got %s", script)
}
if !strings.Contains(script, "systemctl stop gost") {
t.Fatalf("expected script to stop the legacy gost service, got %s", script)
}
if !strings.Contains(script, "systemctl disable gost") {
t.Fatalf("expected script to disable the legacy gost service, got %s", script)
}
if !strings.Contains(script, "rm -f /usr/local/bin/gost") {
t.Fatalf("expected script to remove the legacy gost binary, got %s", script)
}
if !strings.Contains(script, "WorkingDirectory=/etc/gost") {
t.Fatalf("expected script to scope cleanup to the legacy FLVX gost service definition, got %s", script)
}
if !strings.Contains(script, "systemctl start flux_agent") {
t.Fatalf("expected script to restart flux_agent, got %s", script)
}
if strings.Contains(script, "systemctl stop flux_agent && systemctl stop gost 2>/dev/null || true") {
t.Fatalf("expected legacy gost cleanup fallback to be scoped, got %s", script)
}
if runtime.GOARCH == "" {
t.Fatalf("unexpected empty runtime arch")
}
}
func TestStartWebSocketReporterWithConfigPreservesProtocolDefaultsWithoutConfigFile(t *testing.T) {
origDial := wsDial
defer func() { wsDial = origDial }()
origWD, err := os.Getwd()
if err != nil {
t.Fatalf("get working directory: %v", err)
}
t.Cleanup(func() {
_ = os.Chdir(origWD)
})
if err := os.Chdir(t.TempDir()); err != nil {
t.Fatalf("change working directory: %v", err)
}
urls := make(chan string, 1)
wsDial = func(_ *websocket.Dialer, rawURL string) (*websocket.Conn, *http.Response, error) {
select {
case urls <- rawURL:
default:
}
return nil, nil, errors.New("dial failed")
}
reporter := StartWebSocketReporterWithConfig("panel.example.com:443", "abc", 1, 0, 1, "2.0.2")
defer reporter.Stop()
select {
case rawURL := <-urls:
if !strings.Contains(rawURL, "http=1&tls=0&socks=1") {
t.Fatalf("expected reconnect URL to preserve startup protocol values, got %s", rawURL)
}
case <-time.After(2 * time.Second):
t.Fatal("timed out waiting for websocket dial")
}
}
func TestStartWebSocketReporterWithConfigLogsSanitizedURL(t *testing.T) {
origDial := wsDial
defer func() { wsDial = origDial }()
ready := make(chan struct{}, 1)
wsDial = func(_ *websocket.Dialer, rawURL string) (*websocket.Conn, *http.Response, error) {
select {
case ready <- struct{}{}:
default:
}
return nil, nil, errors.New("dial failed")
}
output := captureStdout(t, func() {
reporter := StartWebSocketReporterWithConfig("panel.example.com:443", "abc123", 1, 0, 1, "2.0.2")
select {
case <-ready:
case <-time.After(2 * time.Second):
t.Fatal("timed out waiting for websocket dial")
}
reporter.Stop()
})
if strings.Contains(output, "secret=abc123") {
t.Fatalf("expected logged websocket URL to mask the node secret, got %s", output)
}
if !strings.Contains(output, "secret=%2A%2A%2A") {
t.Fatalf("expected logged websocket URL to include masked secret, got %s", output)
}
}