diff --git a/go-backend/internal/http/handler/control_plane.go b/go-backend/internal/http/handler/control_plane.go index 82942d7..f3db8ac 100644 --- a/go-backend/internal/http/handler/control_plane.go +++ b/go-backend/internal/http/handler/control_plane.go @@ -6,6 +6,7 @@ import ( "fmt" "net" "net/http" + "net/url" "sort" "strconv" "strings" @@ -1612,29 +1613,55 @@ func buildForwarderNodes(targets []string) []map[string]interface{} { } func processServerAddress(serverAddr string) string { - serverAddr = strings.TrimSpace(serverAddr) + serverAddr = normalizeServerAddressInput(serverAddr) if serverAddr == "" { return serverAddr } if strings.HasPrefix(serverAddr, "[") { return serverAddr } + if looksLikeIPv6(serverAddr) { + return "[" + strings.Trim(serverAddr, "[]") + "]" + } + idx := strings.LastIndex(serverAddr, ":") if idx < 0 { - if looksLikeIPv6(serverAddr) { - return "[" + serverAddr + "]" - } return serverAddr } + if strings.Count(serverAddr, ":") != 1 { + return "[" + strings.Trim(serverAddr, "[]") + "]" + } + host := strings.TrimSpace(serverAddr[:idx]) port := strings.TrimSpace(serverAddr[idx+1:]) if host == "" || port == "" { return serverAddr } if looksLikeIPv6(host) { - return "[" + host + "]:" + port + return "[" + strings.Trim(host, "[]") + "]:" + port } - return serverAddr + return host + ":" + port +} + +func normalizeServerAddressInput(serverAddr string) string { + serverAddr = strings.TrimSpace(serverAddr) + if serverAddr == "" { + return serverAddr + } + + if idx := strings.Index(serverAddr, "://"); idx > 0 { + if parsed, err := url.Parse(serverAddr); err == nil { + if host := strings.TrimSpace(parsed.Host); host != "" { + return host + } + } + serverAddr = serverAddr[idx+3:] + } + + if idx := strings.IndexAny(serverAddr, "/?#"); idx >= 0 { + serverAddr = serverAddr[:idx] + } + return strings.TrimSpace(serverAddr) } func looksLikeIPv6(address string) bool { diff --git a/go-backend/internal/http/handler/control_plane_test.go b/go-backend/internal/http/handler/control_plane_test.go index a693e0e..d751868 100644 --- a/go-backend/internal/http/handler/control_plane_test.go +++ b/go-backend/internal/http/handler/control_plane_test.go @@ -420,3 +420,68 @@ func TestBuildForwardServiceConfigs_BindIPAlreadyContainsPort(t *testing.T) { } } } + +func TestProcessServerAddress_StripsURLSchemeAndPath(t *testing.T) { + tests := []struct { + name string + in string + want string + }{ + { + name: "https with path", + in: "https://panel.example.com:8443/api/v1", + want: "panel.example.com:8443", + }, + { + name: "wss with query", + in: "wss://panel.example.com:443/system-info?x=1", + want: "panel.example.com:443", + }, + { + name: "http without port", + in: "http://panel.example.com", + want: "panel.example.com", + }, + { + name: "manual host with trailing path", + in: "panel.example.com:8080/path", + want: "panel.example.com:8080", + }, + } + + for _, tt := range tests { + if got := processServerAddress(tt.in); got != tt.want { + t.Fatalf("%s: expected %q, got %q", tt.name, tt.want, got) + } + } +} + +func TestProcessServerAddress_NormalizesIPv6(t *testing.T) { + tests := []struct { + name string + in string + want string + }{ + { + name: "ipv6 host only", + in: "2001:db8::1", + want: "[2001:db8::1]", + }, + { + name: "ipv6 host and port", + in: "https://[2001:db8::1]:8443/path", + want: "[2001:db8::1]:8443", + }, + { + name: "already bracketed", + in: "[2001:db8::2]:9000", + want: "[2001:db8::2]:9000", + }, + } + + for _, tt := range tests { + if got := processServerAddress(tt.in); got != tt.want { + t.Fatalf("%s: expected %q, got %q", tt.name, tt.want, got) + } + } +} diff --git a/go-gost/main.go b/go-gost/main.go index bf1baa0..42dc9ec 100644 --- a/go-gost/main.go +++ b/go-gost/main.go @@ -109,12 +109,12 @@ func main() { // 加载配置文件 config, err := LoadConfig("config.json") if err != nil { - fmt.Println("❌ 配置加载失败: %v\n", err) + fmt.Printf("❌ 配置加载失败: %v\n", err) fmt.Println("请确保当前目录存在 config.json 文件") os.Exit(1) } - fmt.Println("✅ 配置加载成功 - addr: %s", config.Addr) + fmt.Printf("✅ 配置加载成功 - addr: %s\n", config.Addr) log := xlogger.NewLogger() logger.SetDefault(log) diff --git a/go-gost/x/service/traffic_reporter.go b/go-gost/x/service/traffic_reporter.go index 3e448a1..520f560 100644 --- a/go-gost/x/service/traffic_reporter.go +++ b/go-gost/x/service/traffic_reporter.go @@ -6,7 +6,9 @@ import ( "encoding/json" "fmt" "net/http" + "net/url" "strings" + "sync" "time" "github.com/go-gost/core/observer/stats" @@ -18,6 +20,15 @@ import ( var httpReportURL string var configReportURL string var httpAESCrypto *crypto.AESCrypto // 新增:HTTP上报加密器 +var reportURLPreferenceMutex sync.RWMutex +var preferredUploadURL string +var preferredConfigURL string +var reportDo = func(ctx context.Context, req *http.Request, timeout time.Duration) (*http.Response, error) { + client := &http.Client{ + Timeout: timeout, + } + return client.Do(req.WithContext(ctx)) +} // TrafficReportItem 流量报告项(压缩格式) type TrafficReportItem struct { @@ -27,8 +38,17 @@ type TrafficReportItem struct { } func SetHTTPReportURL(addr string, secret string) { - httpReportURL = "http://" + addr + "/flow/upload?secret=" + secret - configReportURL = "http://" + addr + "/flow/config?secret=" + secret + uploadURLs, configURLs := buildReportURLCandidates(addr, secret) + if len(uploadURLs) > 0 { + httpReportURL = strings.Join(uploadURLs, ",") + } + if len(configURLs) > 0 { + configReportURL = strings.Join(configURLs, ",") + } + reportURLPreferenceMutex.Lock() + preferredUploadURL = "" + preferredConfigURL = "" + reportURLPreferenceMutex.Unlock() // 创建 AES 加密器 var err error @@ -41,8 +61,173 @@ func SetHTTPReportURL(addr string, secret string) { } } +func buildReportURLCandidates(addr string, secret string) (upload []string, config []string) { + normalizedAddr, explicitScheme := normalizeReportAddress(addr) + if normalizedAddr == "" { + normalizedAddr = strings.TrimSpace(addr) + } + + schemes := []string{"https", "http"} + if mappedScheme := mapToHTTPScheme(explicitScheme); mappedScheme == "http" { + schemes = []string{"http", "https"} + } + + upload = []string{ + schemes[0] + "://" + normalizedAddr + "/flow/upload?secret=" + secret, + schemes[1] + "://" + normalizedAddr + "/flow/upload?secret=" + secret, + } + config = []string{ + schemes[0] + "://" + normalizedAddr + "/flow/config?secret=" + secret, + schemes[1] + "://" + normalizedAddr + "/flow/config?secret=" + secret, + } + return upload, config +} + +func normalizeReportAddress(addr string) (string, string) { + raw := strings.TrimSpace(addr) + if raw == "" { + return "", "" + } + + scheme := "" + if idx := strings.Index(raw, "://"); idx > 0 { + scheme = strings.ToLower(strings.TrimSpace(raw[:idx])) + if parsed, err := url.Parse(raw); err == nil { + if host := strings.TrimSpace(parsed.Host); host != "" { + return host, scheme + } + } + raw = raw[idx+3:] + } + + if idx := strings.IndexAny(raw, "/?#"); idx >= 0 { + raw = raw[:idx] + } + return strings.TrimSpace(raw), scheme +} + +func mapToHTTPScheme(scheme string) string { + switch strings.ToLower(strings.TrimSpace(scheme)) { + case "https", "wss": + return "https" + case "http", "ws": + return "http" + default: + return "" + } +} + +func loadPreferredURL(preferred *string) string { + if preferred == nil { + return "" + } + + reportURLPreferenceMutex.RLock() + defer reportURLPreferenceMutex.RUnlock() + return *preferred +} + +func storePreferredURL(preferred *string, value string) { + if preferred == nil { + return + } + + reportURLPreferenceMutex.Lock() + defer reportURLPreferenceMutex.Unlock() + *preferred = value +} + +func prioritizeURLs(urls []string, preferred string) []string { + ordered := append([]string(nil), urls...) + if preferred == "" || len(ordered) < 2 { + return ordered + } + + for i, targetURL := range ordered { + if targetURL == preferred { + if i > 0 { + ordered[0], ordered[i] = ordered[i], ordered[0] + } + break + } + } + + return ordered +} + +func postJSONWithFallback(ctx context.Context, urls []string, requestBody []byte, userAgent string, timeout time.Duration, preferred *string) (bool, error) { + if len(urls) == 0 { + return false, fmt.Errorf("上报URL未设置") + } + + orderedURLs := prioritizeURLs(urls, loadPreferredURL(preferred)) + + var errs []string + for i, targetURL := range orderedURLs { + req, err := http.NewRequest("POST", targetURL, bytes.NewBuffer(requestBody)) + if err != nil { + errs = append(errs, fmt.Sprintf("%s => 创建请求失败: %v", targetURL, err)) + if i < len(orderedURLs)-1 { + fmt.Printf("⚠️ HTTP上报尝试失败,准备回退: %s => 创建请求失败: %v\n", targetURL, err) + } + continue + } + + req.Header.Set("Content-Type", "application/json") + req.Header.Set("User-Agent", userAgent) + + resp, err := reportDo(ctx, req, timeout) + if err != nil { + errs = append(errs, fmt.Sprintf("%s => 请求失败: %v", targetURL, err)) + if i < len(orderedURLs)-1 { + fmt.Printf("⚠️ HTTP上报尝试失败,准备回退: %s => 请求失败: %v\n", targetURL, err) + } + continue + } + + var responseBytes bytes.Buffer + _, readErr := responseBytes.ReadFrom(resp.Body) + resp.Body.Close() + if readErr != nil { + errs = append(errs, fmt.Sprintf("%s => 读取响应失败: %v", targetURL, readErr)) + if i < len(orderedURLs)-1 { + fmt.Printf("⚠️ HTTP上报尝试失败,准备回退: %s => 读取响应失败: %v\n", targetURL, readErr) + } + continue + } + + if resp.StatusCode != http.StatusOK { + errs = append(errs, fmt.Sprintf("%s => HTTP响应错误: %d %s", targetURL, resp.StatusCode, resp.Status)) + if i < len(orderedURLs)-1 { + fmt.Printf("⚠️ HTTP上报尝试失败,准备回退: %s => HTTP响应错误: %d %s\n", targetURL, resp.StatusCode, resp.Status) + } + continue + } + + responseText := strings.TrimSpace(responseBytes.String()) + if responseText == "ok" { + if i > 0 { + fmt.Printf("↪️ HTTP上报已自动回退到: %s\n", targetURL) + } + storePreferredURL(preferred, targetURL) + return true, nil + } + + errs = append(errs, fmt.Sprintf("%s => 服务器响应: %s (期望: ok)", targetURL, responseText)) + if i < len(orderedURLs)-1 { + fmt.Printf("⚠️ HTTP上报尝试失败,准备回退: %s => 服务器响应: %s (期望: ok)\n", targetURL, responseText) + } + } + + return false, fmt.Errorf("发送HTTP请求失败: %s", strings.Join(errs, " | ")) +} + // sendBatchTrafficReport 批量发送多个服务的流量报告到HTTP接口 func sendBatchTrafficReport(ctx context.Context, reportItems []TrafficReportItem) (bool, error) { + if httpReportURL == "" { + return false, fmt.Errorf("流量上报URL未设置") + } + jsonData, err := json.Marshal(reportItems) if err != nil { return false, fmt.Errorf("序列化报告数据失败: %v", err) @@ -73,46 +258,16 @@ func sendBatchTrafficReport(ctx context.Context, reportItems []TrafficReportItem requestBody = jsonData } - req, err := http.NewRequestWithContext(ctx, "POST", httpReportURL, bytes.NewBuffer(requestBody)) - if err != nil { - return false, fmt.Errorf("创建HTTP请求失败: %v", err) - } - - req.Header.Set("Content-Type", "application/json") - req.Header.Set("User-Agent", "GOST-Traffic-Reporter/1.0") - - client := &http.Client{ - Timeout: 5 * time.Second, - } - - resp, err := client.Do(req) - if err != nil { - return false, fmt.Errorf("发送HTTP请求失败: %v", err) - } - defer resp.Body.Close() - - if resp.StatusCode != http.StatusOK { - return false, fmt.Errorf("HTTP响应错误: %d %s", resp.StatusCode, resp.Status) - } - - // 读取响应内容 - var responseBytes bytes.Buffer - _, err = responseBytes.ReadFrom(resp.Body) - if err != nil { - return false, fmt.Errorf("读取响应内容失败: %v", err) - } - - responseText := strings.TrimSpace(responseBytes.String()) - - // 检查响应是否为"ok" - if responseText == "ok" { - return true, nil - } else { - return false, fmt.Errorf("服务器响应: %s (期望: ok)", responseText) - } + return postJSONWithFallback( + ctx, + strings.Split(httpReportURL, ","), + requestBody, + "GOST-Traffic-Reporter/1.0", + 5*time.Second, + &preferredUploadURL, + ) } - // sendConfigReport 发送配置报告到HTTP接口 func sendConfigReport(ctx context.Context) (bool, error) { if configReportURL == "" { @@ -150,43 +305,14 @@ func sendConfigReport(ctx context.Context) (bool, error) { requestBody = configData } - req, err := http.NewRequestWithContext(ctx, "POST", configReportURL, bytes.NewBuffer(requestBody)) - if err != nil { - return false, fmt.Errorf("创建HTTP请求失败: %v", err) - } - - req.Header.Set("Content-Type", "application/json") - req.Header.Set("User-Agent", "Config-Reporter/1.0") - - client := &http.Client{ - Timeout: 10 * time.Second, // 配置上报可以稍长一些 - } - - resp, err := client.Do(req) - if err != nil { - return false, fmt.Errorf("发送HTTP请求失败: %v", err) - } - defer resp.Body.Close() - - if resp.StatusCode != http.StatusOK { - return false, fmt.Errorf("HTTP响应错误: %d %s", resp.StatusCode, resp.Status) - } - - // 读取响应内容 - var responseBytes bytes.Buffer - _, err = responseBytes.ReadFrom(resp.Body) - if err != nil { - return false, fmt.Errorf("读取响应内容失败: %v", err) - } - - responseText := strings.TrimSpace(responseBytes.String()) - - // 检查响应是否为"ok" - if responseText == "ok" { - return true, nil - } else { - return false, fmt.Errorf("服务器响应: %s (期望: ok)", responseText) - } + return postJSONWithFallback( + ctx, + strings.Split(configReportURL, ","), + requestBody, + "Config-Reporter/1.0", + 10*time.Second, + &preferredConfigURL, + ) } // StartConfigReporter 启动配置定时上报器(每10分钟上报一次) diff --git a/go-gost/x/service/traffic_reporter_test.go b/go-gost/x/service/traffic_reporter_test.go new file mode 100644 index 0000000..0fdc651 --- /dev/null +++ b/go-gost/x/service/traffic_reporter_test.go @@ -0,0 +1,150 @@ +package service + +import ( + "context" + "errors" + "io" + "net/http" + "strings" + "testing" + "time" +) + +func TestBuildReportURLCandidatesSecureFirst(t *testing.T) { + upload, config := buildReportURLCandidates("panel.example.com:443", "abc") + + if len(upload) != 2 { + t.Fatalf("expected 2 upload candidates, got %d", len(upload)) + } + if len(config) != 2 { + t.Fatalf("expected 2 config candidates, got %d", len(config)) + } + + if upload[0] != "https://panel.example.com:443/flow/upload?secret=abc" { + t.Fatalf("unexpected upload[0]: %s", upload[0]) + } + if upload[1] != "http://panel.example.com:443/flow/upload?secret=abc" { + t.Fatalf("unexpected upload[1]: %s", upload[1]) + } + if config[0] != "https://panel.example.com:443/flow/config?secret=abc" { + t.Fatalf("unexpected config[0]: %s", config[0]) + } + if config[1] != "http://panel.example.com:443/flow/config?secret=abc" { + t.Fatalf("unexpected config[1]: %s", config[1]) + } +} + +func TestBuildReportURLCandidatesNormalizeSchemeAddr(t *testing.T) { + upload, config := buildReportURLCandidates("https://panel.example.com:8443/path", "abc") + + if upload[0] != "https://panel.example.com:8443/flow/upload?secret=abc" { + t.Fatalf("unexpected upload[0]: %s", upload[0]) + } + if upload[1] != "http://panel.example.com:8443/flow/upload?secret=abc" { + t.Fatalf("unexpected upload[1]: %s", upload[1]) + } + if config[0] != "https://panel.example.com:8443/flow/config?secret=abc" { + t.Fatalf("unexpected config[0]: %s", config[0]) + } + if config[1] != "http://panel.example.com:8443/flow/config?secret=abc" { + t.Fatalf("unexpected config[1]: %s", config[1]) + } +} + +func TestPostJSONWithFallbackUsesHTTPAfterHTTPSFailure(t *testing.T) { + orig := reportDo + defer func() { reportDo = orig }() + + var calls []string + reportDo = func(_ context.Context, req *http.Request, _ time.Duration) (*http.Response, error) { + calls = append(calls, req.URL.String()) + if strings.HasPrefix(req.URL.String(), "https://") { + return nil, errors.New("tls handshake failed") + } + return &http.Response{ + StatusCode: http.StatusOK, + Body: io.NopCloser(strings.NewReader("ok")), + }, nil + } + + ok, err := postJSONWithFallback( + context.Background(), + []string{ + "https://panel.example.com:443/flow/upload?secret=abc", + "http://panel.example.com:443/flow/upload?secret=abc", + }, + []byte(`[]`), + "GOST-Traffic-Reporter/1.0", + 5*time.Second, + nil, + ) + if !ok || err != nil { + t.Fatalf("expected fallback success, ok=%v err=%v", ok, err) + } + if len(calls) != 2 { + t.Fatalf("expected 2 calls, got %d", len(calls)) + } + if !strings.HasPrefix(calls[0], "https://") || !strings.HasPrefix(calls[1], "http://") { + t.Fatalf("unexpected call order: %#v", calls) + } +} + +func TestPostJSONWithFallbackRemembersDetectedURL(t *testing.T) { + orig := reportDo + defer func() { reportDo = orig }() + + targets := []string{ + "https://panel.example.com:443/flow/upload?secret=abc", + "http://panel.example.com:443/flow/upload?secret=abc", + } + + var preferred string + var calls []string + reportDo = func(_ context.Context, req *http.Request, _ time.Duration) (*http.Response, error) { + calls = append(calls, req.URL.String()) + if strings.HasPrefix(req.URL.String(), "https://") { + return nil, errors.New("tls handshake failed") + } + return &http.Response{ + StatusCode: http.StatusOK, + Body: io.NopCloser(strings.NewReader("ok")), + }, nil + } + + ok, err := postJSONWithFallback( + context.Background(), + targets, + []byte(`[]`), + "GOST-Traffic-Reporter/1.0", + 5*time.Second, + &preferred, + ) + if !ok || err != nil { + t.Fatalf("expected first call success, ok=%v err=%v", ok, err) + } + if preferred != targets[1] { + t.Fatalf("expected preferred url to be remembered as %s, got %s", targets[1], preferred) + } + if len(calls) != 2 { + t.Fatalf("expected 2 calls on first attempt, got %d", len(calls)) + } + + calls = nil + ok, err = postJSONWithFallback( + context.Background(), + targets, + []byte(`[]`), + "GOST-Traffic-Reporter/1.0", + 5*time.Second, + &preferred, + ) + if !ok || err != nil { + t.Fatalf("expected second call success, ok=%v err=%v", ok, err) + } + if len(calls) != 1 { + t.Fatalf("expected second call to use remembered url once, got %d calls", len(calls)) + } + if !strings.HasPrefix(calls[0], "http://") { + t.Fatalf("expected remembered http url first, got %s", calls[0]) + } +} diff --git a/go-gost/x/socket/websocket_reporter.go b/go-gost/x/socket/websocket_reporter.go index 743f510..ab03c64 100644 --- a/go-gost/x/socket/websocket_reporter.go +++ b/go-gost/x/socket/websocket_reporter.go @@ -97,20 +97,25 @@ const ( ) type WebSocketReporter struct { - url string - addr string // 保存服务器地址 - secret string // 保存密钥 - version string // 保存版本号 - conn *websocket.Conn - reconnectTime time.Duration - pingInterval time.Duration - configInterval time.Duration - ctx context.Context - cancel context.CancelFunc - connected bool - connecting bool // 新增:正在连接状态 - connMutex sync.Mutex // 新增:连接状态锁 - aesCrypto *crypto.AESCrypto // 新增:AES加密器 + url string + addr string // 保存服务器地址 + secret string // 保存密钥 + version string // 保存版本号 + preferredWSScheme string + conn *websocket.Conn + reconnectTime time.Duration + pingInterval time.Duration + configInterval time.Duration + ctx context.Context + cancel context.CancelFunc + connected bool + connecting bool // 新增:正在连接状态 + connMutex sync.Mutex // 新增:连接状态锁 + aesCrypto *crypto.AESCrypto // 新增:AES加密器 +} + +var wsDial = func(dialer *websocket.Dialer, rawURL string) (*websocket.Conn, *http.Response, error) { + return dialer.Dial(rawURL, nil) } // NewWebSocketReporter 创建一个新的WebSocket报告器 @@ -223,21 +228,14 @@ func (w *WebSocketReporter) connect() error { json.Unmarshal(b, &cfg) } - // 使用最新的配置重新构建 URL - currentURL := "ws://" + w.addr + "/system-info?type=1&secret=" + w.secret + "&version=" + w.version + - "&http=" + strconv.Itoa(cfg.Http) + "&tls=" + strconv.Itoa(cfg.Tls) + "&socks=" + strconv.Itoa(cfg.Socks) - - u, err := url.Parse(currentURL) - if err != nil { - return fmt.Errorf("解析URL失败: %v", err) - } + candidates := buildWebSocketCandidates(w.addr, w.secret, w.version, cfg.Http, cfg.Tls, cfg.Socks, w.preferredWSScheme) dialer := websocket.DefaultDialer dialer.HandshakeTimeout = 10 * time.Second - conn, _, err := dialer.Dial(u.String(), nil) + conn, usedURL, err := dialWebSocketWithFallback(dialer, candidates) if err != nil { - return fmt.Errorf("连接WebSocket失败: %v", err) + return err } // 如果在连接过程中已经有连接了,关闭新连接 @@ -248,6 +246,9 @@ func (w *WebSocketReporter) connect() error { w.conn = conn w.connected = true + if scheme := detectWebSocketScheme(usedURL); scheme != "" { + w.preferredWSScheme = scheme + } _ = conn.SetReadDeadline(time.Now().Add(reporterReadWait)) conn.SetPingHandler(func(appData string) error { _ = conn.SetReadDeadline(time.Now().Add(reporterReadWait)) @@ -265,10 +266,101 @@ func (w *WebSocketReporter) connect() error { return nil }) - fmt.Printf("✅ WebSocket连接建立成功 (http=%d, tls=%d, socks=%d)\n", cfg.Http, cfg.Tls, cfg.Socks) + fmt.Printf("✅ WebSocket连接建立成功 (%s, http=%d, tls=%d, socks=%d)\n", usedURL, cfg.Http, cfg.Tls, cfg.Socks) return nil } +func buildWebSocketCandidates(addr string, secret string, version string, http int, tls int, socks int, preferredScheme string) []string { + normalizedAddr, explicitScheme := normalizeReporterAddress(addr) + if normalizedAddr == "" { + normalizedAddr = strings.TrimSpace(addr) + } + + query := "/system-info?type=1&secret=" + secret + "&version=" + version + + "&http=" + strconv.Itoa(http) + "&tls=" + strconv.Itoa(tls) + "&socks=" + strconv.Itoa(socks) + + schemes := []string{"wss", "ws"} + if mappedScheme := mapToWebSocketScheme(explicitScheme); mappedScheme != "" { + if mappedScheme == "ws" { + schemes = []string{"ws", "wss"} + } + } else if preferredScheme == "ws" { + schemes = []string{"ws", "wss"} + } + + return []string{ + schemes[0] + "://" + normalizedAddr + query, + schemes[1] + "://" + normalizedAddr + query, + } +} + +func normalizeReporterAddress(addr string) (string, string) { + raw := strings.TrimSpace(addr) + if raw == "" { + return "", "" + } + + scheme := "" + if idx := strings.Index(raw, "://"); idx > 0 { + scheme = strings.ToLower(strings.TrimSpace(raw[:idx])) + if parsed, err := url.Parse(raw); err == nil { + if host := strings.TrimSpace(parsed.Host); host != "" { + return host, scheme + } + } + raw = raw[idx+3:] + } + + if idx := strings.IndexAny(raw, "/?#"); idx >= 0 { + raw = raw[:idx] + } + return strings.TrimSpace(raw), scheme +} + +func mapToWebSocketScheme(scheme string) string { + switch strings.ToLower(strings.TrimSpace(scheme)) { + case "wss", "https": + return "wss" + case "ws", "http": + return "ws" + default: + return "" + } +} + +func detectWebSocketScheme(rawURL string) string { + if strings.HasPrefix(rawURL, "wss://") { + return "wss" + } + if strings.HasPrefix(rawURL, "ws://") { + return "ws" + } + return "" +} + +func dialWebSocketWithFallback(dialer *websocket.Dialer, candidates []string) (*websocket.Conn, string, error) { + if len(candidates) == 0 { + return nil, "", fmt.Errorf("WebSocket候选地址为空") + } + + var errs []string + for i, targetURL := range candidates { + conn, _, err := wsDial(dialer, targetURL) + if err == nil { + if i > 0 { + fmt.Printf("↪️ WebSocket已自动回退到: %s\n", targetURL) + } + return conn, targetURL, nil + } + errs = append(errs, fmt.Sprintf("%s => %v", targetURL, err)) + if i < len(candidates)-1 { + fmt.Printf("⚠️ WebSocket连接尝试失败,准备回退: %s => %v\n", targetURL, err) + } + } + + return nil, "", fmt.Errorf("连接WebSocket失败: %s", strings.Join(errs, " | ")) +} + // handleConnection 处理WebSocket连接 func (w *WebSocketReporter) handleConnection() { defer func() { @@ -1290,7 +1382,8 @@ func getMemoryInfo() MemoryInfo { func StartWebSocketReporterWithConfig(addr string, secret string, http int, tls int, socks int, version string) *WebSocketReporter { // 构建初始 WebSocket URL - fullURL := "ws://" + addr + "/system-info?type=1&secret=" + secret + "&version=" + version + "&http=" + strconv.Itoa(http) + "&tls=" + strconv.Itoa(tls) + "&socks=" + strconv.Itoa(socks) + candidates := buildWebSocketCandidates(addr, secret, version, http, tls, socks, "") + fullURL := candidates[0] fmt.Printf("🔗 WebSocket连接URL: %s\n", fullURL) diff --git a/go-gost/x/socket/websocket_reporter_test.go b/go-gost/x/socket/websocket_reporter_test.go new file mode 100644 index 0000000..570b393 --- /dev/null +++ b/go-gost/x/socket/websocket_reporter_test.go @@ -0,0 +1,98 @@ +package socket + +import ( + "errors" + "net/http" + "strings" + "testing" + + "github.com/gorilla/websocket" +) + +func TestBuildWebSocketCandidatesSecureFirst(t *testing.T) { + candidates := buildWebSocketCandidates("panel.example.com:443", "abc", "2.0.2", 1, 0, 1, "") + + if len(candidates) != 2 { + t.Fatalf("expected 2 candidates, got %d", len(candidates)) + } + if !strings.HasPrefix(candidates[0], "wss://") { + t.Fatalf("expected first candidate to start with wss://, got %s", candidates[0]) + } + if !strings.HasPrefix(candidates[1], "ws://") { + t.Fatalf("expected second candidate to start with ws://, got %s", candidates[1]) + } +} + +func TestBuildWebSocketCandidatesUsesPreferredScheme(t *testing.T) { + candidates := buildWebSocketCandidates("panel.example.com:443", "abc", "2.0.2", 1, 0, 1, "ws") + + if len(candidates) != 2 { + t.Fatalf("expected 2 candidates, got %d", len(candidates)) + } + if !strings.HasPrefix(candidates[0], "ws://") { + t.Fatalf("expected preferred ws:// candidate first, got %s", candidates[0]) + } + if !strings.HasPrefix(candidates[1], "wss://") { + t.Fatalf("expected fallback wss:// candidate second, got %s", candidates[1]) + } +} + +func TestBuildWebSocketCandidatesNormalizesSchemePrefixedAddr(t *testing.T) { + candidates := buildWebSocketCandidates("https://panel.example.com:443/path?q=1", "abc", "2.0.2", 0, 0, 0, "") + + if len(candidates) != 2 { + t.Fatalf("expected 2 candidates, got %d", len(candidates)) + } + if !strings.HasPrefix(candidates[0], "wss://panel.example.com:443/") { + t.Fatalf("expected normalized wss candidate, got %s", candidates[0]) + } + if !strings.HasPrefix(candidates[1], "ws://panel.example.com:443/") { + t.Fatalf("expected normalized ws fallback candidate, got %s", candidates[1]) + } +} + +func TestDialWebSocketWithFallbackTriesWSAfterWSSFailure(t *testing.T) { + orig := wsDial + defer func() { wsDial = orig }() + + var attempts []string + wsDial = func(_ *websocket.Dialer, rawURL string) (*websocket.Conn, *http.Response, error) { + attempts = append(attempts, rawURL) + if strings.HasPrefix(rawURL, "wss://") { + return nil, nil, errors.New("tls failed") + } + return &websocket.Conn{}, nil, nil + } + + _, usedURL, err := dialWebSocketWithFallback( + &websocket.Dialer{}, + []string{ + "wss://panel.example.com/system-info?type=1&secret=abc", + "ws://panel.example.com/system-info?type=1&secret=abc", + }, + ) + if err != nil { + t.Fatalf("expected fallback success, got err=%v", err) + } + if !strings.HasPrefix(usedURL, "ws://") { + t.Fatalf("expected fallback ws:// url, got %s", usedURL) + } + if len(attempts) != 2 { + t.Fatalf("expected 2 attempts, got %d", len(attempts)) + } + if !strings.HasPrefix(attempts[0], "wss://") || !strings.HasPrefix(attempts[1], "ws://") { + t.Fatalf("unexpected attempt order: %#v", attempts) + } +} + +func TestDetectWebSocketScheme(t *testing.T) { + if detectWebSocketScheme("wss://panel.example.com/system-info") != "wss" { + t.Fatalf("expected wss detection") + } + if detectWebSocketScheme("ws://panel.example.com/system-info") != "ws" { + t.Fatalf("expected ws detection") + } + if detectWebSocketScheme("http://panel.example.com/system-info") != "" { + t.Fatalf("expected empty detection for non-websocket scheme") + } +}