| 1 | package cli |
| 2 | |
| 3 | import ( |
| 4 | "encoding/json" |
| 5 | "errors" |
| 6 | "net/http" |
| 7 | "net/http/httptest" |
| 8 | "path/filepath" |
| 9 | "sync" |
| 10 | "testing" |
| 11 | |
| 12 | "reasonix/internal/agent" |
| 13 | "reasonix/internal/control" |
| 14 | "reasonix/internal/event" |
| 15 | ) |
| 16 | |
| 17 | type takeoverReturnServers struct { |
| 18 | old, new *httptest.Server |
| 19 | mu sync.Mutex |
| 20 | oldEnd, newEnd []string |
| 21 | } |
| 22 | |
| 23 | func (s *takeoverReturnServers) recordEnd(t *testing.T, r *http.Request, old bool) { |
| 24 | var body struct { |
| 25 | MirrorID string `json:"mirrorId"` |
| 26 | } |
| 27 | if err := json.NewDecoder(r.Body).Decode(&body); err != nil { |
| 28 | t.Error(err) |
| 29 | } |
| 30 | s.mu.Lock() |
| 31 | defer s.mu.Unlock() |
| 32 | if old { |
| 33 | s.oldEnd = append(s.oldEnd, body.MirrorID) |
| 34 | } else { |
| 35 | s.newEnd = append(s.newEnd, body.MirrorID) |
| 36 | } |
| 37 | } |
| 38 | |
| 39 | func newTakeoverReturnServers(t *testing.T, path string) *takeoverReturnServers { |
| 40 | t.Helper() |
| 41 | s := &takeoverReturnServers{} |
| 42 | s.new = httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { |
| 43 | switch r.URL.Path { |
| 44 | case "/auth/token": |
| 45 | w.WriteHeader(http.StatusNoContent) |
| 46 | case "/adopt": |
| 47 | _ = json.NewEncoder(w).Encode(cliTakeoverGrant{SessionPath: path, MirrorID: "mirror-new", ReturnHandoffID: "return-new", SourceWriterID: "serve-new", TargetWriterID: agent.SessionWriterID()}) |
| 48 | case "/mirror-end": |
| 49 | s.recordEnd(t, r, false) |
| 50 | w.WriteHeader(http.StatusNoContent) |
| 51 | default: |
| 52 | w.WriteHeader(http.StatusNotFound) |
| 53 | } |
| 54 | })) |
| 55 | s.old = httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { |
| 56 | if r.URL.Path == "/mirror-end" { |
| 57 | s.recordEnd(t, r, true) |
| 58 | w.WriteHeader(http.StatusNoContent) |
| 59 | return |
| 60 | } |
| 61 | w.WriteHeader(http.StatusUnauthorized) |
| 62 | })) |
| 63 | previous := discoverCLIServesForTakeover |
| 64 | discoverCLIServesForTakeover = func() []cliServeRecord { return []cliServeRecord{{base: s.new.URL, token: "fresh"}} } |
| 65 | t.Cleanup(func() { discoverCLIServesForTakeover = previous; s.old.Close(); s.new.Close() }) |
| 66 | return s |
| 67 | } |
| 68 | |
| 69 | type takeoverReturnEnv struct { |
| 70 | source, target string |
| 71 | servers *takeoverReturnServers |
| 72 | leases *control.SessionLeaseKeeper |
| 73 | manager *cliTakeoverManager |
| 74 | oldBinding *cliTakeoverBinding |
| 75 | } |
| 76 | |
| 77 | func newTakeoverReturnEnv(t *testing.T) *takeoverReturnEnv { |
| 78 | t.Helper() |
| 79 | dir := t.TempDir() |
| 80 | e := &takeoverReturnEnv{source: filepath.Join(dir, "source.jsonl"), target: filepath.Join(dir, "target.jsonl")} |
| 81 | e.servers = newTakeoverReturnServers(t, e.source) |
| 82 | e.leases = control.NewSessionLeaseKeeper() |
| 83 | t.Cleanup(e.leases.Release) |
| 84 | if err := e.leases.Rebind(e.source); err != nil { |
| 85 | t.Fatal(err) |
| 86 | } |
| 87 | e.oldBinding = &cliTakeoverBinding{path: e.source, record: cliServeRecord{base: e.servers.old.URL}, client: e.servers.old.Client(), grant: cliTakeoverGrant{MirrorID: "mirror-old", SourceWriterID: "serve-old", ReturnHandoffID: "return-old"}} |
| 88 | e.manager = newCLITakeoverManager(&takeoverRecordSink{}, e.leases) |
| 89 | e.manager.binding, e.manager.revision = e.oldBinding, 1 |
| 90 | e.manager.Emit(event.Event{Kind: event.Text, Text: "flush before return"}) |
| 91 | return e |
| 92 | } |
| 93 | |
| 94 | func (e *takeoverReturnEnv) assertReturned(t *testing.T) { |
| 95 | t.Helper() |
| 96 | if got := e.leases.HeldPath(); got != agent.CanonicalSessionPath(e.target) { |
| 97 | t.Fatalf("held path = %q, want target", got) |
| 98 | } |
| 99 | info, err := agent.LoadSessionLeaseInfo(e.source) |
| 100 | if err != nil || info == nil || info.HandoffTo != "serve-new" || info.HandoffID != "return-new" { |
| 101 | t.Fatalf("source reservation = %+v, error=%v", info, err) |
| 102 | } |
| 103 | e.servers.mu.Lock() |
| 104 | defer e.servers.mu.Unlock() |
| 105 | if len(e.servers.oldEnd) != 0 || len(e.servers.newEnd) != 1 || e.servers.newEnd[0] != "mirror-new" { |
| 106 | t.Fatalf("mirror-end old=%v new=%v", e.servers.oldEnd, e.servers.newEnd) |
| 107 | } |
| 108 | } |
| 109 | |
| 110 | func (e *takeoverReturnEnv) assertActiveLeaseReturned(t *testing.T) { |
| 111 | t.Helper() |
| 112 | if got := e.leases.HeldPath(); got != "" { |
| 113 | t.Fatalf("held path after return = %q, want none", got) |
| 114 | } |
| 115 | info, err := agent.LoadSessionLeaseInfo(e.source) |
| 116 | if err != nil || info == nil || info.HandoffTo != "serve-new" || info.HandoffID != "return-new" { |
| 117 | t.Fatalf("source reservation = %+v, error=%v", info, err) |
| 118 | } |
| 119 | binding, _, _, _ := e.manager.snapshot() |
| 120 | if binding != nil { |
| 121 | t.Fatalf("active binding survived return: %+v", binding.grant) |
| 122 | } |
| 123 | e.servers.mu.Lock() |
| 124 | defer e.servers.mu.Unlock() |
| 125 | if len(e.servers.oldEnd) != 0 || len(e.servers.newEnd) != 1 || e.servers.newEnd[0] != "mirror-new" { |
| 126 | t.Fatalf("mirror-end old=%v new=%v", e.servers.oldEnd, e.servers.newEnd) |
| 127 | } |
| 128 | } |
| 129 | |
| 130 | func TestCLITakeoverManagerReturnsRefreshedMirrorGeneration(t *testing.T) { |
| 131 | for _, tc := range []struct { |
| 132 | name string |
| 133 | run func(*takeoverReturnEnv) error |
| 134 | }{ |
| 135 | {name: "RebindAway", run: func(e *takeoverReturnEnv) error { |
| 136 | handled, err := e.manager.RebindAway(e.target) |
| 137 | if !handled && err == nil { |
| 138 | return errors.New("RebindAway did not handle the mirror") |
| 139 | } |
| 140 | return err |
| 141 | }}, |
| 142 | {name: "commitPriorMirror", run: func(e *takeoverReturnEnv) error { |
| 143 | previous, err := e.leases.RebindDetaching(e.target) |
| 144 | if err != nil { |
| 145 | return err |
| 146 | } |
| 147 | next := &cliTakeoverBinding{path: e.target, previous: previous, priorMirror: e.oldBinding} |
| 148 | if err := next.commitPrevious(e.manager); err != nil { |
| 149 | return err |
| 150 | } |
| 151 | if next.previous != nil { |
| 152 | return errors.New("detached source keeper retained") |
| 153 | } |
| 154 | return nil |
| 155 | }}, |
| 156 | } { |
| 157 | t.Run(tc.name, func(t *testing.T) { |
| 158 | e := newTakeoverReturnEnv(t) |
| 159 | if err := tc.run(e); err != nil { |
| 160 | t.Fatal(err) |
| 161 | } |
| 162 | e.assertReturned(t) |
| 163 | }) |
| 164 | } |
| 165 | } |
| 166 | |
| 167 | func TestCLITakeoverManagerReturnFailureKeepsRefreshedMirrorActive(t *testing.T) { |
| 168 | e := newTakeoverReturnEnv(t) |
| 169 | wantErr := errors.New("injected reverse reservation failure") |
| 170 | err := e.manager.returnCurrentMirror(e.source, func(current *cliTakeoverBinding) error { |
| 171 | if current.grant.MirrorID != "mirror-new" || current.grant.SourceWriterID != "serve-new" || current.grant.ReturnHandoffID != "return-new" { |
| 172 | t.Fatalf("return binding = %+v", current.grant) |
| 173 | } |
| 174 | return wantErr |
| 175 | }) |
| 176 | if !errors.Is(err, wantErr) { |
| 177 | t.Fatalf("return error = %v", err) |
| 178 | } |
| 179 | binding, _, _, revision := e.manager.snapshot() |
| 180 | e.manager.mu.Lock() |
| 181 | queued := e.manager.queue.Len() |
| 182 | e.manager.mu.Unlock() |
| 183 | if binding == nil || binding.grant.MirrorID != "mirror-new" || revision <= 1 || e.manager.Returned() || queued != 1 { |
| 184 | t.Fatalf("binding=%+v revision=%d returned=%v queued=%d", binding, revision, e.manager.Returned(), queued) |
| 185 | } |
| 186 | if got := e.leases.HeldPath(); got != agent.CanonicalSessionPath(e.source) { |
| 187 | t.Fatalf("source lease moved to %q", got) |
| 188 | } |
| 189 | e.servers.mu.Lock() |
| 190 | defer e.servers.mu.Unlock() |
| 191 | if len(e.servers.oldEnd)+len(e.servers.newEnd) != 0 { |
| 192 | t.Fatalf("mirror ended after reservation failure") |
| 193 | } |
| 194 | } |
| 195 | |
| 196 | func TestCLITakeoverManagerReclaimAndCloseReturnRefreshedGeneration(t *testing.T) { |
| 197 | for _, tc := range []struct { |
| 198 | name string |
| 199 | run func(*takeoverReturnEnv) error |
| 200 | }{ |
| 201 | {name: "reclaim", run: func(e *takeoverReturnEnv) error { |
| 202 | e.manager.reclaiming.Store(true) |
| 203 | return e.manager.returnLeaseFor(e.oldBinding, 1) |
| 204 | }}, |
| 205 | {name: "close", run: func(e *takeoverReturnEnv) error { |
| 206 | return e.manager.Close() |
| 207 | }}, |
| 208 | } { |
| 209 | t.Run(tc.name, func(t *testing.T) { |
| 210 | e := newTakeoverReturnEnv(t) |
| 211 | if err := tc.run(e); err != nil { |
| 212 | t.Fatal(err) |
| 213 | } |
| 214 | e.assertActiveLeaseReturned(t) |
| 215 | }) |
| 216 | } |
| 217 | } |
| 218 |