fix(agent): retire replaced tunnel sessions (#541)

This commit is contained in:
sagit
2026-08-07 14:35:30 +08:00
committed by GitHub
parent 9e6f80019d
commit 0b23d6f7d7
22 changed files with 701 additions and 19 deletions
+41 -1
View File
@@ -2,6 +2,8 @@ package chain
import ( import (
"context" "context"
"errors"
"io"
"github.com/go-gost/core/chain" "github.com/go-gost/core/chain"
"github.com/go-gost/core/hop" "github.com/go-gost/core/hop"
@@ -40,6 +42,7 @@ type chainNamer interface {
type Chain struct { type Chain struct {
name string name string
hops []hop.Hop hops []hop.Hop
ownedHops []hop.Hop
marker selector.Marker marker selector.Marker
metadata metadata.Metadata metadata metadata.Metadata
logger logger.Logger logger logger.Logger
@@ -61,8 +64,15 @@ func NewChain(name string, opts ...ChainOption) *Chain {
} }
} }
func (c *Chain) AddHop(hop hop.Hop) { func (c *Chain) AddHop(hop hop.Hop, owned ...bool) {
c.hops = append(c.hops, hop) c.hops = append(c.hops, hop)
isOwned := true
if len(owned) > 0 {
isOwned = owned[0]
}
if isOwned {
c.ownedHops = append(c.ownedHops, hop)
}
} }
// Metadata implements metadata.Metadatable interface. // Metadata implements metadata.Metadatable interface.
@@ -112,6 +122,36 @@ func (c *Chain) Route(ctx context.Context, network, address string, opts ...chai
return rt return rt
} }
// Retire gracefully drains resources owned by a chain that has been replaced.
func (c *Chain) Retire() {
if c == nil {
return
}
for _, h := range c.ownedHops {
if retirer, ok := h.(interface{ Retire() }); ok {
retirer.Retire()
continue
}
if closer, ok := h.(io.Closer); ok {
_ = closer.Close()
}
}
}
// Close immediately releases all resources owned by the chain.
func (c *Chain) Close() error {
if c == nil {
return nil
}
var errs []error
for _, h := range c.ownedHops {
if closer, ok := h.(io.Closer); ok {
errs = append(errs, closer.Close())
}
}
return errors.Join(errs...)
}
type chainGroup struct { type chainGroup struct {
chains []chain.Chainer chains []chain.Chainer
selector selector.Selector[chain.Chainer] selector selector.Selector[chain.Chainer]
+64
View File
@@ -0,0 +1,64 @@
package chain
import (
"context"
"testing"
corechain "github.com/go-gost/core/chain"
corehop "github.com/go-gost/core/hop"
)
type lifecycleTestHop struct {
selected int
retired int
closed int
}
func (h *lifecycleTestHop) Select(context.Context, ...corehop.SelectOption) *corechain.Node {
h.selected++
return nil
}
func (h *lifecycleTestHop) Retire() {
h.retired++
}
func (h *lifecycleTestHop) Close() error {
h.closed++
return nil
}
func TestChainRoutesThroughSharedHopWithoutOwningLifecycle(t *testing.T) {
hop := &lifecycleTestHop{}
chain := NewChain("shared-hop")
chain.AddHop(hop, false)
if route := chain.Route(context.Background(), "tcp", "example.com:443"); route == nil {
t.Fatal("route is nil")
}
if hop.selected != 1 {
t.Fatalf("shared hop selected %d times, want 1", hop.selected)
}
chain.Retire()
if err := chain.Close(); err != nil {
t.Fatalf("close chain: %v", err)
}
if hop.retired != 0 || hop.closed != 0 {
t.Fatalf("shared hop lifecycle changed: retired=%d closed=%d", hop.retired, hop.closed)
}
}
func TestChainRetiresAndClosesOwnedHop(t *testing.T) {
hop := &lifecycleTestHop{}
chain := NewChain("owned-hop")
chain.AddHop(hop)
chain.Retire()
if err := chain.Close(); err != nil {
t.Fatalf("close chain: %v", err)
}
if hop.retired != 1 || hop.closed != 1 {
t.Fatalf("owned hop lifecycle: retired=%d closed=%d, want 1/1", hop.retired, hop.closed)
}
}
+32 -1
View File
@@ -2,6 +2,8 @@ package chain
import ( import (
"context" "context"
"errors"
"io"
"net" "net"
"github.com/go-gost/core/chain" "github.com/go-gost/core/chain"
@@ -102,5 +104,34 @@ func (tr *Transport) Options() *chain.TransportOptions {
func (tr *Transport) Copy() chain.Transporter { func (tr *Transport) Copy() chain.Transporter {
tr2 := &Transport{} tr2 := &Transport{}
*tr2 = *tr *tr2 = *tr
return tr return tr2
}
// Retire prevents long-lived dialer sessions owned by an obsolete chain from
// accepting new streams while allowing existing streams to drain.
func (tr *Transport) Retire() {
if tr == nil {
return
}
if retirer, ok := tr.dialer.(interface{ Retire() }); ok {
retirer.Retire()
}
if retirer, ok := tr.connector.(interface{ Retire() }); ok {
retirer.Retire()
}
}
// Close immediately releases transport-owned dialer and connector resources.
func (tr *Transport) Close() error {
if tr == nil {
return nil
}
var errs []error
if closer, ok := tr.dialer.(io.Closer); ok {
errs = append(errs, closer.Close())
}
if closer, ok := tr.connector.(io.Closer); ok {
errs = append(errs, closer.Close())
}
return errors.Join(errs...)
} }
+45
View File
@@ -0,0 +1,45 @@
package chain
import (
"context"
"net"
"testing"
corechain "github.com/go-gost/core/chain"
)
type copyTestRoute struct{}
func (copyTestRoute) Dial(context.Context, string, string, ...corechain.DialOption) (net.Conn, error) {
return nil, nil
}
func (copyTestRoute) Bind(context.Context, string, string, ...corechain.BindOption) (net.Listener, error) {
return nil, nil
}
func (copyTestRoute) Nodes() []*corechain.Node {
return nil
}
func TestTransportCopyReturnsIndependentTransport(t *testing.T) {
originalRoute := copyTestRoute{}
replacementRoute := &copyTestRoute{}
original := NewTransport(nil, nil, corechain.RouteTransportOption(originalRoute))
copied, ok := original.Copy().(*Transport)
if !ok {
t.Fatalf("copy type = %T, want *Transport", original.Copy())
}
if copied == original {
t.Fatal("Copy returned the original transport")
}
copied.Options().Route = replacementRoute
if original.Options().Route != originalRoute {
t.Fatal("mutating copied transport changed original route")
}
if copied.Options().Route != replacementRoute {
t.Fatal("copied transport did not retain its independent route")
}
}
+3 -1
View File
@@ -35,16 +35,18 @@ func ParseChain(cfg *config.ChainConfig, log logger.Logger) (chain.Chainer, erro
for _, ch := range cfg.Hops { for _, ch := range cfg.Hops {
var hop hop.Hop var hop hop.Hop
var err error var err error
owned := false
if ch.Nodes != nil || ch.Plugin != nil { if ch.Nodes != nil || ch.Plugin != nil {
if hop, err = hop_parser.ParseHop(ch, log); err != nil { if hop, err = hop_parser.ParseHop(ch, log); err != nil {
return nil, err return nil, err
} }
owned = true
} else { } else {
hop = registry.HopRegistry().Get(ch.Name) hop = registry.HopRegistry().Get(ch.Name)
} }
if hop != nil { if hop != nil {
c.AddHop(hop) c.AddHop(hop, owned)
} }
} }
+5 -2
View File
@@ -19,20 +19,23 @@ func (session *muxSession) Accept() (net.Conn, error) {
} }
func (session *muxSession) Close() error { func (session *muxSession) Close() error {
if session.session == nil { if session == nil || session.session == nil {
return nil return nil
} }
return session.session.Close() return session.session.Close()
} }
func (session *muxSession) IsClosed() bool { func (session *muxSession) IsClosed() bool {
if session.session == nil { if session == nil || session.session == nil {
return true return true
} }
return session.session.IsClosed() return session.session.IsClosed()
} }
func (session *muxSession) NumStreams() int { func (session *muxSession) NumStreams() int {
if session == nil || session.session == nil {
return 0
}
return session.session.NumStreams() return session.session.NumStreams()
} }
+38
View File
@@ -11,6 +11,7 @@ import (
"github.com/go-gost/core/logger" "github.com/go-gost/core/logger"
md "github.com/go-gost/core/metadata" md "github.com/go-gost/core/metadata"
kcp_util "github.com/go-gost/x/internal/util/kcp" kcp_util "github.com/go-gost/x/internal/util/kcp"
"github.com/go-gost/x/internal/util/sessionretire"
mdutil "github.com/go-gost/x/metadata/util" mdutil "github.com/go-gost/x/metadata/util"
"github.com/go-gost/x/registry" "github.com/go-gost/x/registry"
"github.com/xtaci/kcp-go/v5" "github.com/xtaci/kcp-go/v5"
@@ -25,6 +26,7 @@ func init() {
type kcpDialer struct { type kcpDialer struct {
sessions map[string]*muxSession sessions map[string]*muxSession
sessionMutex sync.Mutex sessionMutex sync.Mutex
retired bool
logger logger.Logger logger logger.Logger
md metadata md metadata
options dialer.Options options dialer.Options
@@ -64,6 +66,9 @@ func (d *kcpDialer) Dial(ctx context.Context, addr string, opts ...dialer.DialOp
d.sessionMutex.Lock() d.sessionMutex.Lock()
defer d.sessionMutex.Unlock() defer d.sessionMutex.Unlock()
if d.retired {
return nil, net.ErrClosed
}
session, ok := d.sessions[addr] session, ok := d.sessions[addr]
if session != nil && session.IsClosed() { if session != nil && session.IsClosed() {
@@ -171,3 +176,36 @@ func (d *kcpDialer) initSession(ctx context.Context, addr net.Addr, conn net.Pac
func (d *kcpDialer) Multiplex() bool { func (d *kcpDialer) Multiplex() bool {
return true return true
} }
// Retire drains existing streams and closes their backing sessions once idle.
func (d *kcpDialer) Retire() {
for _, session := range d.detachSessions() {
sessionretire.Gracefully(session)
}
}
// Close immediately releases all cached multiplex sessions.
func (d *kcpDialer) Close() error {
var errs []error
for _, session := range d.detachSessions() {
errs = append(errs, session.Close())
}
return errors.Join(errs...)
}
func (d *kcpDialer) detachSessions() []*muxSession {
if d == nil {
return nil
}
d.sessionMutex.Lock()
d.retired = true
sessions := make([]*muxSession, 0, len(d.sessions))
for _, session := range d.sessions {
if session != nil {
sessions = append(sessions, session)
}
}
d.sessions = make(map[string]*muxSession)
d.sessionMutex.Unlock()
return sessions
}
+17
View File
@@ -0,0 +1,17 @@
package kcp
import (
"context"
"errors"
"net"
"testing"
)
func TestRetiredDialerRejectsNewConnections(t *testing.T) {
dialer := NewDialer().(*kcpDialer)
dialer.Retire()
if _, err := dialer.Dial(context.Background(), "127.0.0.1:1"); !errors.Is(err, net.ErrClosed) {
t.Fatalf("Dial error = %v, want net.ErrClosed", err)
}
}
+14
View File
@@ -20,13 +20,24 @@ func (session *muxSession) Accept() (net.Conn, error) {
} }
func (session *muxSession) Close() error { func (session *muxSession) Close() error {
if session == nil {
return nil
}
if session.session == nil { if session.session == nil {
if session.conn != nil {
conn := session.conn
session.conn = nil
return conn.Close()
}
return nil return nil
} }
return session.session.Close() return session.session.Close()
} }
func (session *muxSession) IsClosed() bool { func (session *muxSession) IsClosed() bool {
if session == nil {
return true
}
if session.session == nil { if session.session == nil {
return true return true
} }
@@ -34,5 +45,8 @@ func (session *muxSession) IsClosed() bool {
} }
func (session *muxSession) NumStreams() int { func (session *muxSession) NumStreams() int {
if session == nil || session.session == nil {
return 0
}
return session.session.NumStreams() return session.session.NumStreams()
} }
+40
View File
@@ -11,6 +11,7 @@ import (
"github.com/go-gost/core/logger" "github.com/go-gost/core/logger"
md "github.com/go-gost/core/metadata" md "github.com/go-gost/core/metadata"
"github.com/go-gost/x/internal/util/mux" "github.com/go-gost/x/internal/util/mux"
"github.com/go-gost/x/internal/util/sessionretire"
"github.com/go-gost/x/registry" "github.com/go-gost/x/registry"
) )
@@ -21,6 +22,7 @@ func init() {
type mtcpDialer struct { type mtcpDialer struct {
sessions map[string]*muxSession sessions map[string]*muxSession
sessionMutex sync.Mutex sessionMutex sync.Mutex
retired bool
logger logger.Logger logger logger.Logger
md metadata md metadata
options dialer.Options options dialer.Options
@@ -55,6 +57,9 @@ func (d *mtcpDialer) Multiplex() bool {
func (d *mtcpDialer) Dial(ctx context.Context, addr string, opts ...dialer.DialOption) (conn net.Conn, err error) { func (d *mtcpDialer) Dial(ctx context.Context, addr string, opts ...dialer.DialOption) (conn net.Conn, err error) {
d.sessionMutex.Lock() d.sessionMutex.Lock()
defer d.sessionMutex.Unlock() defer d.sessionMutex.Unlock()
if d.retired {
return nil, net.ErrClosed
}
session, ok := d.sessions[addr] session, ok := d.sessions[addr]
if session != nil && session.IsClosed() { if session != nil && session.IsClosed() {
@@ -88,6 +93,10 @@ func (d *mtcpDialer) Handshake(ctx context.Context, conn net.Conn, options ...di
d.sessionMutex.Lock() d.sessionMutex.Lock()
defer d.sessionMutex.Unlock() defer d.sessionMutex.Unlock()
if d.retired {
conn.Close()
return nil, net.ErrClosed
}
if d.md.handshakeTimeout > 0 { if d.md.handshakeTimeout > 0 {
conn.SetDeadline(time.Now().Add(d.md.handshakeTimeout)) conn.SetDeadline(time.Now().Add(d.md.handshakeTimeout))
@@ -129,3 +138,34 @@ func (d *mtcpDialer) initSession(ctx context.Context, conn net.Conn) (*muxSessio
} }
return &muxSession{conn: conn, session: session}, nil return &muxSession{conn: conn, session: session}, nil
} }
func (d *mtcpDialer) Retire() {
for _, session := range d.detachSessions() {
sessionretire.Gracefully(session)
}
}
func (d *mtcpDialer) Close() error {
var errs []error
for _, session := range d.detachSessions() {
errs = append(errs, session.Close())
}
return errors.Join(errs...)
}
func (d *mtcpDialer) detachSessions() []*muxSession {
if d == nil {
return nil
}
d.sessionMutex.Lock()
d.retired = true
sessions := make([]*muxSession, 0, len(d.sessions))
for _, session := range d.sessions {
if session != nil {
sessions = append(sessions, session)
}
}
d.sessions = make(map[string]*muxSession)
d.sessionMutex.Unlock()
return sessions
}
+30
View File
@@ -0,0 +1,30 @@
package mtcp
import (
"context"
"errors"
"net"
"testing"
)
func TestRetiredDialerRejectsNewConnections(t *testing.T) {
dialer := NewDialer().(*mtcpDialer)
dialer.Retire()
if _, err := dialer.Dial(context.Background(), "127.0.0.1:1"); !errors.Is(err, net.ErrClosed) {
t.Fatalf("Dial error = %v, want net.ErrClosed", err)
}
}
func TestSessionCloseReleasesPreHandshakeConnection(t *testing.T) {
conn, peer := net.Pipe()
defer peer.Close()
session := &muxSession{conn: conn}
if err := session.Close(); err != nil {
t.Fatalf("Close: %v", err)
}
if !session.IsClosed() {
t.Fatal("pre-handshake session still reports open after Close")
}
}
+14
View File
@@ -20,13 +20,24 @@ func (session *muxSession) Accept() (net.Conn, error) {
} }
func (session *muxSession) Close() error { func (session *muxSession) Close() error {
if session == nil {
return nil
}
if session.session == nil { if session.session == nil {
if session.conn != nil {
conn := session.conn
session.conn = nil
return conn.Close()
}
return nil return nil
} }
return session.session.Close() return session.session.Close()
} }
func (session *muxSession) IsClosed() bool { func (session *muxSession) IsClosed() bool {
if session == nil {
return true
}
if session.session == nil { if session.session == nil {
return true return true
} }
@@ -34,5 +45,8 @@ func (session *muxSession) IsClosed() bool {
} }
func (session *muxSession) NumStreams() int { func (session *muxSession) NumStreams() int {
if session == nil || session.session == nil {
return 0
}
return session.session.NumStreams() return session.session.NumStreams()
} }
+40
View File
@@ -12,6 +12,7 @@ import (
"github.com/go-gost/core/logger" "github.com/go-gost/core/logger"
md "github.com/go-gost/core/metadata" md "github.com/go-gost/core/metadata"
"github.com/go-gost/x/internal/util/mux" "github.com/go-gost/x/internal/util/mux"
"github.com/go-gost/x/internal/util/sessionretire"
"github.com/go-gost/x/registry" "github.com/go-gost/x/registry"
) )
@@ -22,6 +23,7 @@ func init() {
type mtlsDialer struct { type mtlsDialer struct {
sessions map[string]*muxSession sessions map[string]*muxSession
sessionMutex sync.Mutex sessionMutex sync.Mutex
retired bool
logger logger.Logger logger logger.Logger
md metadata md metadata
options dialer.Options options dialer.Options
@@ -56,6 +58,9 @@ func (d *mtlsDialer) Multiplex() bool {
func (d *mtlsDialer) Dial(ctx context.Context, addr string, opts ...dialer.DialOption) (conn net.Conn, err error) { func (d *mtlsDialer) Dial(ctx context.Context, addr string, opts ...dialer.DialOption) (conn net.Conn, err error) {
d.sessionMutex.Lock() d.sessionMutex.Lock()
defer d.sessionMutex.Unlock() defer d.sessionMutex.Unlock()
if d.retired {
return nil, net.ErrClosed
}
session, ok := d.sessions[addr] session, ok := d.sessions[addr]
if session != nil && session.IsClosed() { if session != nil && session.IsClosed() {
@@ -89,6 +94,10 @@ func (d *mtlsDialer) Handshake(ctx context.Context, conn net.Conn, options ...di
d.sessionMutex.Lock() d.sessionMutex.Lock()
defer d.sessionMutex.Unlock() defer d.sessionMutex.Unlock()
if d.retired {
conn.Close()
return nil, net.ErrClosed
}
if d.md.handshakeTimeout > 0 { if d.md.handshakeTimeout > 0 {
conn.SetDeadline(time.Now().Add(d.md.handshakeTimeout)) conn.SetDeadline(time.Now().Add(d.md.handshakeTimeout))
@@ -136,3 +145,34 @@ func (d *mtlsDialer) initSession(ctx context.Context, conn net.Conn) (*muxSessio
} }
return &muxSession{conn: conn, session: session}, nil return &muxSession{conn: conn, session: session}, nil
} }
func (d *mtlsDialer) Retire() {
for _, session := range d.detachSessions() {
sessionretire.Gracefully(session)
}
}
func (d *mtlsDialer) Close() error {
var errs []error
for _, session := range d.detachSessions() {
errs = append(errs, session.Close())
}
return errors.Join(errs...)
}
func (d *mtlsDialer) detachSessions() []*muxSession {
if d == nil {
return nil
}
d.sessionMutex.Lock()
d.retired = true
sessions := make([]*muxSession, 0, len(d.sessions))
for _, session := range d.sessions {
if session != nil {
sessions = append(sessions, session)
}
}
d.sessions = make(map[string]*muxSession)
d.sessionMutex.Unlock()
return sessions
}
+30
View File
@@ -0,0 +1,30 @@
package mtls
import (
"context"
"errors"
"net"
"testing"
)
func TestRetiredDialerRejectsNewConnections(t *testing.T) {
dialer := NewDialer().(*mtlsDialer)
dialer.Retire()
if _, err := dialer.Dial(context.Background(), "127.0.0.1:1"); !errors.Is(err, net.ErrClosed) {
t.Fatalf("Dial error = %v, want net.ErrClosed", err)
}
}
func TestSessionCloseReleasesPreHandshakeConnection(t *testing.T) {
conn, peer := net.Pipe()
defer peer.Close()
session := &muxSession{conn: conn}
if err := session.Close(); err != nil {
t.Fatalf("Close: %v", err)
}
if !session.IsClosed() {
t.Fatal("pre-handshake session still reports open after Close")
}
}
+14
View File
@@ -20,13 +20,24 @@ func (session *muxSession) Accept() (net.Conn, error) {
} }
func (session *muxSession) Close() error { func (session *muxSession) Close() error {
if session == nil {
return nil
}
if session.session == nil { if session.session == nil {
if session.conn != nil {
conn := session.conn
session.conn = nil
return conn.Close()
}
return nil return nil
} }
return session.session.Close() return session.session.Close()
} }
func (session *muxSession) IsClosed() bool { func (session *muxSession) IsClosed() bool {
if session == nil {
return true
}
if session.session == nil { if session.session == nil {
return true return true
} }
@@ -34,5 +45,8 @@ func (session *muxSession) IsClosed() bool {
} }
func (session *muxSession) NumStreams() int { func (session *muxSession) NumStreams() int {
if session == nil || session.session == nil {
return 0
}
return session.session.NumStreams() return session.session.NumStreams()
} }
+40
View File
@@ -12,6 +12,7 @@ import (
"github.com/go-gost/core/logger" "github.com/go-gost/core/logger"
md "github.com/go-gost/core/metadata" md "github.com/go-gost/core/metadata"
"github.com/go-gost/x/internal/util/mux" "github.com/go-gost/x/internal/util/mux"
"github.com/go-gost/x/internal/util/sessionretire"
ws_util "github.com/go-gost/x/internal/util/ws" ws_util "github.com/go-gost/x/internal/util/ws"
"github.com/go-gost/x/registry" "github.com/go-gost/x/registry"
"github.com/gorilla/websocket" "github.com/gorilla/websocket"
@@ -25,6 +26,7 @@ func init() {
type mwsDialer struct { type mwsDialer struct {
sessions map[string]*muxSession sessions map[string]*muxSession
sessionMutex sync.Mutex sessionMutex sync.Mutex
retired bool
tlsEnabled bool tlsEnabled bool
md metadata md metadata
options dialer.Options options dialer.Options
@@ -70,6 +72,9 @@ func (d *mwsDialer) Multiplex() bool {
func (d *mwsDialer) Dial(ctx context.Context, addr string, opts ...dialer.DialOption) (conn net.Conn, err error) { func (d *mwsDialer) Dial(ctx context.Context, addr string, opts ...dialer.DialOption) (conn net.Conn, err error) {
d.sessionMutex.Lock() d.sessionMutex.Lock()
defer d.sessionMutex.Unlock() defer d.sessionMutex.Unlock()
if d.retired {
return nil, net.ErrClosed
}
session, ok := d.sessions[addr] session, ok := d.sessions[addr]
if session != nil && session.IsClosed() { if session != nil && session.IsClosed() {
@@ -108,6 +113,10 @@ func (d *mwsDialer) Handshake(ctx context.Context, conn net.Conn, options ...dia
d.sessionMutex.Lock() d.sessionMutex.Lock()
defer d.sessionMutex.Unlock() defer d.sessionMutex.Unlock()
if d.retired {
conn.Close()
return nil, net.ErrClosed
}
session, ok := d.sessions[opts.Addr] session, ok := d.sessions[opts.Addr]
if session != nil && session.conn != conn { if session != nil && session.conn != conn {
@@ -208,3 +217,34 @@ func (d *mwsDialer) keepAlive(conn ws_util.WebsocketConn) {
conn.SetWriteDeadline(time.Time{}) conn.SetWriteDeadline(time.Time{})
} }
} }
func (d *mwsDialer) Retire() {
for _, session := range d.detachSessions() {
sessionretire.Gracefully(session)
}
}
func (d *mwsDialer) Close() error {
var errs []error
for _, session := range d.detachSessions() {
errs = append(errs, session.Close())
}
return errors.Join(errs...)
}
func (d *mwsDialer) detachSessions() []*muxSession {
if d == nil {
return nil
}
d.sessionMutex.Lock()
d.retired = true
sessions := make([]*muxSession, 0, len(d.sessions))
for _, session := range d.sessions {
if session != nil {
sessions = append(sessions, session)
}
}
d.sessions = make(map[string]*muxSession)
d.sessionMutex.Unlock()
return sessions
}
+30
View File
@@ -0,0 +1,30 @@
package mws
import (
"context"
"errors"
"net"
"testing"
)
func TestRetiredDialerRejectsNewConnections(t *testing.T) {
dialer := NewDialer().(*mwsDialer)
dialer.Retire()
if _, err := dialer.Dial(context.Background(), "127.0.0.1:1"); !errors.Is(err, net.ErrClosed) {
t.Fatalf("Dial error = %v, want net.ErrClosed", err)
}
}
func TestSessionCloseReleasesPreHandshakeConnection(t *testing.T) {
conn, peer := net.Pipe()
defer peer.Close()
session := &muxSession{conn: conn}
if err := session.Close(); err != nil {
t.Fatalf("Close: %v", err)
}
if !session.IsClosed() {
t.Fatal("pre-handshake session still reports open after Close")
}
}
+43 -2
View File
@@ -3,6 +3,7 @@ package hop
import ( import (
"context" "context"
"encoding/json" "encoding/json"
"errors"
"io" "io"
"net" "net"
"sort" "sort"
@@ -92,6 +93,7 @@ type chainHop struct {
nodes []*chain.Node nodes []*chain.Node
mu sync.RWMutex mu sync.RWMutex
cancelFunc context.CancelFunc cancelFunc context.CancelFunc
stopOnce sync.Once
options options options options
} }
@@ -383,7 +385,11 @@ func (p *chainHop) parseNode(r io.Reader) ([]*chain.Node, error) {
return nodes, nil return nodes, nil
} }
func (p *chainHop) Close() error { func (p *chainHop) stopReload() {
if p == nil {
return
}
p.stopOnce.Do(func() {
p.cancelFunc() p.cancelFunc()
if p.options.fileLoader != nil { if p.options.fileLoader != nil {
p.options.fileLoader.Close() p.options.fileLoader.Close()
@@ -391,5 +397,40 @@ func (p *chainHop) Close() error {
if p.options.redisLoader != nil { if p.options.redisLoader != nil {
p.options.redisLoader.Close() p.options.redisLoader.Close()
} }
return nil if p.options.httpLoader != nil {
p.options.httpLoader.Close()
}
})
}
func (p *chainHop) Retire() {
if p == nil {
return
}
p.stopReload()
for _, node := range p.Nodes() {
if node == nil || node.Options().Transport == nil {
continue
}
if retirer, ok := node.Options().Transport.(interface{ Retire() }); ok {
retirer.Retire()
}
}
}
func (p *chainHop) Close() error {
if p == nil {
return nil
}
p.stopReload()
var errs []error
for _, node := range p.Nodes() {
if node == nil || node.Options().Transport == nil {
continue
}
if closer, ok := node.Options().Transport.(io.Closer); ok {
errs = append(errs, closer.Close())
}
}
return errors.Join(errs...)
} }
@@ -0,0 +1,58 @@
package sessionretire
import "time"
const (
defaultIdleGrace = time.Second
defaultPollPeriod = 100 * time.Millisecond
)
// Session is the lifecycle surface shared by the multiplexed dialers.
type Session interface {
Close() error
IsClosed() bool
NumStreams() int
}
// Gracefully closes a retired session after all existing streams have drained.
// A short idle grace covers the Dial/Handshake hand-off used by several dialers.
func Gracefully(session Session) {
if session == nil {
return
}
go waitUntilIdle(session, defaultIdleGrace, defaultPollPeriod)
}
func waitUntilIdle(session Session, idleGrace, pollPeriod time.Duration) {
if session == nil {
return
}
if idleGrace <= 0 {
idleGrace = defaultIdleGrace
}
if pollPeriod <= 0 {
pollPeriod = defaultPollPeriod
}
ticker := time.NewTicker(pollPeriod)
defer ticker.Stop()
var idleSince time.Time
for {
if session.IsClosed() {
_ = session.Close()
return
}
if session.NumStreams() == 0 {
if idleSince.IsZero() {
idleSince = time.Now()
} else if time.Since(idleSince) >= idleGrace {
_ = session.Close()
return
}
} else {
idleSince = time.Time{}
}
<-ticker.C
}
}
@@ -0,0 +1,59 @@
package sessionretire
import (
"sync"
"testing"
"time"
)
type testSession struct {
mu sync.Mutex
streams int
closed bool
}
func (s *testSession) Close() error {
s.mu.Lock()
s.closed = true
s.mu.Unlock()
return nil
}
func (s *testSession) IsClosed() bool {
s.mu.Lock()
defer s.mu.Unlock()
return s.closed
}
func (s *testSession) NumStreams() int {
s.mu.Lock()
defer s.mu.Unlock()
return s.streams
}
func TestWaitUntilIdlePreservesActiveStreams(t *testing.T) {
session := &testSession{streams: 1}
done := make(chan struct{})
go func() {
waitUntilIdle(session, 20*time.Millisecond, time.Millisecond)
close(done)
}()
time.Sleep(30 * time.Millisecond)
if session.IsClosed() {
t.Fatal("active session was closed")
}
session.mu.Lock()
session.streams = 0
session.mu.Unlock()
select {
case <-done:
case <-time.After(250 * time.Millisecond):
t.Fatal("idle session was not closed")
}
if !session.IsClosed() {
t.Fatal("retired session did not close after becoming idle")
}
}
+7 -1
View File
@@ -28,7 +28,13 @@ func (r *chainRegistry) Register(name string, v chain.Chainer) error {
} }
func (r *chainRegistry) replace(name string, v chain.Chainer) { func (r *chainRegistry) replace(name string, v chain.Chainer) {
r.m.Store(name, v) old, loaded := r.m.Swap(name, v)
if !loaded {
return
}
if retirer, ok := old.(interface{ Retire() }); ok {
retirer.Retire()
}
} }
func (r *chainRegistry) Get(name string) chain.Chainer { func (r *chainRegistry) Get(name string) chain.Chainer {
+26
View File
@@ -16,6 +16,15 @@ func (c testChainer) Route(context.Context, string, string, ...chain.RouteOption
return c.route return c.route
} }
type retiringTestChainer struct {
testChainer
retired bool
}
func (c *retiringTestChainer) Retire() {
c.retired = true
}
type testRoute struct { type testRoute struct {
nodes []*chain.Node nodes []*chain.Node
} }
@@ -49,3 +58,20 @@ func TestReplaceChainOverwritesExistingRegistration(t *testing.T) {
t.Fatalf("expected replacement chain route, got %#v", route) t.Fatalf("expected replacement chain route, got %#v", route)
} }
} }
func TestReplaceChainRetiresPreviousRegistration(t *testing.T) {
name := "replace_chain_retire_tdd"
ChainRegistry().Unregister(name)
defer ChainRegistry().Unregister(name)
old := &retiringTestChainer{}
if err := ChainRegistry().Register(name, old); err != nil {
t.Fatalf("register old chain: %v", err)
}
if err := ReplaceChain(name, testChainer{}); err != nil {
t.Fatalf("replace chain: %v", err)
}
if !old.retired {
t.Fatal("previous chain was not retired")
}
}