diff --git a/daemon/internal/api/api.go b/daemon/internal/api/api.go index 18a409c..cec05cd 100644 --- a/daemon/internal/api/api.go +++ b/daemon/internal/api/api.go @@ -760,7 +760,7 @@ func (a *API) streamRun(w http.ResponseWriter, r *http.Request) { return } } - run, ok := a.store.GetRun(id) + _, ok := a.store.GetRun(id) if !ok { writeJSON(w, http.StatusNotFound, map[string]string{"error": "not found"}) return @@ -772,7 +772,7 @@ func (a *API) streamRun(w http.ResponseWriter, r *http.Request) { } sseHeaders(w) - replay, ch, done, cancel := a.hub.SubscribeAfter(id, after) + replay, ch, done, cancel := a.sup.SubscribeRun(id, after) defer cancel() send := func(ev agent.Event) { @@ -793,12 +793,7 @@ func (a *API) streamRun(w http.ResponseWriter, r *http.Request) { send(ev) } send(agent.Event{Type: agent.EventStatus, Meta: map[string]any{"stream": "ready"}}) - // Terminal run: nothing more will be published. Close a phantom topic (freshly created - // by Subscribe after a restart) so it gets reaped, and end the stream. - if done || run.Status != "running" { - if !done { - a.hub.Close(id) - } + if done { return } diff --git a/daemon/internal/orchestrator/transcript.go b/daemon/internal/orchestrator/transcript.go index fe8df43..b6ab80a 100644 --- a/daemon/internal/orchestrator/transcript.go +++ b/daemon/internal/orchestrator/transcript.go @@ -15,6 +15,20 @@ type RunSnapshot struct { stream.Snapshot } +// SubscribeRun keeps the producer's lifetime authoritative for stream completion. +// A terminal record can be saved before its final events/notification are sent; +// subscribers must drain through hub.Close, never close that topic themselves. +// The start/teardown lock also prevents a new run from being mistaken for an old +// run whose replay buffer expired or disappeared after a daemon restart. +func (s *Supervisor) SubscribeRun(id string, after int64) (replay []agent.Event, ch <-chan agent.Event, done bool, cancel func()) { + s.mu.Lock() + defer s.mu.Unlock() + if _, active := s.cancels[id]; active { + return s.hub.SubscribeAfter(id, after) + } + return s.hub.SubscribeRetainedAfter(id, after) +} + func (s *Supervisor) Snapshot(id string) (RunSnapshot, bool) { run, ok := s.store.GetRun(id) if !ok { diff --git a/daemon/internal/stream/hub.go b/daemon/internal/stream/hub.go index d2805bb..28e5329 100644 --- a/daemon/internal/stream/hub.go +++ b/daemon/internal/stream/hub.go @@ -128,7 +128,23 @@ func (h *Hub) Subscribe(id string) (replay []agent.Event, ch <-chan agent.Event, // SubscribeAfter atomically replays events newer than after and subscribes to live // updates. Sequence numbers survive reconnects for the lifetime of the run topic. func (h *Hub) SubscribeAfter(id string, after int64) (replay []agent.Event, ch <-chan agent.Event, done bool, cancel func()) { - t := h.get(id) + return h.get(id).subscribeAfter(after) +} + +// SubscribeRetainedAfter subscribes without creating a topic. Completed runs can +// outlive their replay buffer (or a daemon restart); they must not acquire an +// empty live channel that no producer will ever close. +func (h *Hub) SubscribeRetainedAfter(id string, after int64) (replay []agent.Event, ch <-chan agent.Event, done bool, cancel func()) { + h.mu.Lock() + t := h.topics[id] + h.mu.Unlock() + if t == nil { + return nil, nil, true, func() {} + } + return t.subscribeAfter(after) +} + +func (t *topic) subscribeAfter(after int64) (replay []agent.Event, ch <-chan agent.Event, done bool, cancel func()) { t.mu.Lock() defer t.mu.Unlock() start := min(max(after, 0), int64(len(t.buf))) diff --git a/daemon/sdk/run.go b/daemon/sdk/run.go index fed77bc..37afa23 100644 --- a/daemon/sdk/run.go +++ b/daemon/sdk/run.go @@ -69,7 +69,7 @@ func (r *Run) Stream(ctx context.Context, opts ...StreamOption) iter.Seq[Event] o(&cfg) } return func(yield func(Event) bool) { - replay, ch, done, cancel := r.core.hub.SubscribeAfter(r.data.ID, cfg.after) + replay, ch, done, cancel := r.core.sup.SubscribeRun(r.data.ID, cfg.after) defer cancel() if cfg.openSentinel { @@ -83,11 +83,7 @@ func (r *Run) Stream(ctx context.Context, opts ...StreamOption) iter.Seq[Event] return } } - record, exists := r.core.store.GetRun(r.data.ID) - if done || !exists || record.Status != "running" { - if !done { - r.core.hub.Close(r.data.ID) - } + if done { return } for { @@ -111,7 +107,7 @@ func (r *Run) Stream(ctx context.Context, opts ...StreamOption) iter.Seq[Event] // invoke. Prefer Stream; reach for this only to multiplex a run's events with other channels in a // select. func (r *Run) Subscribe() (replay []Event, ch <-chan Event, done bool, cancel func()) { - return r.core.hub.Subscribe(r.data.ID) + return r.core.sup.SubscribeRun(r.data.ID, 0) } // ---- control --------------------------------------------------------------- diff --git a/daemon/sdk/run_stream_test.go b/daemon/sdk/run_stream_test.go new file mode 100644 index 0000000..0361c6f --- /dev/null +++ b/daemon/sdk/run_stream_test.go @@ -0,0 +1,176 @@ +package mindwire + +import ( + "context" + "encoding/json" + "net/http" + "net/http/httptest" + "strings" + "sync" + "testing" + "time" + + "github.com/oblien/mindwire/daemon/internal/api" + "github.com/oblien/mindwire/daemon/internal/session" +) + +// Completion happens inside the consumer callback, after subscription but before +// Stream starts reading its live channel. No scheduler timing or sleeps required. +func TestStreamDrainsEventsWhenRunFinishesDuringReplay(t *testing.T) { + for _, sentinel := range []bool{false, true} { + name := "replay" + if sentinel { + name = "open sentinel" + } + t.Run(name, func(t *testing.T) { + c := newFakeClient(t, nil) + ctx, cancel := context.WithTimeout(t.Context(), 5*time.Second) + defer cancel() + run, err := c.Turn(ctx, TurnRequest{ChatID: "stream-completion", Message: "partial"}) + if err != nil { + t.Fatal(err) + } + // Make the first text event part of the next subscription's replay. + for range run.Stream(ctx) { + break + } + var options []StreamOption + if sentinel { + options = append(options, WithOpenSentinel()) + } + var text string + var sequence int64 + finished := false + for event := range run.Stream(ctx, options...) { + if !finished { + finished = true + if err := run.Cancel(); err != nil { + t.Fatal(err) + } + c.core.sup.Wait() + } + if event.Type == EventText { + text += event.Text + if event.Sequence != sequence+1 { + t.Fatalf("event sequence jumped from %d to %d", sequence, event.Sequence) + } + sequence = event.Sequence + } + } + if err := ctx.Err(); err != nil { + t.Fatal(err) + } + if text != "Partial reply" || sequence != 2 { + t.Fatalf("completion dropped queued events: text=%q sequence=%d", text, sequence) + } + }) + } +} + +type streamTestNotifier func(context.Context, Notification) error + +func (f streamTestNotifier) Notify(ctx context.Context, n Notification) error { return f(ctx, n) } + +type streamFlushRecorder struct { + *httptest.ResponseRecorder + onFlush func() +} + +func (r *streamFlushRecorder) Flush() { r.ResponseRecorder.Flush(); r.onFlush() } + +// A terminal status is persisted before notification delivery. Both transports +// must leave the producer alive and deliver its final event after replay. +func TestStreamIncludesEventsPublishedAfterTerminalStatus(t *testing.T) { + for _, transport := range []string{"sdk", "http"} { + t.Run(transport, func(t *testing.T) { + reached, release := make(chan struct{}), make(chan struct{}) + var once sync.Once + unblock := func() { once.Do(func() { close(release) }) } + defer unblock() + c := newFakeClient(t, streamTestNotifier(func(ctx context.Context, _ Notification) error { + close(reached) + select { + case <-release: + return nil + case <-ctx.Done(): + return ctx.Err() + } + })) + ctx, cancel := context.WithTimeout(t.Context(), 5*time.Second) + defer cancel() + run, err := c.Turn(ctx, TurnRequest{ChatID: "finishing", Message: "script"}) + if err != nil { + t.Fatal(err) + } + select { + case <-reached: + case <-ctx.Done(): + t.Fatal(ctx.Err()) + } + if record, err := run.Refresh(); err != nil || record.Status != "done" { + t.Fatalf("terminal record not saved before notification: %+v %v", record, err) + } + var events []Event + if transport == "sdk" { + for event := range run.Stream(ctx) { + unblock() + events = append(events, event) + } + } else { + mux := http.NewServeMux() + api.New(c.core.store, c.core.hub, c.core.sup).Register(mux) + response := &streamFlushRecorder{ResponseRecorder: httptest.NewRecorder(), onFlush: unblock} + mux.ServeHTTP(response, httptest.NewRequestWithContext(ctx, "GET", "/runs/"+run.ID()+"/stream", nil)) + if response.Code != http.StatusOK { + t.Fatalf("stream HTTP %d", response.Code) + } + for _, line := range strings.Split(response.Body.String(), "\n") { + if data, ok := strings.CutPrefix(line, "data: "); ok { + var event Event + if err := json.Unmarshal([]byte(data), &event); err != nil { + t.Fatal(err) + } + events = append(events, event) + } + } + } + if ctx.Err() != nil { + t.Fatal(ctx.Err()) + } + if len(events) == 0 { + t.Fatal("stream returned no events") + } + last := events[len(events)-1] + if last.Type != EventStatus || last.Meta["notify"] != "sent ✓" || last.Replay || last.Sequence != 5 { + t.Fatalf("stream lost its post-terminal event: %+v", events) + } + }) + } +} + +func TestStreamClosesAfterRestartWithoutCreatingLiveTopic(t *testing.T) { + c := newFakeClient(t, nil) + if err := c.core.store.SaveRun(session.Run{ID: "restored", ChatID: "chat", Agent: "fake", Status: "done"}); err != nil { + t.Fatal(err) + } + run, err := c.Run("restored") + if err != nil { + t.Fatal(err) + } + ctx, cancel := context.WithTimeout(t.Context(), time.Second) + defer cancel() + for event := range run.Stream(ctx) { + t.Fatalf("restored stream invented an event: %+v", event) + } + if ctx.Err() != nil { + t.Fatal("restored stream waited on a topic with no producer") + } + replay, ch, done, cancelSubscription := run.Subscribe() + defer cancelSubscription() + if len(replay) != 0 || ch != nil || !done { + t.Fatal("restored run subscription did not terminate") + } + if _, retained := c.core.hub.Snapshot(run.ID()); retained { + t.Fatal("restored stream created a live topic") + } +}