diff --git a/core/xray.go b/core/xray.go index 3e785a823..794ce8fff 100644 --- a/core/xray.go +++ b/core/xray.go @@ -4,7 +4,6 @@ import ( "context" "reflect" "sync" - "sync/atomic" "github.com/xtls/xray-core/common" "github.com/xtls/xray-core/common/errors" @@ -85,7 +84,7 @@ type Instance struct { features []features.Feature pendingResolutions []resolution pendingOptionalResolutions []resolution - running atomic.Bool + running bool resolveLock sync.Mutex ctx context.Context @@ -93,7 +92,9 @@ type Instance struct { // Instance state func (server *Instance) IsRunning() bool { - return server.running.Load() + server.statusLock.Lock() + defer server.statusLock.Unlock() + return server.running } func AddInboundHandler(server *Instance, config *InboundHandlerConfig) error { @@ -263,7 +264,7 @@ func (s *Instance) Close() error { s.statusLock.Lock() defer s.statusLock.Unlock() - s.running.Store(false) + s.running = false var errs []interface{} for _, f := range s.features { @@ -321,7 +322,7 @@ func (s *Instance) RequireFeatures(callback interface{}, optional bool) error { // AddFeature registers a feature into current Instance. func (s *Instance) AddFeature(feature features.Feature) error { - if s.running.Load() { + if s.IsRunning() { if err := feature.Start(); err != nil { errors.LogInfoInner(s.ctx, err, "failed to start feature") } @@ -389,7 +390,7 @@ func (s *Instance) Start() error { s.statusLock.Lock() defer s.statusLock.Unlock() - s.running.Store(true) + s.running = true for _, f := range s.features { if err := f.Start(); err != nil { return err diff --git a/transport/internet/hysteria/config.go b/transport/internet/hysteria/config.go index 24d4a98be..bf87d604c 100644 --- a/transport/internet/hysteria/config.go +++ b/transport/internet/hysteria/config.go @@ -79,7 +79,6 @@ const ( StatusNull status = iota StatusActive StatusInactive - StatusClosed ) const protocolName = "hysteria" diff --git a/transport/internet/hysteria/dialer.go b/transport/internet/hysteria/dialer.go index dedc02fd6..29b8a5d18 100644 --- a/transport/internet/hysteria/dialer.go +++ b/transport/internet/hysteria/dialer.go @@ -29,6 +29,7 @@ import ( type client struct { sync.Mutex + instance *core.Instance dest net.Destination config *Config tlsConfig *gotls.Config @@ -36,17 +37,13 @@ type client struct { finalMask *finalmask.FinalMask quicParams *internet.QuicParams - conn *quic.Conn - tr *quic.Transport - pktConn net.PacketConn - udpSM *udpSessionManager - instance *core.Instance + conn *quic.Conn + tr *quic.Transport + pktConn net.PacketConn + udpSM *udpSessionManager } func (c *client) status() status { - if c.instance != nil && !c.instance.IsRunning() { - return StatusClosed - } if c.conn == nil { return StatusNull } @@ -54,16 +51,17 @@ func (c *client) status() status { case <-c.conn.Context().Done(): return StatusInactive default: - return StatusActive + if c.instance == nil || c.instance.IsRunning() { + return StatusActive + } + return StatusInactive } } func (c *client) close() { - if c.conn != nil { - c.conn.CloseWithError(closeErrCodeOK, "") - } - common.CloseIfExists(c.tr) - common.CloseIfExists(c.pktConn) + c.conn.CloseWithError(closeErrCodeOK, "") + c.tr.Close() + c.pktConn.Close() c.conn = nil c.tr = nil c.pktConn = nil @@ -71,9 +69,11 @@ func (c *client) close() { } func (c *client) dial(ctx context.Context) error { - switch c.status() { - case StatusClosed: + if c.instance != nil && !c.instance.IsRunning() { return errors.New("client is closed") + } + + switch c.status() { case StatusActive: return nil case StatusInactive: @@ -264,18 +264,13 @@ func (c *client) udp(ctx context.Context) (stat.Connection, error) { return c.udpSM.udp() } -func (c *client) clean() (shouldDelete bool) { +func (c *client) clean() bool { c.Lock() defer c.Unlock() - switch c.status() { - case StatusClosed: + if c.status() == StatusInactive { c.close() - return true - case StatusInactive: - c.close() - return false } - return false + return c.status() == StatusNull } type dialerConf struct { @@ -291,21 +286,13 @@ type clientManager struct { func (m *clientManager) clean() { ticker := time.NewTicker(idleCleanupInterval) for range ticker.C { - var toDelete []dialerConf - m.RLock() + m.Lock() for k, c := range m.m { if c.clean() { - toDelete = append(toDelete, k) - } - } - m.RUnlock() - if len(toDelete) > 0 { - m.Lock() - for _, k := range toDelete { delete(m.m, k) } - m.Unlock() } + m.Unlock() } } @@ -331,6 +318,7 @@ func Dial(ctx context.Context, dest net.Destination, streamSettings *internet.Me }) dialerConfKey := dialerConf{dest, streamSettings} + manager.RLock() c := manager.m[dialerConfKey] manager.RUnlock()