mirror of
https://github.com/Sagit-chu/flvx.git
synced 2026-10-11 19:56:37 +08:00
fix(agent): retire replaced tunnel sessions (#541)
This commit is contained in:
@@ -19,20 +19,23 @@ func (session *muxSession) Accept() (net.Conn, error) {
|
||||
}
|
||||
|
||||
func (session *muxSession) Close() error {
|
||||
if session.session == nil {
|
||||
if session == nil || session.session == nil {
|
||||
return nil
|
||||
}
|
||||
return session.session.Close()
|
||||
}
|
||||
|
||||
func (session *muxSession) IsClosed() bool {
|
||||
if session.session == nil {
|
||||
if session == nil || session.session == nil {
|
||||
return true
|
||||
}
|
||||
return session.session.IsClosed()
|
||||
}
|
||||
|
||||
func (session *muxSession) NumStreams() int {
|
||||
if session == nil || session.session == nil {
|
||||
return 0
|
||||
}
|
||||
return session.session.NumStreams()
|
||||
}
|
||||
|
||||
|
||||
@@ -11,6 +11,7 @@ import (
|
||||
"github.com/go-gost/core/logger"
|
||||
md "github.com/go-gost/core/metadata"
|
||||
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"
|
||||
"github.com/go-gost/x/registry"
|
||||
"github.com/xtaci/kcp-go/v5"
|
||||
@@ -25,6 +26,7 @@ func init() {
|
||||
type kcpDialer struct {
|
||||
sessions map[string]*muxSession
|
||||
sessionMutex sync.Mutex
|
||||
retired bool
|
||||
logger logger.Logger
|
||||
md metadata
|
||||
options dialer.Options
|
||||
@@ -64,6 +66,9 @@ func (d *kcpDialer) Dial(ctx context.Context, addr string, opts ...dialer.DialOp
|
||||
|
||||
d.sessionMutex.Lock()
|
||||
defer d.sessionMutex.Unlock()
|
||||
if d.retired {
|
||||
return nil, net.ErrClosed
|
||||
}
|
||||
|
||||
session, ok := d.sessions[addr]
|
||||
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 {
|
||||
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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -20,13 +20,24 @@ func (session *muxSession) Accept() (net.Conn, error) {
|
||||
}
|
||||
|
||||
func (session *muxSession) Close() error {
|
||||
if session == nil {
|
||||
return nil
|
||||
}
|
||||
if session.session == nil {
|
||||
if session.conn != nil {
|
||||
conn := session.conn
|
||||
session.conn = nil
|
||||
return conn.Close()
|
||||
}
|
||||
return nil
|
||||
}
|
||||
return session.session.Close()
|
||||
}
|
||||
|
||||
func (session *muxSession) IsClosed() bool {
|
||||
if session == nil {
|
||||
return true
|
||||
}
|
||||
if session.session == nil {
|
||||
return true
|
||||
}
|
||||
@@ -34,5 +45,8 @@ func (session *muxSession) IsClosed() bool {
|
||||
}
|
||||
|
||||
func (session *muxSession) NumStreams() int {
|
||||
if session == nil || session.session == nil {
|
||||
return 0
|
||||
}
|
||||
return session.session.NumStreams()
|
||||
}
|
||||
|
||||
@@ -11,6 +11,7 @@ import (
|
||||
"github.com/go-gost/core/logger"
|
||||
md "github.com/go-gost/core/metadata"
|
||||
"github.com/go-gost/x/internal/util/mux"
|
||||
"github.com/go-gost/x/internal/util/sessionretire"
|
||||
"github.com/go-gost/x/registry"
|
||||
)
|
||||
|
||||
@@ -21,6 +22,7 @@ func init() {
|
||||
type mtcpDialer struct {
|
||||
sessions map[string]*muxSession
|
||||
sessionMutex sync.Mutex
|
||||
retired bool
|
||||
logger logger.Logger
|
||||
md metadata
|
||||
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) {
|
||||
d.sessionMutex.Lock()
|
||||
defer d.sessionMutex.Unlock()
|
||||
if d.retired {
|
||||
return nil, net.ErrClosed
|
||||
}
|
||||
|
||||
session, ok := d.sessions[addr]
|
||||
if session != nil && session.IsClosed() {
|
||||
@@ -88,6 +93,10 @@ func (d *mtcpDialer) Handshake(ctx context.Context, conn net.Conn, options ...di
|
||||
|
||||
d.sessionMutex.Lock()
|
||||
defer d.sessionMutex.Unlock()
|
||||
if d.retired {
|
||||
conn.Close()
|
||||
return nil, net.ErrClosed
|
||||
}
|
||||
|
||||
if d.md.handshakeTimeout > 0 {
|
||||
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
|
||||
}
|
||||
|
||||
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
|
||||
}
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -20,13 +20,24 @@ func (session *muxSession) Accept() (net.Conn, error) {
|
||||
}
|
||||
|
||||
func (session *muxSession) Close() error {
|
||||
if session == nil {
|
||||
return nil
|
||||
}
|
||||
if session.session == nil {
|
||||
if session.conn != nil {
|
||||
conn := session.conn
|
||||
session.conn = nil
|
||||
return conn.Close()
|
||||
}
|
||||
return nil
|
||||
}
|
||||
return session.session.Close()
|
||||
}
|
||||
|
||||
func (session *muxSession) IsClosed() bool {
|
||||
if session == nil {
|
||||
return true
|
||||
}
|
||||
if session.session == nil {
|
||||
return true
|
||||
}
|
||||
@@ -34,5 +45,8 @@ func (session *muxSession) IsClosed() bool {
|
||||
}
|
||||
|
||||
func (session *muxSession) NumStreams() int {
|
||||
if session == nil || session.session == nil {
|
||||
return 0
|
||||
}
|
||||
return session.session.NumStreams()
|
||||
}
|
||||
|
||||
@@ -12,6 +12,7 @@ import (
|
||||
"github.com/go-gost/core/logger"
|
||||
md "github.com/go-gost/core/metadata"
|
||||
"github.com/go-gost/x/internal/util/mux"
|
||||
"github.com/go-gost/x/internal/util/sessionretire"
|
||||
"github.com/go-gost/x/registry"
|
||||
)
|
||||
|
||||
@@ -22,6 +23,7 @@ func init() {
|
||||
type mtlsDialer struct {
|
||||
sessions map[string]*muxSession
|
||||
sessionMutex sync.Mutex
|
||||
retired bool
|
||||
logger logger.Logger
|
||||
md metadata
|
||||
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) {
|
||||
d.sessionMutex.Lock()
|
||||
defer d.sessionMutex.Unlock()
|
||||
if d.retired {
|
||||
return nil, net.ErrClosed
|
||||
}
|
||||
|
||||
session, ok := d.sessions[addr]
|
||||
if session != nil && session.IsClosed() {
|
||||
@@ -89,6 +94,10 @@ func (d *mtlsDialer) Handshake(ctx context.Context, conn net.Conn, options ...di
|
||||
|
||||
d.sessionMutex.Lock()
|
||||
defer d.sessionMutex.Unlock()
|
||||
if d.retired {
|
||||
conn.Close()
|
||||
return nil, net.ErrClosed
|
||||
}
|
||||
|
||||
if d.md.handshakeTimeout > 0 {
|
||||
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
|
||||
}
|
||||
|
||||
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
|
||||
}
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -20,13 +20,24 @@ func (session *muxSession) Accept() (net.Conn, error) {
|
||||
}
|
||||
|
||||
func (session *muxSession) Close() error {
|
||||
if session == nil {
|
||||
return nil
|
||||
}
|
||||
if session.session == nil {
|
||||
if session.conn != nil {
|
||||
conn := session.conn
|
||||
session.conn = nil
|
||||
return conn.Close()
|
||||
}
|
||||
return nil
|
||||
}
|
||||
return session.session.Close()
|
||||
}
|
||||
|
||||
func (session *muxSession) IsClosed() bool {
|
||||
if session == nil {
|
||||
return true
|
||||
}
|
||||
if session.session == nil {
|
||||
return true
|
||||
}
|
||||
@@ -34,5 +45,8 @@ func (session *muxSession) IsClosed() bool {
|
||||
}
|
||||
|
||||
func (session *muxSession) NumStreams() int {
|
||||
if session == nil || session.session == nil {
|
||||
return 0
|
||||
}
|
||||
return session.session.NumStreams()
|
||||
}
|
||||
|
||||
@@ -12,6 +12,7 @@ import (
|
||||
"github.com/go-gost/core/logger"
|
||||
md "github.com/go-gost/core/metadata"
|
||||
"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"
|
||||
"github.com/go-gost/x/registry"
|
||||
"github.com/gorilla/websocket"
|
||||
@@ -25,6 +26,7 @@ func init() {
|
||||
type mwsDialer struct {
|
||||
sessions map[string]*muxSession
|
||||
sessionMutex sync.Mutex
|
||||
retired bool
|
||||
tlsEnabled bool
|
||||
md metadata
|
||||
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) {
|
||||
d.sessionMutex.Lock()
|
||||
defer d.sessionMutex.Unlock()
|
||||
if d.retired {
|
||||
return nil, net.ErrClosed
|
||||
}
|
||||
|
||||
session, ok := d.sessions[addr]
|
||||
if session != nil && session.IsClosed() {
|
||||
@@ -108,6 +113,10 @@ func (d *mwsDialer) Handshake(ctx context.Context, conn net.Conn, options ...dia
|
||||
|
||||
d.sessionMutex.Lock()
|
||||
defer d.sessionMutex.Unlock()
|
||||
if d.retired {
|
||||
conn.Close()
|
||||
return nil, net.ErrClosed
|
||||
}
|
||||
|
||||
session, ok := d.sessions[opts.Addr]
|
||||
if session != nil && session.conn != conn {
|
||||
@@ -208,3 +217,34 @@ func (d *mwsDialer) keepAlive(conn ws_util.WebsocketConn) {
|
||||
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
|
||||
}
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user