返回 DeepSeek-Reasonix
session_display_pager_event_progress_test.go
根目录 / internal / agent / session_display_pager_event_progress_test.go
1 package agent
2
3 import (
4 "context"
5 "database/sql"
6 "encoding/json"
7 "errors"
8 "fmt"
9 "os"
10 "os/exec"
11 "path/filepath"
12 "strings"
13 "testing"
14
15 "reasonix/internal/fileops"
16 "reasonix/internal/historywork"
17 "reasonix/internal/projectiondb"
18 "reasonix/internal/provider"
19 "reasonix/internal/store"
20 )
21
22 func eventResumeOptions(t *testing.T, source, cache string) (projectiondb.OpenOptions, string, int64) {
23 t.Helper()
24 info, err := os.Stat(source)
25 if err != nil {
26 t.Fatal(err)
27 }
28 target, version := fileops.DiskSnapshot(source, info)
29 event, err := os.Stat(store.SessionEventLog(source))
30 if err != nil {
31 t.Fatal(err)
32 }
33 eventTarget, eventVersion := fileops.DiskSnapshot(store.SessionEventLog(source), event)
34 fingerprint := fmt.Sprintf("%s:%s:event:%s:%s:schema1", target.Key, version, eventTarget.Key, eventVersion)
35 return projectiondb.OpenOptions{Path: cache, Migrations: displayPagerMigrations, RequireDisk: true, MaxOpenConns: 1, ResumeKey: "event-v1:" + fingerprint}, fingerprint, info.Size()
36 }
37
38 func TestEventPagerResumesCommittedProgress(t *testing.T) {
39 for _, mode := range []string{"messages", "projection", "publication", "repeated", "process-restart", "fields-after", "damaged", "missing", "changed"} {
40 t.Run(mode, func(t *testing.T) { checkEventPagerResume(t, mode) })
41 }
42 }
43
44 func checkEventPagerResume(t *testing.T, mode string) {
45 t.Helper()
46 messages := make([]provider.Message, historywork.BatchEntries*3)
47 for i := range messages {
48 messages[i] = provider.Message{Role: provider.RoleUser, Content: fmt.Sprintf("question %d %s", i, strings.Repeat("a", 4096))}
49 }
50 encode := func() string {
51 body, err := json.Marshal(messages)
52 if err != nil {
53 t.Fatal(err)
54 }
55 if mode == "fields-after" {
56 return `{"messages":` + string(body) + `,"schema_version":1,"type":"replace"}` + "\n"
57 }
58 return `{"schema_version":1,"type":"replace","messages":` + string(body) + "}\n"
59 }
60 text := encode()
61 source, cache := eventDisplayFixture(t, text)
62 opts, fingerprint, size := eventResumeOptions(t, source, cache)
63 interrupt := func(wantPhase string, wantCount int) {
64 ctx, cancel := context.WithCancel(t.Context())
65 defer cancel()
66 err := projectiondb.Rebuild(ctx, opts, func(ctx context.Context, db *sql.DB) error {
67 return buildEventDisplayPagerObserved(ctx, db, source, fingerprint, size, func(phase string, count int) {
68 if phase == wantPhase && count == wantCount {
69 cancel()
70 }
71 })
72 })
73 if !errors.Is(err, context.Canceled) {
74 t.Fatalf("interruption did not retain progress: %v", err)
75 }
76 }
77 switch mode {
78 case "publication":
79 ctx, cancel := context.WithCancel(t.Context())
80 err := projectiondb.Rebuild(ctx, opts, func(ctx context.Context, db *sql.DB) error {
81 err := buildEventDisplayPager(ctx, db, source, fingerprint, size)
82 cancel()
83 if err != nil {
84 return err
85 }
86 return ctx.Err()
87 })
88 cancel()
89 if !errors.Is(err, context.Canceled) {
90 t.Fatalf("publication cancellation: %v", err)
91 }
92 case "process-restart":
93 executable, err := os.Executable()
94 if err != nil {
95 t.Fatal(err)
96 }
97 child := exec.CommandContext(t.Context(), executable, "-test.run=^TestEventPagerCrashHelper$")
98 child.Env = append(os.Environ(), "REASONIX_EVENT_CRASH_SOURCE="+source, "REASONIX_EVENT_CRASH_CACHE="+cache)
99 if output, err := child.CombinedOutput(); err != nil {
100 t.Fatalf("child failed: %v\n%s", err, output)
101 }
102 case "projection":
103 interrupt("projection", historywork.BatchEntries)
104 default:
105 interrupt("messages", historywork.BatchEntries)
106 }
107 if mode == "repeated" {
108 interrupt("messages", 2*historywork.BatchEntries)
109 }
110 if mode == "damaged" || mode == "missing" {
111 pending := opts
112 pending.Path += ".rebuild-pending"
113 handle, err := projectiondb.Open(t.Context(), pending)
114 if err != nil {
115 t.Fatal(err)
116 }
117 query := `UPDATE metadata SET value='broken' WHERE key='event_scan_progress'`
118 if mode == "missing" {
119 query = `DELETE FROM metadata WHERE key='event_scan_progress'`
120 }
121 _, err = handle.DB.Exec(query)
122 _ = handle.DB.Close()
123 if err != nil {
124 t.Fatal(err)
125 }
126 }
127 if mode == "changed" {
128 info, err := os.Stat(store.SessionEventLog(source))
129 if err != nil {
130 t.Fatal(err)
131 }
132 for i := range messages {
133 messages[i].Content = strings.ReplaceAll(messages[i].Content, "a", "b")
134 }
135 text = encode()
136 if err := os.WriteFile(store.SessionEventLog(source), []byte(text), 0600); err != nil {
137 t.Fatal(err)
138 }
139 if err := os.Chtimes(store.SessionEventLog(source), info.ModTime(), info.ModTime()); err != nil {
140 t.Fatal(err)
141 }
142 }
143 var meter historywork.Coordinator
144 pager, err := OpenDisplayPager(meter.Context(t.Context()), source, cache)
145 if err != nil {
146 t.Fatal(err)
147 }
148 defer pager.Close()
149 digest, err := ContentDigestForMessages(messages)
150 if err != nil || pager.Header.ContentDigest != digest || pager.Header.MessageCount != len(messages) {
151 t.Fatalf("resumed projection changed semantic content: %+v %v", pager.Header, err)
152 }
153 read := meter.Diagnostics().InstrumentedReadBytes
154 if mode == "projection" && read >= int64(len(text))*3/4 {
155 t.Fatalf("projection prefix decoded again: %d of %d", read, len(text))
156 }
157 if (mode == "messages" || mode == "process-restart" || mode == "repeated") && read >= int64(len(text))*19/10 {
158 t.Fatalf("message scan prefix decoded again: %d of %d", read, len(text))
159 }
160 page, err := pager.EventMessages(len(messages)-1, len(messages))
161 if err != nil || len(page) != 1 || page[0].Content != messages[len(messages)-1].Content {
162 t.Fatalf("wrong resumed source locations: %+v %v", page, err)
163 }
164 after, err := os.ReadFile(store.SessionEventLog(source))
165 if err != nil || string(after) != text {
166 t.Fatalf("authoritative source was changed: %v", err)
167 }
168 }
169
170 func TestEventPagerCrashHelper(t *testing.T) {
171 source, cache := os.Getenv("REASONIX_EVENT_CRASH_SOURCE"), os.Getenv("REASONIX_EVENT_CRASH_CACHE")
172 if source == "" || cache == "" {
173 return
174 }
175 opts, fingerprint, size := eventResumeOptions(t, source, cache)
176 err := projectiondb.Rebuild(t.Context(), opts, func(ctx context.Context, db *sql.DB) error {
177 return buildEventDisplayPagerObserved(ctx, db, source, fingerprint, size, func(phase string, count int) {
178 if phase == "messages" && count == historywork.BatchEntries {
179 os.Exit(0)
180 }
181 })
182 })
183 t.Fatalf("child did not exit: %v", err)
184 }
185
186 func TestEventPagerRejectsSourceChangeBeforePublication(t *testing.T) {
187 for _, mode := range []string{"event-rewrite", "event-replace", "checkpoint-rewrite"} {
188 t.Run(mode, func(t *testing.T) {
189 source, cache := eventDisplayFixture(t, `{"schema_version":1,"type":"replace","messages":[{"role":"user","content":"old"}]}`)
190 opts, fingerprint, size := eventResumeOptions(t, source, cache)
191 err := projectiondb.Rebuild(t.Context(), opts, func(ctx context.Context, db *sql.DB) error {
192 return buildEventDisplayPagerObserved(ctx, db, source, fingerprint, size, func(phase string, _ int) {
193 if phase != "projection" {
194 return
195 }
196 path := store.SessionEventLog(source)
197 if mode == "checkpoint-rewrite" {
198 path = source
199 }
200 info, _ := os.Stat(path)
201 body, _ := os.ReadFile(path)
202 if mode == "event-replace" {
203 if err := os.Rename(path, filepath.Join(filepath.Dir(path), "previous")); err != nil {
204 t.Fatal(err)
205 }
206 }
207 if err := os.WriteFile(path, []byte(strings.ReplaceAll(string(body), "old", "new")), 0600); err != nil {
208 t.Fatal(err)
209 }
210 if err := os.Chtimes(path, info.ModTime(), info.ModTime()); err != nil {
211 t.Fatal(err)
212 }
213 })
214 })
215 if !errors.Is(err, ErrDisplaySourceChanged) {
216 t.Fatalf("changed source was published: %v", err)
217 }
218 })
219 }
220 }
221
222 func TestEventPagerResumedArrayPreservesJSONSemantics(t *testing.T) {
223 for _, mode := range []string{"duplicate-messages", "append-replace", "bad-separator", "trailing-comma", "broken-tail", "last-element"} {
224 t.Run(mode, func(t *testing.T) {
225 messages := make([]provider.Message, historywork.BatchEntries)
226 for i := range messages {
227 messages[i] = provider.Message{Role: provider.RoleUser, Content: fmt.Sprint(i)}
228 }
229 body, _ := json.Marshal(messages)
230 prefix := `{"schema_version":1,"type":"replace","messages":`
231 text := prefix + string(body) + "}\n"
232 want := len(messages)
233 switch mode {
234 case "duplicate-messages":
235 text = prefix + string(body) + `,"messages":[{"role":"user","content":"last field wins"}]}`
236 want = 1
237 case "append-replace":
238 text += `{"schema_version":1,"type":"append","message_index":128,"messages":[{"role":"assistant","content":"append"}]}` + "\n" +
239 `{"schema_version":1,"type":"replace","messages":[{"role":"user","content":"replacement"}]}`
240 want = 1
241 case "bad-separator":
242 text = prefix + string(body[:len(body)-1]) + ` {"role":"user","content":"missing comma"}]}`
243 case "trailing-comma":
244 text = prefix + string(body[:len(body)-1]) + `,]}`
245 case "broken-tail":
246 text = prefix + string(body)
247 }
248 source, cache := eventDisplayFixture(t, text)
249 opts, fingerprint, size := eventResumeOptions(t, source, cache)
250 ctx, cancel := context.WithCancel(t.Context())
251 err := projectiondb.Rebuild(ctx, opts, func(ctx context.Context, db *sql.DB) error {
252 return buildEventDisplayPagerObserved(ctx, db, source, fingerprint, size, func(phase string, count int) {
253 if phase == "messages" && count == historywork.BatchEntries {
254 cancel()
255 }
256 })
257 })
258 cancel()
259 if !errors.Is(err, context.Canceled) {
260 t.Fatalf("expected checkpoint cancellation: %v", err)
261 }
262 pager, err := OpenDisplayPager(t.Context(), source, cache)
263 if mode == "bad-separator" || mode == "trailing-comma" || mode == "broken-tail" {
264 if err == nil {
265 pager.Close()
266 t.Fatal("invalid tail was published")
267 }
268 return
269 }
270 if err != nil {
271 t.Fatal(err)
272 }
273 defer pager.Close()
274 if pager.Header.MessageCount != want {
275 t.Fatalf("resumed event semantics changed: %d != %d", pager.Header.MessageCount, want)
276 }
277 })
278 }
279 }
280
280 lines GO