From f5b7f3d35af47325059d72ea0cf16103233ac933 Mon Sep 17 00:00:00 2001 From: mingmxren Date: Sun, 1 Mar 2026 11:58:29 +0800 Subject: [PATCH] fix(session): enforce startnew persistence order and self-heal stale index --- pkg/session/manager.go | 230 +++++++++++++++++++++++++----------- pkg/session/manager_test.go | 127 ++++++++++++++++++++ 2 files changed, 291 insertions(+), 66 deletions(-) diff --git a/pkg/session/manager.go b/pkg/session/manager.go index 8497e74fa..fdce4cffb 100644 --- a/pkg/session/manager.go +++ b/pkg/session/manager.go @@ -103,10 +103,37 @@ func (sm *SessionManager) StartNew(scopeKey string) (string, error) { } newSessionKey := scopeKey + "#" + strconv.Itoa(newOrdinal) + + created := false + if _, ok := sm.sessions[newSessionKey]; !ok { + sm.sessions[newSessionKey] = &Session{ + Key: newSessionKey, + Messages: []providers.Message{}, + Created: now, + Updated: now, + } + created = true + } + if err := sm.saveSessionLocked(newSessionKey); err != nil { + if created { + delete(sm.sessions, newSessionKey) + } + return "", err + } + + prevActive := scope.ActiveSessionKey + prevOrdered := append([]string(nil), scope.OrderedSessions...) + prevUpdated := scope.UpdatedAt scope.ActiveSessionKey = newSessionKey - scope.OrderedSessions = append([]string{newSessionKey}, scope.OrderedSessions...) + scope.OrderedSessions = prependSessionUnique(scope.OrderedSessions, newSessionKey) scope.UpdatedAt = now if err := sm.saveIndexLocked(); err != nil { + scope.ActiveSessionKey = prevActive + scope.OrderedSessions = prevOrdered + scope.UpdatedAt = prevUpdated + if created { + delete(sm.sessions, newSessionKey) + } return "", err } return newSessionKey, nil @@ -379,16 +406,6 @@ func (sm *SessionManager) Save(key string) error { return nil } - filename := sanitizeFilename(key) - - // filepath.IsLocal rejects empty names, "..", absolute paths, and - // OS-reserved device names (NUL, COM1 … on Windows). - // The extra checks reject "." and any directory separators so that - // the session file is always written directly inside sm.storage. - if filename == "." || !filepath.IsLocal(filename) || strings.ContainsAny(filename, `/\`) { - return os.ErrInvalid - } - // Snapshot under read lock, then perform slow file I/O after unlock. sm.mu.RLock() stored, ok := sm.sessions[key] @@ -397,60 +414,10 @@ func (sm *SessionManager) Save(key string) error { return nil } - snapshot := Session{ - Key: stored.Key, - Summary: stored.Summary, - Created: stored.Created, - Updated: stored.Updated, - } - if len(stored.Messages) > 0 { - snapshot.Messages = make([]providers.Message, len(stored.Messages)) - copy(snapshot.Messages, stored.Messages) - } else { - snapshot.Messages = []providers.Message{} - } + snapshot := cloneSession(stored) sm.mu.RUnlock() - data, err := json.MarshalIndent(snapshot, "", " ") - if err != nil { - return err - } - - sessionPath := filepath.Join(sm.storage, filename+".json") - tmpFile, err := os.CreateTemp(sm.storage, "session-*.tmp") - if err != nil { - return err - } - - tmpPath := tmpFile.Name() - cleanup := true - defer func() { - if cleanup { - _ = os.Remove(tmpPath) - } - }() - - if _, err := tmpFile.Write(data); err != nil { - _ = tmpFile.Close() - return err - } - if err := tmpFile.Chmod(0o644); err != nil { - _ = tmpFile.Close() - return err - } - if err := tmpFile.Sync(); err != nil { - _ = tmpFile.Close() - return err - } - if err := tmpFile.Close(); err != nil { - return err - } - - if err := os.Rename(tmpPath, sessionPath); err != nil { - return err - } - cleanup = false - return nil + return sm.writeSessionSnapshot(snapshot) } func (sm *SessionManager) loadSessions() error { @@ -515,20 +482,53 @@ func (sm *SessionManager) loadIndex() error { loaded.Scopes = make(map[string]*scopeIndex) } + changed := false for scopeKey, scope := range loaded.Scopes { if scope == nil { delete(loaded.Scopes, scopeKey) + changed = true continue } - if len(scope.OrderedSessions) == 0 { - scope.OrderedSessions = []string{scopeKey} + + filtered := make([]string, 0, len(scope.OrderedSessions)) + seen := make(map[string]struct{}, len(scope.OrderedSessions)) + for _, sessionKey := range scope.OrderedSessions { + if sessionKey == "" { + changed = true + continue + } + if _, exists := sm.sessions[sessionKey]; !exists { + changed = true + continue + } + if _, dup := seen[sessionKey]; dup { + changed = true + continue + } + seen[sessionKey] = struct{}{} + filtered = append(filtered, sessionKey) } - if scope.ActiveSessionKey == "" { + + if len(filtered) == 0 { + delete(loaded.Scopes, scopeKey) + changed = true + continue + } + + if len(filtered) != len(scope.OrderedSessions) { + changed = true + } + scope.OrderedSessions = filtered + if _, ok := seen[scope.ActiveSessionKey]; !ok { scope.ActiveSessionKey = scope.OrderedSessions[0] + changed = true } } sm.index = loaded + if changed { + return sm.saveIndexLocked() + } return nil } @@ -612,6 +612,104 @@ func (sm *SessionManager) ensureScopeLocked(scopeKey string, now time.Time) (*sc return scope, changed } +func prependSessionUnique(ordered []string, sessionKey string) []string { + next := make([]string, 0, len(ordered)+1) + next = append(next, sessionKey) + for _, existing := range ordered { + if existing == sessionKey { + continue + } + next = append(next, existing) + } + return next +} + +func cloneSession(stored *Session) Session { + snapshot := Session{ + Key: stored.Key, + Summary: stored.Summary, + Created: stored.Created, + Updated: stored.Updated, + } + if len(stored.Messages) > 0 { + snapshot.Messages = make([]providers.Message, len(stored.Messages)) + copy(snapshot.Messages, stored.Messages) + } else { + snapshot.Messages = []providers.Message{} + } + return snapshot +} + +func (sm *SessionManager) saveSessionLocked(key string) error { + if sm.storage == "" { + return nil + } + + stored, ok := sm.sessions[key] + if !ok { + return fmt.Errorf("session %q not found", key) + } + + return sm.writeSessionSnapshot(cloneSession(stored)) +} + +func (sm *SessionManager) writeSessionSnapshot(snapshot Session) error { + if sm.storage == "" { + return nil + } + + filename := sanitizeFilename(snapshot.Key) + + // filepath.IsLocal rejects empty names, "..", absolute paths, and + // OS-reserved device names (NUL, COM1 … on Windows). + // The extra checks reject "." and any directory separators so that + // the session file is always written directly inside sm.storage. + if filename == "." || !filepath.IsLocal(filename) || strings.ContainsAny(filename, `/\`) { + return os.ErrInvalid + } + + data, err := json.MarshalIndent(snapshot, "", " ") + if err != nil { + return err + } + + sessionPath := filepath.Join(sm.storage, filename+".json") + tmpFile, err := os.CreateTemp(sm.storage, "session-*.tmp") + if err != nil { + return err + } + + tmpPath := tmpFile.Name() + cleanup := true + defer func() { + if cleanup { + _ = os.Remove(tmpPath) + } + }() + + if _, err := tmpFile.Write(data); err != nil { + _ = tmpFile.Close() + return err + } + if err := tmpFile.Chmod(0o644); err != nil { + _ = tmpFile.Close() + return err + } + if err := tmpFile.Sync(); err != nil { + _ = tmpFile.Close() + return err + } + if err := tmpFile.Close(); err != nil { + return err + } + + if err := os.Rename(tmpPath, sessionPath); err != nil { + return err + } + cleanup = false + return nil +} + func (sm *SessionManager) deleteSessionFile(sessionKey string) error { if sm.storage == "" { return nil diff --git a/pkg/session/manager_test.go b/pkg/session/manager_test.go index ade3b1e25..9df61cae8 100644 --- a/pkg/session/manager_test.go +++ b/pkg/session/manager_test.go @@ -133,6 +133,48 @@ func TestStartNew_CreatesMonotonicSessionKeys(t *testing.T) { } } +func TestStartNew_PersistsSessionFileWithoutManualSave(t *testing.T) { + dir := t.TempDir() + sm := NewSessionManager(dir) + scope := "agent:main:telegram:direct:user1" + + if _, err := sm.ResolveActive(scope); err != nil { + t.Fatal(err) + } + + s2, err := sm.StartNew(scope) + if err != nil { + t.Fatal(err) + } + + sessionPath := filepath.Join(dir, sanitizeFilename(s2)+".json") + if _, err := os.Stat(sessionPath); err != nil { + t.Fatalf("expected %s to exist: %v", sessionPath, err) + } + + indexPath := filepath.Join(dir, sessionIndexFilename) + raw, err := os.ReadFile(indexPath) + if err != nil { + t.Fatalf("read index: %v", err) + } + + var idx sessionIndex + if err := json.Unmarshal(raw, &idx); err != nil { + t.Fatalf("unmarshal index: %v", err) + } + + scoped := idx.Scopes[scope] + if scoped == nil { + t.Fatalf("expected scope %q in index", scope) + } + if scoped.ActiveSessionKey != s2 { + t.Fatalf("active=%q, want %q", scoped.ActiveSessionKey, s2) + } + if len(scoped.OrderedSessions) == 0 || scoped.OrderedSessions[0] != s2 { + t.Fatalf("ordered_sessions=%v, want first=%q", scoped.OrderedSessions, s2) + } +} + func TestListAndResume_ByScopeOrdinal(t *testing.T) { sm := NewSessionManager(t.TempDir()) scope := "agent:main:telegram:direct:user1" @@ -229,3 +271,88 @@ func TestPrune_RemovesOldestFromMemoryAndDisk(t *testing.T) { t.Fatalf("list order after prune = %+v", list) } } + +func TestLoadIndex_SelfHealsStaleReferences(t *testing.T) { + dir := t.TempDir() + scopeA := "agent:main:telegram:direct:user1" + scopeB := "agent:main:telegram:direct:user2" + validNewest := scopeA + "#3" + validOlder := scopeA + "#2" + + seed := NewSessionManager(dir) + seed.AddMessage(validNewest, "user", "hello") + seed.AddMessage(validOlder, "user", "hello") + if err := seed.Save(validNewest); err != nil { + t.Fatalf("save %q: %v", validNewest, err) + } + if err := seed.Save(validOlder); err != nil { + t.Fatalf("save %q: %v", validOlder, err) + } + + stale := sessionIndex{ + Version: 1, + Scopes: map[string]*scopeIndex{ + scopeA: { + ActiveSessionKey: scopeA + "#999", + OrderedSessions: []string{ + validNewest, + validNewest, + scopeA + "#404", + validOlder, + }, + }, + scopeB: { + ActiveSessionKey: scopeB, + OrderedSessions: []string{scopeB}, + }, + }, + } + + raw, err := json.MarshalIndent(stale, "", " ") + if err != nil { + t.Fatalf("marshal stale index: %v", err) + } + if err := os.WriteFile(filepath.Join(dir, sessionIndexFilename), raw, 0o644); err != nil { + t.Fatalf("write stale index: %v", err) + } + + reloaded := NewSessionManager(dir) + + indexRaw, err := os.ReadFile(filepath.Join(dir, sessionIndexFilename)) + if err != nil { + t.Fatalf("read healed index: %v", err) + } + + var healed sessionIndex + if err := json.Unmarshal(indexRaw, &healed); err != nil { + t.Fatalf("unmarshal healed index: %v", err) + } + + scopeAHealed := healed.Scopes[scopeA] + if scopeAHealed == nil { + t.Fatalf("expected scope %q in healed index", scopeA) + } + if scopeAHealed.ActiveSessionKey != validNewest { + t.Fatalf("active=%q, want %q", scopeAHealed.ActiveSessionKey, validNewest) + } + if len(scopeAHealed.OrderedSessions) != 2 { + t.Fatalf("ordered_sessions=%v, want len=2", scopeAHealed.OrderedSessions) + } + if scopeAHealed.OrderedSessions[0] != validNewest || scopeAHealed.OrderedSessions[1] != validOlder { + t.Fatalf("ordered_sessions=%v, want [%s %s]", scopeAHealed.OrderedSessions, validNewest, validOlder) + } + if _, exists := healed.Scopes[scopeB]; exists { + t.Fatalf("expected stale scope %q removed", scopeB) + } + + list, err := reloaded.List(scopeA) + if err != nil { + t.Fatal(err) + } + if len(list) != 2 { + t.Fatalf("len(list)=%d, want 2", len(list)) + } + if !list[0].Active || list[0].SessionKey != validNewest { + t.Fatalf("list[0]=%+v, want active newest", list[0]) + } +}