| 1 | package control |
| 2 | |
| 3 | import ( |
| 4 | "context" |
| 5 | "errors" |
| 6 | "path/filepath" |
| 7 | "reflect" |
| 8 | "sync" |
| 9 | "testing" |
| 10 | "time" |
| 11 | |
| 12 | "reasonix/internal/event" |
| 13 | ) |
| 14 | |
| 15 | func TestAskExactResolutionDeliversOnceAfterReplay(t *testing.T) { |
| 16 | dir := t.TempDir() |
| 17 | asks := make(chan event.Event, 1) |
| 18 | answers := make(chan []event.AskAnswer, 1) |
| 19 | release := make(chan struct{}) |
| 20 | c := newOwnedTestController(t, Options{ |
| 21 | SessionDir: dir, SessionPath: filepath.Join(dir, "session.jsonl"), |
| 22 | Sink: event.FuncSink(func(e event.Event) { |
| 23 | if e.Kind == event.AskRequest { |
| 24 | asks <- e |
| 25 | } |
| 26 | }), |
| 27 | }) |
| 28 | t.Cleanup(func() { c.Cancel(); waitIdle(t, c); c.Close() }) |
| 29 | c.SetTurnEventRoutingMetadata("runtime-ask", "") |
| 30 | c.runGuarded(func(ctx context.Context) error { |
| 31 | got, err := c.Ask(ctx, askProbeQuestions()) |
| 32 | if err != nil { |
| 33 | return err |
| 34 | } |
| 35 | answers <- got |
| 36 | select { |
| 37 | case <-release: |
| 38 | case <-ctx.Done(): |
| 39 | } |
| 40 | return nil |
| 41 | }) |
| 42 | var request event.Event |
| 43 | select { |
| 44 | case request = <-asks: |
| 45 | case <-time.After(5 * time.Second): |
| 46 | t.Fatal("Ask was not published") |
| 47 | } |
| 48 | identity := PromptIdentity{PromptID: request.Ask.ID, TurnID: request.TurnID, RuntimeEpoch: "runtime-ask", Kind: PromptAsk} |
| 49 | if identity.TurnID == "" || identity.TurnID != c.RuntimeStatus().TurnID { |
| 50 | t.Fatalf("Ask has no current turn identity: %+v", identity) |
| 51 | } |
| 52 | var replay event.Event |
| 53 | c.ReplayPendingPromptsTo(event.FuncSink(func(e event.Event) { replay = e })) |
| 54 | if replay.Ask.ID != identity.PromptID || replay.TurnID != identity.TurnID { |
| 55 | t.Fatalf("replayed Ask changed identity: %+v", replay) |
| 56 | } |
| 57 | want := []event.AskAnswer{{QuestionID: request.Ask.Questions[0].ID, Selected: []string{"custom answer"}}} |
| 58 | answer := PromptAnswer{Questions: want} |
| 59 | stale := identity |
| 60 | stale.RuntimeEpoch = "old-runtime" |
| 61 | if err := c.ResolvePromptExact(stale, answer); !errors.Is(err, ErrPromptStaleRuntime) { |
| 62 | t.Fatalf("stale runtime = %v", err) |
| 63 | } |
| 64 | stale = identity |
| 65 | stale.TurnID = "old-turn" |
| 66 | if err := c.ResolvePromptExact(stale, answer); !errors.Is(err, ErrPromptStaleTurn) { |
| 67 | t.Fatalf("stale turn = %v", err) |
| 68 | } |
| 69 | start := make(chan struct{}) |
| 70 | results := make(chan error, 2) |
| 71 | var wg sync.WaitGroup |
| 72 | for range 2 { |
| 73 | wg.Go(func() { <-start; results <- c.ResolvePromptExact(identity, answer) }) |
| 74 | } |
| 75 | close(start) |
| 76 | wg.Wait() |
| 77 | close(results) |
| 78 | var succeeded int |
| 79 | for err := range results { |
| 80 | switch err { |
| 81 | case nil: |
| 82 | succeeded++ |
| 83 | default: |
| 84 | t.Fatalf("current Ask answer rejected: %v", err) |
| 85 | } |
| 86 | } |
| 87 | if succeeded != 2 { |
| 88 | t.Fatalf("answer results: %d idempotent successes, want 2", succeeded) |
| 89 | } |
| 90 | conflict := PromptAnswer{Questions: []event.AskAnswer{{QuestionID: request.Ask.Questions[0].ID, Selected: []string{"different"}}}} |
| 91 | if err := c.ResolvePromptExact(identity, conflict); !errors.Is(err, ErrPromptAlreadyResolved) { |
| 92 | t.Fatalf("conflicting late answer = %v, want ErrPromptAlreadyResolved", err) |
| 93 | } |
| 94 | select { |
| 95 | case got := <-answers: |
| 96 | if !reflect.DeepEqual(got, want) { |
| 97 | t.Fatalf("Ask returned %+v, want %+v", got, want) |
| 98 | } |
| 99 | case <-time.After(5 * time.Second): |
| 100 | t.Fatal("accepted answer did not unblock Ask") |
| 101 | } |
| 102 | if pending := c.PendingPromptIdentities(); len(pending) != 0 { |
| 103 | t.Fatalf("resolved Ask remains pending: %+v", pending) |
| 104 | } |
| 105 | records, err := c.TurnEventsAfter(0) |
| 106 | if err != nil { |
| 107 | t.Fatal(err) |
| 108 | } |
| 109 | answered := 0 |
| 110 | for _, record := range records { |
| 111 | if record.Kind == "prompt_answered" { |
| 112 | answered++ |
| 113 | } |
| 114 | } |
| 115 | if answered != 1 { |
| 116 | t.Fatalf("durable answers = %d, want exactly one", answered) |
| 117 | } |
| 118 | close(release) |
| 119 | waitIdle(t, c) |
| 120 | } |
| 121 | |
| 122 | func TestAskExactSkipCancelsTurn(t *testing.T) { |
| 123 | dir := t.TempDir() |
| 124 | asks := make(chan event.Event, 1) |
| 125 | done := make(chan event.Event, 1) |
| 126 | c := newOwnedTestController(t, Options{ |
| 127 | SessionDir: dir, SessionPath: filepath.Join(dir, "session.jsonl"), |
| 128 | Sink: event.FuncSink(func(e event.Event) { |
| 129 | switch e.Kind { |
| 130 | case event.AskRequest: |
| 131 | asks <- e |
| 132 | case event.TurnDone: |
| 133 | done <- e |
| 134 | } |
| 135 | }), |
| 136 | }) |
| 137 | t.Cleanup(func() { c.Cancel(); waitIdle(t, c); c.Close() }) |
| 138 | c.runner = &askBlockingRunner{c: c} |
| 139 | c.SetTurnEventRoutingMetadata("runtime-skip", "") |
| 140 | c.Send("ask user") |
| 141 | var request event.Event |
| 142 | select { |
| 143 | case request = <-asks: |
| 144 | case <-time.After(5 * time.Second): |
| 145 | t.Fatal("Ask was not published") |
| 146 | } |
| 147 | identity := PromptIdentity{PromptID: request.Ask.ID, TurnID: request.TurnID, RuntimeEpoch: "runtime-skip", Kind: PromptAsk} |
| 148 | if err := c.ResolvePromptExact(identity, PromptAnswer{}); err != nil { |
| 149 | t.Fatalf("skip Ask: %v", err) |
| 150 | } |
| 151 | if terminal := waitTurnDoneEvent(t, done); !terminal.Cancelled { |
| 152 | t.Fatalf("empty answer did not cancel the turn: %+v", terminal) |
| 153 | } |
| 154 | waitIdle(t, c) |
| 155 | if c.PendingPrompt() || len(c.PendingPromptIdentities()) != 0 { |
| 156 | t.Fatal("skipped Ask remains pending") |
| 157 | } |
| 158 | } |
| 159 |