| 1 | package agent |
| 2 | |
| 3 | import ( |
| 4 | "errors" |
| 5 | "strings" |
| 6 | "testing" |
| 7 | |
| 8 | "reasonix/internal/provider" |
| 9 | ) |
| 10 | |
| 11 | func TestForkHeadStartsANewHeadAndKeepsTheOldChain(t *testing.T) { |
| 12 | path := dagTestSession(t) |
| 13 | s := dagSavedSession(t, path, "q1", "a1", "q2", "a2") |
| 14 | msgs := s.Snapshot() |
| 15 | forkAt := msgs[2].ID // a1 |
| 16 | head, err := s.ForkHead(path, forkAt, HeadKindFork, "alt") |
| 17 | if err != nil || head == "" || head == SessionMainHead { |
| 18 | t.Fatalf("ForkHead = %q err=%v", head, err) |
| 19 | } |
| 20 | if got := strings.Join(dagContents(s.Snapshot()), ","); got != "sys,q1,a1" { |
| 21 | t.Fatalf("in-memory transcript after fork = %s", got) |
| 22 | } |
| 23 | if ref, ok := s.Head(); !ok || ref.HeadID != head || ref.LeafID != forkAt { |
| 24 | t.Fatalf("head ref after fork = %+v ok=%v", ref, ok) |
| 25 | } |
| 26 | if err := s.Save(path); err != nil { |
| 27 | t.Fatalf("no-op save after fork: %v", err) |
| 28 | } |
| 29 | s.Add(provider.Message{Role: provider.RoleUser, Content: "q2-alt"}) |
| 30 | if err := s.Save(path); err != nil { |
| 31 | t.Fatal(err) |
| 32 | } |
| 33 | st := dagReplay(t, path) |
| 34 | if got := strings.Join(dagChain(st, SessionMainHead), ","); got != "sys,q1,a1,q2,a2" { |
| 35 | t.Fatalf("main chain after fork = %s", got) |
| 36 | } |
| 37 | if got := strings.Join(dagChain(st, head), ","); got != "sys,q1,a1,q2-alt" { |
| 38 | t.Fatalf("fork chain = %s", got) |
| 39 | } |
| 40 | heads, err := ListSessionHeads(path) |
| 41 | if err != nil || len(heads) != 2 || heads[1].ID != head || heads[1].Kind != HeadKindFork || heads[1].Name != "alt" || heads[1].ForkFrom != forkAt || !heads[1].Selected { |
| 42 | t.Fatalf("heads = %+v err=%v", heads, err) |
| 43 | } |
| 44 | reloaded, err := LoadSession(path) |
| 45 | if err != nil { |
| 46 | t.Fatal(err) |
| 47 | } |
| 48 | if ref, _ := reloaded.Head(); ref.HeadID != head { |
| 49 | t.Fatalf("reload must land on the selected fork, got %+v", ref) |
| 50 | } |
| 51 | idx, err := ReadSessionHeadIndex(path) |
| 52 | if err != nil || idx == nil || !idx.Current(path) || idx.SelectedHead != head || len(idx.Heads) != 2 { |
| 53 | t.Fatalf("head index after fork = %+v err=%v", idx, err) |
| 54 | } |
| 55 | meta, _, _ := LoadBranchMeta(path) |
| 56 | if meta.HeadID != head || meta.HeadCount != 2 { |
| 57 | t.Fatalf("meta mirror after fork = head %q count %d", meta.HeadID, meta.HeadCount) |
| 58 | } |
| 59 | } |
| 60 | |
| 61 | func TestSwitchHeadMovesTheSessionBackAndForth(t *testing.T) { |
| 62 | path := dagTestSession(t) |
| 63 | s := dagSavedSession(t, path, "q1", "a1") |
| 64 | fork, err := s.ForkHead(path, s.Snapshot()[1].ID, HeadKindRewind, "") |
| 65 | if err != nil { |
| 66 | t.Fatal(err) |
| 67 | } |
| 68 | s.Add(provider.Message{Role: provider.RoleAssistant, Content: "a1-rewound"}) |
| 69 | if err := s.Save(path); err != nil { |
| 70 | t.Fatal(err) |
| 71 | } |
| 72 | if err := s.SwitchHead(path, SessionMainHead); err != nil { |
| 73 | t.Fatalf("SwitchHead main: %v", err) |
| 74 | } |
| 75 | if got := strings.Join(dagContents(s.Snapshot()), ","); got != "sys,q1,a1" { |
| 76 | t.Fatalf("transcript on main = %s", got) |
| 77 | } |
| 78 | s.Add(provider.Message{Role: provider.RoleUser, Content: "q2-main"}) |
| 79 | if err := s.Save(path); err != nil { |
| 80 | t.Fatal(err) |
| 81 | } |
| 82 | st := dagReplay(t, path) |
| 83 | if got := strings.Join(dagChain(st, SessionMainHead), ","); got != "sys,q1,a1,q2-main" { |
| 84 | t.Fatalf("main chain = %s", got) |
| 85 | } |
| 86 | if got := strings.Join(dagChain(st, fork), ","); got != "sys,q1,a1-rewound" { |
| 87 | t.Fatalf("rewind chain = %s", got) |
| 88 | } |
| 89 | if st.selectedHead() != SessionMainHead { |
| 90 | t.Fatalf("selected = %q, want main after switch", st.selectedHead()) |
| 91 | } |
| 92 | if err := s.SwitchHead(path, "nope"); !errors.Is(err, ErrSessionHeadUnknown) { |
| 93 | t.Fatalf("unknown head err = %v", err) |
| 94 | } |
| 95 | if err := s.SwitchHead(path, SessionMainHead); err != nil { |
| 96 | t.Fatalf("switching to the current head must be a no-op: %v", err) |
| 97 | } |
| 98 | } |
| 99 | |
| 100 | func TestHeadMarkersOnDiskSelectRetireRename(t *testing.T) { |
| 101 | path := dagTestSession(t) |
| 102 | s := dagSavedSession(t, path, "q1", "a1") |
| 103 | fork, err := s.ForkHead(path, s.Snapshot()[1].ID, HeadKindFork, "side") |
| 104 | if err != nil { |
| 105 | t.Fatal(err) |
| 106 | } |
| 107 | if err := RenameSessionHead(path, fork, "renamed"); err != nil { |
| 108 | t.Fatal(err) |
| 109 | } |
| 110 | if err := RetireSessionHead(path, fork); err == nil { |
| 111 | t.Fatal("retiring the selected head must be refused") |
| 112 | } |
| 113 | if err := SelectSessionHead(path, SessionMainHead); err != nil { |
| 114 | t.Fatal(err) |
| 115 | } |
| 116 | if err := RetireSessionHead(path, fork); err != nil { |
| 117 | t.Fatalf("retire: %v", err) |
| 118 | } |
| 119 | if err := SelectSessionHead(path, fork); err == nil { |
| 120 | t.Fatal("a retired head must not become the selection") |
| 121 | } |
| 122 | heads, err := ListSessionHeads(path) |
| 123 | if err != nil || len(heads) != 2 { |
| 124 | t.Fatalf("heads = %+v err=%v", heads, err) |
| 125 | } |
| 126 | if !heads[1].Retired || heads[1].Name != "renamed" || heads[1].Selected || !heads[0].Selected { |
| 127 | t.Fatalf("head rows = %+v", heads) |
| 128 | } |
| 129 | reloaded, err := LoadSession(path) |
| 130 | if err != nil { |
| 131 | t.Fatal(err) |
| 132 | } |
| 133 | if ref, _ := reloaded.Head(); ref.HeadID != SessionMainHead { |
| 134 | t.Fatalf("reload after select = %+v", ref) |
| 135 | } |
| 136 | if err := RetireSessionHead(path, "missing"); !errors.Is(err, ErrSessionHeadUnknown) { |
| 137 | t.Fatalf("unknown head err = %v", err) |
| 138 | } |
| 139 | meta, _, _ := LoadBranchMeta(path) |
| 140 | if meta.HeadID != SessionMainHead || meta.HeadCount != 2 { |
| 141 | t.Fatalf("meta mirror = head %q count %d", meta.HeadID, meta.HeadCount) |
| 142 | } |
| 143 | } |
| 144 | |
| 145 | func TestHeadOperationsRefuseSchemaOneSessions(t *testing.T) { |
| 146 | useSchemaOneLog(t) |
| 147 | path := dagTestSession(t) |
| 148 | s := dagSavedSession(t, path, "q1") |
| 149 | if _, err := s.ForkHead(path, "", HeadKindFork, ""); !errors.Is(err, ErrSessionNotDAG) { |
| 150 | t.Fatalf("ForkHead on schema 1 err = %v", err) |
| 151 | } |
| 152 | if err := SelectSessionHead(path, SessionMainHead); !errors.Is(err, ErrSessionNotDAG) { |
| 153 | t.Fatalf("SelectSessionHead on schema 1 err = %v", err) |
| 154 | } |
| 155 | } |
| 156 | |
| 157 | func TestHeadListMarksCoveredHeads(t *testing.T) { |
| 158 | path := dagTestSession(t) |
| 159 | s := dagSavedSession(t, path, "q1", "a1") |
| 160 | fork, err := s.ForkHead(path, s.Snapshot()[2].ID, HeadKindFork, "") |
| 161 | if err != nil { |
| 162 | t.Fatal(err) |
| 163 | } |
| 164 | covered := func() map[string]bool { |
| 165 | t.Helper() |
| 166 | heads, err := ListSessionHeads(path) |
| 167 | if err != nil { |
| 168 | t.Fatal(err) |
| 169 | } |
| 170 | out := map[string]bool{} |
| 171 | for _, h := range heads { |
| 172 | out[h.ID] = h.Covered |
| 173 | } |
| 174 | return out |
| 175 | } |
| 176 | if got := covered(); !got[SessionMainHead] || got[fork] { |
| 177 | t.Fatalf("tip fork selected: covered = %v, want the parent covered and the selection never flagged", got) |
| 178 | } |
| 179 | s.Add(provider.Message{Role: provider.RoleUser, Content: "q2-alt"}) |
| 180 | if err := s.Save(path); err != nil { |
| 181 | t.Fatal(err) |
| 182 | } |
| 183 | if got := covered(); !got[SessionMainHead] || got[fork] { |
| 184 | t.Fatalf("after the fork grew: covered = %v", got) |
| 185 | } |
| 186 | if err := SelectSessionHead(path, SessionMainHead); err != nil { |
| 187 | t.Fatal(err) |
| 188 | } |
| 189 | if got := covered(); got[SessionMainHead] || got[fork] { |
| 190 | t.Fatalf("main selected: covered = %v, want the diverged fork kept", got) |
| 191 | } |
| 192 | } |
| 193 |