Merge pull request #14 from Sagit-chu/opencode/cosmic-pixel

fix(gost): implement failover for multi-node forwarding rules
This commit is contained in:
sagit
2026-02-04 12:30:49 +08:00
committed by GitHub
8 changed files with 396 additions and 202 deletions
+20
View File
@@ -109,3 +109,23 @@ func LoggerFromContext(ctx context.Context) logger.Logger {
v, _ := ctx.Value(keyLogger).(logger.Logger) v, _ := ctx.Value(keyLogger).(logger.Logger)
return v 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
}
+41 -8
View File
@@ -176,16 +176,40 @@ func (h *forwardHandler) Handle(ctx context.Context, conn net.Conn, opts ...hand
} }
} }
// 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
}
}
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{} target := &chain.Node{}
if h.hop != nil { if h.hop != nil {
target = h.hop.Select(ctx, target = h.hop.Select(selectCtx,
hop.ProtocolSelectOption(proto), hop.ProtocolSelectOption(proto),
) )
} }
if target == nil { if target == nil {
err := errors.New("node not available") if lastErr != nil {
return err return lastErr
} }
return errors.New("node not available")
}
// Track this node as tried
triedNodes = append(triedNodes, target.Addr)
addr := target.Addr addr := target.Addr
if opts := target.Options(); opts != nil { if opts := target.Options(); opts != nil {
@@ -203,24 +227,33 @@ func (h *forwardHandler) Handle(ctx context.Context, conn net.Conn, opts ...hand
ro.Host = addr ro.Host = addr
var buf bytes.Buffer var buf bytes.Buffer
cc, err := h.options.Router.Dial(ctxvalue.ContextWithBuffer(ctx, &buf), network, addr) cc, err = h.options.Router.Dial(ctxvalue.ContextWithBuffer(ctx, &buf), network, addr)
ro.Route = buf.String() ro.Route = buf.String()
if err != nil { if err != nil {
// TODO: the router itself may be failed due to the failed node in the router, // Mark node as failed for future selections
// the dead marker may be a wrong operation.
if marker := target.Marker(); marker != nil { if marker := target.Marker(); marker != nil {
marker.Mark() marker.Mark()
} }
return err lastErr = err
// Try next node
continue
} }
// Success - reset marker and proceed
if marker := target.Marker(); marker != nil { if marker := target.Marker(); marker != nil {
marker.Reset() marker.Reset()
} }
defer cc.Close() defer cc.Close()
xnet.Transport(conn, cc) xnet.Transport(conn, cc)
return nil return nil
}
// All retries exhausted
if lastErr != nil {
return lastErr
}
return errors.New("all nodes failed")
} }
func (h *forwardHandler) checkRateLimit(addr net.Addr) bool { func (h *forwardHandler) checkRateLimit(addr net.Addr) bool {
@@ -25,6 +25,12 @@ type metadata struct {
privateKey crypto.PrivateKey privateKey crypto.PrivateKey
alpn string alpn string
mitmBypass bypass.Bypass 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) { 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.alpn = mdutil.GetString(md, "mitm.alpn")
h.md.mitmBypass = registry.BypassRegistry().Get(mdutil.GetString(md, "mitm.bypass")) 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 return
} }
+46 -11
View File
@@ -204,6 +204,25 @@ func (h *forwardHandler) Handle(ctx context.Context, conn net.Conn, opts ...hand
} }
} }
// 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
}
}
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)
var target *chain.Node var target *chain.Node
if host != "" { if host != "" {
target = &chain.Node{ target = &chain.Node{
@@ -211,16 +230,22 @@ func (h *forwardHandler) Handle(ctx context.Context, conn net.Conn, opts ...hand
} }
} }
if h.hop != nil { if h.hop != nil {
target = h.hop.Select(ctx, target = h.hop.Select(selectCtx,
hop.ProtocolSelectOption(proto), hop.ProtocolSelectOption(proto),
) )
} }
if target == nil { if target == nil {
if lastErr != nil {
return lastErr
}
err := errors.New("node not available") err := errors.New("node not available")
log.Error(err) log.Error(err)
return err return err
} }
// Track this node as tried
triedNodes = append(triedNodes, target.Addr)
if opts := target.Options(); opts != nil { if opts := target.Options(); opts != nil {
switch opts.Network { switch opts.Network {
case "unix": case "unix":
@@ -232,40 +257,50 @@ func (h *forwardHandler) Handle(ctx context.Context, conn net.Conn, opts ...hand
ro.Network = network ro.Network = network
ro.Host = target.Addr ro.Host = target.Addr
log = log.WithFields(map[string]any{ targetLog := log.WithFields(map[string]any{
"node": target.Name, "node": target.Name,
"dst": fmt.Sprintf("%s/%s", target.Addr, network), "dst": fmt.Sprintf("%s/%s", target.Addr, network),
}) })
log.Debugf("%s >> %s", conn.RemoteAddr(), target.Addr) targetLog.Debugf("%s >> %s", conn.RemoteAddr(), target.Addr)
var buf bytes.Buffer var buf bytes.Buffer
cc, err := h.options.Router.Dial(ctxvalue.ContextWithBuffer(ctx, &buf), network, target.Addr) cc, err = h.options.Router.Dial(ctxvalue.ContextWithBuffer(ctx, &buf), network, target.Addr)
ro.Route = buf.String() ro.Route = buf.String()
if err != nil { if err != nil {
log.Error(err) targetLog.Error(err)
// TODO: the router itself may be failed due to the failed node in the router, // Mark node as failed for future selections
// the dead marker may be a wrong operation.
if marker := target.Marker(); marker != nil { if marker := target.Marker(); marker != nil {
marker.Mark() marker.Mark()
} }
return err lastErr = err
// Try next node
continue
} }
defer cc.Close()
// Success - reset marker and proceed
if marker := target.Marker(); marker != nil { if marker := target.Marker(); marker != nil {
marker.Reset() marker.Reset()
} }
defer cc.Close()
cc = proxyproto.WrapClientConn(h.md.proxyProtocol, conn.RemoteAddr(), convertAddr(conn.LocalAddr()), cc) cc = proxyproto.WrapClientConn(h.md.proxyProtocol, conn.RemoteAddr(), convertAddr(conn.LocalAddr()), cc)
t := time.Now() t := time.Now()
log.Infof("%s <-> %s", conn.RemoteAddr(), target.Addr) targetLog.Infof("%s <-> %s", conn.RemoteAddr(), target.Addr)
xnet.Transport(conn, cc) xnet.Transport(conn, cc)
log.WithFields(map[string]any{ targetLog.WithFields(map[string]any{
"duration": time.Since(t), "duration": time.Since(t),
}).Infof("%s >-< %s", conn.RemoteAddr(), target.Addr) }).Infof("%s >-< %s", conn.RemoteAddr(), target.Addr)
return nil return nil
}
// All retries exhausted
if lastErr != nil {
return lastErr
}
return errors.New("all nodes failed")
} }
func (h *forwardHandler) checkRateLimit(addr net.Addr) bool { func (h *forwardHandler) checkRateLimit(addr net.Addr) bool {
@@ -26,6 +26,12 @@ type metadata struct {
privateKey crypto.PrivateKey privateKey crypto.PrivateKey
alpn string alpn string
mitmBypass bypass.Bypass 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) { 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.alpn = mdutil.GetString(md, "mitm.alpn")
h.md.mitmBypass = registry.BypassRegistry().Get(mdutil.GetString(md, "mitm.bypass")) 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 return
} }
+20 -3
View File
@@ -18,6 +18,7 @@ import (
"github.com/go-gost/core/selector" "github.com/go-gost/core/selector"
"github.com/go-gost/x/config" "github.com/go-gost/x/config"
node_parser "github.com/go-gost/x/config/parsing/node" node_parser "github.com/go-gost/x/config/parsing/node"
ctxvalue "github.com/go-gost/x/ctx"
"github.com/go-gost/x/internal/loader" "github.com/go-gost/x/internal/loader"
) )
@@ -141,11 +142,25 @@ func (p *chainHop) Select(ctx context.Context, opts ...hop.SelectOption) *chain.
return nil 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 var nodes []*chain.Node
for _, node := range p.Nodes() { for _, node := range p.Nodes() {
if node == nil { if node == nil {
continue 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 // node level bypass
if node.Options().Bypass != nil && if node.Options().Bypass != nil &&
node.Options().Bypass.Contains(ctx, options.Network, options.Addr, bypass.WithHostOpton(options.Host)) { 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 { if len(nodes) == 0 {
return nil return nil
} }
if len(nodes) == 1 {
return nodes[0]
}
sort.Slice(nodes, func(i, j int) bool { sort.Slice(nodes, func(i, j int) bool {
return nodes[i].Options().Priority > nodes[j].Options().Priority 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] 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 { if s := p.options.selector; s != nil {
return s.Select(ctx, nodes...) return s.Select(ctx, nodes...)
} }
// Fallback: return first node if no selector configured
return nodes[0] return nodes[0]
} }
+86 -17
View File
@@ -247,11 +247,27 @@ func (h *Sniffer) dial(ctx context.Context, conn net.Conn, req *http.Request, ho
} }
} }
// Determine max retry attempts
maxRetries := 1
if nl, ok := ho.Hop.(hop.NodeList); ok {
maxRetries = len(nl.Nodes())
}
if maxRetries <= 0 {
maxRetries = 1
}
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 = &chain.Node{ node = &chain.Node{
Addr: host, Addr: host,
} }
if ho.Hop != nil { if ho.Hop != nil {
node = ho.Hop.Select(ctx, node = ho.Hop.Select(selectCtx,
hop.ClientIPSelectOption(net.ParseIP(ro.ClientIP)), hop.ClientIPSelectOption(net.ParseIP(ro.ClientIP)),
hop.ProtocolSelectOption(sniffing.ProtoHTTP), hop.ProtocolSelectOption(sniffing.ProtoHTTP),
hop.HostSelectOption(host), hop.HostSelectOption(host),
@@ -262,6 +278,13 @@ func (h *Sniffer) dial(ctx context.Context, conn net.Conn, req *http.Request, ho
) )
} }
if node == nil { 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) ho.Log.Warnf("node for %s not found", host)
res.StatusCode = http.StatusBadGateway res.StatusCode = http.StatusBadGateway
ro.HTTP.StatusCode = res.StatusCode ro.HTTP.StatusCode = res.StatusCode
@@ -269,6 +292,9 @@ func (h *Sniffer) dial(ctx context.Context, conn net.Conn, req *http.Request, ho
return nil, nil, errors.New("node not available") return nil, nil, errors.New("node not available")
} }
// Track this node as tried
triedNodes = append(triedNodes, node.Addr)
ro.Host = node.Addr ro.Host = node.Addr
ho.Log = ho.Log.WithFields(map[string]any{ ho.Log = ho.Log.WithFields(map[string]any{
"node": node.Name, "node": node.Name,
@@ -278,15 +304,16 @@ func (h *Sniffer) dial(ctx context.Context, conn net.Conn, req *http.Request, ho
cc, err = dial(ctx, "tcp", node.Addr) cc, err = dial(ctx, "tcp", node.Addr)
if err != nil { if err != nil {
// TODO: the router itself may be failed due to the failed node in the router, // Mark node as failed for future selections
// the dead marker may be a wrong operation.
if marker := node.Marker(); marker != nil { if marker := node.Marker(); marker != nil {
marker.Mark() marker.Mark()
} }
ho.Log.Warnf("connect to node %s(%s) failed: %v", node.Name, node.Addr, err) ho.Log.Warnf("connect to node %s(%s) failed: %v, trying next node", node.Name, node.Addr, err)
res.Write(conn) lastErr = err
return continue
} }
// Success - reset marker
if marker := node.Marker(); marker != nil { if marker := node.Marker(); marker != nil {
marker.Reset() marker.Reset()
} }
@@ -304,7 +331,16 @@ func (h *Sniffer) dial(ctx context.Context, conn net.Conn, req *http.Request, ho
}) })
cc = tls.Client(cc, cfg) cc = tls.Client(cc, cfg)
} }
return return node, cc, nil
}
// 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 { func (h *Sniffer) serveH2(ctx context.Context, conn net.Conn, ho *HandleOptions) error {
@@ -847,24 +883,48 @@ func (h *Sniffer) dialTLS(ctx context.Context, host string, ho *HandleOptions) (
return return
} }
ro := ho.RecorderObject
// Determine max retry attempts
maxRetries := 1
if nl, ok := ho.Hop.(hop.NodeList); ok {
maxRetries = len(nl.Nodes())
}
if maxRetries <= 0 {
maxRetries = 1
}
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 != "" { if host != "" {
node = &chain.Node{ node = &chain.Node{
Addr: host, Addr: host,
} }
} }
ro := ho.RecorderObject
if ho.Hop != nil { if ho.Hop != nil {
node = ho.Hop.Select(ctx, node = ho.Hop.Select(selectCtx,
hop.ClientIPSelectOption(net.ParseIP(ro.ClientIP)), hop.ClientIPSelectOption(net.ParseIP(ro.ClientIP)),
hop.HostSelectOption(host), hop.HostSelectOption(host),
hop.ProtocolSelectOption(sniffing.ProtoTLS), hop.ProtocolSelectOption(sniffing.ProtoTLS),
) )
} }
if node == nil { if node == nil {
err = errors.New("node not available") if lastErr != nil {
return 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")
}
// Track this node as tried
triedNodes = append(triedNodes, node.Addr)
addr := node.Addr addr := node.Addr
if opts := node.Options(); opts != nil { if opts := node.Options(); opts != nil {
@@ -888,15 +948,16 @@ func (h *Sniffer) dialTLS(ctx context.Context, host string, ho *HandleOptions) (
cc, err = dial(ctx, ro.Network, addr) cc, err = dial(ctx, ro.Network, addr)
if err != nil { if err != nil {
// TODO: the router itself may be failed due to the failed node in the router, // Mark node as failed for future selections
// the dead marker may be a wrong operation.
if marker := node.Marker(); marker != nil { if marker := node.Marker(); marker != nil {
marker.Mark() marker.Mark()
} }
ho.Log.Warnf("connect to node %s(%s) failed: %v", node.Name, node.Addr, err) ho.Log.Warnf("connect to node %s(%s) failed: %v, trying next node", node.Name, node.Addr, err)
return lastErr = err
continue
} }
// Success - reset marker
if marker := node.Marker(); marker != nil { if marker := node.Marker(); marker != nil {
marker.Reset() marker.Reset()
} }
@@ -914,7 +975,15 @@ func (h *Sniffer) dialTLS(ctx context.Context, host string, ho *HandleOptions) (
}) })
cc = tls.Client(cc, cfg) cc = tls.Client(cc, cfg)
} }
return return node, cc, nil
}
// 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 { func (h *Sniffer) terminateTLS(ctx context.Context, conn, cc net.Conn, clientHello *dissector.ClientHelloInfo, ho *HandleOptions) error {
+5 -4
View File
@@ -5,8 +5,8 @@ import (
"time" "time"
"github.com/go-gost/core/metadata" "github.com/go-gost/core/metadata"
mdutil "github.com/go-gost/x/metadata/util"
"github.com/go-gost/core/selector" "github.com/go-gost/core/selector"
mdutil "github.com/go-gost/x/metadata/util"
) )
type failFilter[T any] struct { 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. // 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 { func (f *failFilter[T]) Filter(ctx context.Context, vs ...T) []T {
if len(vs) <= 1 {
return vs
}
var l []T var l []T
for _, v := range vs { for _, v := range vs {
maxFails := f.maxFails maxFails := f.maxFails