返回 DeepSeek-Reasonix
checkpoint.go
根目录 / internal / transcript / checkpoint.go
1 package transcript
2
3 import (
4 "encoding/json"
5 "errors"
6 "fmt"
7 "os"
8 "reflect"
9
10 "reasonix/internal/eventwire"
11 "reasonix/internal/fileutil"
12 "reasonix/internal/store"
13 )
14
15 // Checkpoint is the durable display state, separate from provider messages.
16 // Its digest binds it to the terminal transcript that produced its coverage.
17 type Checkpoint struct {
18 Version int `json:"version"`
19 Identity Identity `json:"identity"`
20 CoveredThroughSeq uint64 `json:"coveredThroughSeq"`
21 TranscriptDigest string `json:"transcriptDigest"`
22 ProviderCount int `json:"providerCount"`
23 Records []Message `json:"records"`
24 Runtime Runtime `json:"runtime"`
25 ActiveAttempts []ActiveAttempt `json:"activeAttempts"`
26 Completion *eventwire.CompletionSummary `json:"completion,omitempty"`
27 }
28
29 // ToolResultRepairStats reports legacy identity recovery without including
30 // tool arguments or result bodies in diagnostics.
31 type ToolResultRepairStats struct {
32 Repaired int
33 Missing int
34 Conflicts int
35 }
36
37 // NeedsToolResultRepair keeps the common restore path from rebuilding
38 // canonical history when every tool-result display row already has identity.
39 func NeedsToolResultRepair(records []Message) bool {
40 for _, message := range records {
41 if message.Role == "tool" && message.MessageID == "" && message.ToolCallID != "" {
42 return true
43 }
44 }
45 return false
46 }
47
48 func checkpointMessageIDs(records []Message) map[string]bool {
49 occupied := make(map[string]bool)
50 for _, message := range records {
51 if message.MessageID != "" {
52 occupied[message.MessageID] = true
53 }
54 }
55 return occupied
56 }
57
58 // RepairCheckpointToolResults joins legacy display rows that lost MessageID
59 // with authoritative persisted history. ToolCallID is the only cross-stream
60 // join key; known turn boundaries must also agree. Ambiguous or conflicting
61 // rows are deliberately left unchanged.
62 func RepairCheckpointToolResults(records, canonical []Message) ([]Message, ToolResultRepairStats) {
63 repaired := append([]Message(nil), records...)
64 byCall := make(map[string][]Message)
65 occupied := checkpointMessageIDs(records)
66 historyTurn := 0
67 for _, message := range canonical {
68 if message.Role == "user" {
69 historyTurn++
70 }
71 if message.HistoryTurn == 0 {
72 message.HistoryTurn = historyTurn
73 }
74 if message.Role == "tool" && message.ToolCallID != "" && message.MessageID != "" {
75 byCall[message.ToolCallID] = append(byCall[message.ToolCallID], message)
76 }
77 }
78 stats := ToolResultRepairStats{}
79 for index := range repaired {
80 legacy := &repaired[index]
81 if legacy.Role != "tool" || legacy.MessageID != "" || legacy.ToolCallID == "" {
82 continue
83 }
84 candidates := make([]Message, 0, len(byCall[legacy.ToolCallID]))
85 for _, candidate := range byCall[legacy.ToolCallID] {
86 if legacy.TurnID != "" && candidate.TurnID != "" && legacy.TurnID != candidate.TurnID {
87 continue
88 }
89 if legacy.HistoryTurn != 0 && candidate.HistoryTurn != 0 && legacy.HistoryTurn != candidate.HistoryTurn {
90 continue
91 }
92 candidates = append(candidates, candidate)
93 }
94 if len(candidates) == 0 {
95 stats.Missing++
96 continue
97 }
98 if len(candidates) != 1 {
99 stats.Conflicts++
100 continue
101 }
102 formal := candidates[0]
103 if occupied[formal.MessageID] {
104 stats.Conflicts++
105 continue
106 }
107 // Keep the checkpoint row's display location and event-formatted result,
108 // while restoring fields owned by the persisted message.
109 legacy.MessageID = formal.MessageID
110 occupied[formal.MessageID] = true
111 if legacy.RecordID == "" {
112 legacy.RecordID = formal.RecordID
113 }
114 if legacy.ToolName == "" {
115 legacy.ToolName = formal.ToolName
116 }
117 if legacy.CreatedAt == 0 {
118 legacy.CreatedAt = formal.CreatedAt
119 }
120 if legacy.Execution == nil {
121 legacy.Execution = formal.Execution
122 }
123 legacy.ToolResultArchived = legacy.ToolResultArchived || formal.ToolResultArchived
124 if len(legacy.PresentedFiles) == 0 {
125 legacy.PresentedFiles = append(legacy.PresentedFiles, formal.PresentedFiles...)
126 }
127 if legacy.ReadCompletion == nil {
128 legacy.ReadCompletion = formal.ReadCompletion
129 }
130 if legacy.Diagnostic == nil {
131 legacy.Diagnostic = formal.Diagnostic
132 }
133 stats.Repaired++
134 }
135 return repaired, stats
136 }
137
138 func (p *Projection) Checkpoint(digest string) (Checkpoint, error) {
139 p.mu.Lock()
140 defer p.mu.Unlock()
141 runtime, attempts := p.runtimeLocked()
142 state := Checkpoint{Version: ProtocolVersion, Identity: p.identity, CoveredThroughSeq: p.covered,
143 TranscriptDigest: digest, Records: p.buffer.Messages(), Runtime: runtime, ActiveAttempts: attempts, Completion: p.buffer.completion}
144 // Detach mutable metadata while sharing immutable strings. Encoding and
145 // decoding the full transcript here duplicates large bodies under p.mu;
146 // SaveCheckpoint already owns the required encoding outside that lock.
147 owned := mapContentStrings(reflect.ValueOf(state), nil, func(text string, _ []string) string { return text }).Interface().(Checkpoint)
148 return owned, nil
149 }
150
151 func RestoreCheckpoint(state Checkpoint, identity Identity) (*Projection, error) {
152 if state.Version != ProtocolVersion || state.Identity.SessionID != identity.SessionID ||
153 state.Identity.HeadID != identity.HeadID || state.Identity.RewriteEpoch != identity.RewriteEpoch {
154 return nil, errors.New("transcript checkpoint identity mismatch")
155 }
156 records := repairCheckpointRecordIdentities(state.Records, state.CoveredThroughSeq)
157 p, err := NewProjection(identity, records, state.CoveredThroughSeq)
158 if err != nil {
159 return nil, err
160 }
161 b, err := json.Marshal(state)
162 if err != nil {
163 return nil, err
164 }
165 var owned Checkpoint
166 if err = json.Unmarshal(b, &owned); err != nil {
167 return nil, err
168 }
169 p.runtime = owned.Runtime
170 if p.runtime.StartedAt > 0 {
171 p.startedTurnID = p.runtime.TurnID
172 }
173 for _, attempt := range owned.ActiveAttempts {
174 p.attempts[attempt.ID] = attempt
175 }
176 for _, prompt := range owned.Runtime.PendingEvents {
177 id := prompt.PromptID
178 if id != "" {
179 p.prompts[id] = prompt
180 }
181 }
182 p.buffer.completion = owned.Completion
183 return p, nil
184 }
185
186 // Older builds could persist display-only rows without an identity when a
187 // frame was published outside the active turn. Repair only that legacy shape;
188 // non-empty duplicate identities remain corruption and are rejected by
189 // NewProjection. The generated value is deterministic for this checkpoint so
190 // repeated recovery cannot reshuffle mounted rows.
191 func repairCheckpointRecordIdentities(records []Message, covered uint64) []Message {
192 repaired := append([]Message(nil), records...)
193 used := make(map[string]bool, len(repaired))
194 for _, record := range repaired {
195 if record.RecordID != "" {
196 used[record.RecordID] = true
197 }
198 }
199 for index := range repaired {
200 if repaired[index].RecordID != "" {
201 continue
202 }
203 switch {
204 case repaired[index].Role == "tool" && repaired[index].ToolCallID != "":
205 repaired[index].RecordID = "tool:" + repaired[index].ToolCallID
206 case repaired[index].MessageID != "":
207 repaired[index].RecordID = "m:" + repaired[index].MessageID
208 default:
209 base := fmt.Sprintf("view:checkpoint:%d:%d", covered, index)
210 repaired[index].RecordID = base
211 for suffix := 1; used[repaired[index].RecordID]; suffix++ {
212 repaired[index].RecordID = fmt.Sprintf("%s:%d", base, suffix)
213 }
214 }
215 used[repaired[index].RecordID] = true
216 }
217 return repaired
218 }
219
220 func SaveCheckpoint(sessionPath string, state Checkpoint) error {
221 path := store.SessionTranscriptProjection(sessionPath)
222 if path == "" {
223 return nil
224 }
225 b, err := json.Marshal(state)
226 if err != nil {
227 return err
228 }
229 return fileutil.AtomicWriteFile(path, b, 0o600)
230 }
231
232 func LoadCheckpoint(sessionPath string) (Checkpoint, bool, error) {
233 path := store.SessionTranscriptProjection(sessionPath)
234 if path == "" {
235 return Checkpoint{}, false, nil
236 }
237 b, err := os.ReadFile(path)
238 if errors.Is(err, os.ErrNotExist) {
239 return Checkpoint{}, false, nil
240 }
241 if err != nil {
242 return Checkpoint{}, false, err
243 }
244 var state Checkpoint
245 if err = json.Unmarshal(b, &state); err != nil {
246 return Checkpoint{}, false, err
247 }
248 if state.Version != ProtocolVersion {
249 return Checkpoint{}, false, errors.New("unsupported transcript checkpoint version")
250 }
251 return state, true, nil
252 }
253
253 lines GO