From 74bf20d349433557d744f0ad19d4326daf69b85c Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E4=B8=96=E7=95=8C?= Date: Fri, 13 Mar 2026 21:31:23 +0800 Subject: [PATCH] ccm,ocm: fix reverse session shutdown race --- service/ccm/credential_external.go | 24 +++++++++++++++++++----- service/ccm/reverse.go | 5 ++++- service/ocm/credential_external.go | 24 +++++++++++++++++++----- service/ocm/reverse.go | 5 ++++- 4 files changed, 46 insertions(+), 12 deletions(-) diff --git a/service/ccm/credential_external.go b/service/ccm/credential_external.go index 141c893c3..a0350a9fd 100644 --- a/service/ccm/credential_external.go +++ b/service/ccm/credential_external.go @@ -50,6 +50,7 @@ type externalCredential struct { reverse bool reverseSession *yamux.Session reverseAccess sync.RWMutex + closed bool reverseContext context.Context reverseCancel context.CancelFunc connectorDialer N.Dialer @@ -542,12 +543,16 @@ func (c *externalCredential) httpTransport() *http.Client { } func (c *externalCredential) close() { + var session *yamux.Session c.reverseAccess.Lock() - if c.reverseCancel != nil { - c.reverseCancel() + if !c.closed { + c.closed = true + if c.reverseCancel != nil { + c.reverseCancel() + } + session = c.reverseSession + c.reverseSession = nil } - session := c.reverseSession - c.reverseSession = nil c.reverseAccess.Unlock() if session != nil { session.Close() @@ -567,14 +572,19 @@ func (c *externalCredential) getReverseSession() *yamux.Session { return c.reverseSession } -func (c *externalCredential) setReverseSession(session *yamux.Session) { +func (c *externalCredential) setReverseSession(session *yamux.Session) bool { c.reverseAccess.Lock() + if c.closed { + c.reverseAccess.Unlock() + return false + } old := c.reverseSession c.reverseSession = session c.reverseAccess.Unlock() if old != nil { old.Close() } + return true } func (c *externalCredential) clearReverseSession(session *yamux.Session) { @@ -593,6 +603,10 @@ func (c *externalCredential) getReverseContext() context.Context { func (c *externalCredential) resetReverseContext() { c.reverseAccess.Lock() + if c.closed { + c.reverseAccess.Unlock() + return + } c.reverseCancel() c.reverseContext, c.reverseCancel = context.WithCancel(context.Background()) c.reverseAccess.Unlock() diff --git a/service/ccm/reverse.go b/service/ccm/reverse.go index ae00df79f..625e55a9d 100644 --- a/service/ccm/reverse.go +++ b/service/ccm/reverse.go @@ -110,7 +110,10 @@ func (s *Service) handleReverseConnect(w http.ResponseWriter, r *http.Request) { return } - receiverCredential.setReverseSession(session) + if !receiverCredential.setReverseSession(session) { + session.Close() + return + } s.logger.Info("reverse connection established for ", receiverCredential.tagName(), " from ", r.RemoteAddr) go func() { diff --git a/service/ocm/credential_external.go b/service/ocm/credential_external.go index 0d6e6b4b1..5c42350d7 100644 --- a/service/ocm/credential_external.go +++ b/service/ocm/credential_external.go @@ -52,6 +52,7 @@ type externalCredential struct { reverse bool reverseSession *yamux.Session reverseAccess sync.RWMutex + closed bool reverseContext context.Context reverseCancel context.CancelFunc connectorDialer N.Dialer @@ -595,12 +596,16 @@ func (c *externalCredential) ocmGetBaseURL() string { } func (c *externalCredential) close() { + var session *yamux.Session c.reverseAccess.Lock() - if c.reverseCancel != nil { - c.reverseCancel() + if !c.closed { + c.closed = true + if c.reverseCancel != nil { + c.reverseCancel() + } + session = c.reverseSession + c.reverseSession = nil } - session := c.reverseSession - c.reverseSession = nil c.reverseAccess.Unlock() if session != nil { session.Close() @@ -620,14 +625,19 @@ func (c *externalCredential) getReverseSession() *yamux.Session { return c.reverseSession } -func (c *externalCredential) setReverseSession(session *yamux.Session) { +func (c *externalCredential) setReverseSession(session *yamux.Session) bool { c.reverseAccess.Lock() + if c.closed { + c.reverseAccess.Unlock() + return false + } old := c.reverseSession c.reverseSession = session c.reverseAccess.Unlock() if old != nil { old.Close() } + return true } func (c *externalCredential) clearReverseSession(session *yamux.Session) { @@ -646,6 +656,10 @@ func (c *externalCredential) getReverseContext() context.Context { func (c *externalCredential) resetReverseContext() { c.reverseAccess.Lock() + if c.closed { + c.reverseAccess.Unlock() + return + } c.reverseCancel() c.reverseContext, c.reverseCancel = context.WithCancel(context.Background()) c.reverseAccess.Unlock() diff --git a/service/ocm/reverse.go b/service/ocm/reverse.go index 25cf017e3..906778df5 100644 --- a/service/ocm/reverse.go +++ b/service/ocm/reverse.go @@ -110,7 +110,10 @@ func (s *Service) handleReverseConnect(w http.ResponseWriter, r *http.Request) { return } - receiverCredential.setReverseSession(session) + if !receiverCredential.setReverseSession(session) { + session.Close() + return + } s.logger.Info("reverse connection established for ", receiverCredential.tagName(), " from ", r.RemoteAddr) go func() {