Răsfoiți Sursa

webapi: close more shutdown holes

Mike 13 ore în urmă
părinte
comite
29530dc73d

+ 9 - 0
server/webapi/handlers/session.go

@@ -224,6 +224,10 @@ func (h *SessionHandler) StartSession(w http.ResponseWriter, r *http.Request) {
 
 	if err = instance.Session().RunOnce(h.FnSessInit(instance)); err != nil {
 		h.Logger.ErrorContext(context.Background(), "failed to init session", "err", err.Error())
+		// Nothing owns the instance yet, so close it here. Left open, it stays
+		// registered in the session manager, counting against the user's
+		// concurrent-instance budget until the server restarts.
+		instance.CloseInstance()
 		h.sendError(w, r, http.StatusInternalServerError, "internal server error")
 		return
 	}
@@ -246,6 +250,7 @@ func (h *SessionHandler) StartSession(w http.ResponseWriter, r *http.Request) {
 
 	if err := h.OServiceService.ClientOnline(ctx, wire.BOS, wire.SNAC_0x01_0x02_OServiceClientOnline{}, instance); err != nil {
 		h.Logger.ErrorContext(ctx, "failed to set client online", "err", err.Error())
+		instance.CloseInstance()
 		h.sendError(w, r, http.StatusInternalServerError, "internal server error")
 		return
 	}
@@ -263,6 +268,10 @@ func (h *SessionHandler) StartSession(w http.ResponseWriter, r *http.Request) {
 	session, err := h.SessionManager.CreateSession(screenName, apiKey.DevID, events, instance, baseURL, h.Logger)
 	if err != nil {
 		h.Logger.ErrorContext(ctx, "failed to create session", "err", err.Error())
+		// CreateSession refuses once the manager is shut down, so this is the
+		// path a startSession racing shutdown takes. The WebAPISession that
+		// would have owned the instance was never created.
+		instance.CloseInstance()
 		h.sendError(w, r, http.StatusInternalServerError, "failed to create session")
 		return
 	}

+ 12 - 2
server/webapi/server.go

@@ -405,10 +405,20 @@ func (s *Server) Shutdown(ctx context.Context) error {
 	s.logger.Debug("Initiating graceful shutdown...")
 	s.shutdownCancel() // stop the session reaper so ListenAndServe's errgroup can drain
 
-	s.sessionManager.Shutdown()
+	var errs []error
+	if err := s.sessionManager.Shutdown(ctx); err != nil {
+		errs = append(errs, fmt.Errorf("draining webapi sessions: %w", err))
+	}
 
 	for _, srv := range s.servers {
-		_ = srv.Shutdown(ctx)
+		if err := srv.Shutdown(ctx); err != nil {
+			errs = append(errs, fmt.Errorf("stopping webapi listener %s: %w", srv.Addr, err))
+		}
+	}
+
+	if err := errors.Join(errs...); err != nil {
+		s.logger.Error("shutdown incomplete", "err", err.Error())
+		return err
 	}
 	s.logger.Info("shutdown complete")
 	return nil

+ 49 - 13
state/webapi_session.go

@@ -88,6 +88,12 @@ type WebAPISession struct {
 	imLogMu      sync.Mutex
 	logger       *slog.Logger // Logger for debugging
 	listeners    sync.WaitGroup
+
+	ctx    context.Context
+	cancel context.CancelFunc
+
+	closeMu sync.Mutex
+	closed  bool
 }
 
 // IsExpired checks if the session has expired.
@@ -143,7 +149,7 @@ func (s *WebAPISession) InvalidateAliases() {
 // every event naming a buddy has to repeat it.
 func (s *WebAPISession) aliasFor(buddy IdentScreenName) string {
 	// Runs on the SNAC listener goroutine, which has no request context.
-	return s.Aliases(context.Background())[buddy.String()]
+	return s.Aliases(s.ctx)[buddy.String()]
 }
 
 // Touch updates the last accessed time and extends expiration if needed.
@@ -172,7 +178,12 @@ func (s *WebAPISession) StartListeningToOSCARSession() {
 		return
 	}
 
-	// Start goroutine to listen for OSCAR messages
+	s.closeMu.Lock()
+	defer s.closeMu.Unlock()
+	if s.closed {
+		return
+	}
+
 	s.listeners.Add(1)
 	go func() {
 		defer s.listeners.Done()
@@ -197,8 +208,18 @@ func (s *WebAPISession) StartListeningToOSCARSession() {
 // the OSCAR instance, and waits for the listener goroutine to unwind. Safe to
 // call more than once.
 func (s *WebAPISession) Close() {
+	s.closeMu.Lock()
+	if s.closed {
+		s.closeMu.Unlock()
+		return
+	}
+	s.closed = true
+	s.closeMu.Unlock()
+
 	s.EventQueue.Close()
 	s.OSCARSession.CloseInstance()
+
+	s.cancel()
 	s.listeners.Wait()
 }
 
@@ -236,7 +257,7 @@ func (s *WebAPISession) handleOServiceMessage(msg wire.SNACMessage) {
 	if s.MyInfoRefresher == nil {
 		return
 	}
-	data, err := s.MyInfoRefresher(context.Background())
+	data, err := s.MyInfoRefresher(s.ctx)
 	if err != nil {
 		s.logger.Error("failed to refresh myInfo after user-info update", "err", err)
 		return
@@ -461,7 +482,7 @@ func (s *WebAPISession) handleFeedbagMessage(msg wire.SNACMessage) {
 		s.InvalidateAliases()
 
 		if s.BuddyListRefresher != nil {
-			groups, err := s.BuddyListRefresher(context.Background())
+			groups, err := s.BuddyListRefresher(s.ctx)
 			if err != nil {
 				s.logger.Error("failed to refresh buddy list after feedbag change", "err", err)
 			} else {
@@ -475,7 +496,7 @@ func (s *WebAPISession) handleFeedbagMessage(msg wire.SNACMessage) {
 					if item.ClassID == wire.FeedbagClassIDPermit ||
 						item.ClassID == wire.FeedbagClassIDDeny ||
 						item.ClassID == wire.FeedbagClassIdPdinfo {
-						pdd, err := s.PermitDenyRefresher(context.Background())
+						pdd, err := s.PermitDenyRefresher(s.ctx)
 						if err != nil {
 							s.logger.Error("failed to refresh permit/deny after feedbag change", "err", err)
 						} else {
@@ -532,7 +553,10 @@ func (m *WebAPISessionManager) CreateSession(screenName DisplayScreenName, devID
 	}
 
 	now := time.Now()
+	sessCtx, sessCancel := context.WithCancel(context.Background())
 	session := &WebAPISession{
+		ctx:             sessCtx,
+		cancel:          sessCancel,
 		AimSID:          aimsid,
 		ScreenName:      screenName,
 		OSCARSession:    oscarSession,
@@ -665,12 +689,13 @@ func (m *WebAPISessionManager) reapExpired() {
 // Shutdown drains and closes all sessions, stops the reaper started by Run, and
 // blocks further CreateSession calls. It does not depend on the caller
 // cancelling Run's context. Safe to call more than once, though only the first
-// call waits for the drain.
-func (m *WebAPISessionManager) Shutdown() {
+// call waits for the drain. The drain is bounded by ctx: Shutdown returns
+// ctx.Err() rather than block forever on a listener that ignores cancellation.
+func (m *WebAPISessionManager) Shutdown(ctx context.Context) error {
 	m.mu.Lock()
 	if m.closed {
 		m.mu.Unlock()
-		return
+		return nil
 	}
 	m.closed = true
 	close(m.stopCh)
@@ -683,12 +708,23 @@ func (m *WebAPISessionManager) Shutdown() {
 	m.sessions = make(map[string]*WebAPISession)
 	m.mu.Unlock()
 
-	// Tear down outside the lock: CloseInstance fans out to buddy-departed
-	// broadcasts and signout, which we don't want to run under m.mu.
-	for _, session := range sessions {
-		session.Close()
+	drained := make(chan struct{})
+	go func() {
+		defer close(drained)
+		// Tear down outside the lock: CloseInstance fans out to buddy-departed
+		// broadcasts and signout, which we don't want to run under m.mu.
+		for _, session := range sessions {
+			session.Close()
+		}
+		m.reaperWG.Wait()
+	}()
+
+	select {
+	case <-drained:
+		return nil
+	case <-ctx.Done():
+		return ctx.Err()
 	}
-	m.reaperWG.Wait()
 }
 
 // generateSessionID creates a cryptographically secure session ID.

+ 56 - 7
state/webapi_session_test.go

@@ -271,10 +271,10 @@ func TestWebAPISession_TempBuddiesIndependence(t *testing.T) {
 func TestWebAPISessionManager_ShutdownIdempotent(t *testing.T) {
 	mgr := NewWebAPISessionManager()
 
-	mgr.Shutdown()
+	_ = mgr.Shutdown(context.Background())
 
 	assert.NotPanics(t, func() {
-		mgr.Shutdown()
+		_ = mgr.Shutdown(context.Background())
 	})
 }
 
@@ -284,7 +284,7 @@ func TestWebAPISessionManager_ShutdownIdempotent(t *testing.T) {
 func TestWebAPISessionManager_CreateAfterShutdown(t *testing.T) {
 	mgr := NewWebAPISessionManager()
 
-	mgr.Shutdown()
+	_ = mgr.Shutdown(context.Background())
 
 	sess, err := mgr.CreateSession(DisplayScreenName("testuser"), "dev", []string{"presence"}, nil, "", nil)
 	assert.Nil(t, sess)
@@ -306,7 +306,7 @@ func TestWebAPISessionManager_ShutdownDrainsAndClosesSessions(t *testing.T) {
 	s2, err := mgr.CreateSession(DisplayScreenName("bob"), "dev", []string{"presence"}, inst2, "", slog.Default())
 	assert.NoError(t, err)
 
-	mgr.Shutdown()
+	mgr.Shutdown(context.Background())
 
 	// Maps drained: the collect loop ran over both sessions.
 	assert.Empty(t, mgr.sessions)
@@ -387,7 +387,7 @@ func TestWebAPISessionManager_ShutdownWithoutReaper(t *testing.T) {
 	done := make(chan struct{})
 	go func() {
 		defer close(done)
-		mgr.Shutdown()
+		mgr.Shutdown(context.Background())
 	}()
 
 	select {
@@ -414,7 +414,7 @@ func TestWebAPISessionManager_ShutdownJoinsReaper(t *testing.T) {
 	done := make(chan struct{})
 	go func() {
 		defer close(done)
-		mgr.Shutdown()
+		mgr.Shutdown(context.Background())
 	}()
 
 	select {
@@ -435,7 +435,7 @@ func TestWebAPISessionManager_ShutdownJoinsReaper(t *testing.T) {
 // with Shutdown never starts, so it cannot reap an already-drained manager.
 func TestWebAPISessionManager_RunAfterShutdown(t *testing.T) {
 	mgr := NewWebAPISessionManager()
-	mgr.Shutdown()
+	mgr.Shutdown(context.Background())
 
 	done := make(chan struct{})
 	go func() {
@@ -843,3 +843,52 @@ func TestWebAPISession_PushesMyInfoOnUserInfoUpdate(t *testing.T) {
 		assert.Equal(t, 0, *refreshes)
 	})
 }
+
+// TestWebAPISessionManager_ShutdownBoundedByContext verifies that Shutdown
+// honors its context instead of blocking indefinitely. A listener goroutine that
+// ignores cancellation must not be able to hold the whole server open: main
+// budgets a few seconds for every server's shutdown combined, so an unbounded
+// wait here means the process never exits.
+func TestWebAPISessionManager_ShutdownBoundedByContext(t *testing.T) {
+	mgr := NewWebAPISessionManager()
+
+	inst := NewSession().AddInstance()
+	sess, err := mgr.CreateSession("alice", "dev", []string{"presence"}, inst, "", slog.Default())
+	assert.NoError(t, err)
+
+	// Stand in for a listener wedged somewhere that never observes cancellation.
+	release := make(chan struct{})
+	defer close(release)
+	sess.listeners.Add(1)
+	go func() {
+		defer sess.listeners.Done()
+		<-release
+	}()
+
+	ctx, cancel := context.WithTimeout(context.Background(), 100*time.Millisecond)
+	defer cancel()
+
+	start := time.Now()
+	err = mgr.Shutdown(ctx)
+	elapsed := time.Since(start)
+
+	assert.ErrorIs(t, err, context.DeadlineExceeded)
+	assert.Less(t, elapsed, 2*time.Second, "Shutdown must give up at its deadline, not wait on the stuck listener")
+}
+
+// TestWebAPISession_CloseCancelsSessionContext verifies that Close cancels the
+// context handed to the refresher callbacks. The listener runs feedbag queries
+// through it, and without cancellation Close's wait lasts as long as the query.
+func TestWebAPISession_CloseCancelsSessionContext(t *testing.T) {
+	mgr := NewWebAPISessionManager()
+
+	inst := NewSession().AddInstance()
+	sess, err := mgr.CreateSession("alice", "dev", []string{"presence"}, inst, "", slog.Default())
+	assert.NoError(t, err)
+
+	assert.NoError(t, sess.ctx.Err(), "session context should be live before Close")
+
+	sess.Close()
+
+	assert.ErrorIs(t, sess.ctx.Err(), context.Canceled)
+}

+ 0 - 39
state/zz_deadlock_probe_test.go

@@ -1,39 +0,0 @@
-package state
-
-import (
-	"testing"
-	"time"
-)
-
-// Demonstrates that ScaleWarningAndRateLimit sends on warningCh while holding
-// the write lock, so a full buffer wedges every reader of the Session.
-func TestProbe_SendUnderWriteLockBlocksReaders(t *testing.T) {
-	s := NewSession()
-	s.SetIdentScreenName(NewIdentScreenName("probe"))
-
-	// 1. Fill the 1-slot warningCh buffer. Nobody is draining it.
-	s.ScaleWarningAndRateLimit(10, 3)
-
-	// 2. Next call blocks on the send -- while holding mutex.Lock().
-	blocked := make(chan struct{})
-	go func() {
-		close(blocked)
-		s.ScaleWarningAndRateLimit(10, 3)
-	}()
-	<-blocked
-	time.Sleep(200 * time.Millisecond) // let it reach the send
-
-	// 3. Any reader now blocks forever. This is icbm.go:719 TLVUserInfo().
-	done := make(chan struct{})
-	go func() {
-		defer close(done)
-		s.DisplayScreenName()
-	}()
-
-	select {
-	case <-done:
-		t.Log("OK: reader acquired RLock")
-	case <-time.After(2 * time.Second):
-		t.Fatal("DEADLOCK: reader blocked on RLock because the writer is parked on a channel send")
-	}
-}