| 1 | package agent |
| 2 | |
| 3 | import ( |
| 4 | "context" |
| 5 | "crypto/sha256" |
| 6 | "database/sql" |
| 7 | "encoding" |
| 8 | "encoding/json" |
| 9 | "errors" |
| 10 | "fmt" |
| 11 | "os" |
| 12 | "time" |
| 13 | |
| 14 | "reasonix/internal/fileops" |
| 15 | "reasonix/internal/provider" |
| 16 | "reasonix/internal/store" |
| 17 | ) |
| 18 | |
| 19 | // Only display-relevant head state is retained. Message bodies and execution |
| 20 | // state (writers, open turns, compaction) never enter this disposable cache. |
| 21 | type dagPagerHead struct { |
| 22 | ID, Leaf, System string |
| 23 | CreatedAt, LastActivity time.Time |
| 24 | LastOffset int64 |
| 25 | Retired bool |
| 26 | } |
| 27 | |
| 28 | type dagPagerProgress struct { |
| 29 | Version int |
| 30 | SourceKey, Phase string |
| 31 | Offset int64 |
| 32 | Records int |
| 33 | Heads []dagPagerHead `json:"-"` |
| 34 | HeadPositions map[string]int `json:"-"` |
| 35 | HeadCount int |
| 36 | Selected string |
| 37 | HeadID, NextID string |
| 38 | ChainCount int |
| 39 | NextPosition int |
| 40 | Header SessionDisplayIndex |
| 41 | Users int |
| 42 | Hash []byte |
| 43 | } |
| 44 | |
| 45 | func (p *dagPagerProgress) snapshot(st *sessionDAGState) { |
| 46 | p.Selected = st.selected |
| 47 | p.Heads = make([]dagPagerHead, 0, len(st.headOrder)) |
| 48 | for _, id := range st.headOrder { |
| 49 | h := st.heads[id] |
| 50 | p.Heads = append(p.Heads, snapshotDAGPagerHead(h)) |
| 51 | } |
| 52 | p.HeadCount = len(p.Heads) |
| 53 | } |
| 54 | |
| 55 | func snapshotDAGPagerHead(h *sessionDAGHead) dagPagerHead { |
| 56 | saved := dagPagerHead{ID: h.id, Leaf: h.leaf, CreatedAt: h.createdAt, LastActivity: h.lastActivity, LastOffset: h.lastOffset, Retired: h.retired} |
| 57 | if h.system != nil { |
| 58 | saved.System = h.system.ID |
| 59 | } |
| 60 | return saved |
| 61 | } |
| 62 | |
| 63 | func (p *dagPagerProgress) state(source string) *sessionDAGState { |
| 64 | st := newSessionDAGState(source) |
| 65 | st.heads, st.headOrder, st.selected = map[string]*sessionDAGHead{}, nil, p.Selected |
| 66 | for _, saved := range p.Heads { |
| 67 | h := &sessionDAGHead{id: saved.ID, leaf: saved.Leaf, createdAt: saved.CreatedAt, lastActivity: saved.LastActivity, lastOffset: saved.LastOffset, retired: saved.Retired} |
| 68 | if saved.System != "" { |
| 69 | h.system = &provider.Message{ID: saved.System} |
| 70 | } |
| 71 | st.heads[h.id] = h |
| 72 | st.headOrder = append(st.headOrder, h.id) |
| 73 | } |
| 74 | return st |
| 75 | } |
| 76 | |
| 77 | func validateDAGPagerFile(source string, file *os.File, fingerprint, head string) error { |
| 78 | key, err := displayImportFileKey(store.SessionEventLog(source), file) |
| 79 | if err != nil { |
| 80 | return err |
| 81 | } |
| 82 | info, err := os.Stat(source) |
| 83 | if err != nil { |
| 84 | return err |
| 85 | } |
| 86 | target, version := fileops.DiskSnapshot(source, info) |
| 87 | info, err = file.Stat() |
| 88 | if err != nil { |
| 89 | return err |
| 90 | } |
| 91 | if fmt.Sprintf("%s:%s:event:%s:dag:%d:%d:%s", target.Key, version, key, info.Size(), info.ModTime().UnixNano(), head) != fingerprint { |
| 92 | return ErrDisplaySourceChanged |
| 93 | } |
| 94 | return nil |
| 95 | } |
| 96 | |
| 97 | func restoreDAGPagerProgress(ctx context.Context, db *sql.DB, source, fingerprint string, size int64) (*dagPagerProgress, error) { |
| 98 | _, err := db.ExecContext(ctx, `CREATE TABLE IF NOT EXISTS dag_nodes(id TEXT PRIMARY KEY,parent TEXT NOT NULL,location BLOB NOT NULL); |
| 99 | CREATE TABLE IF NOT EXISTS dag_overlays(kind TEXT NOT NULL,id TEXT NOT NULL,location BLOB NOT NULL,PRIMARY KEY(kind,id)); |
| 100 | CREATE TABLE IF NOT EXISTS dag_chain(position INTEGER PRIMARY KEY,id TEXT UNIQUE NOT NULL); |
| 101 | CREATE TABLE IF NOT EXISTS dag_locations(position INTEGER PRIMARY KEY,location BLOB NOT NULL); |
| 102 | CREATE TABLE IF NOT EXISTS dag_heads(position INTEGER PRIMARY KEY,id TEXT UNIQUE NOT NULL,record BLOB NOT NULL);`) |
| 103 | if err != nil { |
| 104 | return nil, err |
| 105 | } |
| 106 | var raw string |
| 107 | err = db.QueryRowContext(ctx, `SELECT value FROM metadata WHERE key='dag_progress'`).Scan(&raw) |
| 108 | if err != nil && !errors.Is(err, sql.ErrNoRows) { |
| 109 | return nil, err |
| 110 | } |
| 111 | var p dagPagerProgress |
| 112 | valid := err == nil && json.Unmarshal([]byte(raw), &p) == nil && validDAGPagerHeader(&p, fingerprint, size) |
| 113 | info, err := os.Stat(store.SessionEventLog(source)) |
| 114 | if err != nil { |
| 115 | return nil, err |
| 116 | } |
| 117 | headsValid, err := restoreDAGPagerHeads(ctx, db, &p) |
| 118 | if err != nil { |
| 119 | return nil, err |
| 120 | } |
| 121 | valid = valid && headsValid && p.Offset <= info.Size() && validDAGPagerPhase(&p) |
| 122 | for _, table := range []struct { |
| 123 | name string |
| 124 | rows int |
| 125 | }{{"dag_chain", p.ChainCount}, {"entries", p.Header.MessageCount}, {"dag_locations", p.Header.MessageCount}} { |
| 126 | var count, first, last int |
| 127 | if err := db.QueryRowContext(ctx, `SELECT COUNT(*),COALESCE(MIN(position),0),COALESCE(MAX(position),-1) FROM `+table.name).Scan(&count, &first, &last); err != nil { |
| 128 | return nil, err |
| 129 | } |
| 130 | valid = valid && count == table.rows && first == 0 && last == count-1 |
| 131 | } |
| 132 | if valid { |
| 133 | return &p, nil |
| 134 | } |
| 135 | // A parser position and its derived rows are an atomic unit. Never adopt |
| 136 | // leftover rows when their durable progress proof is missing or invalid. |
| 137 | _, err = db.ExecContext(ctx, `DELETE FROM dag_nodes; DELETE FROM dag_overlays; DELETE FROM dag_chain; |
| 138 | DELETE FROM dag_locations; DELETE FROM dag_heads; DELETE FROM entries; DELETE FROM metadata`) |
| 139 | p = dagPagerProgress{Version: 1, SourceKey: fingerprint, Phase: "scan", |
| 140 | Header: SessionDisplayIndex{SchemaVersion: SessionDisplayIndexSchemaVersion, TranscriptSize: size, ListingPreviewKnown: true}} |
| 141 | p.snapshot(newSessionDAGState(source)) |
| 142 | return &p, err |
| 143 | } |
| 144 | |
| 145 | func restoreDAGPagerHeads(ctx context.Context, db *sql.DB, p *dagPagerProgress) (bool, error) { |
| 146 | rows, err := db.QueryContext(ctx, `SELECT position,id,record FROM dag_heads ORDER BY position`) |
| 147 | if err != nil { |
| 148 | return false, err |
| 149 | } |
| 150 | defer rows.Close() |
| 151 | valid := true |
| 152 | for rows.Next() { |
| 153 | var position int |
| 154 | var id string |
| 155 | var raw []byte |
| 156 | var h dagPagerHead |
| 157 | if err := rows.Scan(&position, &id, &raw); err != nil { |
| 158 | return false, err |
| 159 | } |
| 160 | if json.Unmarshal(raw, &h) != nil || position != len(p.Heads) || id != h.ID { |
| 161 | valid = false |
| 162 | } |
| 163 | p.Heads = append(p.Heads, h) |
| 164 | } |
| 165 | return valid && len(p.Heads) == p.HeadCount, rows.Err() |
| 166 | } |
| 167 | |
| 168 | func validDAGPagerHeader(p *dagPagerProgress, fingerprint string, size int64) bool { |
| 169 | return p.Version == 1 && p.SourceKey == fingerprint && p.Offset >= 0 && p.Records >= 0 && p.ChainCount >= 0 && |
| 170 | p.Header.SchemaVersion == SessionDisplayIndexSchemaVersion && p.Header.TranscriptSize == size && p.Header.Entries == nil && |
| 171 | p.Header.MessageCount >= 0 && p.Users >= 0 && p.Users <= p.Header.MessageCount && p.Header.AuthoredTurns >= 0 && p.Header.AuthoredTurns <= p.Header.MessageCount |
| 172 | } |
| 173 | |
| 174 | func validDAGPagerPhase(p *dagPagerProgress) bool { |
| 175 | seen := map[string]bool{} |
| 176 | for _, h := range p.Heads { |
| 177 | if h.ID == "" || seen[h.ID] || h.LastOffset < 0 || h.LastOffset > p.Offset { |
| 178 | return false |
| 179 | } |
| 180 | seen[h.ID] = true |
| 181 | } |
| 182 | if !seen[SessionMainHead] { |
| 183 | return false |
| 184 | } |
| 185 | switch p.Phase { |
| 186 | case "scan": |
| 187 | return p.ChainCount == 0 && p.Header.MessageCount == 0 |
| 188 | case "chain": |
| 189 | return seen[p.HeadID] && p.Header.MessageCount == 0 |
| 190 | case "projection", "done": |
| 191 | if !seen[p.HeadID] || p.NextID != "" || p.NextPosition < -1 || p.NextPosition >= p.ChainCount { |
| 192 | return false |
| 193 | } |
| 194 | processed := p.ChainCount - p.NextPosition - 1 |
| 195 | if p.Header.MessageCount < processed || p.Header.MessageCount > processed+1 { |
| 196 | return false |
| 197 | } |
| 198 | hash := sha256.New() |
| 199 | if hash.(encoding.BinaryUnmarshaler).UnmarshalBinary(p.Hash) != nil { |
| 200 | return false |
| 201 | } |
| 202 | return p.Phase != "done" || p.NextPosition == -1 && p.Header.ContentDigest == fmt.Sprintf("%x", hash.Sum(nil)) |
| 203 | default: |
| 204 | return false |
| 205 | } |
| 206 | } |
| 207 | |
| 208 | func commitDAGPagerProgress(ctx context.Context, tx *sql.Tx, p *dagPagerProgress) error { |
| 209 | body, err := json.Marshal(p) |
| 210 | if err == nil { |
| 211 | _, err = tx.ExecContext(ctx, `INSERT OR REPLACE INTO metadata VALUES('dag_progress',?)`, string(body)) |
| 212 | } |
| 213 | if err != nil { |
| 214 | return err |
| 215 | } |
| 216 | return tx.Commit() |
| 217 | } |
| 218 |