From ea46a30f1120dc13ab4b430be1e18bdc018ca86f Mon Sep 17 00:00:00 2001 From: Awuqing <3184394176@qq.com> Date: Sun, 9 Aug 2026 02:30:42 +0800 Subject: [PATCH] feat(cluster): support restricted-network agent connectivity --- server/cmd/backupx/agent.go | 16 ++- server/config.example.yaml | 3 + server/internal/agent/agent.go | 7 +- server/internal/agent/client.go | 52 ++++++++- server/internal/agent/client_test.go | 93 +++++++++++++++ server/internal/agent/config.go | 109 ++++++++++++++++-- server/internal/agent/config_test.go | 50 +++++++- server/internal/config/config.go | 25 ++++ server/internal/config/config_test.go | 44 +++++++ server/internal/http/agent_handler.go | 10 +- .../internal/http/forwarded_headers_test.go | 49 ++++++++ server/internal/http/middleware.go | 41 +++++++ server/internal/http/router.go | 4 + 13 files changed, 481 insertions(+), 22 deletions(-) create mode 100644 server/internal/agent/client_test.go create mode 100644 server/internal/http/forwarded_headers_test.go diff --git a/server/cmd/backupx/agent.go b/server/cmd/backupx/agent.go index 2cccee9..418a48a 100644 --- a/server/cmd/backupx/agent.go +++ b/server/cmd/backupx/agent.go @@ -24,7 +24,10 @@ func runAgent(args []string) { configPath := fs.String("config", "", "path to agent config YAML (optional)") master := fs.String("master", "", "master URL, e.g. http://master.example.com:8340") token := fs.String("token", "", "agent authentication token") + tokenFile := fs.String("token-file", "", "read the agent authentication token from a file") tempDir := fs.String("temp-dir", "", "local temp directory for backup artifacts") + proxyURL := fs.String("proxy-url", "", "HTTP(S) or SOCKS5 proxy used to reach the master") + caCertFile := fs.String("ca-cert", "", "PEM CA certificate used to verify the master") insecureTLS := fs.Bool("insecure-tls", false, "skip TLS verification (testing only)") if err := fs.Parse(args); err != nil { @@ -36,10 +39,21 @@ func runAgent(args []string) { fmt.Fprintf(os.Stderr, "agent: load config: %v\n", err) os.Exit(2) } - cfg.MergeWithFlags(*master, *token, *tempDir) + cfg.ApplyOverrides(agent.Overrides{ + Master: *master, + Token: *token, + TokenFile: *tokenFile, + TempDir: *tempDir, + ProxyURL: *proxyURL, + CACertFile: *caCertFile, + }) if *insecureTLS { cfg.InsecureSkipTLSVerify = true } + if err := cfg.ResolveToken(); err != nil { + fmt.Fprintf(os.Stderr, "agent: %v\n", err) + os.Exit(2) + } if err := cfg.Validate(); err != nil { fmt.Fprintf(os.Stderr, "agent: %v\n", err) os.Exit(2) diff --git a/server/config.example.yaml b/server/config.example.yaml index 351b378..82da115 100644 --- a/server/config.example.yaml +++ b/server/config.example.yaml @@ -4,6 +4,9 @@ server: port: 8340 mode: "release" # debug | release external_url: "" # 可选:Master 对 Agent 可达的 URL,例如 https://backup.example.com + trusted_proxies: # 仅这些代理可提供 X-Forwarded-For;跨容器代理需加入其网段 + - "127.0.0.1" + - "::1" web_root: "" # 前端静态目录;留空自动探测(./web、/opt/backupx/web 等)。 # 命中后后端直接托管 Web 控制台,无需额外 nginx 反向代理。 diff --git a/server/internal/agent/agent.go b/server/internal/agent/agent.go index 93bedf3..cf5d4d7 100644 --- a/server/internal/agent/agent.go +++ b/server/internal/agent/agent.go @@ -28,10 +28,16 @@ type Agent struct { // New 构造 Agent。 func New(cfg *Config, version string) (*Agent, error) { + if err := cfg.ResolveToken(); err != nil { + return nil, err + } if err := cfg.Validate(); err != nil { return nil, err } client := NewMasterClient(cfg.Master, cfg.Token, cfg.InsecureSkipTLSVerify) + if err := client.ConfigureTransport(cfg.ProxyURL, cfg.CACertFile); err != nil { + return nil, fmt.Errorf("configure master connection: %w", err) + } executor := NewExecutor(client, cfg.TempDir) return &Agent{ cfg: cfg, @@ -93,7 +99,6 @@ func (a *Agent) heartbeatLoop(ctx context.Context, interval time.Duration) { func (a *Agent) heartbeatOnce(ctx context.Context) error { hostname, _ := os.Hostname() req := HeartbeatRequest{ - Token: a.cfg.Token, Hostname: hostname, IPAddress: detectLocalIP(), AgentVersion: a.version, diff --git a/server/internal/agent/client.go b/server/internal/agent/client.go index 64ba166..43f2b50 100644 --- a/server/internal/agent/client.go +++ b/server/internal/agent/client.go @@ -4,11 +4,14 @@ import ( "bytes" "context" "crypto/tls" + "crypto/x509" "encoding/json" "errors" "fmt" "io" "net/http" + "net/url" + "os" "strings" "time" ) @@ -22,23 +25,64 @@ type MasterClient struct { // NewMasterClient 构造 Master 客户端。 func NewMasterClient(baseURL, token string, insecureTLS bool) *MasterClient { - transport := &http.Transport{} - if insecureTLS { - transport.TLSClientConfig = &tls.Config{InsecureSkipVerify: true} + transport := http.DefaultTransport.(*http.Transport).Clone() + tlsConfig := &tls.Config{MinVersion: tls.VersionTLS12} + if transport.TLSClientConfig != nil { + tlsConfig = transport.TLSClientConfig.Clone() + tlsConfig.MinVersion = tls.VersionTLS12 } + // 仅用于用户显式开启的测试模式。生产环境应配置受信 CA。 + tlsConfig.InsecureSkipVerify = insecureTLS // #nosec G402 + transport.TLSClientConfig = tlsConfig return &MasterClient{ baseURL: strings.TrimRight(baseURL, "/"), token: token, httpClient: &http.Client{ Timeout: 120 * time.Second, Transport: transport, + // Agent Token 是自定义认证头。禁止自动重定向,避免代理或错误 + // 配置把它转发到另一个主机;Master URL 必须直接指向 API。 + CheckRedirect: func(_ *http.Request, _ []*http.Request) error { + return http.ErrUseLastResponse + }, }, } } +// ConfigureTransport 应用显式代理和私有 CA。默认 Transport 已保留 +// ProxyFromEnvironment,因此 ProxyURL 留空时 HTTP_PROXY/HTTPS_PROXY/NO_PROXY 生效。 +func (c *MasterClient) ConfigureTransport(proxyURL, caCertFile string) error { + transport, ok := c.httpClient.Transport.(*http.Transport) + if !ok { + return errors.New("agent http transport has unexpected type") + } + if strings.TrimSpace(proxyURL) != "" { + parsedProxy, err := url.Parse(strings.TrimSpace(proxyURL)) + if err != nil { + return fmt.Errorf("parse proxy URL: %w", err) + } + transport.Proxy = http.ProxyURL(parsedProxy) + } + if strings.TrimSpace(caCertFile) == "" { + return nil + } + pemData, err := os.ReadFile(strings.TrimSpace(caCertFile)) + if err != nil { + return fmt.Errorf("read CA certificate: %w", err) + } + roots, err := x509.SystemCertPool() + if err != nil || roots == nil { + roots = x509.NewCertPool() + } + if !roots.AppendCertsFromPEM(pemData) { + return errors.New("CA certificate file does not contain a valid PEM certificate") + } + transport.TLSClientConfig.RootCAs = roots + return nil +} + // HeartbeatRequest Agent 上报心跳的请求 type HeartbeatRequest struct { - Token string `json:"token"` Hostname string `json:"hostname,omitempty"` IPAddress string `json:"ipAddress,omitempty"` AgentVersion string `json:"agentVersion,omitempty"` diff --git a/server/internal/agent/client_test.go b/server/internal/agent/client_test.go new file mode 100644 index 0000000..bad09ca --- /dev/null +++ b/server/internal/agent/client_test.go @@ -0,0 +1,93 @@ +package agent + +import ( + "context" + "io" + "net/http" + "net/http/httptest" + "net/url" + "os" + "path/filepath" + "strings" + "testing" +) + +func TestMasterClientKeepsEnvironmentProxySupport(t *testing.T) { + client := NewMasterClient("https://master.example.com", "token", false) + transport := client.httpClient.Transport.(*http.Transport) + if transport.Proxy == nil { + t.Fatal("default transport proxy function must be preserved") + } +} + +func TestMasterClientConfiguresExplicitProxy(t *testing.T) { + client := NewMasterClient("https://master.example.com", "token", false) + if err := client.ConfigureTransport("socks5h://127.0.0.1:1080", ""); err != nil { + t.Fatal(err) + } + transport := client.httpClient.Transport.(*http.Transport) + requestURL, _ := url.Parse("https://master.example.com") + proxyURL, err := transport.Proxy(&http.Request{URL: requestURL}) + if err != nil { + t.Fatal(err) + } + if proxyURL == nil || proxyURL.String() != "socks5h://127.0.0.1:1080" { + t.Fatalf("proxy URL = %v", proxyURL) + } +} + +func TestMasterClientRejectsInvalidCACertificate(t *testing.T) { + path := filepath.Join(t.TempDir(), "invalid.pem") + if err := os.WriteFile(path, []byte("not a certificate"), 0o600); err != nil { + t.Fatal(err) + } + client := NewMasterClient("https://master.example.com", "token", false) + if err := client.ConfigureTransport("", path); err == nil { + t.Fatal("expected invalid CA certificate error") + } +} + +func TestMasterClientDoesNotForwardTokenThroughRedirects(t *testing.T) { + receivedToken := "" + target := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + receivedToken = r.Header.Get("X-Agent-Token") + w.WriteHeader(http.StatusOK) + })) + defer target.Close() + redirector := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + http.Redirect(w, r, target.URL, http.StatusTemporaryRedirect) + })) + defer redirector.Close() + + client := NewMasterClient(redirector.URL, "secret-agent-token", false) + if _, err := client.Heartbeat(context.Background(), HeartbeatRequest{}); err == nil { + t.Fatal("redirect response should not be accepted as a Master API response") + } + if receivedToken != "" { + t.Fatalf("Agent token leaked through redirect: %q", receivedToken) + } +} + +func TestHeartbeatSendsTokenOnlyInAuthenticationHeader(t *testing.T) { + requestBody := "" + receivedHeader := "" + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + body, _ := io.ReadAll(r.Body) + requestBody = string(body) + receivedHeader = r.Header.Get("X-Agent-Token") + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"data":{"status":"ok","nodeId":1,"name":"node"}}`)) + })) + defer server.Close() + + client := NewMasterClient(server.URL, "secret-agent-token", false) + if _, err := client.Heartbeat(context.Background(), HeartbeatRequest{Hostname: "node"}); err != nil { + t.Fatal(err) + } + if receivedHeader != "secret-agent-token" { + t.Fatalf("authentication header = %q", receivedHeader) + } + if strings.Contains(requestBody, "secret-agent-token") || strings.Contains(requestBody, `"token"`) { + t.Fatalf("heartbeat body exposed the Agent token: %s", requestBody) + } +} diff --git a/server/internal/agent/config.go b/server/internal/agent/config.go index 03b0a34..df63d4b 100644 --- a/server/internal/agent/config.go +++ b/server/internal/agent/config.go @@ -10,8 +10,10 @@ package agent import ( "errors" "fmt" + "net/url" "os" "strings" + "time" "gopkg.in/yaml.v3" ) @@ -22,16 +24,34 @@ type Config struct { Master string `yaml:"master"` // Token 节点认证令牌(在 Master 创建节点时生成) Token string `yaml:"token"` + // TokenFile 从文件读取节点认证令牌;适合 systemd 凭据和容器 secret。 + // Token 与 TokenFile 同时设置时优先使用 Token。 + TokenFile string `yaml:"tokenFile"` // HeartbeatInterval 心跳间隔,默认 15s HeartbeatInterval string `yaml:"heartbeatInterval"` // PollInterval 命令轮询间隔,默认 5s PollInterval string `yaml:"pollInterval"` // TempDir 备份临时目录,默认 /var/lib/backupx-agent/tmp TempDir string `yaml:"tempDir"` + // ProxyURL Agent 访问 Master 使用的显式代理。留空时遵循 + // HTTP_PROXY、HTTPS_PROXY 与 NO_PROXY;支持 http(s) 和 socks5(h)。 + ProxyURL string `yaml:"proxyUrl"` + // CACertFile 私有 CA 的 PEM 文件路径,用于安全连接内网 HTTPS Master。 + CACertFile string `yaml:"caCertFile"` // InsecureSkipTLSVerify 测试环境允许跳过 TLS 证书校验 InsecureSkipTLSVerify bool `yaml:"insecureSkipTlsVerify"` } +// Overrides 表示命令行显式提供的 Agent 配置覆盖项。 +type Overrides struct { + Master string + Token string + TokenFile string + TempDir string + ProxyURL string + CACertFile string +} + // LoadConfigFile 从 YAML 文件加载 Agent 配置。 func LoadConfigFile(path string) (*Config, error) { data, err := os.ReadFile(path) @@ -50,42 +70,109 @@ func LoadConfigFile(path string) (*Config, error) { // 支持的环境变量: // - BACKUPX_AGENT_MASTER Master URL // - BACKUPX_AGENT_TOKEN 节点认证令牌 +// - BACKUPX_AGENT_TOKEN_FILE 节点认证令牌文件 // - BACKUPX_AGENT_HEARTBEAT 心跳间隔(如 15s) // - BACKUPX_AGENT_POLL 命令轮询间隔(如 5s) // - BACKUPX_AGENT_TEMP_DIR 临时目录 +// - BACKUPX_AGENT_PROXY_URL 显式 HTTP(S)/SOCKS5 代理 +// - BACKUPX_AGENT_CA_CERT_FILE 私有 CA PEM 文件 // - BACKUPX_AGENT_INSECURE_TLS true / 1 跳过 TLS 校验 func LoadConfigFromEnv() (*Config, error) { cfg := &Config{ Master: strings.TrimSpace(os.Getenv("BACKUPX_AGENT_MASTER")), Token: strings.TrimSpace(os.Getenv("BACKUPX_AGENT_TOKEN")), + TokenFile: strings.TrimSpace(os.Getenv("BACKUPX_AGENT_TOKEN_FILE")), HeartbeatInterval: strings.TrimSpace(os.Getenv("BACKUPX_AGENT_HEARTBEAT")), PollInterval: strings.TrimSpace(os.Getenv("BACKUPX_AGENT_POLL")), TempDir: strings.TrimSpace(os.Getenv("BACKUPX_AGENT_TEMP_DIR")), + ProxyURL: strings.TrimSpace(os.Getenv("BACKUPX_AGENT_PROXY_URL")), + CACertFile: strings.TrimSpace(os.Getenv("BACKUPX_AGENT_CA_CERT_FILE")), InsecureSkipTLSVerify: strings.EqualFold(os.Getenv("BACKUPX_AGENT_INSECURE_TLS"), "true") || os.Getenv("BACKUPX_AGENT_INSECURE_TLS") == "1", } return applyConfigDefaults(cfg) } -// MergeWithFlags 把命令行覆盖值合并入配置(非空覆盖)。 -func (c *Config) MergeWithFlags(master, token, tempDir string) { - if strings.TrimSpace(master) != "" { - c.Master = master +// ApplyOverrides 把命令行覆盖值合并入配置(非空覆盖)。 +func (c *Config) ApplyOverrides(overrides Overrides) { + if strings.TrimSpace(overrides.Master) != "" { + c.Master = strings.TrimSpace(overrides.Master) } - if strings.TrimSpace(token) != "" { - c.Token = token + tokenProvided := strings.TrimSpace(overrides.Token) != "" + if strings.TrimSpace(overrides.TokenFile) != "" { + c.TokenFile = strings.TrimSpace(overrides.TokenFile) + if !tokenProvided { + // An explicit --token-file must override a token inherited from YAML. + c.Token = "" + } } - if strings.TrimSpace(tempDir) != "" { - c.TempDir = tempDir + if tokenProvided { + c.Token = strings.TrimSpace(overrides.Token) } + if strings.TrimSpace(overrides.TempDir) != "" { + c.TempDir = strings.TrimSpace(overrides.TempDir) + } + if strings.TrimSpace(overrides.ProxyURL) != "" { + c.ProxyURL = strings.TrimSpace(overrides.ProxyURL) + } + if strings.TrimSpace(overrides.CACertFile) != "" { + c.CACertFile = strings.TrimSpace(overrides.CACertFile) + } +} + +// ResolveToken 在所有配置源合并完成后读取 token 文件。 +func (c *Config) ResolveToken() error { + if strings.TrimSpace(c.Token) != "" || strings.TrimSpace(c.TokenFile) == "" { + c.Token = strings.TrimSpace(c.Token) + return nil + } + data, err := os.ReadFile(strings.TrimSpace(c.TokenFile)) + if err != nil { + return fmt.Errorf("read agent token file: %w", err) + } + c.Token = strings.TrimSpace(string(data)) + if c.Token == "" { + return errors.New("agent token file is empty") + } + return nil } // Validate 校验必填字段。 func (c *Config) Validate() error { + masterURL, err := url.Parse(strings.TrimSpace(c.Master)) if strings.TrimSpace(c.Master) == "" { return errors.New("master url is required (set via --master, BACKUPX_AGENT_MASTER or config file)") } + if err != nil || (masterURL.Scheme != "http" && masterURL.Scheme != "https") || masterURL.Host == "" || masterURL.User != nil || masterURL.RawQuery != "" || masterURL.Fragment != "" { + return errors.New("master url must be an absolute http(s) URL without credentials, query or fragment") + } if strings.TrimSpace(c.Token) == "" { - return errors.New("token is required (set via --token, BACKUPX_AGENT_TOKEN or config file)") + return errors.New("token is required (set via --token, --token-file, environment or config file)") + } + if c.ProxyURL != "" { + proxyURL, proxyErr := url.Parse(c.ProxyURL) + if proxyErr != nil || proxyURL.Host == "" { + return errors.New("proxy url must be an absolute URL") + } + switch proxyURL.Scheme { + case "http", "https", "socks5", "socks5h": + default: + return errors.New("proxy url scheme must be http, https, socks5 or socks5h") + } + if proxyURL.RawQuery != "" || proxyURL.Fragment != "" || (proxyURL.Path != "" && proxyURL.Path != "/") { + return errors.New("proxy url must not contain a path, query or fragment") + } + } + if c.CACertFile != "" && c.InsecureSkipTLSVerify { + return errors.New("ca cert file and insecure TLS cannot be enabled together") + } + for name, value := range map[string]string{ + "heartbeat interval": c.HeartbeatInterval, + "poll interval": c.PollInterval, + } { + duration, durationErr := time.ParseDuration(value) + if durationErr != nil || duration <= 0 { + return fmt.Errorf("%s must be a positive duration", name) + } } return nil } @@ -101,5 +188,9 @@ func applyConfigDefaults(cfg *Config) (*Config, error) { cfg.TempDir = "/var/lib/backupx-agent/tmp" } cfg.Master = strings.TrimRight(strings.TrimSpace(cfg.Master), "/") + cfg.Token = strings.TrimSpace(cfg.Token) + cfg.TokenFile = strings.TrimSpace(cfg.TokenFile) + cfg.ProxyURL = strings.TrimSpace(cfg.ProxyURL) + cfg.CACertFile = strings.TrimSpace(cfg.CACertFile) return cfg, nil } diff --git a/server/internal/agent/config_test.go b/server/internal/agent/config_test.go index f1e5292..619235f 100644 --- a/server/internal/agent/config_test.go +++ b/server/internal/agent/config_test.go @@ -11,9 +11,12 @@ func TestLoadConfigFile(t *testing.T) { path := filepath.Join(dir, "agent.yaml") content := `master: http://master.example.com:8340/ token: abc123 +tokenFile: /run/secrets/backupx_agent_token heartbeatInterval: 20s pollInterval: 3s tempDir: /var/backupx-agent +proxyUrl: socks5h://127.0.0.1:1080 +caCertFile: /etc/backupx-agent/ca.pem insecureSkipTlsVerify: true ` if err := os.WriteFile(path, []byte(content), 0644); err != nil { @@ -35,6 +38,9 @@ insecureSkipTlsVerify: true if !cfg.InsecureSkipTLSVerify { t.Errorf("insecure should be true") } + if cfg.ProxyURL != "socks5h://127.0.0.1:1080" || cfg.CACertFile != "/etc/backupx-agent/ca.pem" { + t.Errorf("connection options not loaded: %+v", cfg) + } } func TestLoadConfigDefaults(t *testing.T) { @@ -64,8 +70,15 @@ func TestConfigValidate(t *testing.T) { {"valid", Config{Master: "http://m", Token: "t"}, false}, {"missing master", Config{Token: "t"}, true}, {"missing token", Config{Master: "http://m"}, true}, + {"invalid master scheme", Config{Master: "ssh://m", Token: "t"}, true}, + {"master credentials rejected", Config{Master: "https://user:pass@m", Token: "t"}, true}, + {"valid socks proxy", Config{Master: "https://m", Token: "t", ProxyURL: "socks5h://127.0.0.1:1080"}, false}, + {"invalid proxy", Config{Master: "https://m", Token: "t", ProxyURL: "ftp://proxy"}, true}, + {"proxy path rejected", Config{Master: "https://m", Token: "t", ProxyURL: "http://proxy/connect"}, true}, + {"invalid heartbeat", Config{Master: "https://m", Token: "t", HeartbeatInterval: "never"}, true}, } for _, c := range cases { + _, _ = applyConfigDefaults(&c.cfg) err := c.cfg.Validate() if (err != nil) != c.wantErr { t.Errorf("%s: err=%v wantErr=%v", c.name, err, c.wantErr) @@ -75,7 +88,7 @@ func TestConfigValidate(t *testing.T) { func TestMergeWithFlags(t *testing.T) { cfg := &Config{Master: "http://old", Token: "old"} - cfg.MergeWithFlags("http://new", "", "/tmp/x") + cfg.ApplyOverrides(Overrides{Master: "http://new", TempDir: "/tmp/x", ProxyURL: "http://proxy:3128"}) if cfg.Master != "http://new" { t.Errorf("master not overridden: %q", cfg.Master) } @@ -85,17 +98,50 @@ func TestMergeWithFlags(t *testing.T) { if cfg.TempDir != "/tmp/x" { t.Errorf("tempDir: %q", cfg.TempDir) } + if cfg.ProxyURL != "http://proxy:3128" { + t.Errorf("proxyUrl: %q", cfg.ProxyURL) + } +} + +func TestTokenFileOverrideReplacesConfiguredToken(t *testing.T) { + path := filepath.Join(t.TempDir(), "agent.token") + if err := os.WriteFile(path, []byte("file-token\n"), 0o600); err != nil { + t.Fatal(err) + } + cfg := &Config{Token: "yaml-token", TokenFile: "/old/token"} + cfg.ApplyOverrides(Overrides{TokenFile: " " + path + " "}) + if err := cfg.ResolveToken(); err != nil { + t.Fatal(err) + } + if cfg.Token != "file-token" || cfg.TokenFile != path { + t.Fatalf("token file override was not applied: %+v", cfg) + } } func TestLoadConfigFromEnv(t *testing.T) { t.Setenv("BACKUPX_AGENT_MASTER", "http://env-master") t.Setenv("BACKUPX_AGENT_TOKEN", "env-token") + t.Setenv("BACKUPX_AGENT_PROXY_URL", "http://env-proxy:8080") t.Setenv("BACKUPX_AGENT_INSECURE_TLS", "true") cfg, err := LoadConfigFromEnv() if err != nil { t.Fatal(err) } - if cfg.Master != "http://env-master" || cfg.Token != "env-token" || !cfg.InsecureSkipTLSVerify { + if cfg.Master != "http://env-master" || cfg.Token != "env-token" || cfg.ProxyURL != "http://env-proxy:8080" || !cfg.InsecureSkipTLSVerify { t.Errorf("env not picked up: %+v", cfg) } } + +func TestResolveTokenFile(t *testing.T) { + path := filepath.Join(t.TempDir(), "agent.token") + if err := os.WriteFile(path, []byte(" file-token\n"), 0o600); err != nil { + t.Fatal(err) + } + cfg := &Config{TokenFile: path} + if err := cfg.ResolveToken(); err != nil { + t.Fatal(err) + } + if cfg.Token != "file-token" { + t.Fatalf("token = %q", cfg.Token) + } +} diff --git a/server/internal/config/config.go b/server/internal/config/config.go index 05fcb85..33871fe 100644 --- a/server/internal/config/config.go +++ b/server/internal/config/config.go @@ -2,6 +2,8 @@ package config import ( "fmt" + "net" + "net/url" "strings" "time" @@ -21,6 +23,9 @@ type ServerConfig struct { Port int `mapstructure:"port"` Mode string `mapstructure:"mode"` ExternalURL string `mapstructure:"external_url"` + // TrustedProxies 限定可提供 X-Forwarded-For 等头部的反向代理地址。 + // 默认仅信任本机代理;空列表表示不信任任何代理头。 + TrustedProxies []string `mapstructure:"trusted_proxies"` // WebRoot 指向前端构建产物目录。留空时后端会按部署惯例自动探测 // (./web、./web/dist、/opt/backupx/web 等)。探测命中后后端直接托管 // 前端 SPA,无需额外的 nginx 反向代理即可访问 Web 控制台。 @@ -91,6 +96,25 @@ func Load(configPath string) (Config, error) { if cfg.Server.Mode == "" { cfg.Server.Mode = "release" } + cfg.Server.ExternalURL = strings.TrimRight(strings.TrimSpace(cfg.Server.ExternalURL), "/") + if cfg.Server.ExternalURL != "" { + externalURL, parseErr := url.Parse(cfg.Server.ExternalURL) + if parseErr != nil || (externalURL.Scheme != "http" && externalURL.Scheme != "https") || externalURL.Host == "" || externalURL.User != nil || externalURL.RawQuery != "" || externalURL.Fragment != "" { + return Config{}, fmt.Errorf("server.external_url must be an absolute http(s) URL without credentials, query or fragment") + } + } + if len(cfg.Server.TrustedProxies) == 1 && strings.Contains(cfg.Server.TrustedProxies[0], ",") { + cfg.Server.TrustedProxies = strings.Split(cfg.Server.TrustedProxies[0], ",") + } + for index := range cfg.Server.TrustedProxies { + proxy := strings.TrimSpace(cfg.Server.TrustedProxies[index]) + cfg.Server.TrustedProxies[index] = proxy + if net.ParseIP(proxy) == nil { + if _, _, parseErr := net.ParseCIDR(proxy); parseErr != nil { + return Config{}, fmt.Errorf("server.trusted_proxies contains invalid IP or CIDR %q", proxy) + } + } + } if cfg.Database.Path == "" { cfg.Database.Path = "./data/backupx.db" } @@ -142,6 +166,7 @@ func applyDefaults(v *viper.Viper) { v.SetDefault("server.port", 8340) v.SetDefault("server.mode", "release") v.SetDefault("server.external_url", "") + v.SetDefault("server.trusted_proxies", []string{"127.0.0.1", "::1"}) v.SetDefault("server.web_root", "") v.SetDefault("database.path", "./data/backupx.db") v.SetDefault("security.jwt_expire", "24h") diff --git a/server/internal/config/config_test.go b/server/internal/config/config_test.go index 7fe7d7e..3e59894 100644 --- a/server/internal/config/config_test.go +++ b/server/internal/config/config_test.go @@ -21,6 +21,25 @@ func TestLoadUsesDefaultsWithoutConfigFile(t *testing.T) { if cfg.Database.Path != "./data/backupx.db" { t.Fatalf("expected default database path, got %s", cfg.Database.Path) } + if len(cfg.Server.TrustedProxies) != 2 { + t.Fatalf("expected loopback trusted proxies, got %#v", cfg.Server.TrustedProxies) + } +} + +func TestLoadRejectsInvalidExternalURLAndTrustedProxy(t *testing.T) { + tests := []string{ + "server:\n external_url: \"ssh://master.example.com\"\n", + "server:\n trusted_proxies: [\"not-an-ip\"]\n", + } + for _, content := range tests { + configPath := filepath.Join(t.TempDir(), "config.yaml") + if err := os.WriteFile(configPath, []byte(content), 0o600); err != nil { + t.Fatal(err) + } + if _, err := Load(configPath); err == nil { + t.Fatalf("expected invalid configuration to fail: %s", content) + } + } } func TestLoadReadsServerExternalURLFromFile(t *testing.T) { @@ -52,3 +71,28 @@ func TestLoadReadsServerExternalURLFromEnv(t *testing.T) { t.Fatalf("expected external URL from env, got %q", cfg.Server.ExternalURL) } } + +func TestLoadReadsTrustedProxiesFromEnv(t *testing.T) { + t.Setenv("BACKUPX_SERVER_TRUSTED_PROXIES", "127.0.0.1,172.18.0.0/16") + cfg, err := Load("") + if err != nil { + t.Fatal(err) + } + if len(cfg.Server.TrustedProxies) != 2 || cfg.Server.TrustedProxies[1] != "172.18.0.0/16" { + t.Fatalf("trusted proxies = %#v", cfg.Server.TrustedProxies) + } +} + +func TestLoadAllowsTrustedProxiesToBeDisabled(t *testing.T) { + configPath := filepath.Join(t.TempDir(), "config.yaml") + if err := os.WriteFile(configPath, []byte("server:\n trusted_proxies: []\n"), 0o600); err != nil { + t.Fatal(err) + } + cfg, err := Load(configPath) + if err != nil { + t.Fatal(err) + } + if len(cfg.Server.TrustedProxies) != 0 { + t.Fatalf("trusted proxies should be disabled, got %#v", cfg.Server.TrustedProxies) + } +} diff --git a/server/internal/http/agent_handler.go b/server/internal/http/agent_handler.go index 23d28c1..7912390 100644 --- a/server/internal/http/agent_handler.go +++ b/server/internal/http/agent_handler.go @@ -23,7 +23,7 @@ func NewAgentHandler(agentService *service.AgentService, nodeService *service.No return &AgentHandler{agentService: agentService, nodeService: nodeService, restoreService: restoreService} } -// extractToken 从请求头或 JSON body 中提取 Agent Token。 +// extractToken 从认证请求头中提取 Agent Token。 func extractToken(c *gin.Context) string { if t := strings.TrimSpace(c.GetHeader("X-Agent-Token")); t != "" { return t @@ -46,10 +46,10 @@ func (h *AgentHandler) Heartbeat(c *gin.Context) { Arch string `json:"arch"` } _ = c.ShouldBindJSON(&input) - // token 优先走 body(向后兼容),否则从 header 读 - token := input.Token + // 新版 Agent 只通过请求头发送 Token;JSON body 仅保留旧版本兼容。 + token := extractToken(c) if token == "" { - token = extractToken(c) + token = input.Token } if token == "" { c.JSON(stdhttp.StatusBadRequest, gin.H{"code": "INVALID_INPUT", "message": "missing token"}) @@ -72,7 +72,7 @@ func (h *AgentHandler) Heartbeat(c *gin.Context) { }) } -// Poll Agent 长轮询获取下一条待执行命令。 +// Poll Agent 获取下一条待执行命令;Agent 按配置间隔主动轮询。 // 无命令时返回 {command: null}。 func (h *AgentHandler) Poll(c *gin.Context) { node, err := h.agentService.AuthenticatedNode(c.Request.Context(), extractToken(c)) diff --git a/server/internal/http/forwarded_headers_test.go b/server/internal/http/forwarded_headers_test.go new file mode 100644 index 0000000..41369b4 --- /dev/null +++ b/server/internal/http/forwarded_headers_test.go @@ -0,0 +1,49 @@ +package http + +import ( + stdhttp "net/http" + "net/http/httptest" + "testing" + + "github.com/gin-gonic/gin" +) + +func TestForwardedHeadersMiddlewareRejectsUntrustedHeaders(t *testing.T) { + gin.SetMode(gin.TestMode) + engine := gin.New() + engine.Use(ForwardedHeadersMiddleware([]string{"127.0.0.1", "10.0.0.0/8"})) + engine.GET("/master-url", func(c *gin.Context) { + c.String(stdhttp.StatusOK, resolveMasterURL(c, "")) + }) + + request := httptest.NewRequest(stdhttp.MethodGet, "http://master.example.com/master-url", nil) + request.RemoteAddr = "203.0.113.10:54321" + request.Header.Set("X-Forwarded-Host", "attacker.example.com") + request.Header.Set("X-Forwarded-Proto", "https") + recorder := httptest.NewRecorder() + engine.ServeHTTP(recorder, request) + + if recorder.Body.String() != "http://master.example.com" { + t.Fatalf("untrusted forwarding headers changed URL: %q", recorder.Body.String()) + } +} + +func TestForwardedHeadersMiddlewareAcceptsTrustedProxy(t *testing.T) { + gin.SetMode(gin.TestMode) + engine := gin.New() + engine.Use(ForwardedHeadersMiddleware([]string{"10.0.0.0/8"})) + engine.GET("/master-url", func(c *gin.Context) { + c.String(stdhttp.StatusOK, resolveMasterURL(c, "")) + }) + + request := httptest.NewRequest(stdhttp.MethodGet, "http://backupx:8340/master-url", nil) + request.RemoteAddr = "10.10.0.5:43210" + request.Header.Set("X-Forwarded-Host", "backup.example.com") + request.Header.Set("X-Forwarded-Proto", "https") + recorder := httptest.NewRecorder() + engine.ServeHTTP(recorder, request) + + if recorder.Body.String() != "https://backup.example.com" { + t.Fatalf("trusted forwarding headers were ignored: %q", recorder.Body.String()) + } +} diff --git a/server/internal/http/middleware.go b/server/internal/http/middleware.go index 7e79a0b..f9991b6 100644 --- a/server/internal/http/middleware.go +++ b/server/internal/http/middleware.go @@ -3,6 +3,7 @@ package http import ( "context" stdhttp "net/http" + "net/netip" "strings" "backupx/server/internal/apperror" @@ -11,6 +12,46 @@ import ( "github.com/gin-gonic/gin" ) +// ForwardedHeadersMiddleware 只允许配置中的反向代理提供转发头。 +// Gin 的 trusted_proxies 保护 ClientIP;这里同步保护安装命令使用的 +// X-Forwarded-Host 与 X-Forwarded-Proto,避免直连请求伪造 Agent 地址。 +func ForwardedHeadersMiddleware(trustedProxies []string) gin.HandlerFunc { + trustedPrefixes := make([]netip.Prefix, 0, len(trustedProxies)) + for _, raw := range trustedProxies { + raw = strings.TrimSpace(raw) + if prefix, err := netip.ParsePrefix(raw); err == nil { + trustedPrefixes = append(trustedPrefixes, prefix) + continue + } + if addr, err := netip.ParseAddr(raw); err == nil { + trustedPrefixes = append(trustedPrefixes, netip.PrefixFrom(addr, addr.BitLen())) + } + } + + return func(c *gin.Context) { + remote, err := netip.ParseAddrPort(c.Request.RemoteAddr) + trusted := false + if err == nil { + remoteAddr := remote.Addr().Unmap() + for _, prefix := range trustedPrefixes { + if prefix.Contains(remoteAddr) { + trusted = true + break + } + } + } + if !trusted { + for _, header := range []string{ + "Forwarded", "X-Forwarded-For", "X-Forwarded-Host", + "X-Forwarded-Port", "X-Forwarded-Proto", "X-Real-IP", + } { + c.Request.Header.Del(header) + } + } + c.Next() + } +} + // CORSMiddleware handles Cross-Origin Resource Sharing for the API. func CORSMiddleware() gin.HandlerFunc { return func(c *gin.Context) { diff --git a/server/internal/http/router.go b/server/internal/http/router.go index f003fb2..2c02949 100644 --- a/server/internal/http/router.go +++ b/server/internal/http/router.go @@ -61,7 +61,11 @@ type RouterDependencies struct { func NewRouter(deps RouterDependencies) *gin.Engine { gin.SetMode(deps.Config.Server.Mode) engine := gin.New() + if err := engine.SetTrustedProxies(deps.Config.Server.TrustedProxies); err != nil { + panic("invalid trusted proxy configuration: " + err.Error()) + } engine.Use(gin.Recovery()) + engine.Use(ForwardedHeadersMiddleware(deps.Config.Server.TrustedProxies)) engine.Use(CORSMiddleware()) engine.Use(requestLogger(deps.Logger))