| 1 | package serve |
| 2 | |
| 3 | import ( |
| 4 | "context" |
| 5 | "encoding/json" |
| 6 | "io" |
| 7 | "net/http" |
| 8 | "net/http/httptest" |
| 9 | "path/filepath" |
| 10 | "slices" |
| 11 | "strings" |
| 12 | "testing" |
| 13 | |
| 14 | "reasonix/internal/boot" |
| 15 | "reasonix/internal/config" |
| 16 | "reasonix/internal/control" |
| 17 | "reasonix/internal/provider" |
| 18 | "reasonix/internal/session" |
| 19 | ) |
| 20 | |
| 21 | // appendForkTurn commits one turn into a live exclusive session. A completed |
| 22 | // turn is the only forkable boundary; the open variant exists to be refused. |
| 23 | func appendForkTurn(t *testing.T, service *session.Service, ref session.SessionRef, turnID string, open bool) { |
| 24 | t.Helper() |
| 25 | runtime, ok := service.Runtime(ref) |
| 26 | if !ok { |
| 27 | t.Fatal("current session has no runtime") |
| 28 | } |
| 29 | reply, _ := json.Marshal(map[string]any{"message": provider.Message{ID: "m1", Role: provider.RoleUser, Content: "one"}}) |
| 30 | events := []session.Event{ |
| 31 | {Kind: "turn/start"}, {Kind: "message/complete", Payload: reply}, |
| 32 | {Kind: "turn/end", Payload: json.RawMessage(`{"status":"completed"}`)}, |
| 33 | } |
| 34 | if open { |
| 35 | events = []session.Event{{Kind: "turn/start"}} |
| 36 | } |
| 37 | if _, err := runtime.Session().Append(t.Context(), session.Batch{OperationID: turnID, TurnID: turnID, Events: events}); err != nil { |
| 38 | t.Fatal(err) |
| 39 | } |
| 40 | } |
| 41 | |
| 42 | func forkResponseBody(t *testing.T, resp *http.Response) string { |
| 43 | t.Helper() |
| 44 | body, err := io.ReadAll(resp.Body) |
| 45 | if err != nil { |
| 46 | t.Fatal(err) |
| 47 | } |
| 48 | return string(body) |
| 49 | } |
| 50 | |
| 51 | func forkCreateJSON(ref session.SessionRef, target session.ForkTarget, operationID string) string { |
| 52 | boundary := target.BoundarySequence |
| 53 | if boundary == 0 { |
| 54 | boundary = 1 |
| 55 | } |
| 56 | body, _ := json.Marshal(map[string]any{"sourceSessionId": ref.SessionID, "turnId": target.TurnID, |
| 57 | "boundarySequence": boundary, "operationId": operationID}) |
| 58 | return string(body) |
| 59 | } |
| 60 | |
| 61 | func postForkJSON(t *testing.T, baseURL string, ref session.SessionRef, body string) *http.Response { |
| 62 | t.Helper() |
| 63 | req, _ := http.NewRequest(http.MethodPost, baseURL+"/fork-session", strings.NewReader(body)) |
| 64 | req.Header.Set("Content-Type", "application/json") |
| 65 | req.Header.Set(expectedSessionIDHeader, ref.SessionID) |
| 66 | resp, err := http.DefaultClient.Do(req) |
| 67 | if err != nil { |
| 68 | t.Fatal(err) |
| 69 | } |
| 70 | return resp |
| 71 | } |
| 72 | |
| 73 | func getForkTargets(t *testing.T, baseURL string, refs ...session.SessionRef) forkTargetsResponse { |
| 74 | t.Helper() |
| 75 | req, _ := http.NewRequest(http.MethodGet, baseURL+"/fork-targets", nil) |
| 76 | if len(refs) > 0 { |
| 77 | req.Header.Set(expectedSessionIDHeader, refs[0].SessionID) |
| 78 | } |
| 79 | resp, err := http.DefaultClient.Do(req) |
| 80 | if err != nil { |
| 81 | t.Fatal(err) |
| 82 | } |
| 83 | defer resp.Body.Close() |
| 84 | if resp.StatusCode != http.StatusOK { |
| 85 | t.Fatalf("fork targets status = %d: %s", resp.StatusCode, forkResponseBody(t, resp)) |
| 86 | } |
| 87 | var payload forkTargetsResponse |
| 88 | if err := json.NewDecoder(resp.Body).Decode(&payload); err != nil { |
| 89 | t.Fatal(err) |
| 90 | } |
| 91 | return payload |
| 92 | } |
| 93 | |
| 94 | func TestForkTargetsRouteEncodesEmptyTargetSetAsArray(t *testing.T) { |
| 95 | srv, _, _, ref := newExclusiveSessionServe(t) |
| 96 | server := httptest.NewServer(srv.Handler()) |
| 97 | defer server.Close() |
| 98 | |
| 99 | req, _ := http.NewRequest(http.MethodGet, server.URL+"/fork-targets", nil) |
| 100 | req.Header.Set(expectedSessionIDHeader, ref.SessionID) |
| 101 | resp, err := http.DefaultClient.Do(req) |
| 102 | if err != nil { |
| 103 | t.Fatal(err) |
| 104 | } |
| 105 | defer resp.Body.Close() |
| 106 | if resp.StatusCode != http.StatusOK { |
| 107 | t.Fatalf("fork targets status = %d", resp.StatusCode) |
| 108 | } |
| 109 | var raw map[string]json.RawMessage |
| 110 | if err := json.Unmarshal([]byte(forkResponseBody(t, resp)), &raw); err != nil { |
| 111 | t.Fatal(err) |
| 112 | } |
| 113 | if got := string(raw["targets"]); got != "[]" { |
| 114 | t.Fatalf("empty targets encode as %s, want []", got) |
| 115 | } |
| 116 | // A null array decodes into a nil slice, so the decoded value is the check |
| 117 | // that a client can always map or measure the list. |
| 118 | payload := getForkTargets(t, server.URL, ref) |
| 119 | if payload.Targets == nil || len(payload.Targets) != 0 || payload.Verifiable { |
| 120 | t.Fatalf("empty target set = %+v", payload) |
| 121 | } |
| 122 | } |
| 123 | |
| 124 | func TestForkRoutesRequireAndEnforceSessionFence(t *testing.T) { |
| 125 | srv, _, service, ref := newExclusiveSessionServe(t) |
| 126 | appendForkTurn(t, service, ref, "turn-1", false) |
| 127 | server := httptest.NewServer(srv.Handler()) |
| 128 | defer server.Close() |
| 129 | |
| 130 | get, _ := http.NewRequest(http.MethodGet, server.URL+"/fork-targets", nil) |
| 131 | missing, err := http.DefaultClient.Do(get) |
| 132 | if err != nil { |
| 133 | t.Fatal(err) |
| 134 | } |
| 135 | defer missing.Body.Close() |
| 136 | var refusal forkErrorResponse |
| 137 | if err := json.NewDecoder(missing.Body).Decode(&refusal); err != nil || missing.StatusCode != http.StatusBadRequest || |
| 138 | refusal.Code != "fork_unavailable" || refusal.Reason != session.ForkStaleSource { |
| 139 | t.Fatalf("missing fence status=%d refusal=%+v err=%v", missing.StatusCode, refusal, err) |
| 140 | } |
| 141 | |
| 142 | target := getForkTargets(t, server.URL, ref).Targets[0] |
| 143 | staleRef := ref |
| 144 | staleRef.SessionID = "another-session" |
| 145 | stale := postForkJSON(t, server.URL, staleRef, forkCreateJSON(ref, target, "stale-fence")) |
| 146 | defer stale.Body.Close() |
| 147 | refusal = forkErrorResponse{} |
| 148 | if err := json.NewDecoder(stale.Body).Decode(&refusal); err != nil || stale.StatusCode != http.StatusConflict || |
| 149 | refusal.Code != "fork_unavailable" || refusal.Reason != session.ForkStaleSource { |
| 150 | t.Fatalf("stale fence status=%d refusal=%+v err=%v", stale.StatusCode, refusal, err) |
| 151 | } |
| 152 | } |
| 153 | |
| 154 | func TestForkRoutesReportLegacySessionWithoutTurnRecords(t *testing.T) { |
| 155 | srv := newBrokerTestServer(t, boot.Options{}) |
| 156 | server := httptest.NewServer(srv.Handler()) |
| 157 | defer server.Close() |
| 158 | |
| 159 | payload := getForkTargets(t, server.URL) |
| 160 | if payload.Targets == nil || len(payload.Targets) != 0 || payload.Verifiable { |
| 161 | t.Fatalf("legacy target set = %+v", payload) |
| 162 | } |
| 163 | resp := postRuntimeJSON(t, server.URL+"/fork-session", `{"sourceSessionId":"legacy","turnId":"turn-1","boundarySequence":1,"operationId":"legacy-op"}`) |
| 164 | defer resp.Body.Close() |
| 165 | if resp.StatusCode != http.StatusNotImplemented { |
| 166 | t.Fatalf("legacy fork session status = %d: %s", resp.StatusCode, forkResponseBody(t, resp)) |
| 167 | } |
| 168 | } |
| 169 | |
| 170 | func TestForkTargetsRouteListsCompletedAndOpenTurns(t *testing.T) { |
| 171 | srv, _, service, ref := newExclusiveSessionServe(t) |
| 172 | appendForkTurn(t, service, ref, "turn-1", false) |
| 173 | appendForkTurn(t, service, ref, "turn-2", true) |
| 174 | server := httptest.NewServer(srv.Handler()) |
| 175 | defer server.Close() |
| 176 | |
| 177 | payload := getForkTargets(t, server.URL, ref) |
| 178 | if !payload.Verifiable || len(payload.Targets) != 2 { |
| 179 | t.Fatalf("target set = %+v", payload) |
| 180 | } |
| 181 | completed, open := payload.Targets[0], payload.Targets[1] |
| 182 | if completed.TurnID != "turn-1" || completed.TurnNumber != 1 || !completed.Available || completed.Reason != "" { |
| 183 | t.Fatalf("completed target = %+v", completed) |
| 184 | } |
| 185 | if open.TurnID != "turn-2" || open.TurnNumber != 2 || open.Available || open.Reason != session.ForkTurnOpen { |
| 186 | t.Fatalf("open target = %+v", open) |
| 187 | } |
| 188 | } |
| 189 | |
| 190 | func TestForkSessionRouteRejectsMissingOrEmptyTurnID(t *testing.T) { |
| 191 | srv, _, _, ref := newExclusiveSessionServe(t) |
| 192 | server := httptest.NewServer(srv.Handler()) |
| 193 | defer server.Close() |
| 194 | |
| 195 | for _, body := range []string{`{}`, `{"turnId":""}`, `{"turnId":" "}`, `{`, `{"turnId":"t1","name":`} { |
| 196 | resp := postForkJSON(t, server.URL, ref, body) |
| 197 | if resp.StatusCode != http.StatusBadRequest { |
| 198 | t.Fatalf("body %q status = %d: %s", body, resp.StatusCode, forkResponseBody(t, resp)) |
| 199 | } |
| 200 | resp.Body.Close() |
| 201 | } |
| 202 | } |
| 203 | |
| 204 | func TestForkSessionRouteCreatesChildWithoutSwitchingParent(t *testing.T) { |
| 205 | srv, ctrl, service, ref := newExclusiveSessionServe(t) |
| 206 | appendForkTurn(t, service, ref, "turn-1", false) |
| 207 | parentPath := ctrl.SessionPath() |
| 208 | server := httptest.NewServer(srv.Handler()) |
| 209 | defer server.Close() |
| 210 | target := getForkTargets(t, server.URL, ref).Targets[0] |
| 211 | |
| 212 | resp := postForkJSON(t, server.URL, ref, forkCreateJSON(ref, target, "create-op")) |
| 213 | if resp.StatusCode != http.StatusOK { |
| 214 | t.Fatalf("fork session status = %d: %s", resp.StatusCode, forkResponseBody(t, resp)) |
| 215 | } |
| 216 | var created forkSessionResponse |
| 217 | if err := json.NewDecoder(resp.Body).Decode(&created); err != nil { |
| 218 | t.Fatal(err) |
| 219 | } |
| 220 | resp.Body.Close() |
| 221 | if created.SessionID == "" || created.SessionID == ref.SessionID || created.TurnID != "turn-1" || created.TurnNumber != 1 { |
| 222 | t.Fatalf("created child = %+v", created) |
| 223 | } |
| 224 | // The response never moves the serve session: the parent keeps its identity |
| 225 | // and its path, and the X-Reasonix-Session-ID header still names the parent. |
| 226 | if after, ok := ctrl.SessionRef(); !ok || after != ref || ctrl.SessionPath() != parentPath { |
| 227 | t.Fatalf("parent identity = %+v (ok=%v path=%q)", after, ok, ctrl.SessionPath()) |
| 228 | } |
| 229 | if got := resp.Header.Get(sessionIDHeader); got != "" { |
| 230 | t.Fatalf("fork response session id header = %q, want none", got) |
| 231 | } |
| 232 | child, err := service.Query().Snapshot(t.Context(), session.SessionRef{HostID: created.HostID, SessionID: created.SessionID}) |
| 233 | if err != nil { |
| 234 | t.Fatal(err) |
| 235 | } |
| 236 | inherited := false |
| 237 | for _, message := range child.Projection.Messages { |
| 238 | if message.ID == "m1" { |
| 239 | inherited = true |
| 240 | } |
| 241 | } |
| 242 | if len(child.Projection.Turns) != 1 || !inherited { |
| 243 | t.Fatalf("child projection = %+v", child.Projection) |
| 244 | } |
| 245 | |
| 246 | retry := postForkJSON(t, server.URL, ref, forkCreateJSON(ref, target, "create-op")) |
| 247 | defer retry.Body.Close() |
| 248 | if retry.StatusCode != http.StatusOK { |
| 249 | t.Fatalf("retry status = %d: %s", retry.StatusCode, forkResponseBody(t, retry)) |
| 250 | } |
| 251 | var again forkSessionResponse |
| 252 | if err := json.NewDecoder(retry.Body).Decode(&again); err != nil { |
| 253 | t.Fatal(err) |
| 254 | } |
| 255 | if again.SessionID != created.SessionID { |
| 256 | t.Fatalf("retry created %q, want the existing child %q", again.SessionID, created.SessionID) |
| 257 | } |
| 258 | } |
| 259 | |
| 260 | func TestForkSessionRouteReportsUnavailableTurnReason(t *testing.T) { |
| 261 | srv, _, service, ref := newExclusiveSessionServe(t) |
| 262 | appendForkTurn(t, service, ref, "turn-1", false) |
| 263 | appendForkTurn(t, service, ref, "turn-2", true) |
| 264 | server := httptest.NewServer(srv.Handler()) |
| 265 | defer server.Close() |
| 266 | targets := getForkTargets(t, server.URL, ref) |
| 267 | |
| 268 | resp := postForkJSON(t, server.URL, ref, forkCreateJSON(ref, targets.Targets[1], "open-op")) |
| 269 | defer resp.Body.Close() |
| 270 | body := forkResponseBody(t, resp) |
| 271 | if resp.StatusCode < 400 || resp.StatusCode >= 500 { |
| 272 | t.Fatalf("refused fork status = %d, want a 4xx: %s", resp.StatusCode, body) |
| 273 | } |
| 274 | var refusal forkErrorResponse |
| 275 | if err := json.Unmarshal([]byte(body), &refusal); err != nil || refusal.Code != "fork_unavailable" || refusal.Reason != session.ForkTurnOpen { |
| 276 | t.Fatalf("refusal %q does not carry structured reason %q: %+v err=%v", body, session.ForkTurnOpen, refusal, err) |
| 277 | } |
| 278 | } |
| 279 | |
| 280 | func TestServerAdvertisesSessionForkTargetsOnlyForExclusiveSessions(t *testing.T) { |
| 281 | service, err := session.NewService("serve", session.NewFilesystemPersistence(filepath.Join(t.TempDir(), "sessions-v4"))) |
| 282 | if err != nil { |
| 283 | t.Fatal(err) |
| 284 | } |
| 285 | t.Cleanup(func() { |
| 286 | if err := service.Shutdown(context.Background()); err != nil { |
| 287 | t.Errorf("shutdown session service: %v", err) |
| 288 | } |
| 289 | }) |
| 290 | ctrl := control.New(control.Options{SessionService: service, ExclusiveSession: true}) |
| 291 | defer ctrl.Close() |
| 292 | srv := New(ctrl, NewBroadcaster(), config.ServeConfig{AuthMode: "token", Token: "secret"}) |
| 293 | if !slices.Contains(srv.capabilities(), capabilityForkTargetsV1) { |
| 294 | t.Fatalf("exclusive capabilities = %v", srv.capabilities()) |
| 295 | } |
| 296 | if legacy := newBrokerTestServer(t, boot.Options{}); slices.Contains(legacy.capabilities(), capabilityForkTargetsV1) { |
| 297 | t.Fatalf("legacy capabilities = %v", legacy.capabilities()) |
| 298 | } |
| 299 | |
| 300 | server := httptest.NewServer(srv.Handler()) |
| 301 | defer server.Close() |
| 302 | resp, err := http.Post(server.URL+"/auth/token", "application/json", strings.NewReader(`{"token":"secret"}`)) |
| 303 | if err != nil { |
| 304 | t.Fatal(err) |
| 305 | } |
| 306 | defer resp.Body.Close() |
| 307 | if resp.StatusCode != http.StatusNoContent { |
| 308 | t.Fatalf("handshake status = %d", resp.StatusCode) |
| 309 | } |
| 310 | advertised := strings.Split(resp.Header.Get(capabilitiesHeader), ",") |
| 311 | if !slices.Contains(advertised, capabilityForkTargetsV1) { |
| 312 | t.Fatalf("capabilities header = %v", advertised) |
| 313 | } |
| 314 | } |
| 315 |