返回 DeepSeek-Reasonix
save_dag_test.go
根目录 / internal / agent / save_dag_test.go
1 package agent
2
3 import (
4 "errors"
5 "os"
6 "path/filepath"
7 "strings"
8 "sync"
9 "testing"
10
11 "reasonix/internal/fileutil"
12 "reasonix/internal/provider"
13 "reasonix/internal/store"
14 )
15
16 func dagSavedSession(t *testing.T, path string, contents ...string) *Session {
17 t.Helper()
18 s := NewSession("sys")
19 for i, c := range contents {
20 role := provider.RoleUser
21 if i%2 == 1 {
22 role = provider.RoleAssistant
23 }
24 s.Add(provider.Message{Role: role, Content: c})
25 }
26 if err := s.Save(path); err != nil {
27 t.Fatalf("save: %v", err)
28 }
29 return s
30 }
31
32 func dagEntryTypes(t *testing.T, path string) []string {
33 t.Helper()
34 b, err := os.ReadFile(store.SessionEventLog(path))
35 if err != nil {
36 t.Fatal(err)
37 }
38 var types []string
39 for line := range strings.SplitSeq(strings.TrimSpace(string(b)), "\n") {
40 _, rest, _ := strings.Cut(line, `"type":"`)
41 typ, _, _ := strings.Cut(rest, `"`)
42 types = append(types, typ)
43 }
44 return types
45 }
46
47 func assertNoTranscriptCopies(t *testing.T, path string) {
48 t.Helper()
49 entries, _ := os.ReadDir(filepath.Dir(path))
50 for _, entry := range entries {
51 if store.IsSessionTranscriptName(entry.Name()) && entry.Name() != filepath.Base(path) {
52 t.Fatalf("unexpected transcript copy %s", entry.Name())
53 }
54 }
55 }
56
57 func TestDAGSaveCreatesSchemaTwoLogAndAppendsDelta(t *testing.T) {
58 path := dagTestSession(t)
59 s := dagSavedSession(t, path, "q1", "a1")
60 probe, err := probeSessionEventLog(path)
61 if err != nil || !probe.dag {
62 t.Fatalf("probe = %+v err=%v", probe, err)
63 }
64 if got := dagEntryTypes(t, path); strings.Join(got, ",") != "log,writer,message,message,message" {
65 t.Fatalf("entries = %v", got)
66 }
67 ref, ok := s.Head()
68 if !ok || ref.HeadID != SessionMainHead || ref.LeafID != s.LeafID() || ref.LogGeneration != 1 {
69 t.Fatalf("head = %+v ok=%v", ref, ok)
70 }
71 if b, err := os.ReadFile(path); err != nil || strings.Count(string(b), "\n") != 3 {
72 t.Fatalf("checkpoint cache: %v %q", err, b)
73 }
74 idx, err := ReadSessionHeadIndex(path)
75 if err != nil || idx == nil || !idx.Current(path) || idx.MessageCount != 3 || idx.SelectedHead != SessionMainHead {
76 t.Fatalf("index = %+v err=%v", idx, err)
77 }
78 meta, _, err := LoadBranchMeta(path)
79 if err != nil || meta.HeadID != SessionMainHead || meta.LogSchema != 2 || meta.HeadCount != 1 || meta.Revision == 0 {
80 t.Fatalf("meta = %+v err=%v", meta, err)
81 }
82 s.Add(provider.Message{Role: provider.RoleUser, Content: "q2"})
83 if err := s.Save(path); err != nil {
84 t.Fatal(err)
85 }
86 if got := dagEntryTypes(t, path); strings.Join(got, ",") != "log,writer,message,message,message,message" {
87 t.Fatalf("entries after append = %v", got)
88 }
89 if err := s.Save(path); err != nil {
90 t.Fatal(err)
91 }
92 if got := len(dagEntryTypes(t, path)); got != 6 {
93 t.Fatalf("no-op save appended: %d entries", got)
94 }
95 loaded, err := LoadSession(path)
96 if err != nil || len(loaded.Messages) != 4 || loaded.LeafID() != s.LeafID() {
97 t.Fatalf("reload: err=%v len=%d", err, len(loaded.Messages))
98 }
99 assertNoTranscriptCopies(t, path)
100 }
101
102 func TestDAGSaveDisabledByEnvKeepsSchemaOne(t *testing.T) {
103 useSchemaOneLog(t)
104 path := dagTestSession(t)
105 dagSavedSession(t, path, "q1")
106 probe, err := probeSessionEventLog(path)
107 if err != nil || probe.dag || !probe.native {
108 t.Fatalf("probe = %+v err=%v", probe, err)
109 }
110 }
111
112 func TestLoadedLegacyCheckpointStaysOnLegacyWriter(t *testing.T) {
113 path := filepath.Join(t.TempDir(), "legacy.jsonl")
114 body := `{"role":"user","content":"old question"}` + "\n"
115 if err := os.WriteFile(path, []byte(body), 0o600); err != nil {
116 t.Fatal(err)
117 }
118 s, err := LoadSession(path)
119 if err != nil {
120 t.Fatal(err)
121 }
122 s.Add(provider.Message{Role: provider.RoleAssistant, Content: "continued answer"})
123 if err := s.Save(path); err != nil {
124 t.Fatal(err)
125 }
126 probe, err := probeSessionEventLog(path)
127 if err != nil {
128 t.Fatal(err)
129 }
130 if probe.dag || probe.size != 0 {
131 t.Fatalf("legacy save changed format: %+v", probe)
132 }
133 if _, err := os.Stat(store.SessionEventLog(path)); !os.IsNotExist(err) {
134 t.Fatalf("checkpoint-only continuation created a duplicate log: %v", err)
135 }
136 loaded, err := LoadSession(path)
137 if err != nil || len(loaded.Messages) != 2 || loaded.Messages[1].Content != "continued answer" {
138 t.Fatalf("continued legacy session = %+v, err=%v", loaded, err)
139 }
140 }
141
142 func TestDAGSaveLocalMetadataBecomesPatch(t *testing.T) {
143 path := dagTestSession(t)
144 s := dagSavedSession(t, path, "q1", "a1")
145 msgs := s.Snapshot()
146 msgs[1].Edited = true
147 msgs[1].WorkDurationMs = 42
148 s.ReplaceLocalMetadata(msgs)
149 if err := s.SaveRewrite(path); err != nil {
150 t.Fatal(err)
151 }
152 types := dagEntryTypes(t, path)
153 if types[len(types)-1] != sessionDAGTypePatch {
154 t.Fatalf("entries = %v", types)
155 }
156 loaded, err := LoadSession(path)
157 if err != nil || !loaded.Messages[1].Edited || loaded.Messages[1].WorkDurationMs != 42 || loaded.Messages[1].ID != msgs[1].ID {
158 t.Fatalf("reload = %+v err=%v", loaded.Messages[1], err)
159 }
160 if reasons := s.DrainContentRewriteReasons(); len(reasons) != 0 {
161 t.Fatalf("local metadata save queued cache reasons %v", reasons)
162 }
163 }
164
165 func TestDAGSaveSystemPromptRefreshKeepsLaterIDs(t *testing.T) {
166 path := dagTestSession(t)
167 s := dagSavedSession(t, path, "q1", "a1")
168 before := s.Snapshot()
169 s.SetLeadingSystemPrompt("sys-v2")
170 if err := s.SaveRewrite(path); err != nil {
171 t.Fatal(err)
172 }
173 types := dagEntryTypes(t, path)
174 if types[len(types)-1] != sessionDAGTypeSystem {
175 t.Fatalf("entries = %v", types)
176 }
177 loaded, err := LoadSession(path)
178 if err != nil || loaded.Messages[0].Content != "sys-v2" {
179 t.Fatalf("reload = %+v err=%v", loaded.Messages, err)
180 }
181 for i := range before {
182 if loaded.Messages[i].ID != before[i].ID {
183 t.Fatalf("message %d id changed across system refresh", i)
184 }
185 }
186 }
187
188 func TestDAGSaveTruncationRewindsWithoutErasingBytes(t *testing.T) {
189 path := dagTestSession(t)
190 s := dagSavedSession(t, path, "q1", "a1", "q2", "a2")
191 logBefore, _ := os.ReadFile(store.SessionEventLog(path))
192 msgs := s.Snapshot()
193 s.Rewrite(msgs[:3], "rewind_truncate")
194 if err := s.SaveRewrite(path); err != nil {
195 t.Fatal(err)
196 }
197 logAfter, _ := os.ReadFile(store.SessionEventLog(path))
198 if !strings.HasPrefix(string(logAfter), string(logBefore)) {
199 t.Fatal("rewind must not rewrite earlier bytes")
200 }
201 types := dagEntryTypes(t, path)
202 if types[len(types)-1] != sessionDAGTypeRewind {
203 t.Fatalf("entries = %v", types)
204 }
205 loaded, err := LoadSession(path)
206 if err != nil || len(loaded.Messages) != 3 || loaded.LeafID() != msgs[2].ID {
207 t.Fatalf("reload len=%d leaf=%q err=%v", len(loaded.Messages), loaded.LeafID(), err)
208 }
209 s.Add(provider.Message{Role: provider.RoleAssistant, Content: "a2-new"})
210 if err := s.Save(path); err != nil {
211 t.Fatal(err)
212 }
213 loaded, err = LoadSession(path)
214 if err != nil || strings.Join(dagContents(loaded.Messages), ",") != "sys,q1,a1,a2-new" {
215 t.Fatalf("after re-append: %v err=%v", dagContents(loaded.Messages), err)
216 }
217 }
218
219 func TestLoadedSchemaOneLogStaysNativeEvenUnderLease(t *testing.T) {
220 t.Setenv(SessionLogSchemaEnv, "v1")
221 path := dagTestSession(t)
222 v1 := dagSavedSession(t, path, "q1", "a1")
223 if err := os.Unsetenv(SessionLogSchemaEnv); err != nil {
224 t.Fatal(err)
225 }
226 loaded, err := LoadSession(path)
227 if err != nil {
228 t.Fatal(err)
229 }
230 loaded.Add(provider.Message{Role: provider.RoleUser, Content: "q2"})
231 if err := loaded.Save(path); err != nil {
232 t.Fatal(err)
233 }
234 if probe, _ := probeSessionEventLog(path); probe.dag {
235 t.Fatal("unleased writer must not upgrade an existing schema-1 log")
236 }
237 lease, err := TryAcquireSessionLease(path)
238 if err != nil {
239 t.Fatal(err)
240 }
241 defer lease.Release()
242 loaded.Add(provider.Message{Role: provider.RoleAssistant, Content: "a2"})
243 if err := loaded.Save(path); err != nil {
244 t.Fatal(err)
245 }
246 probe, _ := probeSessionEventLog(path)
247 if probe.dag {
248 t.Fatal("a lease must not implicitly upgrade an existing schema-1 log")
249 }
250 again, err := LoadSession(path)
251 if err != nil || strings.Join(dagContents(again.Messages), ",") != "sys,q1,a1,q2,a2" {
252 t.Fatalf("after native append: %v err=%v", dagContents(again.Messages), err)
253 }
254 for i := range v1.Messages {
255 if again.Messages[i].ID != loaded.Messages[i].ID {
256 t.Fatalf("message %d id changed across native append", i)
257 }
258 }
259 }
260
261 func TestDAGSaveConcurrentWritersForkInsteadOfConflicting(t *testing.T) {
262 path := dagTestSession(t)
263 a := dagSavedSession(t, path, "q1", "a1")
264 b, err := LoadSession(path)
265 if err != nil {
266 t.Fatal(err)
267 }
268 a.Add(provider.Message{Role: provider.RoleUser, Content: "q2-from-a"})
269 if err := a.Save(path); err != nil {
270 t.Fatal(err)
271 }
272 b.Add(provider.Message{Role: provider.RoleUser, Content: "q2-from-b"})
273 if err := b.Save(path); err != nil {
274 t.Fatalf("second writer must not conflict: %v", err)
275 }
276 refA, _ := a.Head()
277 refB, _ := b.Head()
278 if refA.HeadID != SessionMainHead || refB.HeadID == SessionMainHead || refB.HeadID == "" {
279 t.Fatalf("heads a=%+v b=%+v", refA, refB)
280 }
281 events := b.DrainHeadEvents()
282 if len(events) != 1 || events[0].Kind != HeadEventForkedConcurrent || events[0].HeadID != refB.HeadID {
283 t.Fatalf("events = %+v", events)
284 }
285 heads, err := ListSessionHeads(path)
286 if err != nil || len(heads) != 2 || heads[1].Kind != HeadKindConcurrent || heads[1].MessageCount != 4 || heads[0].MessageCount != 4 {
287 t.Fatalf("heads = %+v err=%v", heads, err)
288 }
289 st := dagReplay(t, path)
290 if got := dagChain(st, SessionMainHead); strings.Join(got, ",") != "sys,q1,a1,q2-from-a" {
291 t.Fatalf("main chain %v", got)
292 }
293 if got := dagChain(st, refB.HeadID); strings.Join(got, ",") != "sys,q1,a1,q2-from-b" {
294 t.Fatalf("fork chain %v", got)
295 }
296 assertNoTranscriptCopies(t, path)
297 // Each writer keeps extending its own head afterwards.
298 a.Add(provider.Message{Role: provider.RoleAssistant, Content: "a2-from-a"})
299 b.Add(provider.Message{Role: provider.RoleAssistant, Content: "a2-from-b"})
300 if err := a.Save(path); err != nil {
301 t.Fatal(err)
302 }
303 if err := b.Save(path); err != nil {
304 t.Fatal(err)
305 }
306 if len(b.DrainHeadEvents()) != 0 {
307 t.Fatal("continuing on the fork must not fork again")
308 }
309 st = dagReplay(t, path)
310 if len(st.heads) != 2 || len(dagChain(st, SessionMainHead)) != 5 || len(dagChain(st, refB.HeadID)) != 5 {
311 t.Fatalf("heads=%d main=%d fork=%d", len(st.heads), len(dagChain(st, SessionMainHead)), len(dagChain(st, refB.HeadID)))
312 }
313 }
314
315 func TestDAGSaveBehindDiskReportsStalePrefix(t *testing.T) {
316 path := dagTestSession(t)
317 a := dagSavedSession(t, path, "q1", "a1")
318 b, err := LoadSession(path)
319 if err != nil {
320 t.Fatal(err)
321 }
322 a.Add(provider.Message{Role: provider.RoleUser, Content: "q2"})
323 if err := a.Save(path); err != nil {
324 t.Fatal(err)
325 }
326 err = b.Save(path)
327 if !errors.Is(err, ErrSessionSnapshotConflict) {
328 t.Fatalf("behind writer err = %v, want stale prefix conflict", err)
329 }
330 if kind, ok := SnapshotConflictKind(err); !ok || kind != SessionSnapshotConflictStalePrefix {
331 t.Fatalf("kind = %q ok=%v", kind, ok)
332 }
333 if st := dagReplay(t, path); len(st.heads) != 1 || len(dagChain(st, SessionMainHead)) != 4 {
334 t.Fatal("a behind writer must not append or fork")
335 }
336 assertNoTranscriptCopies(t, path)
337 }
338
339 func TestDAGSaveRedactionCompactErasesBytesUnderLease(t *testing.T) {
340 path := dagTestSession(t)
341 lease, err := TryAcquireSessionLease(path)
342 if err != nil {
343 t.Fatal(err)
344 }
345 defer lease.Release()
346 s := dagSavedSession(t, path, "q1 secret-token", "a1")
347 msgs := s.Snapshot()
348 ids := []string{msgs[0].ID, msgs[1].ID, msgs[2].ID}
349 msgs[1].Content = "q1 [redacted]"
350 s.Rewrite(msgs, "redact")
351 if err := s.SaveRewriteCompact(path); err != nil {
352 t.Fatal(err)
353 }
354 raw, _ := os.ReadFile(store.SessionEventLog(path))
355 if strings.Contains(string(raw), "secret-token") {
356 t.Fatal("redaction left the secret in the log")
357 }
358 st := dagReplay(t, path)
359 if st.generation != 2 {
360 t.Fatalf("generation = %d, want rotation", st.generation)
361 }
362 loaded, err := LoadSession(path)
363 if err != nil || loaded.Messages[1].Content != "q1 [redacted]" {
364 t.Fatalf("reload = %+v err=%v", loaded.Messages, err)
365 }
366 for i, id := range ids {
367 if loaded.Messages[i].ID != id {
368 t.Fatalf("message %d id changed by redaction", i)
369 }
370 }
371 if ref, _ := s.Head(); ref.LogGeneration != 2 {
372 t.Fatalf("session did not follow the rotation: %+v", ref)
373 }
374 }
375
376 func TestDAGSaveOversizeLogRotatesUnderLease(t *testing.T) {
377 path := dagTestSession(t)
378 lease, err := TryAcquireSessionLease(path)
379 if err != nil {
380 t.Fatal(err)
381 }
382 defer lease.Release()
383 big := strings.Repeat("x", 100<<10)
384 s := dagSavedSession(t, path, big+"1", big+"2", big+"3", big+"4", big+"5", big+"6")
385 msgs := s.Snapshot()
386 s.Rewrite(msgs[:2], "rewind_truncate")
387 if err := s.SaveRewrite(path); err != nil {
388 t.Fatal(err)
389 }
390 st := dagReplay(t, path)
391 if st.generation != 2 || len(st.nodes) != 2 {
392 t.Fatalf("generation=%d nodes=%d, want rotated log with only the live chain", st.generation, len(st.nodes))
393 }
394 if info, _ := os.Stat(store.SessionEventLog(path)); info.Size() > int64(len(big))*3 {
395 t.Fatalf("rotated log still %d bytes", info.Size())
396 }
397 loaded, err := LoadSession(path)
398 if err != nil || len(loaded.Messages) != 2 {
399 t.Fatalf("reload len=%d err=%v", len(loaded.Messages), err)
400 }
401 }
402
403 func TestDAGSaveCrashAtAppendRecoversOnNextSave(t *testing.T) {
404 path := dagTestSession(t)
405 s := dagSavedSession(t, path, "q1")
406 s.Add(provider.Message{Role: provider.RoleAssistant, Content: "a1"})
407 fileutil.CrashPoint = func(op, _ string) {
408 if op == "dag-append" {
409 panic("crash:dag-append")
410 }
411 }
412 func() {
413 defer func() {
414 if recover() == nil {
415 t.Fatal("crash point did not fire")
416 }
417 }()
418 _ = s.Save(path)
419 }()
420 fileutil.CrashPoint = nil
421 // The crash happened inside the locked save; a later save from the same
422 // session must still land exactly one copy of the message.
423 if err := s.Save(path); err != nil {
424 t.Fatalf("save after crash: %v", err)
425 }
426 loaded, err := LoadSession(path)
427 if err != nil || strings.Join(dagContents(loaded.Messages), ",") != "sys,q1,a1" {
428 t.Fatalf("after crash: %v err=%v", dagContents(loaded.Messages), err)
429 }
430 }
431
432 func TestDAGSaveConcurrentGoroutinesExtendTheirOwnHeads(t *testing.T) {
433 path := dagTestSession(t)
434 a := dagSavedSession(t, path, "q1", "a1")
435 b, err := LoadSession(path)
436 if err != nil {
437 t.Fatal(err)
438 }
439 const rounds = 15
440 var wg sync.WaitGroup
441 run := func(s *Session, tag string) {
442 defer wg.Done()
443 for i := range rounds {
444 s.Add(provider.Message{Role: provider.RoleUser, Content: tag + string(rune('a'+i))})
445 if err := s.Save(path); err != nil {
446 t.Errorf("%s save %d: %v", tag, i, err)
447 return
448 }
449 }
450 }
451 wg.Add(2)
452 go run(a, "A")
453 go run(b, "B")
454 wg.Wait()
455 st := dagReplay(t, path)
456 if st.damaged || len(st.heads) != 2 {
457 t.Fatalf("damaged=%v heads=%d", st.damaged, len(st.heads))
458 }
459 for _, id := range st.headOrder {
460 if got := len(dagChain(st, id)); got != 3+rounds {
461 t.Fatalf("head %s chain length %d", id, got)
462 }
463 }
464 assertNoTranscriptCopies(t, path)
465 }
466
467 func TestExportSessionSchemaOneWritesReadableSchemaOneSession(t *testing.T) {
468 path := dagTestSession(t)
469 s := dagSavedSession(t, path, "q1", "a1")
470 dst := filepath.Join(t.TempDir(), "export.jsonl")
471 if err := ExportSessionSchemaOne(path, dst); err != nil {
472 t.Fatal(err)
473 }
474 probe, err := probeSessionEventLog(dst)
475 if err != nil || !probe.native || probe.dag {
476 t.Fatalf("export probe = %+v err=%v", probe, err)
477 }
478 exported, err := LoadSession(dst)
479 if err != nil || strings.Join(dagContents(exported.Messages), ",") != strings.Join(dagContents(s.Messages), ",") {
480 t.Fatalf("export reload = %v err=%v", dagContents(exported.Messages), err)
481 }
482 if _, ok := exported.Head(); ok {
483 t.Fatal("exported session must be schema 1")
484 }
485 if err := ExportSessionSchemaOne(path, dst); err == nil {
486 t.Fatal("export must refuse to overwrite an existing destination")
487 }
488 }
489
490 func TestDAGSaveIndependentIdenticalTranscriptsConverge(t *testing.T) {
491 path := dagTestSession(t)
492 a := dagSavedSession(t, path, "q1", "a1")
493 b := NewSession("sys")
494 b.Add(provider.Message{Role: provider.RoleUser, Content: "q1"})
495 b.Add(provider.Message{Role: provider.RoleAssistant, Content: "a1"})
496 b.Add(provider.Message{Role: provider.RoleUser, Content: "q2"})
497 if err := b.Save(path); err != nil {
498 t.Fatal(err)
499 }
500 st := dagReplay(t, path)
501 if len(st.heads) != 1 || strings.Join(dagChain(st, SessionMainHead), ",") != "sys,q1,a1,q2" {
502 t.Fatalf("identical prefix must extend main: heads=%d chain=%v", len(st.heads), dagChain(st, SessionMainHead))
503 }
504 for i := range a.Messages {
505 if b.Messages[i].ID != a.Messages[i].ID {
506 t.Fatalf("message %d: independent writer did not adopt the persisted id", i)
507 }
508 }
509 if ref, _ := b.Head(); ref.HeadID != SessionMainHead || ref.LeafID != b.LeafID() {
510 t.Fatalf("b head = %+v", ref)
511 }
512 }
513
514 func TestDAGSaveUnrelatedWriterForksInsteadOfRewinding(t *testing.T) {
515 path := dagTestSession(t)
516 dagSavedSession(t, path, "q1", "a1", "q2", "a2")
517 b := NewSession("sys")
518 b.Add(provider.Message{Role: provider.RoleUser, Content: "q1"})
519 b.Add(provider.Message{Role: provider.RoleAssistant, Content: "a1"})
520 b.Add(provider.Message{Role: provider.RoleUser, Content: "q2-other"})
521 if err := b.Save(path); err != nil {
522 t.Fatal(err)
523 }
524 st := dagReplay(t, path)
525 ref, _ := b.Head()
526 if len(st.heads) != 2 || ref.HeadID == SessionMainHead {
527 t.Fatalf("unrelated writer must fork: heads=%d ref=%+v", len(st.heads), ref)
528 }
529 if got := dagChain(st, SessionMainHead); strings.Join(got, ",") != "sys,q1,a1,q2,a2" {
530 t.Fatalf("main was rewritten by an unrelated writer: %v", got)
531 }
532 if got := dagChain(st, ref.HeadID); strings.Join(got, ",") != "sys,q1,a1,q2-other" {
533 t.Fatalf("fork chain %v", got)
534 }
535 }
536
536 lines GO