From a98057d06a99a7630b7d942c0cc306587e871c1c Mon Sep 17 00:00:00 2001 From: root Date: Wed, 4 Feb 2026 04:12:38 +0000 Subject: [PATCH] fix(gost): implement failover for multi-node forwarding rules (#12) When a forwarding rule has multiple backend nodes configured, the first node failure would cause the entire forward to fail instead of trying the next available node. Root causes fixed: - FailFilter skipped filtering when only 1 node remained - hop.Select() bypassed selector for single-node hops - Handlers only attempted one node before giving up Changes: - selector/filter.go: Remove len<=1 early return, always filter failed nodes - hop/hop.go: Remove single-node bypass, add ExcludeNodes context support - ctx/value.go: Add ContextWithExcludeNodes/ExcludeNodesFromContext helpers - handler/forward/local: Add maxRetries config, implement retry loop - handler/forward/remote: Add maxRetries config, implement retry loop - forwarder/sniffer.go: Add retry logic to dial() and dialTLS() Closes #12 --- go-gost/x/ctx/value.go | 20 ++ go-gost/x/handler/forward/local/handler.go | 105 ++++--- go-gost/x/handler/forward/local/metadata.go | 9 + go-gost/x/handler/forward/remote/handler.go | 137 +++++---- go-gost/x/handler/forward/remote/metadata.go | 10 + go-gost/x/hop/hop.go | 23 +- go-gost/x/internal/util/forwarder/sniffer.go | 285 ++++++++++++------- go-gost/x/selector/filter.go | 9 +- 8 files changed, 396 insertions(+), 202 deletions(-) diff --git a/go-gost/x/ctx/value.go b/go-gost/x/ctx/value.go index 349aa08..b7161cb 100644 --- a/go-gost/x/ctx/value.go +++ b/go-gost/x/ctx/value.go @@ -109,3 +109,23 @@ func LoggerFromContext(ctx context.Context) logger.Logger { v, _ := ctx.Value(keyLogger).(logger.Logger) return v } + +// excludeNodesKey saves the list of node addresses to exclude during selection. +// This is used for failover retry logic - when a node fails, it gets added to +// the exclude list so the next Select() call will skip it. +type excludeNodesKey struct{} + +var ( + keyExcludeNodes = &excludeNodesKey{} +) + +// ContextWithExcludeNodes returns a context with the list of node addresses to exclude. +func ContextWithExcludeNodes(ctx context.Context, nodes []string) context.Context { + return context.WithValue(ctx, keyExcludeNodes, nodes) +} + +// ExcludeNodesFromContext returns the list of node addresses to exclude from selection. +func ExcludeNodesFromContext(ctx context.Context) []string { + v, _ := ctx.Value(keyExcludeNodes).([]string) + return v +} diff --git a/go-gost/x/handler/forward/local/handler.go b/go-gost/x/handler/forward/local/handler.go index b64466f..4f7881b 100644 --- a/go-gost/x/handler/forward/local/handler.go +++ b/go-gost/x/handler/forward/local/handler.go @@ -176,51 +176,84 @@ func (h *forwardHandler) Handle(ctx context.Context, conn net.Conn, opts ...hand } } - target := &chain.Node{} - if h.hop != nil { - target = h.hop.Select(ctx, - hop.ProtocolSelectOption(proto), - ) - } - if target == nil { - err := errors.New("node not available") - return err + // Determine max retry attempts + maxRetries := h.md.maxRetries + if maxRetries <= 0 { + // Default: try all available nodes + if nl, ok := h.hop.(hop.NodeList); ok { + maxRetries = len(nl.Nodes()) + } + if maxRetries <= 0 { + maxRetries = 1 + } } - addr := target.Addr - if opts := target.Options(); opts != nil { - switch opts.Network { - case "unix": - network = opts.Network - default: - if _, _, err := net.SplitHostPort(addr); err != nil { - addr += ":0" + var triedNodes []string + var lastErr error + var cc net.Conn + + for attempt := 0; attempt < maxRetries; attempt++ { + // Select a target node, excluding previously tried nodes + selectCtx := ctxvalue.ContextWithExcludeNodes(ctx, triedNodes) + target := &chain.Node{} + if h.hop != nil { + target = h.hop.Select(selectCtx, + hop.ProtocolSelectOption(proto), + ) + } + if target == nil { + if lastErr != nil { + return lastErr + } + return errors.New("node not available") + } + + // Track this node as tried + triedNodes = append(triedNodes, target.Addr) + + addr := target.Addr + if opts := target.Options(); opts != nil { + switch opts.Network { + case "unix": + network = opts.Network + default: + if _, _, err := net.SplitHostPort(addr); err != nil { + addr += ":0" + } } } - } - ro.Network = network - ro.Host = addr + ro.Network = network + ro.Host = addr - var buf bytes.Buffer - cc, err := h.options.Router.Dial(ctxvalue.ContextWithBuffer(ctx, &buf), network, addr) - ro.Route = buf.String() - if err != nil { - // TODO: the router itself may be failed due to the failed node in the router, - // the dead marker may be a wrong operation. - if marker := target.Marker(); marker != nil { - marker.Mark() + var buf bytes.Buffer + cc, err = h.options.Router.Dial(ctxvalue.ContextWithBuffer(ctx, &buf), network, addr) + ro.Route = buf.String() + if err != nil { + // Mark node as failed for future selections + if marker := target.Marker(); marker != nil { + marker.Mark() + } + lastErr = err + // Try next node + continue } - return err - } - if marker := target.Marker(); marker != nil { - marker.Reset() - } - defer cc.Close() - xnet.Transport(conn, cc) + // Success - reset marker and proceed + if marker := target.Marker(); marker != nil { + marker.Reset() + } + defer cc.Close() - return nil + xnet.Transport(conn, cc) + return nil + } + + // All retries exhausted + if lastErr != nil { + return lastErr + } + return errors.New("all nodes failed") } func (h *forwardHandler) checkRateLimit(addr net.Addr) bool { diff --git a/go-gost/x/handler/forward/local/metadata.go b/go-gost/x/handler/forward/local/metadata.go index 13ecb87..1307eea 100644 --- a/go-gost/x/handler/forward/local/metadata.go +++ b/go-gost/x/handler/forward/local/metadata.go @@ -25,6 +25,12 @@ type metadata struct { privateKey crypto.PrivateKey alpn string mitmBypass bypass.Bypass + + // maxRetries specifies the maximum number of failover retry attempts. + // When a target node fails, the handler will try the next available node. + // 0 means use the total number of available nodes (try all nodes once). + // Default: 0 (try all available nodes) + maxRetries int } func (h *forwardHandler) parseMetadata(md mdata.Metadata) (err error) { @@ -56,5 +62,8 @@ func (h *forwardHandler) parseMetadata(md mdata.Metadata) (err error) { h.md.alpn = mdutil.GetString(md, "mitm.alpn") h.md.mitmBypass = registry.BypassRegistry().Get(mdutil.GetString(md, "mitm.bypass")) + // maxRetries: 0 means try all available nodes (default behavior) + h.md.maxRetries = mdutil.GetInt(md, "maxRetries", "retry.max") + return } diff --git a/go-gost/x/handler/forward/remote/handler.go b/go-gost/x/handler/forward/remote/handler.go index 9d45497..35f604a 100644 --- a/go-gost/x/handler/forward/remote/handler.go +++ b/go-gost/x/handler/forward/remote/handler.go @@ -204,68 +204,103 @@ func (h *forwardHandler) Handle(ctx context.Context, conn net.Conn, opts ...hand } } - var target *chain.Node - if host != "" { - target = &chain.Node{ - Addr: host, + // Determine max retry attempts + maxRetries := h.md.maxRetries + if maxRetries <= 0 { + // Default: try all available nodes + if nl, ok := h.hop.(hop.NodeList); ok { + maxRetries = len(nl.Nodes()) } - } - if h.hop != nil { - target = h.hop.Select(ctx, - hop.ProtocolSelectOption(proto), - ) - } - if target == nil { - err := errors.New("node not available") - log.Error(err) - return err - } - - if opts := target.Options(); opts != nil { - switch opts.Network { - case "unix": - network = opts.Network - default: + if maxRetries <= 0 { + maxRetries = 1 } } - ro.Network = network - ro.Host = target.Addr + var triedNodes []string + var lastErr error + var cc net.Conn - log = log.WithFields(map[string]any{ - "node": target.Name, - "dst": fmt.Sprintf("%s/%s", target.Addr, network), - }) + for attempt := 0; attempt < maxRetries; attempt++ { + // Select a target node, excluding previously tried nodes + selectCtx := ctxvalue.ContextWithExcludeNodes(ctx, triedNodes) + var target *chain.Node + if host != "" { + target = &chain.Node{ + Addr: host, + } + } + if h.hop != nil { + target = h.hop.Select(selectCtx, + hop.ProtocolSelectOption(proto), + ) + } + if target == nil { + if lastErr != nil { + return lastErr + } + err := errors.New("node not available") + log.Error(err) + return err + } - log.Debugf("%s >> %s", conn.RemoteAddr(), target.Addr) + // Track this node as tried + triedNodes = append(triedNodes, target.Addr) - var buf bytes.Buffer - cc, err := h.options.Router.Dial(ctxvalue.ContextWithBuffer(ctx, &buf), network, target.Addr) - ro.Route = buf.String() - if err != nil { - log.Error(err) - // TODO: the router itself may be failed due to the failed node in the router, - // the dead marker may be a wrong operation. + if opts := target.Options(); opts != nil { + switch opts.Network { + case "unix": + network = opts.Network + default: + } + } + + ro.Network = network + ro.Host = target.Addr + + targetLog := log.WithFields(map[string]any{ + "node": target.Name, + "dst": fmt.Sprintf("%s/%s", target.Addr, network), + }) + + targetLog.Debugf("%s >> %s", conn.RemoteAddr(), target.Addr) + + var buf bytes.Buffer + cc, err = h.options.Router.Dial(ctxvalue.ContextWithBuffer(ctx, &buf), network, target.Addr) + ro.Route = buf.String() + if err != nil { + targetLog.Error(err) + // Mark node as failed for future selections + if marker := target.Marker(); marker != nil { + marker.Mark() + } + lastErr = err + // Try next node + continue + } + + // Success - reset marker and proceed if marker := target.Marker(); marker != nil { - marker.Mark() + marker.Reset() } - return err - } - defer cc.Close() - if marker := target.Marker(); marker != nil { - marker.Reset() + defer cc.Close() + + cc = proxyproto.WrapClientConn(h.md.proxyProtocol, conn.RemoteAddr(), convertAddr(conn.LocalAddr()), cc) + + t := time.Now() + targetLog.Infof("%s <-> %s", conn.RemoteAddr(), target.Addr) + xnet.Transport(conn, cc) + targetLog.WithFields(map[string]any{ + "duration": time.Since(t), + }).Infof("%s >-< %s", conn.RemoteAddr(), target.Addr) + + return nil } - cc = proxyproto.WrapClientConn(h.md.proxyProtocol, conn.RemoteAddr(), convertAddr(conn.LocalAddr()), cc) - - t := time.Now() - log.Infof("%s <-> %s", conn.RemoteAddr(), target.Addr) - xnet.Transport(conn, cc) - log.WithFields(map[string]any{ - "duration": time.Since(t), - }).Infof("%s >-< %s", conn.RemoteAddr(), target.Addr) - - return nil + // All retries exhausted + if lastErr != nil { + return lastErr + } + return errors.New("all nodes failed") } func (h *forwardHandler) checkRateLimit(addr net.Addr) bool { diff --git a/go-gost/x/handler/forward/remote/metadata.go b/go-gost/x/handler/forward/remote/metadata.go index 6b1e541..67c8902 100644 --- a/go-gost/x/handler/forward/remote/metadata.go +++ b/go-gost/x/handler/forward/remote/metadata.go @@ -26,6 +26,12 @@ type metadata struct { privateKey crypto.PrivateKey alpn string mitmBypass bypass.Bypass + + // maxRetries specifies the maximum number of failover retry attempts. + // When a target node fails, the handler will try the next available node. + // 0 means use the total number of available nodes (try all nodes once). + // Default: 0 (try all available nodes) + maxRetries int } func (h *forwardHandler) parseMetadata(md mdata.Metadata) (err error) { @@ -57,5 +63,9 @@ func (h *forwardHandler) parseMetadata(md mdata.Metadata) (err error) { } h.md.alpn = mdutil.GetString(md, "mitm.alpn") h.md.mitmBypass = registry.BypassRegistry().Get(mdutil.GetString(md, "mitm.bypass")) + + // maxRetries: 0 means try all available nodes (default behavior) + h.md.maxRetries = mdutil.GetInt(md, "maxRetries", "retry.max") + return } diff --git a/go-gost/x/hop/hop.go b/go-gost/x/hop/hop.go index a7cc481..60076eb 100644 --- a/go-gost/x/hop/hop.go +++ b/go-gost/x/hop/hop.go @@ -18,6 +18,7 @@ import ( "github.com/go-gost/core/selector" "github.com/go-gost/x/config" node_parser "github.com/go-gost/x/config/parsing/node" + ctxvalue "github.com/go-gost/x/ctx" "github.com/go-gost/x/internal/loader" ) @@ -141,11 +142,25 @@ func (p *chainHop) Select(ctx context.Context, opts ...hop.SelectOption) *chain. return nil } + // Get list of nodes to exclude (for failover retry) + excludeNodes := ctxvalue.ExcludeNodesFromContext(ctx) + excludeSet := make(map[string]bool) + for _, addr := range excludeNodes { + excludeSet[addr] = true + } + var nodes []*chain.Node for _, node := range p.Nodes() { if node == nil { continue } + + // Skip nodes in the exclude list (failover retry) + if excludeSet[node.Addr] || excludeSet[node.Name] { + log.Debugf("node %s(%s) excluded for failover retry", node.Name, node.Addr) + continue + } + // node level bypass if node.Options().Bypass != nil && node.Options().Bypass.Contains(ctx, options.Network, options.Addr, bypass.WithHostOpton(options.Host)) { @@ -177,9 +192,6 @@ func (p *chainHop) Select(ctx context.Context, opts ...hop.SelectOption) *chain. if len(nodes) == 0 { return nil } - if len(nodes) == 1 { - return nodes[0] - } sort.Slice(nodes, func(i, j int) bool { return nodes[i].Options().Priority > nodes[j].Options().Priority @@ -189,9 +201,14 @@ func (p *chainHop) Select(ctx context.Context, opts ...hop.SelectOption) *chain. return nodes[0] } + // Always go through selector for proper FailFilter evaluation, + // even when there's only one node. This ensures failed nodes + // can be filtered out properly. if s := p.options.selector; s != nil { return s.Select(ctx, nodes...) } + + // Fallback: return first node if no selector configured return nodes[0] } diff --git a/go-gost/x/internal/util/forwarder/sniffer.go b/go-gost/x/internal/util/forwarder/sniffer.go index bbc0f72..4d9c153 100644 --- a/go-gost/x/internal/util/forwarder/sniffer.go +++ b/go-gost/x/internal/util/forwarder/sniffer.go @@ -247,64 +247,100 @@ func (h *Sniffer) dial(ctx context.Context, conn net.Conn, req *http.Request, ho } } - node = &chain.Node{ - Addr: host, + // Determine max retry attempts + maxRetries := 1 + if nl, ok := ho.Hop.(hop.NodeList); ok { + maxRetries = len(nl.Nodes()) } - if ho.Hop != nil { - node = ho.Hop.Select(ctx, - hop.ClientIPSelectOption(net.ParseIP(ro.ClientIP)), - hop.ProtocolSelectOption(sniffing.ProtoHTTP), - hop.HostSelectOption(host), - hop.MethodSelectOption(req.Method), - hop.PathSelectOption(req.URL.Path), - hop.QuerySelectOption(req.URL.Query()), - hop.HeaderSelectOption(req.Header), - ) - } - if node == nil { - ho.Log.Warnf("node for %s not found", host) - res.StatusCode = http.StatusBadGateway - ro.HTTP.StatusCode = res.StatusCode - res.Write(conn) - return nil, nil, errors.New("node not available") + if maxRetries <= 0 { + maxRetries = 1 } - ro.Host = node.Addr - ho.Log = ho.Log.WithFields(map[string]any{ - "node": node.Name, - "dst": node.Addr, - }) - ho.Log.Debugf("find node for host %s -> %s(%s)", host, node.Name, node.Addr) + var triedNodes []string + var lastErr error - cc, err = dial(ctx, "tcp", node.Addr) - if err != nil { - // TODO: the router itself may be failed due to the failed node in the router, - // the dead marker may be a wrong operation. - if marker := node.Marker(); marker != nil { - marker.Mark() + for attempt := 0; attempt < maxRetries; attempt++ { + // Select a node, excluding previously tried nodes + selectCtx := ctxvalue.ContextWithExcludeNodes(ctx, triedNodes) + + node = &chain.Node{ + Addr: host, } - ho.Log.Warnf("connect to node %s(%s) failed: %v", node.Name, node.Addr, err) - res.Write(conn) - return - } - if marker := node.Marker(); marker != nil { - marker.Reset() - } - - if tlsSettings := node.Options().TLS; tlsSettings != nil { - cfg := &tls.Config{ - ServerName: tlsSettings.ServerName, - InsecureSkipVerify: !tlsSettings.Secure, + if ho.Hop != nil { + node = ho.Hop.Select(selectCtx, + hop.ClientIPSelectOption(net.ParseIP(ro.ClientIP)), + hop.ProtocolSelectOption(sniffing.ProtoHTTP), + hop.HostSelectOption(host), + hop.MethodSelectOption(req.Method), + hop.PathSelectOption(req.URL.Path), + hop.QuerySelectOption(req.URL.Query()), + hop.HeaderSelectOption(req.Header), + ) } - tls_util.SetTLSOptions(cfg, &config.TLSOptions{ - MinVersion: tlsSettings.Options.MinVersion, - MaxVersion: tlsSettings.Options.MaxVersion, - CipherSuites: tlsSettings.Options.CipherSuites, - ALPN: tlsSettings.Options.ALPN, + if node == nil { + if lastErr != nil { + ho.Log.Warnf("node for %s not found after retries", host) + res.StatusCode = http.StatusBadGateway + ro.HTTP.StatusCode = res.StatusCode + res.Write(conn) + return nil, nil, lastErr + } + ho.Log.Warnf("node for %s not found", host) + res.StatusCode = http.StatusBadGateway + ro.HTTP.StatusCode = res.StatusCode + res.Write(conn) + return nil, nil, errors.New("node not available") + } + + // Track this node as tried + triedNodes = append(triedNodes, node.Addr) + + ro.Host = node.Addr + ho.Log = ho.Log.WithFields(map[string]any{ + "node": node.Name, + "dst": node.Addr, }) - cc = tls.Client(cc, cfg) + ho.Log.Debugf("find node for host %s -> %s(%s)", host, node.Name, node.Addr) + + cc, err = dial(ctx, "tcp", node.Addr) + if err != nil { + // Mark node as failed for future selections + if marker := node.Marker(); marker != nil { + marker.Mark() + } + ho.Log.Warnf("connect to node %s(%s) failed: %v, trying next node", node.Name, node.Addr, err) + lastErr = err + continue + } + + // Success - reset marker + if marker := node.Marker(); marker != nil { + marker.Reset() + } + + if tlsSettings := node.Options().TLS; tlsSettings != nil { + cfg := &tls.Config{ + ServerName: tlsSettings.ServerName, + InsecureSkipVerify: !tlsSettings.Secure, + } + tls_util.SetTLSOptions(cfg, &config.TLSOptions{ + MinVersion: tlsSettings.Options.MinVersion, + MaxVersion: tlsSettings.Options.MaxVersion, + CipherSuites: tlsSettings.Options.CipherSuites, + ALPN: tlsSettings.Options.ALPN, + }) + cc = tls.Client(cc, cfg) + } + return node, cc, nil } - return + + // All retries exhausted + ho.Log.Warnf("all nodes failed for host %s", host) + res.Write(conn) + if lastErr != nil { + return nil, nil, lastErr + } + return nil, nil, errors.New("all nodes failed") } func (h *Sniffer) serveH2(ctx context.Context, conn net.Conn, ho *HandleOptions) error { @@ -847,74 +883,107 @@ func (h *Sniffer) dialTLS(ctx context.Context, host string, ho *HandleOptions) ( return } - if host != "" { - node = &chain.Node{ - Addr: host, - } - } - ro := ho.RecorderObject - if ho.Hop != nil { - node = ho.Hop.Select(ctx, - hop.ClientIPSelectOption(net.ParseIP(ro.ClientIP)), - hop.HostSelectOption(host), - hop.ProtocolSelectOption(sniffing.ProtoTLS), - ) + + // Determine max retry attempts + maxRetries := 1 + if nl, ok := ho.Hop.(hop.NodeList); ok { + maxRetries = len(nl.Nodes()) } - if node == nil { - err = errors.New("node not available") - return + if maxRetries <= 0 { + maxRetries = 1 } - addr := node.Addr - if opts := node.Options(); opts != nil { - switch opts.Network { - case "unix": - ro.Network = opts.Network - default: - if _, _, err := net.SplitHostPort(addr); err != nil { - addr += ":443" + var triedNodes []string + var lastErr error + + for attempt := 0; attempt < maxRetries; attempt++ { + // Select a node, excluding previously tried nodes + selectCtx := ctxvalue.ContextWithExcludeNodes(ctx, triedNodes) + + node = nil + if host != "" { + node = &chain.Node{ + Addr: host, } } - } - ro.Host = addr - - ho.Log = ho.Log.WithFields(map[string]any{ - "host": host, - "node": node.Name, - "dst": fmt.Sprintf("%s/%s", addr, ro.Network), - }) - ho.Log.Debugf("find node for host %s -> %s(%s)", host, node.Name, addr) - - cc, err = dial(ctx, ro.Network, addr) - if err != nil { - // TODO: the router itself may be failed due to the failed node in the router, - // the dead marker may be a wrong operation. - if marker := node.Marker(); marker != nil { - marker.Mark() + if ho.Hop != nil { + node = ho.Hop.Select(selectCtx, + hop.ClientIPSelectOption(net.ParseIP(ro.ClientIP)), + hop.HostSelectOption(host), + hop.ProtocolSelectOption(sniffing.ProtoTLS), + ) } - ho.Log.Warnf("connect to node %s(%s) failed: %v", node.Name, node.Addr, err) - return - } - - if marker := node.Marker(); marker != nil { - marker.Reset() - } - - if tlsSettings := node.Options().TLS; tlsSettings != nil { - cfg := &tls.Config{ - ServerName: tlsSettings.ServerName, - InsecureSkipVerify: !tlsSettings.Secure, + if node == nil { + if lastErr != nil { + ho.Log.Warnf("node for %s not found after retries", host) + return nil, nil, lastErr + } + ho.Log.Warnf("node for %s not found", host) + return nil, nil, errors.New("node not available") } - tls_util.SetTLSOptions(cfg, &config.TLSOptions{ - MinVersion: tlsSettings.Options.MinVersion, - MaxVersion: tlsSettings.Options.MaxVersion, - CipherSuites: tlsSettings.Options.CipherSuites, - ALPN: tlsSettings.Options.ALPN, + + // Track this node as tried + triedNodes = append(triedNodes, node.Addr) + + addr := node.Addr + if opts := node.Options(); opts != nil { + switch opts.Network { + case "unix": + ro.Network = opts.Network + default: + if _, _, err := net.SplitHostPort(addr); err != nil { + addr += ":443" + } + } + } + ro.Host = addr + + ho.Log = ho.Log.WithFields(map[string]any{ + "host": host, + "node": node.Name, + "dst": fmt.Sprintf("%s/%s", addr, ro.Network), }) - cc = tls.Client(cc, cfg) + ho.Log.Debugf("find node for host %s -> %s(%s)", host, node.Name, addr) + + cc, err = dial(ctx, ro.Network, addr) + if err != nil { + // Mark node as failed for future selections + if marker := node.Marker(); marker != nil { + marker.Mark() + } + ho.Log.Warnf("connect to node %s(%s) failed: %v, trying next node", node.Name, node.Addr, err) + lastErr = err + continue + } + + // Success - reset marker + if marker := node.Marker(); marker != nil { + marker.Reset() + } + + if tlsSettings := node.Options().TLS; tlsSettings != nil { + cfg := &tls.Config{ + ServerName: tlsSettings.ServerName, + InsecureSkipVerify: !tlsSettings.Secure, + } + tls_util.SetTLSOptions(cfg, &config.TLSOptions{ + MinVersion: tlsSettings.Options.MinVersion, + MaxVersion: tlsSettings.Options.MaxVersion, + CipherSuites: tlsSettings.Options.CipherSuites, + ALPN: tlsSettings.Options.ALPN, + }) + cc = tls.Client(cc, cfg) + } + return node, cc, nil } - return + + // All retries exhausted + ho.Log.Warnf("all nodes failed for host %s", host) + if lastErr != nil { + return nil, nil, lastErr + } + return nil, nil, errors.New("all nodes failed") } func (h *Sniffer) terminateTLS(ctx context.Context, conn, cc net.Conn, clientHello *dissector.ClientHelloInfo, ho *HandleOptions) error { diff --git a/go-gost/x/selector/filter.go b/go-gost/x/selector/filter.go index 465ba3f..69befe3 100644 --- a/go-gost/x/selector/filter.go +++ b/go-gost/x/selector/filter.go @@ -5,8 +5,8 @@ import ( "time" "github.com/go-gost/core/metadata" - mdutil "github.com/go-gost/x/metadata/util" "github.com/go-gost/core/selector" + mdutil "github.com/go-gost/x/metadata/util" ) type failFilter[T any] struct { @@ -24,10 +24,11 @@ func FailFilter[T any](maxFails int, timeout time.Duration) selector.Filter[T] { } // Filter filters dead objects. +// Note: We intentionally do NOT skip filtering when len(vs) <= 1. +// This ensures that even a single dead node gets filtered out, +// allowing the caller to know that no healthy nodes are available +// and potentially trigger failover behavior. func (f *failFilter[T]) Filter(ctx context.Context, vs ...T) []T { - if len(vs) <= 1 { - return vs - } var l []T for _, v := range vs { maxFails := f.maxFails