fix(session): enforce startnew persistence order and self-heal stale index

This commit is contained in:
mingmxren 2026-03-01 11:58:29 +08:00
parent 34916332d8
commit f5b7f3d35a
2 changed files with 291 additions and 66 deletions

View file

@ -103,10 +103,37 @@ func (sm *SessionManager) StartNew(scopeKey string) (string, error) {
} }
newSessionKey := scopeKey + "#" + strconv.Itoa(newOrdinal) 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.ActiveSessionKey = newSessionKey
scope.OrderedSessions = append([]string{newSessionKey}, scope.OrderedSessions...) scope.OrderedSessions = prependSessionUnique(scope.OrderedSessions, newSessionKey)
scope.UpdatedAt = now scope.UpdatedAt = now
if err := sm.saveIndexLocked(); err != nil { if err := sm.saveIndexLocked(); err != nil {
scope.ActiveSessionKey = prevActive
scope.OrderedSessions = prevOrdered
scope.UpdatedAt = prevUpdated
if created {
delete(sm.sessions, newSessionKey)
}
return "", err return "", err
} }
return newSessionKey, nil return newSessionKey, nil
@ -379,16 +406,6 @@ func (sm *SessionManager) Save(key string) error {
return nil 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. // Snapshot under read lock, then perform slow file I/O after unlock.
sm.mu.RLock() sm.mu.RLock()
stored, ok := sm.sessions[key] stored, ok := sm.sessions[key]
@ -397,60 +414,10 @@ func (sm *SessionManager) Save(key string) error {
return nil return nil
} }
snapshot := Session{ snapshot := cloneSession(stored)
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{}
}
sm.mu.RUnlock() sm.mu.RUnlock()
data, err := json.MarshalIndent(snapshot, "", " ") return sm.writeSessionSnapshot(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) loadSessions() error { func (sm *SessionManager) loadSessions() error {
@ -515,20 +482,53 @@ func (sm *SessionManager) loadIndex() error {
loaded.Scopes = make(map[string]*scopeIndex) loaded.Scopes = make(map[string]*scopeIndex)
} }
changed := false
for scopeKey, scope := range loaded.Scopes { for scopeKey, scope := range loaded.Scopes {
if scope == nil { if scope == nil {
delete(loaded.Scopes, scopeKey) delete(loaded.Scopes, scopeKey)
changed = true
continue 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 scope.ActiveSessionKey == "" { 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 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] scope.ActiveSessionKey = scope.OrderedSessions[0]
changed = true
} }
} }
sm.index = loaded sm.index = loaded
if changed {
return sm.saveIndexLocked()
}
return nil return nil
} }
@ -612,6 +612,104 @@ func (sm *SessionManager) ensureScopeLocked(scopeKey string, now time.Time) (*sc
return scope, changed 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 { func (sm *SessionManager) deleteSessionFile(sessionKey string) error {
if sm.storage == "" { if sm.storage == "" {
return nil return nil

View file

@ -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) { func TestListAndResume_ByScopeOrdinal(t *testing.T) {
sm := NewSessionManager(t.TempDir()) sm := NewSessionManager(t.TempDir())
scope := "agent:main:telegram:direct:user1" scope := "agent:main:telegram:direct:user1"
@ -229,3 +271,88 @@ func TestPrune_RemovesOldestFromMemoryAndDisk(t *testing.T) {
t.Fatalf("list order after prune = %+v", list) 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])
}
}