diff --git a/cmd/root/run.go b/cmd/root/run.go index c6a0e2da36..3d44219383 100644 --- a/cmd/root/run.go +++ b/cmd/root/run.go @@ -1163,7 +1163,6 @@ func (f *runExecFlags) runLeanTUI(ctx context.Context, rt runtime.Runtime, sess opts = append(opts, app.WithTitleGenerator(gen)) } a := app.New(ctx, rt, sess, opts...) - a.Start(ctx) firstMessage, err := readInitialMessage(args) if err != nil { diff --git a/cmd/root/run_listen.go b/cmd/root/run_listen.go index 908244cd6d..63ca7fbc70 100644 --- a/cmd/root/run_listen.go +++ b/cmd/root/run_listen.go @@ -53,7 +53,7 @@ func (f *runExecFlags) recallCoordinatorOpt(ctx context.Context, rt runtime.Runt // plane's per-session event log (GET /api/sessions/:id/events). func registerAppEventSource(sm *server.SessionManager, sessionID string, a *app.App) { sm.RegisterEventSource(sessionID, func(ctx context.Context, send func(any)) { - a.SubscribeWith(ctx, func(msg tea.Msg) { + a.SubscribeReliable(ctx, func(msg tea.Msg) { if ev, ok := msg.(runtime.Event); ok { send(ev) } diff --git a/docs/features/api-server/index.md b/docs/features/api-server/index.md index 549636cbca..2d0a13dbc9 100644 --- a/docs/features/api-server/index.md +++ b/docs/features/api-server/index.md @@ -64,7 +64,7 @@ For an agent loaded from a remote HTTP(S) configuration source, endpoints that n | `GET` | `/api/sessions/:id` | Get a session by ID (messages, tokens, permissions) | | `GET` | `/api/sessions/:id/status` | Lightweight runtime state (streaming, title, agent, tokens). Requires an attached runtime. | | `GET` | `/api/sessions/:id/snapshot` | Full state in one call (stored fields + runtime state + `last_event_seq`) for gapless resync — see [Reconnecting without gaps](#reconnecting-without-gaps). | -| `GET` | `/api/sessions/:id/events` | Live session event stream (SSE) with sequence numbers and replay. Available for a run attached via [`--listen`](#listen), or once a session has raised at least one out-of-band event (e.g. a background job's elicitation, answered via `POST .../elicitation`), which creates a session-scoped event log on demand carrying such out-of-band events — see [Session event stream](#session-event-stream-and-reconnection) for what each kind of log contains. | +| `GET` | `/api/sessions/:id/events` | Live session event stream (SSE) with sequence numbers and replay. Available for a run attached via [`--listen`](#listen), or once a session has raised at least one out-of-band event (e.g. a background job's elicitation or an idle-session recall), which creates a session-scoped event log on demand carrying such out-of-band events — see [Session event stream](#session-event-stream-and-reconnection) for what each kind of log contains. | | `DELETE` | `/api/sessions/:id` | Delete a session | | `PATCH` | `/api/sessions/:id/title` | Update session title | | `PATCH` | `/api/sessions/:id/permissions` | Update session permissions | @@ -277,6 +277,17 @@ $ curl -X POST http://127.0.0.1:8080/api/sessions/$SID/followup \ ## Session event stream and reconnection +The Go remote client reconnects sequenced `/events` streams from the last received +ID. After a replay gap, the remote TUI waits for an idle snapshot and recovers +missing plain assistant text by message ID. Tool results, reasoning, and transient +interaction events are not reconstructed; a warning reports this limitation. +Ambiguous histories fail visibly instead of replaying answers twice. + +Foreground POST run streams are not automatically resubmitted after interruption: +retrying a run could repeat tool side effects. A stream ending without a root +`stream_stopped` event reports an incomplete-response error. Streaming requests +are not subject to the metadata client's 30-second total timeout. + `GET /api/sessions/:id/events` is a **Server-Sent Events** stream of the session's runtime events — `stream_started`, `agent_choice`, `tool_call`, `session_title`, `token_usage`, `stream_stopped`, and so on. Unlike the @@ -284,7 +295,7 @@ per-request stream returned by the agent-execution endpoint, it is session-scoped and survives across turns, so a client can watch a session for its whole lifetime. It is available for a run attached via [`--listen`](#listen), and — since a session-scoped event log is created on -demand the first time a session raises an out-of-band event, such as an +demand the first time a session runs an idle recall or raises an out-of-band event, such as an `elicitation_request` from a background job — for any API-created session that has produced at least one (see the [Sessions endpoint table](#sessions) above). The two kinds of log differ in diff --git a/pkg/app/app.go b/pkg/app/app.go index 99a39c1b8a..3420bb61c4 100644 --- a/pkg/app/app.go +++ b/pkg/app/app.go @@ -57,10 +57,11 @@ type App struct { snapshotController builtins.SnapshotController // Drives /undo, /snapshots, /reset; nil for runtimes that don't capture snapshots streamGuard sync.Locker // Held for the duration of every direct RunStream call; nil when not attached to a SessionManager (see WithStreamGuard) - startOnce sync.Once - subsMu sync.Mutex - subs []chan tea.Msg - fanoutOnce sync.Once + eventGeneration atomic.Uint64 + startOnce sync.Once + subsMu sync.Mutex + subs []*eventSubscriber + fanoutOnce sync.Once } // Opt is an option for creating a new App. @@ -176,6 +177,7 @@ func (a *App) Start(ctx context.Context) { // a.session concurrently, and it re-emits startup info for the new // session itself. sess := a.session + startupCtx := a.eventContext(ctx) go func() { startupEvents := make(chan runtime.Event, 10) go func() { @@ -183,7 +185,7 @@ func (a *App) Start(ctx context.Context) { a.runtime.EmitStartupInfo(ctx, sess, runtime.NewChannelSink(startupEvents)) }() for event := range startupEvents { - a.sendEvent(ctx, event) + a.sendEvent(startupCtx, event) } }() @@ -196,9 +198,22 @@ func (a *App) Start(ctx context.Context) { // Forward events surfaced from detached background work (token usage // from background agent tasks) so the sidebar and agent inspector can // account for background agents' context usage. - a.runtime.OnBackgroundEvent(func(event runtime.Event) { - a.sendEvent(ctx, event) - }) + backgroundCtx := ctx + if _, ok := a.runtime.(interface{ RetireBackgroundEvents() }); ok { + backgroundCtx = a.eventContext(ctx) + } + if contextual, ok := a.runtime.(interface { + OnBackgroundEventWithContext(handler func(context.Context, runtime.Event)) + }); ok { + contextual.OnBackgroundEventWithContext(func(origin context.Context, event runtime.Event) { + if origin == nil { + origin = ctx + } + a.sendEvent(context.WithoutCancel(origin), event) + }) + } else { + a.runtime.OnBackgroundEvent(func(event runtime.Event) { a.sendEvent(backgroundCtx, event) }) + } // Forward elicitation requests raised anywhere in the runtime — // including background-job (run_background_agent) sub-sessions whose @@ -212,9 +227,18 @@ func (a *App) Start(ctx context.Context) { // don't mirror it (RemoteRuntime, whose OnElicitationRequest below is // a no-op) deliver elicitations only through that RunStream copy, // which those loops forward unfiltered (#3584 review). - a.runtime.OnElicitationRequest(func(event runtime.Event) { - a.sendEvent(ctx, event) - }) + if contextual, ok := a.runtime.(interface { + OnElicitationRequestWithContext(handler func(context.Context, runtime.Event)) + }); ok { + contextual.OnElicitationRequestWithContext(func(origin context.Context, event runtime.Event) { + if origin == nil { + origin = ctx + } + a.sendEvent(origin, event) + }) + } else { + a.runtime.OnElicitationRequest(func(event runtime.Event) { a.sendEvent(ctx, event) }) + } }) } @@ -419,6 +443,9 @@ func (a *App) SkillCommandFork(_ context.Context, input string) (skillName, task // opens the child; the sub-session's first user message is the expanded // SKILL.md body. Companion of SkillCommandFork. func (a *App) RunSkillFork(ctx context.Context, cancel context.CancelFunc, skillName, task string, _ []messages.Attachment) { + if _, ok := ctx.Value(eventGenerationKey{}).(uint64); !ok { + ctx = a.eventContext(ctx) + } a.cancel = cancel // Snapshot the session like Run does: the goroutines below outlive any // concurrent ReplaceSession and must keep working against this session. @@ -561,6 +588,9 @@ func (a *App) EmitStartupInfo(ctx context.Context, events chan runtime.Event) { // Run one agent loop func (a *App) Run(ctx context.Context, cancel context.CancelFunc, message string, attachments []messages.Attachment) { + if _, ok := ctx.Value(eventGenerationKey{}).(uint64); !ok { + ctx = a.eventContext(ctx) + } a.cancel = cancel sess := a.session @@ -750,7 +780,7 @@ func (a *App) sendEvent(ctx context.Context, event tea.Msg) { default: } select { - case a.events <- event: + case a.events <- a.stampEvent(ctx, event): case <-ctx.Done(): case <-a.eventsDone: } @@ -805,6 +835,9 @@ func mustSkipMirroredElicitation(rt runtime.Runtime) bool { // suppress the pre-StreamStarted re-emitted user message; Run and // RunWithMessage pass nil. func (a *App) forwardRunStreamEvents(ctx context.Context, sess *session.Session, ch <-chan runtime.Event, filter func(event runtime.Event) (forward bool)) { + if _, ok := ctx.Value(eventGenerationKey{}).(uint64); !ok { + ctx = a.eventContext(ctx) + } skipMirroredElicitation := mustSkipMirroredElicitation(a.runtime) // sawRootStop/sawRootError/agentName drive the #4136 fallback below: the @@ -952,6 +985,9 @@ func (a *App) processInlineAttachment(att messages.Attachment, textBuilder *stri // re-emission; genuine user messages injected mid-run (steer / follow-up) // arrive after StreamStarted and are forwarded normally. func (a *App) Retry(ctx context.Context, cancel context.CancelFunc) { + if _, ok := ctx.Value(eventGenerationKey{}).(uint64); !ok { + ctx = a.eventContext(ctx) + } a.cancel = cancel sess := a.session @@ -980,6 +1016,9 @@ func (a *App) Retry(ctx context.Context, cancel context.CancelFunc) { // RunWithMessage runs the agent loop with a pre-constructed message. // This is used for special cases like image attachments. func (a *App) RunWithMessage(ctx context.Context, cancel context.CancelFunc, msg *session.Message) { + if _, ok := ctx.Value(eventGenerationKey{}).(uint64); !ok { + ctx = a.eventContext(ctx) + } a.cancel = cancel sess := a.session @@ -1012,6 +1051,9 @@ func (a *App) RunWithMessage(ctx context.Context, cancel context.CancelFunc, msg } func (a *App) RunBangCommand(ctx context.Context, command string) { + if _, ok := ctx.Value(eventGenerationKey{}).(uint64); !ok { + ctx = a.eventContext(ctx) + } command = strings.TrimSpace(command) if command == "" { a.sendEvent(ctx, runtime.ShellOutput("Error: empty command")) @@ -1078,8 +1120,9 @@ func (a *App) InjectUserMessage(ctx context.Context, content string) { // Slow subscribers drop events rather than block the bus. func (a *App) SubscribeWith(ctx context.Context, send func(tea.Msg)) { ch := make(chan tea.Msg, subscriberBufferSize) - a.addSubscriber(ch) - defer a.removeSubscriber(ch) + sub := &eventSubscriber{ch: ch} + a.addSubscriber(sub) + defer a.removeSubscriber(sub) a.fanoutOnce.Do(a.startFanOut) @@ -1095,38 +1138,100 @@ func (a *App) SubscribeWith(ctx context.Context, send func(tea.Msg)) { } } +// SubscribeReliable delivers every event in order, buffering without a size +// limit while send is blocked. Use it for the TUI: dropped deltas cannot recover. +// Cancellation discards pending events; send must unblock when ctx is canceled. +func (a *App) SubscribeReliable(ctx context.Context, send func(tea.Msg), opts ...SubscribeOption) { + queue := newEventQueue() + sub := &eventSubscriber{queue: queue} + for _, opt := range opts { + opt(sub) + } + a.addSubscriber(sub) + if sub.registered != nil { + sub.registered() + } + cleanup := func() { + queue.close() + a.removeSubscriber(sub) + } + finished := make(chan struct{}) + defer close(finished) + defer cleanup() + go func() { + select { + case <-ctx.Done(): + case <-a.eventsDone: + case <-finished: + return + } + cleanup() + }() + + a.fanoutOnce.Do(a.startFanOut) + + for { + msg, ok := queue.next(ctx, a.eventsDone) + if !ok { + return + } + send(msg) + } +} + const subscriberBufferSize = 1024 -func (a *App) addSubscriber(ch chan tea.Msg) { +func (a *App) addSubscriber(sub *eventSubscriber) { a.subsMu.Lock() defer a.subsMu.Unlock() - a.subs = append(a.subs, ch) + a.subs = append(a.subs, sub) } -func (a *App) removeSubscriber(ch chan tea.Msg) { +func (a *App) removeSubscriber(sub *eventSubscriber) { a.subsMu.Lock() defer a.subsMu.Unlock() - a.subs = slices.DeleteFunc(a.subs, func(c chan tea.Msg) bool { return c == ch }) + a.subs = slices.DeleteFunc(a.subs, func(s *eventSubscriber) bool { return s == sub }) } // startFanOut runs once per App. It throttles the raw events channel and -// scatters every message to all currently-registered subscribers. Sends are -// non-blocking; if a subscriber's buffer is full the event is dropped for -// that subscriber so one slow consumer cannot stall the others. -// -// Turn-boundary events are the exception: dropping a stream_started or -// stream_stopped skews a consumer's turn accounting for good (the SSE replay -// buffer never sees the event, so reconnecting cannot recover it). For those, -// the oldest pending message — almost always a content delta, which the next -// delta supersedes — is evicted to make room instead. +// scatters every message to all currently-registered subscribers. Reliable +// subscribers queue every event; best-effort subscribers drop on overflow. +// Turn boundaries evict the oldest pending best-effort delivery to preserve +// turn accounting even when that subscriber falls behind. func (a *App) startFanOut() { throttled := a.throttleEvents(a.ctx(), a.events) go func() { for msg := range throttled { + generation := a.eventGeneration.Load() + if stamped, ok := msg.(generationEvent); ok { + generation = stamped.generation + if stamped.generation != a.eventGeneration.Load() { + continue + } + msg = stamped.inner + } a.subsMu.Lock() subs := slices.Clone(a.subs) a.subsMu.Unlock() - for _, ch := range subs { + for _, sub := range subs { + delivery := msg + if sub.prepareGeneration != nil { + delivery = sub.prepareGeneration(msg, generation) + if delivery == nil { + continue + } + } + if sub.prepare != nil { + delivery = sub.prepare(msg) + if delivery == nil { + continue + } + } + if sub.queue != nil { + sub.queue.pushGeneration(delivery, generation) + continue + } + ch := sub.ch select { case ch <- msg: default: @@ -1155,8 +1260,9 @@ func (a *App) startFanOut() { // isTurnBoundaryEvent reports whether msg is one of the events consumers use // to track turn state (running/waiting/failed/paused) and identity (title). -// These are low-frequency and irrecoverable when lost, unlike the content -// deltas that dominate the stream, so the fan-out prefers them on overflow. +// These are low-frequency and keep best-effort consumers' turn accounting +// usable on overflow. Content deltas are also irrecoverable; TUI consumers +// must use SubscribeReliable instead. func isTurnBoundaryEvent(msg tea.Msg) bool { switch msg.(type) { case *runtime.StreamStartedEvent, @@ -1279,6 +1385,7 @@ func (a *App) ResumeElicitation(ctx context.Context, action tools.ElicitationAct } func (a *App) NewSession() { + a.RetireEvents() if a.cancel != nil { a.cancel() a.cancel = nil @@ -1318,6 +1425,9 @@ func (a *App) NewSession() { // reEmitStartupInfo resets and re-emits startup info (agent, team, tools) // through the events channel so the sidebar updates. func (a *App) reEmitStartupInfo(ctx context.Context) { + if _, ok := ctx.Value(eventGenerationKey{}).(uint64); !ok { + ctx = a.eventContext(ctx) + } a.runtime.ResetStartupInfo() // Snapshot before handing off to the background goroutine so a later // ReplaceSession cannot race with this read. @@ -1342,7 +1452,7 @@ func (a *App) pumpToEvents(ctx context.Context, emit func(runtime.EventSink)) { continue } select { - case a.events <- event: + case a.events <- a.stampEvent(ctx, event): case <-ctx.Done(): case <-a.eventsDone: default: @@ -1437,6 +1547,9 @@ type liveSessionCompactor interface { // runtime cannot target live sessions (e.g. remote runtimes), or the // runtime's rejection for unknown/finished sessions and duplicate requests. func (a *App) CompactLiveSession(ctx context.Context, sessionID, additionalPrompt string) error { + if _, ok := ctx.Value(eventGenerationKey{}).(uint64); !ok { + ctx = a.eventContext(ctx) + } compactor, ok := a.runtime.(liveSessionCompactor) if !ok { return fmt.Errorf("targeted session compaction: %w", runtime.ErrUnsupported) @@ -1717,6 +1830,9 @@ func (a *App) IsReadOnly() bool { } func (a *App) CompactSession(ctx context.Context, cancel context.CancelFunc, additionalPrompt string) { + if _, ok := ctx.Value(eventGenerationKey{}).(uint64); !ok { + ctx = a.eventContext(ctx) + } a.cancel = cancel sess := a.session @@ -1769,6 +1885,7 @@ func (a *App) SessionStore() session.Store { // so the sidebar displays the agent and tool information. // If the session has stored model overrides, they are applied to the runtime. func (a *App) ReplaceSession(ctx context.Context, sess *session.Session) { + a.RetireEvents() if a.cancel != nil { a.cancel() a.cancel = nil @@ -1867,6 +1984,9 @@ func (a *App) throttleEvents(ctx context.Context, in <-chan tea.Msg) <-chan tea. // shouldThrottle determines if an event should be buffered/throttled func (a *App) shouldThrottle(msg tea.Msg) bool { + if stamped, ok := msg.(generationEvent); ok { + msg = stamped.inner + } switch msg.(type) { case *runtime.AgentChoiceEvent: return true @@ -1895,6 +2015,22 @@ func (a *App) mergeEvents(events []tea.Msg) []tea.Msg { result := make([]tea.Msg, 0, len(events)) for i := 0; i < len(events); i++ { + if first, ok := events[i].(generationEvent); ok { + run := []tea.Msg{first.inner} + n := i + 1 + for ; n < len(events); n++ { + next, ok := events[n].(generationEvent) + if !ok || next.generation != first.generation { + break + } + run = append(run, next.inner) + } + for _, merged := range a.mergeEvents(run) { + result = append(result, generationEvent{generation: first.generation, inner: merged}) + } + i = n - 1 + continue + } switch ev := events[i].(type) { case *runtime.AgentChoiceEvent: merged, consumed := mergeAgentChoiceRun(ev, events[i+1:]) @@ -2144,6 +2280,7 @@ func (a *App) generateTitle(ctx context.Context, sess *session.Session, userMess // RegenerateSessionTitle triggers AI-based title regeneration for the current session. // Returns ErrTitleGenerating if a title generation is already in progress. func (a *App) RegenerateSessionTitle(ctx context.Context) error { + ctx = a.eventContext(ctx) if a.session == nil { return errors.New("no active session") } diff --git a/pkg/app/app_test.go b/pkg/app/app_test.go index e8b8f1a45f..1f23fe3adc 100644 --- a/pkg/app/app_test.go +++ b/pkg/app/app_test.go @@ -238,6 +238,7 @@ func TestApp_Retry_SuppressesReEmittedUserMessage(t *testing.T) { for !sawStreamStopped { select { case ev := <-events: + ev = unwrappedTestEvent(ev) switch e := ev.(type) { case *runtime.UserMessageEvent: userMessages = append(userMessages, e.Message) @@ -297,6 +298,7 @@ func TestApp_Start_ForwardsBackgroundEvents(t *testing.T) { select { case msg := <-events: + msg = unwrappedTestEvent(msg) assert.Equal(t, usage, msg, "the background event must reach the app's event stream unchanged") case <-time.After(2 * time.Second): t.Fatal("timed out waiting for the forwarded background event") @@ -362,6 +364,7 @@ func TestApp_Start_ForwardsElicitationRequests(t *testing.T) { select { case msg := <-events: + msg = unwrappedTestEvent(msg) assert.Equal(t, ev, msg, "the elicitation request must reach the app's event stream unchanged") case <-time.After(2 * time.Second): t.Fatal("timed out waiting for the forwarded elicitation request") @@ -566,6 +569,7 @@ func TestApp_UpdateSessionTitle(t *testing.T) { // Check that an event was emitted select { case event := <-events: + event = unwrappedTestEvent(event) titleEvent, ok := event.(*runtime.SessionTitleEvent) require.True(t, ok, "should emit SessionTitleEvent") assert.Equal(t, "New Title", titleEvent.Title) @@ -719,6 +723,7 @@ func TestApp_SubscribeWith_FanOutToMultipleSubscribers(t *testing.T) { for _, ch := range []chan tea.Msg{a, b} { select { case msg := <-ch: + msg = unwrappedTestEvent(msg) ev, ok := msg.(*runtime.SessionTitleEvent) require.True(t, ok) assert.Equal(t, "hello", ev.Title) @@ -805,6 +810,7 @@ func TestApp_InjectUserMessage(t *testing.T) { select { case msg := <-events: + msg = unwrappedTestEvent(msg) sendMsg, ok := msg.(messages.SendMsg) require.True(t, ok, "should emit a SendMsg, got %T", msg) assert.Equal(t, "do the thing", sendMsg.Content) @@ -999,6 +1005,7 @@ func TestApp_CompactLiveSession_BridgesEventsIntoStream(t *testing.T) { select { case msg := <-app.events: + msg = unwrappedTestEvent(msg) evt, ok := msg.(*runtime.SessionCompactionEvent) require.True(t, ok, "expected SessionCompactionEvent, got %T", msg) assert.Equal(t, "child-1", evt.SessionID) @@ -1096,6 +1103,7 @@ func countElicitationDeliveries(t *testing.T, events <-chan tea.Msg) int { for { select { case msg := <-events: + msg = unwrappedTestEvent(msg) if _, ok := msg.(*runtime.ElicitationRequestEvent); ok { n++ } @@ -1310,6 +1318,7 @@ func collectUntilQuiet(t *testing.T, events <-chan tea.Msg) []tea.Msg { for { select { case msg := <-events: + msg = unwrappedTestEvent(msg) collected = append(collected, msg) case <-time.After(collectUntilQuietWindow): return collected @@ -1486,3 +1495,10 @@ func TestForwardRunStreamEvents_SynthesizesRootStreamStopped(t *testing.T) { func (m *mockRuntime) ReadSkillContent(context.Context, *session.Session, string) (string, error) { return "", nil } + +func unwrappedTestEvent(msg tea.Msg) tea.Msg { + if event, ok := msg.(generationEvent); ok { + return event.inner + } + return msg +} diff --git a/pkg/app/event_generation.go b/pkg/app/event_generation.go new file mode 100644 index 0000000000..9de097df34 --- /dev/null +++ b/pkg/app/event_generation.go @@ -0,0 +1,50 @@ +package app + +import ( + "context" + + tea "charm.land/bubbletea/v2" + + "github.com/docker/docker-agent/pkg/runtime" +) + +type eventGenerationKey struct{} + +type generationEvent struct { + generation uint64 + inner tea.Msg +} + +// RetireEvents discards deliveries belonging to a replaced conversation. +func (a *App) RetireEvents() { + generation := a.eventGeneration.Add(1) + a.subsMu.Lock() + for _, sub := range a.subs { + if sub.queue != nil { + sub.queue.retire(generation) + } + } + a.subsMu.Unlock() + if retire, ok := a.runtime.(interface{ RetireBackgroundEvents() }); ok { + retire.RetireBackgroundEvents() + ctx := a.eventContext(a.ctx()) + a.runtime.OnBackgroundEvent(func(event runtime.Event) { a.sendEvent(ctx, event) }) + } +} + +// IsEventGeneration checks producer identity at the final consumer boundary. +func (a *App) IsEventGeneration(generation uint64) bool { + return generation == a.eventGeneration.Load() +} + +func (a *App) eventContext(ctx context.Context) context.Context { + return context.WithValue(ctx, eventGenerationKey{}, a.eventGeneration.Load()) +} + +func (a *App) stampEvent(ctx context.Context, event tea.Msg) tea.Msg { + generation, ok := ctx.Value(eventGenerationKey{}).(uint64) + if !ok { + generation = a.eventGeneration.Load() + } + return generationEvent{generation: generation, inner: event} +} diff --git a/pkg/app/event_generation_test.go b/pkg/app/event_generation_test.go new file mode 100644 index 0000000000..87f26ac66e --- /dev/null +++ b/pkg/app/event_generation_test.go @@ -0,0 +1,156 @@ +package app + +import ( + "context" + "testing" + "testing/synctest" + + tea "charm.land/bubbletea/v2" + "github.com/stretchr/testify/require" + + "github.com/docker/docker-agent/pkg/runtime" + "github.com/docker/docker-agent/pkg/session" +) + +func TestReplacementDiscardsOldThrottledContentAndLateStop(t *testing.T) { + t.Parallel() + synctest.Test(t, func(t *testing.T) { + ctx, cancel := context.WithCancel(t.Context()) + defer cancel() + a := New(ctx, &mockRuntime{}, session.New()) + var got []tea.Msg + ready := make(chan struct{}) + go a.SubscribeReliable(ctx, func(msg tea.Msg) { got = append(got, msg) }, WithSubscriptionReady(func() { close(ready) })) + <-ready + oldCtx := a.eventContext(ctx) + a.sendEvent(oldCtx, runtime.AgentChoice("root", a.session.ID, "OLD-ANSWER", "old")) + synctest.Wait() + require.Empty(t, got, "fixture must leave text in the throttle buffer") + a.ReplaceSession(ctx, session.New()) + a.sendEvent(context.WithoutCancel(oldCtx), runtime.StreamStopped("old-session", "root", "canceled")) + want := runtime.AgentChoice("root", a.session.ID, "NEW-ANSWER", "new") + a.sendEvent(ctx, want) + a.sendEvent(ctx, runtime.StreamStopped(a.session.ID, "root", "normal")) + synctest.Wait() + require.Len(t, got, 2) + require.Same(t, want, got[0]) + require.Equal(t, "normal", got[1].(*runtime.StreamStoppedEvent).Reason) + }) +} + +func TestRetireEventsDiscardsRawAndThrottledDeliveries(t *testing.T) { + t.Parallel() + synctest.Test(t, func(t *testing.T) { + ctx, cancel := context.WithCancel(t.Context()) + defer cancel() + a := New(ctx, &mockRuntime{}, session.New()) + oldCtx := a.eventContext(ctx) + a.sendEvent(oldCtx, runtime.AgentChoice("root", a.session.ID, "RAW-OLD-ANSWER", "old")) + a.RetireEvents() + var got []tea.Msg + ready := make(chan struct{}) + go a.SubscribeReliable(ctx, func(msg tea.Msg) { got = append(got, msg) }, WithSubscriptionReady(func() { close(ready) })) + <-ready + synctest.Wait() + marker := runtime.StreamStopped(a.session.ID, "root", "normal") + a.sendEvent(ctx, marker) + synctest.Wait() + require.Equal(t, []tea.Msg{marker}, got) + }) +} + +func TestRoutingRetainsOriginWhenRetirementRacesMapping(t *testing.T) { + t.Parallel() + synctest.Test(t, func(t *testing.T) { + ctx, cancel := context.WithCancel(t.Context()) + defer cancel() + a := New(ctx, &mockRuntime{}, session.New()) + mapping := make(chan struct{}) + resume := make(chan struct{}) + ready := make(chan struct{}) + type delivery struct{ generation uint64 } + var got []tea.Msg + go a.SubscribeReliable(ctx, func(msg tea.Msg) { got = append(got, msg) }, WithSubscriptionReady(func() { close(ready) }), WithGenerationEventMapper(func(_ tea.Msg, generation uint64) tea.Msg { + close(mapping) + select { + case <-resume: + case <-ctx.Done(): + return nil + } + return delivery{generation: generation} + })) + <-ready + a.sendEvent(ctx, runtime.StreamStopped(a.session.ID, "root", "normal")) + <-mapping + a.RetireEvents() + close(resume) + synctest.Wait() + require.Len(t, got, 1) + require.False(t, a.IsEventGeneration(got[0].(delivery).generation), "the consumer must reject an event accepted just before retirement") + }) +} + +type contextualEventsRuntime struct { + mockRuntime + + background func(context.Context, runtime.Event) + elicitation func(context.Context, runtime.Event) +} + +func (r *contextualEventsRuntime) OnBackgroundEventWithContext(handler func(context.Context, runtime.Event)) { + r.background = handler +} + +func (r *contextualEventsRuntime) OnElicitationRequestWithContext(handler func(context.Context, runtime.Event)) { + r.elicitation = handler +} + +func TestDetachedCallbacksKeepOriginAcrossReplacement(t *testing.T) { + t.Parallel() + synctest.Test(t, func(t *testing.T) { + ctx, cancel := context.WithCancel(t.Context()) + defer cancel() + rt := &contextualEventsRuntime{} + a := New(ctx, rt, session.New()) + a.Start(ctx) + var got []tea.Msg + ready := make(chan struct{}) + go a.SubscribeReliable(ctx, func(msg tea.Msg) { got = append(got, msg) }, WithSubscriptionReady(func() { close(ready) })) + <-ready + old := a.eventContext(ctx) + a.ReplaceSession(ctx, session.New()) + rt.background(old, runtime.StreamStopped("old-child", "root", "normal")) + rt.elicitation(old, runtime.ElicitationRequest("OLD REQUEST", "form", nil, "", "old", "", "old-child", nil, "root")) + current := a.eventContext(ctx) + request := runtime.ElicitationRequest("CURRENT REQUEST", "form", nil, "", "new", "", "new-child", nil, "root") + rt.elicitation(current, request) + synctest.Wait() + require.Equal(t, []tea.Msg{request}, got) + }) +} + +func TestCanceledDetachedAccountingStillDeliversWithOriginalGeneration(t *testing.T) { + t.Parallel() + synctest.Test(t, func(t *testing.T) { + ctx, cancel := context.WithCancel(t.Context()) + defer cancel() + rt := &contextualEventsRuntime{} + a := New(ctx, rt, session.New()) + a.Start(ctx) + var got []tea.Msg + ready := make(chan struct{}) + go a.SubscribeReliable(ctx, func(msg tea.Msg) { got = append(got, msg) }, WithSubscriptionReady(func() { close(ready) })) + <-ready + origin, stop := context.WithCancel(a.eventContext(ctx)) + stop() + for range 100 { + rt.background(origin, runtime.NewTokenUsageEvent("child", "root", &runtime.Usage{})) + } + synctest.Wait() + require.Len(t, got, 100) + a.RetireEvents() + rt.background(origin, runtime.NewTokenUsageEvent("child", "root", &runtime.Usage{})) + synctest.Wait() + require.Len(t, got, 100) + }) +} diff --git a/pkg/app/fanout_test.go b/pkg/app/fanout_test.go index 08e8fba74a..c0e68cc383 100644 --- a/pkg/app/fanout_test.go +++ b/pkg/app/fanout_test.go @@ -34,7 +34,7 @@ func TestFanOut_TurnBoundaryEventEvictsPendingDelta(t *testing.T) { // A one-slot subscriber makes the overflow deterministic. The subscriber // never reads, standing in for a consumer that fell behind. ch := make(chan tea.Msg, 1) - app.addSubscriber(ch) + app.addSubscriber(&eventSubscriber{ch: ch}) app.fanoutOnce.Do(app.startFanOut) // Fill the subscriber's buffer with a droppable event. @@ -75,11 +75,11 @@ func TestFanOut_DroppableEventIsDroppedOnOverflow(t *testing.T) { } ch := make(chan tea.Msg, 1) - app.addSubscriber(ch) + app.addSubscriber(&eventSubscriber{ch: ch}) // The witness is registered after ch, so once a message reaches it the // fan-out has already made its keep-or-drop decision for ch. witness := make(chan tea.Msg, 16) - app.addSubscriber(witness) + app.addSubscriber(&eventSubscriber{ch: witness}) app.fanoutOnce.Do(app.startFanOut) first := runtime.NewTokenUsageEvent("sess", "root", &runtime.Usage{}) diff --git a/pkg/app/reliable_subscriber_test.go b/pkg/app/reliable_subscriber_test.go new file mode 100644 index 0000000000..d82964b0f8 --- /dev/null +++ b/pkg/app/reliable_subscriber_test.go @@ -0,0 +1,164 @@ +package app + +import ( + "context" + "testing" + "testing/synctest" + + tea "charm.land/bubbletea/v2" + "github.com/stretchr/testify/require" + + "github.com/docker/docker-agent/pkg/runtime" + "github.com/docker/docker-agent/pkg/session" +) + +func TestReliableSubscriberRetainsFinalResponseWhenStalled(t *testing.T) { + t.Parallel() + + synctest.Test(t, func(t *testing.T) { + ctx, cancel := context.WithCancel(t.Context()) + defer cancel() + a := New(ctx, &mockRuntime{}, session.New()) + resume := make(chan struct{}) + var received []tea.Msg + go a.SubscribeReliable(ctx, func(msg tea.Msg) { + select { + case <-resume: + case <-ctx.Done(): + return + } + received = append(received, msg) + }) + var witness []tea.Msg + go a.SubscribeWith(ctx, func(msg tea.Msg) { witness = append(witness, msg) }) + synctest.Wait() + + want := []tea.Msg{runtime.StreamStarted(a.session.ID, "root")} + for range subscriberBufferSize + 1 { + want = append(want, runtime.NewTokenUsageEvent(a.session.ID, "root", &runtime.Usage{})) + } + want = append(want, + runtime.AgentChoiceReasoning("root", a.session.ID, "thinking", "answer"), + runtime.AgentChoice("root", a.session.ID, "FINAL-RESPONSE", "answer"), + runtime.StreamStopped(a.session.ID, "root", "normal"), + ) + for _, msg := range want { + a.sendEvent(ctx, msg) + } + synctest.Wait() + require.Equal(t, want, witness, "a stalled TUI must not block other subscribers") + + close(resume) + synctest.Wait() + require.Len(t, received, len(want), "all content must reach the TUI before completion") + for i, msg := range want { + require.Equal(t, msg, received[i], "delivery %d must preserve event order", i) + } + }) +} + +func TestReliableSubscriberCancellationReleasesBlockedQueue(t *testing.T) { + t.Parallel() + for _, cancelOwner := range []bool{false, true} { + t.Run(map[bool]string{false: "subscriber", true: "owner"}[cancelOwner], func(t *testing.T) { + t.Parallel() + synctest.Test(t, func(t *testing.T) { + ctx, cancel := context.WithCancel(t.Context()) + defer cancel() + a := New(ctx, &mockRuntime{}, session.New()) + a.Start(ctx) + subCtx, cancelSub := context.WithCancel(t.Context()) + defer cancelSub() + release := make(chan struct{}) + defer close(release) + go func() { + a.SubscribeReliable(subCtx, func(tea.Msg) { <-release }) + }() + synctest.Wait() + require.Len(t, a.subs, 1) + queue := a.subs[0].queue + for range subscriberBufferSize + 2 { + a.sendEvent(ctx, runtime.SessionTitle(a.session.ID, "queued")) + } + synctest.Wait() + require.NotEmpty(t, queue.pending) + if cancelOwner { + cancel() + } else { + cancelSub() + } + synctest.Wait() + require.Empty(t, a.subs, "cleanup must not wait for a blocked callback") + require.Empty(t, queue.pending) + queue.push(runtime.SessionTitle(a.session.ID, "late")) + require.Empty(t, queue.pending, "retired fan-out snapshots cannot retain new events") + }) + }) + } +} + +func TestReliableSubscriberMapsBeforeBuffering(t *testing.T) { + t.Parallel() + synctest.Test(t, func(t *testing.T) { + ctx, cancel := context.WithCancel(t.Context()) + defer cancel() + a := New(ctx, &mockRuntime{}, session.New()) + resume := make(chan struct{}) + mapped := make(chan tea.Msg, 4) + type delivery struct { + generation int + msg tea.Msg + } + generation := 1 + var received []tea.Msg + go a.SubscribeReliable(ctx, func(msg tea.Msg) { + select { + case <-resume: + case <-ctx.Done(): + return + } + received = append(received, msg) + }, WithEventMapper(func(msg tea.Msg) tea.Msg { + result := delivery{generation, msg} + mapped <- result + return result + })) + synctest.Wait() + first := runtime.SessionTitle(a.session.ID, "first") + second := runtime.SessionTitle(a.session.ID, "second") + a.sendEvent(ctx, first) + a.sendEvent(ctx, second) + synctest.Wait() + require.Equal(t, delivery{1, first}, <-mapped) + require.Equal(t, delivery{1, second}, <-mapped) + generation = 2 + third := runtime.SessionTitle(a.session.ID, "replacement") + a.sendEvent(ctx, third) + synctest.Wait() + require.Equal(t, delivery{2, third}, <-mapped) + close(resume) + synctest.Wait() + require.Equal(t, []tea.Msg{delivery{1, first}, delivery{1, second}, delivery{2, third}}, received) + }) +} + +func TestReliableSubscriptionReadyPrecedesFirstProducedEvent(t *testing.T) { + t.Parallel() + synctest.Test(t, func(t *testing.T) { + ctx, cancel := context.WithCancel(t.Context()) + defer cancel() + a := New(ctx, &mockRuntime{}, session.New()) + go a.SubscribeWith(ctx, func(tea.Msg) {}) + synctest.Wait() + ready := make(chan struct{}) + var received []tea.Msg + go a.SubscribeReliable(ctx, func(msg tea.Msg) { received = append(received, msg) }, WithSubscriptionReady(func() { close(ready) })) + <-ready + want := runtime.AgentChoice("root", a.session.ID, "FAST-FINAL-ANSWER", "answer") + a.sendEvent(ctx, want) + a.sendEvent(ctx, runtime.StreamStopped(a.session.ID, "root", "normal")) + synctest.Wait() + require.Len(t, received, 2) + require.Same(t, want, received[0]) + }) +} diff --git a/pkg/app/session_lifecycle_test.go b/pkg/app/session_lifecycle_test.go index 5dbb6bed49..60ef40d1ef 100644 --- a/pkg/app/session_lifecycle_test.go +++ b/pkg/app/session_lifecycle_test.go @@ -142,7 +142,7 @@ func TestAppRunKeepsWorkScopedToOriginalSession(t *testing.T) { app.ReplaceSession(t.Context(), newSession) close(rt.release) - event := <-app.events + event := unwrappedTestEvent(<-app.events) stop, ok := event.(*runtime.StreamStoppedEvent) require.True(t, ok) assert.Equal(t, oldSession.ID, stop.SessionID) @@ -275,7 +275,7 @@ func TestGenerateTitleKeepsOriginalSession(t *testing.T) { assert.Equal(t, "Original title", oldSession.TitleSnapshot()) assert.Empty(t, newSession.TitleSnapshot()) - titleEvent, ok := (<-app.events).(*runtime.SessionTitleEvent) + titleEvent, ok := unwrappedTestEvent(<-app.events).(*runtime.SessionTitleEvent) require.True(t, ok) assert.Equal(t, oldSession.ID, titleEvent.SessionID) }) diff --git a/pkg/app/subscriber.go b/pkg/app/subscriber.go new file mode 100644 index 0000000000..f66a7dfda5 --- /dev/null +++ b/pkg/app/subscriber.go @@ -0,0 +1,126 @@ +package app + +import ( + "context" + "sync" + + tea "charm.land/bubbletea/v2" +) + +// eventSubscriber buffers reliable deliveries independently of slow consumers. +// Best-effort subscribers keep their bounded channel and overflow policy. +type eventSubscriber struct { + ch chan tea.Msg + queue *eventQueue + prepare func(tea.Msg) tea.Msg + registered func() + prepareGeneration func(tea.Msg, uint64) tea.Msg +} + +// SubscribeOption configures reliable event delivery. +type SubscribeOption func(*eventSubscriber) + +// WithEventMapper stamps routing metadata before an event is queued. +// The mapper runs on the fan-out goroutine and must not block; nil skips delivery. +func WithEventMapper(mapper func(tea.Msg) tea.Msg) SubscribeOption { + return func(sub *eventSubscriber) { + sub.prepare = mapper + } +} + +// WithGenerationEventMapper preserves producer identity through routing queues. +func WithGenerationEventMapper(mapper func(tea.Msg, uint64) tea.Msg) SubscribeOption { + return func(sub *eventSubscriber) { sub.prepareGeneration = mapper } +} + +// WithSubscriptionReady signals after registration, before any events are delivered. +func WithSubscriptionReady(ready func()) SubscribeOption { + return func(sub *eventSubscriber) { sub.registered = ready } +} + +type queuedEvent struct { + generation uint64 + msg tea.Msg +} + +type eventQueue struct { + mu sync.Mutex + pending []queuedEvent + ready chan struct{} + closed bool +} + +func newEventQueue() *eventQueue { + return &eventQueue{ready: make(chan struct{}, 1)} +} + +func (q *eventQueue) push(msg tea.Msg) { q.pushGeneration(msg, 0) } + +func (q *eventQueue) pushGeneration(msg tea.Msg, generation uint64) { + q.mu.Lock() + defer q.mu.Unlock() + if q.closed { + return + } + q.pending = append(q.pending, queuedEvent{generation: generation, msg: msg}) + select { + case q.ready <- struct{}{}: + default: + } +} + +func (q *eventQueue) next(ctx context.Context, done <-chan struct{}) (tea.Msg, bool) { + for { + select { + case <-ctx.Done(): + return nil, false + case <-done: + return nil, false + default: + } + + q.mu.Lock() + if len(q.pending) > 0 { + msg := q.pending[0].msg + q.pending[0] = queuedEvent{} + q.pending = q.pending[1:] + if len(q.pending) == 0 { + q.pending = nil + } + q.mu.Unlock() + return msg, true + } + q.mu.Unlock() + + select { + case <-ctx.Done(): + return nil, false + case <-done: + return nil, false + case <-q.ready: + } + } +} + +func (q *eventQueue) retire(generation uint64) { + q.mu.Lock() + defer q.mu.Unlock() + kept := q.pending[:0] + for _, event := range q.pending { + if event.generation >= generation { + kept = append(kept, event) + } + } + clear(q.pending[len(kept):]) + q.pending = kept + if len(kept) == 0 { + q.pending = nil + } +} + +func (q *eventQueue) close() { + q.mu.Lock() + defer q.mu.Unlock() + q.closed = true + q.pending = nil +} diff --git a/pkg/app/subscriber_test.go b/pkg/app/subscriber_test.go new file mode 100644 index 0000000000..ef1d20f040 --- /dev/null +++ b/pkg/app/subscriber_test.go @@ -0,0 +1,40 @@ +package app + +import ( + "context" + "testing" + "testing/synctest" + + tea "charm.land/bubbletea/v2" + "github.com/stretchr/testify/require" +) + +func TestEventQueueWakeupAndDrain(t *testing.T) { + t.Parallel() + synctest.Test(t, func(t *testing.T) { + ctx, cancel := context.WithCancel(t.Context()) + defer cancel() + queue := newEventQueue() + received := make(chan tea.Msg, 1) + done := make(chan struct{}) + go func() { + defer close(done) + for { + msg, ok := queue.next(ctx, nil) + if !ok { + return + } + received <- msg + } + }() + for i := range 10 { + synctest.Wait() + queue.push(i) + require.Equal(t, i, <-received) + synctest.Wait() + require.Nil(t, queue.pending, "draining must release the backlog storage") + } + cancel() + <-done + }) +} diff --git a/pkg/chat/visible_content.go b/pkg/chat/visible_content.go new file mode 100644 index 0000000000..1888f846e5 --- /dev/null +++ b/pkg/chat/visible_content.go @@ -0,0 +1,9 @@ +package chat + +import "strings" + +// VisibleAssistantContent preserves the live stream's XML tool-call suppression. +func VisibleAssistantContent(content string) string { + visible, _, _ := strings.Cut(content, "") + return visible +} diff --git a/pkg/chat/visible_content_test.go b/pkg/chat/visible_content_test.go new file mode 100644 index 0000000000..6e19f6361d --- /dev/null +++ b/pkg/chat/visible_content_test.go @@ -0,0 +1,19 @@ +package chat + +import ( + "testing" + + "github.com/stretchr/testify/assert" +) + +func TestVisibleAssistantContentSuppressesToolPayloads(t *testing.T) { + t.Parallel() + for _, test := range []struct{ content, want string }{ + {"safe answer", "safe answer"}, + {`{"arguments":"PRIVATE`, ""}, + {`safe prefix{"arguments":"PRIVATE`, "safe prefix"}, + {"safehiddenalso hidden", "safe"}, + } { + assert.Equal(t, test.want, VisibleAssistantContent(test.content)) + } +} diff --git a/pkg/leantui/assistant_content_test.go b/pkg/leantui/assistant_content_test.go new file mode 100644 index 0000000000..a8f9caf5f1 --- /dev/null +++ b/pkg/leantui/assistant_content_test.go @@ -0,0 +1,99 @@ +package leantui + +import ( + "strings" + "testing" + + "github.com/charmbracelet/x/ansi" + "github.com/stretchr/testify/require" + + "github.com/docker/docker-agent/pkg/chat" + "github.com/docker/docker-agent/pkg/leantui/ui" + "github.com/docker/docker-agent/pkg/runtime" + "github.com/docker/docker-agent/pkg/session" +) + +func leanTranscriptText(m *model) string { + return ansi.Strip(strings.Join(m.screen.Transcript.Lines(100, 0, false, m.sessionState, nil), "\n")) +} + +func TestCanonicalAssistantContentReconcilesLeanTranscript(t *testing.T) { + t.Parallel() + for _, partial := range []string{"", "FINAL-", "ANSWER", "FINAL-ANSWER"} { + t.Run(partial, func(t *testing.T) { + m := bareModel(40) + m.handleEvent(t.Context(), runtime.StreamStarted("session", "root")) + if partial != "" { + m.handleEvent(t.Context(), runtime.AgentChoice("root", "session", partial, "answer")) + } + msg := session.NewAgentMessage("root", &chat.Message{Role: chat.MessageRoleAssistant, MessageID: "answer", Content: "FINAL-ANSWER"}) + for range 2 { + m.handleEvent(t.Context(), runtime.MessageAdded("session", msg, "root")) + } + require.Equal(t, 1, strings.Count(leanTranscriptText(m), "FINAL-ANSWER")) + }) + } +} + +func TestLeanRetryMessageIDsKeepAnswerVisible(t *testing.T) { + t.Parallel() + m := bareModel(40) + m.handleEvent(t.Context(), runtime.AgentChoice("root", "session", "[diagnostic](", "attempt-one")) + m.handleEvent(t.Context(), runtime.AgentChoice("root", "session", "FINAL-ANSWER)", "attempt-two")) + require.Contains(t, leanTranscriptText(m), "FINAL-ANSWER") +} + +func TestLeanInterleavedSessionsDoNotMergeMarkdown(t *testing.T) { + t.Parallel() + m := bareModel(40) + m.handleEvent(t.Context(), runtime.AgentChoice("root", "child", "[diagnostic](", "answer")) + m.handleEvent(t.Context(), runtime.AgentChoice("root", "parent", "FINAL-ANSWER)", "answer")) + require.Contains(t, leanTranscriptText(m), "FINAL-ANSWER") +} + +func TestRetiredLeanEventsCannotAffectReplacement(t *testing.T) { + t.Parallel() + m := bareModel(40) + old := leanEvent{generation: 0, inner: runtime.AgentChoice("root", "same-session", "OLD-ANSWER", "old")} + m.eventGeneration.Add(1) + m.handleEvent(t.Context(), old) + m.handleEvent(t.Context(), leanEvent{generation: 1, inner: runtime.AgentChoice("root", "same-session", "NEW-ANSWER", "new")}) + require.NotContains(t, leanTranscriptText(m), "OLD-ANSWER") + require.Contains(t, leanTranscriptText(m), "NEW-ANSWER") +} + +func TestLeanCanonicalContentCannotRevealSuppressedToolXML(t *testing.T) { + t.Parallel() + m := bareModel(40) + msg := session.NewAgentMessage("root", &chat.Message{Role: chat.MessageRoleAssistant, MessageID: "answer", Content: `safe prefix{"arguments":"PRIVATE-TOOL-ARGUMENT`}) + m.handleEvent(t.Context(), runtime.AgentChoice("root", "session", "safe prefix", "answer")) + m.handleEvent(t.Context(), runtime.MessageAdded("session", msg, "root")) + require.Contains(t, leanTranscriptText(m), "safe prefix") + require.NotContains(t, leanTranscriptText(m), "PRIVATE-TOOL-ARGUMENT") +} + +func TestLeanRecoveryClearsNestedBusyWithoutDrainingQueue(t *testing.T) { + t.Parallel() + m := bareModel(40) + for range 3 { + m.handleEvent(t.Context(), runtime.StreamStarted("session", "root")) + } + m.queue = []ui.PendingUserMessage{{Content: "must not run"}} + m.handleEvent(t.Context(), runtime.SessionRecovered("")) + require.Zero(t, m.streamDepth) + require.False(t, m.busy) + require.Len(t, m.queue, 1) +} + +func TestLeanRecoveryRestoresRootContextUsage(t *testing.T) { + t.Parallel() + m := bareModel(40) + m.handleEvent(t.Context(), runtime.StreamStarted("root", "root")) + m.handleEvent(t.Context(), runtime.NewTokenUsageEvent("root", "root", &runtime.Usage{ContextLength: 100})) + m.handleEvent(t.Context(), runtime.StreamStarted("child", "child")) + m.handleEvent(t.Context(), runtime.NewTokenUsageEvent("child", "child", &runtime.Usage{ContextLength: 999})) + m.handleEvent(t.Context(), runtime.SessionRecovered("")) + m.handleEvent(t.Context(), runtime.StreamStarted("root", "root")) + m.handleEvent(t.Context(), runtime.StreamStopped("root", "root", "normal")) + require.Equal(t, int64(100), m.status.ContextLength) +} diff --git a/pkg/leantui/events.go b/pkg/leantui/events.go index 0f09509c54..854336d2da 100644 --- a/pkg/leantui/events.go +++ b/pkg/leantui/events.go @@ -4,6 +4,7 @@ import ( "context" "time" + "github.com/docker/docker-agent/pkg/chat" "github.com/docker/docker-agent/pkg/leantui/ui" "github.com/docker/docker-agent/pkg/runtime" "github.com/docker/docker-agent/pkg/sound" @@ -17,6 +18,12 @@ import ( // handleEvent applies a single runtime event emitted by the App to the model, // updating the conversation, tool state, status footer, or busy state. func (m *model) handleEvent(ctx context.Context, ev any) { + if routed, ok := ev.(leanEvent); ok { + if (routed.valid != nil && !routed.valid()) || routed.generation != m.eventGeneration.Load() { + return + } + ev = routed.inner + } switch e := ev.(type) { case fileCompletionsLoaded: m.screen.Autocomplete.SetFiles(e) @@ -28,6 +35,7 @@ func (m *model) handleEvent(ctx context.Context, ev any) { m.submitFollowUp(ctx, e.Content) } case *runtime.StreamStartedEvent: + m.contentIdentity.Finish(m.contentSession(e.SessionID)) if m.streamDepth == 0 { m.streamStartTime = time.Now() } @@ -36,7 +44,23 @@ func (m *model) handleEvent(ctx context.Context, ev any) { m.trackStreamStarted(e.SessionID) case *runtime.UserMessageEvent: m.handleUserMessageEvent(e) + case *runtime.SessionRecoveredEvent: + if e.SessionID != "" && e.SessionID != m.contentSession("") { + return + } + m.screen.Transcript.FlushPending() + m.screen.Transcript.FinalizeTools(tuitypes.ToolStatusError, m.sessionState) + m.streamDepth = 0 + m.usage.RecoverIdle(m.contentSession(e.SessionID)) + m.applyUsageSnapshot() + m.busy = false + m.runCancel = nil + m.cancelMarkerPending = false + m.screen.Confirm = nil + m.status.Compacting = false + m.contentIdentity.Finish(m.contentSession(e.SessionID)) case *runtime.StreamStoppedEvent: + m.contentIdentity.Finish(m.contentSession(e.SessionID)) m.trackStreamStopped() m.streamDepth = max(0, m.streamDepth-1) if m.streamDepth > 0 { @@ -45,9 +69,17 @@ func (m *model) handleEvent(ctx context.Context, ev any) { m.notifyStreamStopped(ctx, e.Reason) m.handleStreamStopped(ctx) case *runtime.AgentChoiceReasoningEvent: - m.screen.Transcript.AppendReasoning(e.Content) + m.screen.Transcript.AppendReasoningContent(m.contentIdentity.Resolve(m.contentSession(e.SessionID), e.MessageID), e.Content) case *runtime.AgentChoiceEvent: - m.screen.Transcript.AppendAssistant(e.Content) + m.screen.Transcript.AppendAssistantContent(m.contentIdentity.Resolve(m.contentSession(e.SessionID), e.MessageID), e.Content) + case *runtime.MessageAddedEvent: + if e.Message == nil || e.Message.Implicit || e.Message.Message.Role != chat.MessageRoleAssistant { + return + } + sessionID := m.contentSession(e.SessionID) + identity := m.contentIdentity.Resolve(sessionID, e.Message.Message.MessageID) + m.screen.Transcript.ReconcileAssistantContent(identity, chat.VisibleAssistantContent(e.Message.Message.Content)) + m.contentIdentity.Finish(sessionID) case *runtime.PartialToolCallEvent: m.screen.Transcript.FlushPending() toolDef := tools.Tool{Name: e.ToolCall.Function.Name} @@ -245,3 +277,10 @@ func (m *model) notifyStreamStopped(ctx context.Context, reason string) { } } } + +func (m *model) contentSession(sessionID string) string { + if sessionID == "" && m.app != nil && m.app.Session() != nil { + return m.app.Session().ID + } + return sessionID +} diff --git a/pkg/leantui/leantui.go b/pkg/leantui/leantui.go index 2f17ad0033..231b148de6 100644 --- a/pkg/leantui/leantui.go +++ b/pkg/leantui/leantui.go @@ -6,6 +6,7 @@ import ( "log/slog" "os" "strings" + "sync/atomic" "time" tea "charm.land/bubbletea/v2" @@ -19,6 +20,7 @@ import ( "github.com/docker/docker-agent/pkg/tui/components/editor/completions" "github.com/docker/docker-agent/pkg/tui/messages" "github.com/docker/docker-agent/pkg/tui/service" + "github.com/docker/docker-agent/pkg/tui/streamcontent" "github.com/docker/docker-agent/pkg/userconfig" ) @@ -93,14 +95,19 @@ func Run(ctx context.Context, cfg Config) error { } go readKeys(term.Reader(), keys, done) + subscriptionReady := make(chan struct{}) go func() { - m.app.SubscribeWith(loopCtx, func(msg tea.Msg) { + m.app.SubscribeReliable(loopCtx, func(msg tea.Msg) { select { case events <- msg: case <-done: } - }) + }, app.WithSubscriptionReady(func() { close(subscriptionReady) }), app.WithGenerationEventMapper(func(msg tea.Msg, generation uint64) tea.Msg { + return leanEvent{generation: m.eventGeneration.Load(), inner: msg, valid: func() bool { return m.app.IsEventGeneration(generation) }} + })) }() + <-subscriptionReady + m.app.Start(loopCtx) go func() { for { w, h, ok := term.Resized() @@ -191,10 +198,17 @@ func readKeys(r io.Reader, keys chan<- ui.Key, done <-chan struct{}) { } } +type leanEvent struct { + generation uint64 + inner tea.Msg + valid func() bool +} + type model struct { - app *app.App - term *ui.Terminal - r *ui.Renderer + eventGeneration atomic.Uint64 + app *app.App + term *ui.Terminal + r *ui.Renderer width int height int @@ -205,6 +219,7 @@ type model struct { sessionState *service.SessionState usage *ui.UsageTracker + contentIdentity streamcontent.Tracker streamDepth int streamStartTime time.Time playSound func(context.Context, sound.Event) diff --git a/pkg/leantui/ui/renderer.go b/pkg/leantui/ui/renderer.go index 38b513cfa0..c8f6b2043d 100644 --- a/pkg/leantui/ui/renderer.go +++ b/pkg/leantui/ui/renderer.go @@ -91,6 +91,11 @@ func (r *Renderer) Frame(newLines []string, cursorLine, cursorCol int) { // Changes to a live block may start above the viewport. Repaint only the // visible rows rather than clearing terminal scrollback on every update. if first < r.viewportTop || newViewportTop < r.viewportTop { + if first < newViewportTop { + // Scrollback cannot be edited; replay the changed suffix without erasing history. + r.redrawSuffix(newLines, first, cursorLine, cursorCol) + return + } r.repaintVisible(newLines, cursorLine, cursorCol) return } @@ -174,6 +179,42 @@ func (r *Renderer) repaintVisible(newLines []string, cursorLine, cursorCol int) r.prev = newLines } +// redrawSuffix archives only changed offscreen rows, then restores the visible tail. +// Scrollback is immutable; replaying the whole suffix would duplicate old answers. +func (r *Renderer) redrawSuffix(newLines []string, first, cursorLine, cursorCol int) { + top := max(0, len(newLines)-r.height) + oldEnd, newEnd := len(r.prev), len(newLines) + for oldEnd > first && newEnd > first && r.prev[oldEnd-1] == newLines[newEnd-1] { + oldEnd-- + newEnd-- + } + var changed []string + for i := first; i < min(newEnd, top); i++ { + if i < r.viewportTop && i < len(r.prev) && r.prev[i] == newLines[i] { + continue + } + changed = append(changed, newLines[i]) + } + var b strings.Builder + b.WriteString(seqSyncStart) + b.WriteString(seqHideCursor) + b.WriteString("\x1b[2J\x1b[H") + rows := append(changed, newLines[top:]...) + for i, line := range rows { + if i > 0 { + b.WriteString("\r\n") + } + b.WriteString(seqEraseLine) + b.WriteString(line) + } + r.viewportTop = top + r.cursorRow = r.moveCursor(&b, len(newLines)-1, cursorLine, cursorCol) + b.WriteString(seqShowCursor) + b.WriteString(seqSyncEnd) + r.write(b.String()) + r.prev = newLines +} + // fullRedraw repaints every line. When wipe is set it also clears the screen // and scrollback first (used on resize); otherwise it assumes a clean line and // streams the content out, letting the terminal scroll as needed. diff --git a/pkg/leantui/ui/renderer_scrollback_test.go b/pkg/leantui/ui/renderer_scrollback_test.go new file mode 100644 index 0000000000..f2c7f2484a --- /dev/null +++ b/pkg/leantui/ui/renderer_scrollback_test.go @@ -0,0 +1,109 @@ +package ui + +import ( + "fmt" + "regexp" + "slices" + "strconv" + "strings" + "testing" + + "github.com/stretchr/testify/require" +) + +// ASCII-only screen fixture covers the cursor/erase sequences used by Renderer. +type scrollbackScreen struct { + rows []string + history []string + row, col int +} + +var rendererCSIPattern = regexp.MustCompile(`^\x1b\[([?0-9;]*)([A-Za-z])`) + +func (s *scrollbackScreen) write(output string) { + for output != "" { + if match := rendererCSIPattern.FindStringSubmatch(output); match != nil { + n, _ := strconv.Atoi(match[1]) + if n == 0 { + n = 1 + } + switch match[2] { + case "H": + s.row, s.col = 0, 0 + case "A": + s.row = max(0, s.row-n) + case "B": + s.row = min(len(s.rows)-1, s.row+n) + case "C": + s.col += n + case "K": + s.rows[s.row] = "" + case "J": + if match[1] == "2" { + clear(s.rows) + } + if match[1] == "3" { + s.history = nil + } + } + output = output[len(match[0]):] + continue + } + switch output[0] { + case '\r': + s.col = 0 + case '\n': + if s.row == len(s.rows)-1 { + s.history = append(s.history, s.rows[0]) + copy(s.rows, s.rows[1:]) + s.rows[s.row] = "" + } else { + s.row++ + } + default: + line := s.rows[s.row] + for len(line) <= s.col { + line += " " + } + s.rows[s.row] = line[:s.col] + output[:1] + line[s.col+1:] + s.col++ + } + output = output[1:] + } +} + +func TestRendererRepeatedOffscreenChangesDoNotReplayUnchangedAnswers(t *testing.T) { + t.Parallel() + r, buf := newTestRenderer(3) + screen := scrollbackScreen{rows: make([]string, 3)} + lines := []string{"header", "tool-start", "tool-timer-0", "tool-end", "UNCHANGED-ANSWER", "input", "footer"} + r.Frame(lines, 5, 0) + screen.write(buf.String()) + buf.Reset() + for i := range 10 { + updated := slices.Clone(lines) + updated[2] = fmt.Sprintf("tool-timer-%d", i+1) + r.Frame(updated, 5, 0) + screen.write(buf.String()) + buf.Reset() + lines = updated + } + require.Equal(t, 1, strings.Count(strings.Join(append(slices.Clone(screen.history), screen.rows...), "\n"), "UNCHANGED-ANSWER")) + require.NotContains(t, screen.history, "input") + require.NotContains(t, screen.history, "footer") + require.Equal(t, []string{"UNCHANGED-ANSWER", "input", "footer"}, screen.rows) + require.Len(t, screen.history, 14, "one changed timer row per update, not the full transcript suffix") +} + +func TestRendererPreservesAnswerPushedOffscreenDuringCorrection(t *testing.T) { + t.Parallel() + r, buf := newTestRenderer(3) + screen := scrollbackScreen{rows: make([]string, 3)} + r.Frame([]string{"history", "timer-old", "ANSWER", "input", "footer"}, 3, 0) + screen.write(buf.String()) + buf.Reset() + r.Frame([]string{"history", "timer-new", "ANSWER", "input", "input-line2", "footer"}, 4, 0) + screen.write(buf.String()) + require.Equal(t, 1, strings.Count(strings.Join(append(slices.Clone(screen.history), screen.rows...), "\n"), "ANSWER")) + require.Equal(t, []string{"input", "input-line2", "footer"}, screen.rows) +} diff --git a/pkg/leantui/ui/renderer_test.go b/pkg/leantui/ui/renderer_test.go index abeef59566..cf37c13224 100644 --- a/pkg/leantui/ui/renderer_test.go +++ b/pkg/leantui/ui/renderer_test.go @@ -148,3 +148,46 @@ func TestRendererEraseBelow(t *testing.T) { assert.Contains(t, out, seqShowCursor) assert.Equal(t, 2, r.cursorRow) } + +func TestRendererOffscreenInsertionWritesEntireChangedSuffix(t *testing.T) { + t.Parallel() + r, buf := newTestRenderer(3) + r.Frame([]string{"old-header", "history", "old-tail", "input", "footer"}, 3, 0) + buf.Reset() + lines := []string{"new-header", "FINAL-ANSWER", "new-line-1", "new-line-2", "new-line-3", "new-tail", "input", "footer"} + r.Frame(lines, 6, 0) + assert.Contains(t, buf.String(), "FINAL-ANSWER") + assert.Contains(t, buf.String(), "new-line-1") + assert.Contains(t, buf.String(), "new-tail") + assert.NotContains(t, buf.String(), "\x1b[3J", "immutable terminal history must not be erased") + assert.Equal(t, 5, r.ViewportTop()) + buf.Reset() + r.Frame(lines, 6, 0) + assert.NotContains(t, buf.String(), "FINAL-ANSWER", "identical frames must not duplicate content") +} + +func TestRendererLongAppendWritesAllNewRowsWithoutRedraw(t *testing.T) { + t.Parallel() + r, buf := newTestRenderer(3) + r.Frame([]string{"history", "old-tail", "input", "footer"}, 2, 0) + buf.Reset() + r.Frame([]string{"history", "old-tail", "FINAL-ANSWER", "line2", "line3", "line4", "input", "footer"}, 6, 0) + assert.Contains(t, buf.String(), "FINAL-ANSWER") + assert.Contains(t, buf.String(), "line2") + assert.NotContains(t, buf.String(), "\x1b[2J") + assert.NotContains(t, buf.String(), "\x1b[3J") +} + +func TestRendererOffscreenCanonicalCorrectionWithoutGrowthIsWritten(t *testing.T) { + t.Parallel() + r, buf := newTestRenderer(3) + r.Frame([]string{"ANSWER", "tool1", "tool2", "tool3", "input", "footer"}, 4, 0) + buf.Reset() + r.Frame([]string{"FINAL-ANSWER", "tool1", "tool2", "tool3", "input", "footer"}, 4, 0) + assert.Contains(t, buf.String(), "FINAL-ANSWER") + assert.NotContains(t, buf.String(), "\x1b[3J") + buf.Reset() + r.Frame([]string{"SHORT-FINAL-ANSWER", "tool1", "tool2", "input", "footer"}, 3, 0) + assert.Contains(t, buf.String(), "SHORT-FINAL-ANSWER") + assert.NotContains(t, buf.String(), "\x1b[3J") +} diff --git a/pkg/leantui/ui/transcript.go b/pkg/leantui/ui/transcript.go index 6503627ab5..bfac196311 100644 --- a/pkg/leantui/ui/transcript.go +++ b/pkg/leantui/ui/transcript.go @@ -5,6 +5,7 @@ import ( "github.com/docker/docker-agent/pkg/tools" "github.com/docker/docker-agent/pkg/tui/service" + "github.com/docker/docker-agent/pkg/tui/streamcontent" tuitypes "github.com/docker/docker-agent/pkg/tui/types" ) @@ -33,18 +34,22 @@ type PendingUserMessage struct { // pendingBlock accumulates the text of the block currently being streamed. type pendingBlock struct { - kind blockKind - text strings.Builder + kind blockKind + identity streamcontent.Identity + text strings.Builder } // block is a finalized piece of the conversation. Its lines are rendered lazily // and cached per width, so finalized content is not re-rendered every frame and // only reflows when the terminal is resized. type block struct { - render func(width int) []string - cacheW int - cache []string - cached bool + render func(width int) []string + identity streamcontent.Identity + kind blockKind + text string + cacheW int + cache []string + cached bool } func (b *block) lines(width int) []string { @@ -84,22 +89,26 @@ func (t *Transcript) AddBlock(render func(width int) []string) { t.blocks = append(t.blocks, &block{render: render}) } -func (t *Transcript) appendPending(kind blockKind, content string) { +func (t *Transcript) appendPending(kind blockKind, identity streamcontent.Identity, content string) { if content == "" { return } - if t.pending == nil || t.pending.kind != kind { + if t.pending == nil || t.pending.kind != kind || t.pending.identity != identity { t.FlushPending() - t.pending = &pendingBlock{kind: kind} + t.pending = &pendingBlock{kind: kind, identity: identity} } t.pending.text.WriteString(content) } // AppendReasoning appends streamed reasoning text. -func (t *Transcript) AppendReasoning(content string) { t.appendPending(blockReasoning, content) } +func (t *Transcript) AppendReasoning(content string) { + t.appendPending(blockReasoning, streamcontent.Identity{}, content) +} // AppendAssistant appends streamed assistant text. -func (t *Transcript) AppendAssistant(content string) { t.appendPending(blockAssistant, content) } +func (t *Transcript) AppendAssistant(content string) { + t.appendPending(blockAssistant, streamcontent.Identity{}, content) +} // FlushPending finalizes the in-progress streamed block into the conversation. func (t *Transcript) FlushPending() { @@ -108,13 +117,79 @@ func (t *Transcript) FlushPending() { } text := t.pending.text.String() kind := t.pending.kind + identity := t.pending.identity t.pending = nil - switch kind { - case blockReasoning: - t.AddBlock(func(w int) []string { return RenderReasoningLines(text, w) }) - case blockAssistant: - t.AddBlock(func(w int) []string { return RenderAssistantLines(text, w) }) + t.blocks = append(t.blocks, textBlock(kind, identity, text)) +} + +func textBlock(kind blockKind, identity streamcontent.Identity, text string) *block { + b := &block{kind: kind, identity: identity, text: text} + b.render = func(w int) []string { + if kind == blockReasoning { + return RenderReasoningLines(b.text, w) + } + return RenderAssistantLines(b.text, w) + } + return b +} + +// AppendAssistantContent separates logical messages even within one stream. +func (t *Transcript) AppendAssistantContent(identity streamcontent.Identity, content string) { + t.appendPending(blockAssistant, identity, content) +} + +// AppendReasoningContent shares the assistant message's logical identity. +func (t *Transcript) AppendReasoningContent(identity streamcontent.Identity, content string) { + t.appendPending(blockReasoning, identity, content) +} + +// ReconcileAssistantContent repairs incomplete live text from the saved message. +func (t *Transcript) ReconcileAssistantContent(identity streamcontent.Identity, content string) { + if content == "" { + return + } + var matching []*block + var delivered strings.Builder + for _, b := range t.blocks { + if b.identity == identity && b.kind == blockAssistant { + matching = append(matching, b) + delivered.WriteString(b.text) + } + } + pending := t.pending != nil && t.pending.kind == blockAssistant && t.pending.identity == identity + if pending { + delivered.WriteString(t.pending.text.String()) + } + if delivered.String() == content { + return + } + if len(matching) == 0 && !pending { + t.AppendAssistantContent(identity, content) + return + } + if suffix, ok := strings.CutPrefix(content, delivered.String()); ok { + if pending { + t.pending.text.WriteString(suffix) + } else { + b := matching[len(matching)-1] + b.text += suffix + b.cached = false + } + return + } + for i, b := range matching { + b.text = "" + if i == 0 { + b.text = content + } + b.cached = false + } + if pending { + t.pending.text.Reset() + if len(matching) == 0 { + t.pending.text.WriteString(content) + } } } diff --git a/pkg/leantui/ui/usage.go b/pkg/leantui/ui/usage.go index a7ef21522f..d14a79801e 100644 --- a/pkg/leantui/ui/usage.go +++ b/pkg/leantui/ui/usage.go @@ -32,6 +32,14 @@ func (u *UsageTracker) Reset() { u.stack = nil } +// RecoverIdle restores root context selection without discarding accounted usage. +func (u *UsageTracker) RecoverIdle(sessionID string) { + u.stack = nil + if sessionID != "" { + u.rootSessionID = sessionID + } +} + // StreamStarted pushes a newly-started session onto the active stack, adopting // the first one as the root session. func (u *UsageTracker) StreamStarted(sessionID string) { diff --git a/pkg/leantui/update.go b/pkg/leantui/update.go index 721f96b0d4..61881da6d7 100644 --- a/pkg/leantui/update.go +++ b/pkg/leantui/update.go @@ -337,6 +337,7 @@ func (m *model) handleSlash(ctx context.Context, text string, mode busySubmitMod m.quit() return true case "new": + m.eventGeneration.Add(1) m.app.NewSession() m.resetConversation() m.addNotice("", "Started a new session.", ui.StMuted()) @@ -495,6 +496,7 @@ func (m *model) resumeSession(ctx context.Context, sessionID string) { return } + m.eventGeneration.Add(1) m.app.ReplaceSession(ctx, sess) m.resetConversation() m.screen.Transcript = ui.NewTranscript() @@ -533,13 +535,16 @@ func (m *model) loadSessionTranscript(sess *session.Session) { case chat.MessageRoleUser: m.addUserEcho(content) case chat.MessageRoleAssistant: + content = chat.VisibleAssistantContent(content) if msg.Message.ReasoningContent != "" { reasoning := msg.Message.ReasoningContent m.screen.Transcript.AddBlock(func(w int) []string { return ui.RenderReasoningLines(reasoning, w) }) } if content != "" { - answer := content - m.screen.Transcript.AddBlock(func(w int) []string { return ui.RenderAssistantLines(answer, w) }) + identity := m.contentIdentity.Resolve(sess.ID, msg.Message.MessageID) + m.screen.Transcript.AppendAssistantContent(identity, content) + m.screen.Transcript.FlushPending() + m.contentIdentity.Finish(sess.ID) } for i, toolCall := range msg.Message.ToolCalls { toolDef := tools.Tool{} diff --git a/pkg/runtime/agent_delegation.go b/pkg/runtime/agent_delegation.go index 5b341256c0..222824a9e5 100644 --- a/pkg/runtime/agent_delegation.go +++ b/pkg/runtime/agent_delegation.go @@ -467,7 +467,7 @@ func (r *LocalRuntime) runCollecting(ctx context.Context, parent *session.Sessio var errMsg string events := r.RunStream(ctx, s) for event := range events { - r.forwardBackgroundUsage(event) + r.forwardBackgroundUsage(ctx, event) if ctx.Err() != nil { break } @@ -493,7 +493,7 @@ func (r *LocalRuntime) runCollecting(ctx context.Context, parent *session.Sessio // Drain remaining events so the RunStream goroutine can complete and // close the channel without blocking on a full buffer. for event := range events { - r.forwardBackgroundUsage(event) + r.forwardBackgroundUsage(ctx, event) } // Emit one authoritative final snapshot before the child @@ -511,7 +511,7 @@ func (r *LocalRuntime) runCollecting(ctx context.Context, parent *session.Sessio // usageCtx: the context-limit lookup must still resolve for a // cancelled task. finalUsage.ContextLimit = r.contextLimitForAgentModel(usageCtx, child, r.getEffectiveModelID(usageCtx, child)) - r.emitBackgroundEvent(NewTokenUsageEvent(s.ID, cfg.AgentName, finalUsage)) + r.emitBackgroundEvent(ctx, NewTokenUsageEvent(s.ID, cfg.AgentName, finalUsage)) } // Persist the sub-session unconditionally — the partial transcript is @@ -828,9 +828,9 @@ func (r *LocalRuntime) applyForceHandoff(ctx context.Context, sess *session.Sess } // Preserve accounting during cancellation drains, including nested children. -func (r *LocalRuntime) forwardBackgroundUsage(event Event) { +func (r *LocalRuntime) forwardBackgroundUsage(ctx context.Context, event Event) { switch event.(type) { case *TokenUsageEvent, *EvaluationUsageEvent: - r.emitBackgroundEvent(event) + r.emitBackgroundEvent(ctx, event) } } diff --git a/pkg/runtime/client.go b/pkg/runtime/client.go index fcc3b75a2b..0e5f88ec3d 100644 --- a/pkg/runtime/client.go +++ b/pkg/runtime/client.go @@ -5,12 +5,15 @@ import ( "bytes" "context" "encoding/json" + "errors" "fmt" "io" "log/slog" "net/http" "net/url" "path" + "strconv" + "strings" "time" "github.com/docker/docker-agent/pkg/api" @@ -44,7 +47,7 @@ func WithAuthToken(token string) ClientOption { } } -// WithTimeout sets the HTTP client timeout (deprecated: prefer per-request timeouts) +// WithTimeout sets the non-streaming HTTP timeout (deprecated: prefer per-request timeouts). func WithTimeout(timeout time.Duration) ClientOption { return func(c *Client) { if c.httpClient == nil { @@ -85,6 +88,7 @@ func NewClient(baseURL string, opts ...ClientOption) (*Client, error) { "token_usage": func() Event { return &TokenUsageEvent{} }, "evaluation_usage": func() Event { return &EvaluationUsageEvent{} }, "stream_stopped": func() Event { return &StreamStoppedEvent{} }, + "session_recovered": func() Event { return &SessionRecoveredEvent{} }, "runtime_paused": func() Event { return &PausedEvent{} }, "stream_started": func() Event { return &StreamStartedEvent{} }, "shell": func() Event { return &ShellOutputEvent{} }, @@ -336,7 +340,8 @@ func (c *Client) GetDesktopToken(ctx context.Context) (*api.DesktopTokenResponse // RunAgent executes an agent and returns a channel of streaming events. The // optional model override is persisted on the session's current agent before // the user messages are appended; pass an empty string to leave the existing -// override (if any) untouched. +// override (if any) untouched. EOF without a root StreamStoppedEvent is reported +// as incomplete; a run is never automatically resubmitted. func (c *Client) RunAgent(ctx context.Context, sessionID, agent string, messages []api.Message, model string) (<-chan Event, error) { return c.runAgentWithAgentName(ctx, sessionID, agent, "", messages, model) } @@ -375,7 +380,7 @@ func (c *Client) runAgentWithAgentName(ctx context.Context, sessionID, agent, ag req.Header.Set("Authorization", "Bearer "+c.authToken) } - resp, err := c.httpClient.Do(req) //nolint:bodyclose // body is closed in the goroutine below + resp, err := c.streamingHTTPClient().Do(req) //nolint:bodyclose // body is closed in the goroutine below if err != nil { return nil, fmt.Errorf("performing request: %w", err) } @@ -405,6 +410,7 @@ func (c *Client) runAgentWithAgentName(ctx context.Context, sessionID, agent, ag // above bufio's 64 KiB default so an oversized line does not silently // truncate the stream (bufio.ErrTooLong). scanner.Buffer(make([]byte, 0, bufio.MaxScanTokenSize), maxSSELineBytes) + var sawRootStop, sawError bool for scanner.Scan() { line := scanner.Bytes() if len(line) == 0 || line[0] == ':' { @@ -423,8 +429,8 @@ func (c *Client) runAgentWithAgentName(ctx context.Context, sessionID, agent, ag Type string `json:"type"` } if err := json.Unmarshal(after, &baseEvent); err != nil { - slog.DebugContext(ctx, "event", "error", err) - continue + sendClientEvent(ctx, eventChan, Error(fmt.Sprintf("decoding remote agent event: %v", err))) + return } // Then unmarshal the full event @@ -436,11 +442,19 @@ func (c *Client) runAgentWithAgentName(ctx context.Context, sessionID, agent, ag e := createEvent() if err := json.Unmarshal(after, &e); err != nil { - slog.DebugContext(ctx, "event", "error", err) - continue + sendClientEvent(ctx, eventChan, Error(fmt.Sprintf("decoding remote agent event: %v", err))) + return } - eventChan <- e + switch event := e.(type) { + case *StreamStoppedEvent: + sawRootStop = sawRootStop || event.SessionID == "" || event.SessionID == sessionID + case *ErrorEvent: + sawError = true + } + if !sendClientEvent(ctx, eventChan, e) { + return + } } // Surface a read failure (e.g. an over-long line) instead of ending @@ -448,9 +462,14 @@ func (c *Client) runAgentWithAgentName(ctx context.Context, sessionID, agent, ag // error after the last event that fit. if err := scanner.Err(); err != nil { slog.DebugContext(ctx, "event", "scanner_error", err) - eventChan <- Error(fmt.Sprintf("reading event stream: %v", err)) + if ctx.Err() == nil { + sendClientEvent(ctx, eventChan, Error(fmt.Sprintf("reading event stream: %v", err))) + } return } + if ctx.Err() == nil && !sawRootStop && !sawError { + sendClientEvent(ctx, eventChan, Error("remote agent stream ended before completion; the response may be incomplete")) + } }() return eventChan, nil @@ -493,116 +512,212 @@ func (c *Client) GetAgentToolCount(ctx context.Context, agentFilename, agentName return resp.AvailableTools, nil } -// StreamSessionEvents streams events for a session as they occur via Server-Sent Events. -// The returned channel is closed when ctx is cancelled, the stream's max -// duration is reached, or the server closes the connection. +// GetSessionSnapshot retrieves the recovery state and event-stream cursor. +func (c *Client) GetSessionSnapshot(ctx context.Context, sessionID string) (*api.SessionSnapshotResponse, error) { + var snapshot api.SessionSnapshotResponse + err := c.doRequest(ctx, http.MethodGet, "/api/sessions/"+sessionID+"/snapshot", nil, &snapshot) + return &snapshot, err +} + +const sessionEventGapError = "session event stream has a gap; reload the session snapshot before reconnecting" + +// StreamSessionEvents replays buffered events, then tails the session. Sequenced +// streams reconnect after transport drops. A gap produces an ErrorEvent and +// closes the channel: callers must reload a snapshot before subscribing again. +// Unsequenced legacy streams cannot safely reconnect and close at EOF. func (c *Client) StreamSessionEvents(ctx context.Context, sessionID string) (<-chan Event, error) { - endpoint := fmt.Sprintf("/api/sessions/%s/events", sessionID) + return c.streamSessionEvents(ctx, sessionID, nil) +} - u := *c.baseURL - u.Path = path.Join(u.Path, endpoint) +// StreamSessionEventsSince tails events newer than a snapshot's LastEventSeq. +// It never replays older history; see StreamSessionEvents for gap handling. +func (c *Client) StreamSessionEventsSince(ctx context.Context, sessionID string, since uint64) (<-chan Event, error) { + return c.streamSessionEvents(ctx, sessionID, &since) +} - // Bound the maximum lifetime of a single SSE connection. The cancel - // must be tied to the goroutine consuming the stream, not to this - // function's return: cancelling streamCtx kills the in-flight HTTP - // request, which would turn the stream into a one-shot read. - timeout := c.timeoutFor("streaming") - streamCtx, cancel := context.WithTimeout(ctx, timeout) +// HTTP total timeouts also cover body reads; SSE lifetime belongs to its context. +func (c *Client) streamingHTTPClient() *http.Client { + client := *c.httpClient + client.Timeout = 0 + return &client +} + +func sendClientEvent(ctx context.Context, events chan<- Event, event Event) bool { + select { + case events <- event: + return true + case <-ctx.Done(): + return false + } +} + +func waitEventStreamRetry(ctx context.Context, delay time.Duration) bool { + timer := time.NewTimer(delay) + defer timer.Stop() + select { + case <-timer.C: + return true + case <-ctx.Done(): + return false + } +} + +type sessionEventHTTPError struct { + status int + body string +} + +func (e *sessionEventHTTPError) Error() string { + var response ErrorResponse + if err := json.Unmarshal([]byte(e.body), &response); err == nil && response.Error != "" { + return fmt.Sprintf("API error (%d): %s", e.status, response.Error) + } + return fmt.Sprintf("HTTP error %d: %s", e.status, e.body) +} - req, err := http.NewRequestWithContext(streamCtx, http.MethodGet, u.String(), http.NoBody) +func (c *Client) openSessionEventStream(ctx context.Context, sessionID string, since *uint64) (*http.Response, error) { + u := *c.baseURL + u.Path = path.Join(u.Path, "/api/sessions/"+sessionID+"/events") + // The explicit cursor must win over a query inherited from the base URL. + query := u.Query() + query.Del("since") + u.RawQuery = query.Encode() + req, err := http.NewRequestWithContext(ctx, http.MethodGet, u.String(), http.NoBody) if err != nil { - cancel() return nil, fmt.Errorf("creating request: %w", err) } - req.Header.Set("Accept", "text/event-stream") req.Header.Set("Cache-Control", "no-cache") - + if since != nil { + req.Header.Set("Last-Event-ID", strconv.FormatUint(*since, 10)) + } if c.authToken != "" { req.Header.Set("Authorization", "Bearer "+c.authToken) } - - resp, err := c.httpClient.Do(req) //nolint:bodyclose // body is closed in the goroutine below + resp, err := c.streamingHTTPClient().Do(req) if err != nil { - cancel() return nil, fmt.Errorf("performing request: %w", err) } - if resp.StatusCode >= 400 { - defer cancel() defer resp.Body.Close() - respBody, err := io.ReadAll(resp.Body) + body, err := io.ReadAll(resp.Body) if err != nil { return nil, fmt.Errorf("reading error response body: %w", err) } - - var errResp ErrorResponse - if err := json.Unmarshal(respBody, &errResp); err == nil && errResp.Error != "" { - return nil, fmt.Errorf("API error (%d): %s", resp.StatusCode, errResp.Error) - } - return nil, fmt.Errorf("HTTP error %d: %s", resp.StatusCode, string(respBody)) + return nil, &sessionEventHTTPError{status: resp.StatusCode, body: string(body)} } + return resp, nil +} - eventChan := make(chan Event, defaultEventChannelCapacity) - +func (c *Client) streamSessionEvents(ctx context.Context, sessionID string, since *uint64) (<-chan Event, error) { + resp, err := c.openSessionEventStream(ctx, sessionID, since) //nolint:bodyclose // consumed and closed by readSessionEventStream + if err != nil { + return nil, err + } + events := make(chan Event, defaultEventChannelCapacity) go func() { - defer cancel() - defer close(eventChan) - defer resp.Body.Close() - - scanner := bufio.NewScanner(resp.Body) - // A single SSE line can carry a large tool response; raise the cap - // above bufio's 64 KiB default so an oversized line does not silently - // truncate the stream (bufio.ErrTooLong). - scanner.Buffer(make([]byte, 0, bufio.MaxScanTokenSize), maxSSELineBytes) - for scanner.Scan() { - line := scanner.Bytes() - if len(line) == 0 || line[0] == ':' { - continue + defer close(events) + delay := 250 * time.Millisecond + for { + terminal, readErr := c.readSessionEventStream(ctx, resp, events, &since) + if terminal || ctx.Err() != nil { + return } - - after, ok := bytes.CutPrefix(line, []byte("data: ")) - if !ok { - continue + if errors.Is(readErr, bufio.ErrTooLong) { + sendClientEvent(ctx, events, Error(fmt.Sprintf("reading event stream: %v", readErr))) + return } - - slog.DebugContext(ctx, "received event", "data", string(after)) - - // First unmarshal to get the type - var baseEvent struct { - Type string `json:"type"` + // Replaying without a cursor can duplicate answers on older servers. + if since == nil { + if readErr != nil { + sendClientEvent(ctx, events, Error(fmt.Sprintf("reading event stream: %v", readErr))) + } + return } - if err := json.Unmarshal(after, &baseEvent); err != nil { - slog.DebugContext(ctx, "failed to unmarshal event type", "error", err) - continue + for { + if !waitEventStreamRetry(ctx, delay) { + return + } + resp, err = c.openSessionEventStream(ctx, sessionID, since) //nolint:bodyclose // consumed and closed by readSessionEventStream + if err == nil { + break + } + var httpErr *sessionEventHTTPError + if errors.As(err, &httpErr) && httpErr.status >= 400 && httpErr.status < 500 && httpErr.status != http.StatusTooManyRequests { + sendClientEvent(ctx, events, Error(fmt.Sprintf("reconnecting session event stream: %v", err))) + return + } + delay = min(2*delay, 5*time.Second) } + } + }() + return events, nil +} - // Then unmarshal the full event - createEvent, found := c.registry[baseEvent.Type] - if !found { - slog.DebugContext(ctx, "unknown event type", "type", baseEvent.Type) - continue +func (c *Client) readSessionEventStream(ctx context.Context, resp *http.Response, events chan<- Event, since **uint64) (terminal bool, err error) { + defer resp.Body.Close() + scanner := bufio.NewScanner(resp.Body) + scanner.Buffer(make([]byte, 0, bufio.MaxScanTokenSize), maxSSELineBytes) + var id *uint64 + sawData := false + for scanner.Scan() { + line := scanner.Text() + if raw, ok := strings.CutPrefix(line, "id:"); ok { + seq, err := strconv.ParseUint(strings.TrimSpace(raw), 10, 64) + if err != nil { + sendClientEvent(ctx, events, Error("invalid session event cursor; reload the session snapshot")) + return true, nil } - - e := createEvent() - if err := json.Unmarshal(after, &e); err != nil { - slog.DebugContext(ctx, "failed to unmarshal event", "error", err) - continue + id = &seq + continue + } + data, ok := strings.CutPrefix(line, "data:") + if !ok { + continue + } + sawData = true + var base struct { + Type string `json:"type"` + } + if err := json.Unmarshal([]byte(data), &base); err != nil { + sendClientEvent(ctx, events, Error(fmt.Sprintf("decoding session event: %v", err))) + return true, nil + } + if base.Type == "gap" { + sendClientEvent(ctx, events, Error(sessionEventGapError)) + return true, nil + } + if base.Type == "session_exited" { + return true, nil + } + if id == nil && *since != nil { + sendClientEvent(ctx, events, Error("session event stream has no cursor; cannot safely resume")) + return true, nil + } + if id != nil && *since != nil && *id <= **since { + id = nil + continue + } + if create, ok := c.registry[base.Type]; ok { + event := create() + if err := json.Unmarshal([]byte(data), event); err != nil { + sendClientEvent(ctx, events, Error(fmt.Sprintf("decoding session event: %v", err))) + return true, nil + } + if !sendClientEvent(ctx, events, event) { + return true, nil } - - eventChan <- e } - - // Surface a read failure (e.g. an over-long line) instead of ending - // the stream silently — otherwise the run appears to stop with no - // error after the last event that fit. - if err := scanner.Err(); err != nil { - slog.DebugContext(ctx, "scanner error", "error", err) - eventChan <- Error(fmt.Sprintf("reading event stream: %v", err)) + if id != nil { + *since = id + id = nil } - }() - - return eventChan, nil + } + if *since == nil && !sawData { + // No event was delivered, so resuming from zero cannot duplicate it. + *since = new(uint64) + } + return false, scanner.Err() } // GetSessionTools retrieves tools available in a session. diff --git a/pkg/runtime/client_test.go b/pkg/runtime/client_test.go index 8a6eaf20fd..0ca5620cc4 100644 --- a/pkg/runtime/client_test.go +++ b/pkg/runtime/client_test.go @@ -5,6 +5,7 @@ import ( "fmt" "net/http" "net/http/httptest" + "sync/atomic" "testing" "time" @@ -126,3 +127,246 @@ func TestClient_StreamSessionEvents_StopsWhenContextCancelled(t *testing.T) { } } } + +func TestClient_RunAgentIgnoresTotalHTTPTimeout(t *testing.T) { + t.Parallel() + + proceed := make(chan struct{}) + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "text/event-stream") + fmt.Fprint(w, "data: {\"type\":\"warning\",\"message\":\"first\"}\n\n") + w.(http.Flusher).Flush() + select { + case <-proceed: + fmt.Fprint(w, "data: {\"type\":\"warning\",\"message\":\"last\"}\n\ndata: {\"type\":\"stream_stopped\"}\n\n") + case <-r.Context().Done(): + } + })) + t.Cleanup(srv.Close) + c, err := NewClient(srv.URL, WithHTTPClient(&http.Client{Timeout: 100 * time.Millisecond})) + require.NoError(t, err) + ctx, cancel := context.WithCancel(t.Context()) + t.Cleanup(cancel) + stream, err := c.RunAgent(ctx, "s", "agent.yaml", nil, "") + require.NoError(t, err) + first := awaitEvent[*WarningEvent](t, stream, "first event") + assert.Equal(t, "first", first.Message) + // Let the configured total timeout expire while the stream is healthy. + <-time.After(200 * time.Millisecond) + close(proceed) + var got []Event + for event := range stream { + got = append(got, event) + } + require.Len(t, got, 2) + assert.IsType(t, &StreamStoppedEvent{}, got[1]) + last, ok := got[0].(*WarningEvent) + require.True(t, ok, "got %T", got[0]) + assert.Equal(t, "last", last.Message) + assert.Equal(t, 100*time.Millisecond, c.httpClient.Timeout, "do not mutate the supplied client") +} + +func TestClient_StreamSessionEventsReconnectsFromLastDeliveredID(t *testing.T) { + t.Parallel() + + var calls atomic.Int32 + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + assert.Equal(t, "Bearer secret", r.Header.Get("Authorization")) + w.Header().Set("Content-Type", "text/event-stream") + switch calls.Add(1) { + case 1: + assert.Empty(t, r.Header.Get("Last-Event-ID")) + fmt.Fprint(w, "id: 1\ndata: {\"type\":\"session_title\",\"title\":\"one\"}\n\n") + fmt.Fprint(w, "id: 2\ndata: {\"type\":\"future_event\"}\n\n") + case 2: + assert.Equal(t, "2", r.Header.Get("Last-Event-ID"), "unknown events advance the cursor too") + fmt.Fprint(w, "id: 2\ndata: {\"type\":\"session_title\",\"title\":\"duplicate\"}\n\n") + fmt.Fprint(w, "id: 3\ndata: {\"type\":\"session_title\",\"title\":\"three\"}\n\n") + fmt.Fprint(w, "id: 4\ndata: {\"type\":\"session_exited\"}\n\n") + default: + t.Error("reconnected after session_exited") + } + })) + t.Cleanup(srv.Close) + c, err := NewClient(srv.URL, WithAuthToken("secret")) + require.NoError(t, err) + ctx, cancel := context.WithTimeout(t.Context(), 5*time.Second) + defer cancel() + stream, err := c.StreamSessionEvents(ctx, "s") + require.NoError(t, err) + var titles []string + for event := range stream { + title, ok := event.(*SessionTitleEvent) + require.True(t, ok, "got %T", event) + titles = append(titles, title.Title) + } + assert.Equal(t, []string{"one", "three"}, titles) + assert.Equal(t, int32(2), calls.Load()) +} + +func TestClient_StreamSessionEventsGapRequiresSnapshot(t *testing.T) { + t.Parallel() + + var calls atomic.Int32 + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + calls.Add(1) + assert.Equal(t, "7", r.Header.Get("Last-Event-ID")) + assert.Empty(t, r.URL.Query().Get("since")) + w.Header().Set("Content-Type", "text/event-stream") + fmt.Fprint(w, "data: {\"type\":\"gap\"}\n\nid: 99\ndata: {\"type\":\"session_title\",\"title\":\"partial history\"}\n\n") + })) + t.Cleanup(srv.Close) + c, err := NewClient(srv.URL + "?since=1") + require.NoError(t, err) + stream, err := c.StreamSessionEventsSince(t.Context(), "s", 7) + require.NoError(t, err) + var got []Event + for event := range stream { + got = append(got, event) + } + require.Len(t, got, 1) + failure, ok := got[0].(*ErrorEvent) + require.True(t, ok) + assert.Contains(t, failure.Error, "reload the session snapshot") + assert.Equal(t, int32(1), calls.Load()) +} + +func TestClient_SSECancelWithUnreadFullBuffer(t *testing.T) { + t.Parallel() + + for _, run := range []bool{false, true} { + t.Run(fmt.Sprintf("run=%v", run), func(t *testing.T) { + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "text/event-stream") + for i := range 2 * defaultEventChannelCapacity { + fmt.Fprintf(w, "id: %d\ndata: {\"type\":\"warning\",\"message\":\"x\"}\n\n", i+1) + } + w.(http.Flusher).Flush() + <-r.Context().Done() + })) + t.Cleanup(srv.Close) + c, err := NewClient(srv.URL) + require.NoError(t, err) + ctx, cancel := context.WithCancel(t.Context()) + t.Cleanup(cancel) + var stream <-chan Event + if run { + stream, err = c.RunAgent(ctx, "s", "agent.yaml", nil, "") + } else { + stream, err = c.StreamSessionEvents(ctx, "s") + } + require.NoError(t, err) + require.Eventually(t, func() bool { return len(stream) == defaultEventChannelCapacity }, 2*time.Second, time.Millisecond) + cancel() + done := make(chan struct{}) + go func() { + defer close(done) + for range stream { + } + }() + select { + case <-done: + case <-time.After(2 * time.Second): + t.Fatal("cancelled stream stayed open") + } + }) + } +} + +func TestClient_StreamSessionEventsReconnectsAfterEmptyConnection(t *testing.T) { + t.Parallel() + + var calls atomic.Int32 + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "text/event-stream") + if calls.Add(1) == 1 { + fmt.Fprint(w, ": ping\n\n") + return + } + assert.Equal(t, "0", r.Header.Get("Last-Event-ID")) + fmt.Fprint(w, "id: 1\ndata: {\"type\":\"session_title\",\"title\":\"recovered\"}\n\nid: 2\ndata: {\"type\":\"session_exited\"}\n\n") + })) + t.Cleanup(srv.Close) + c, err := NewClient(srv.URL) + require.NoError(t, err) + ctx, cancel := context.WithTimeout(t.Context(), 3*time.Second) + defer cancel() + stream, err := c.StreamSessionEvents(ctx, "s") + require.NoError(t, err) + var got []Event + for event := range stream { + got = append(got, event) + } + require.Len(t, got, 1) + assert.Equal(t, int32(2), calls.Load()) +} + +func TestClient_StreamSessionEventsIgnoresTotalHTTPTimeout(t *testing.T) { + t.Parallel() + + proceed := make(chan struct{}) + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) { + w.Header().Set("Content-Type", "text/event-stream") + fmt.Fprint(w, "id: 1\ndata: {\"type\":\"session_title\",\"title\":\"first\"}\n\n") + w.(http.Flusher).Flush() + select { + case <-proceed: + fmt.Fprint(w, "id: 2\ndata: {\"type\":\"session_title\",\"title\":\"last\"}\n\nid: 3\ndata: {\"type\":\"session_exited\"}\n\n") + case <-req.Context().Done(): + } + })) + t.Cleanup(srv.Close) + client, err := NewClient(srv.URL, WithTimeout(100*time.Millisecond)) + require.NoError(t, err) + ctx, cancel := context.WithTimeout(t.Context(), 3*time.Second) + defer cancel() + stream, err := client.StreamSessionEvents(ctx, "s") + require.NoError(t, err) + first := awaitEvent[*SessionTitleEvent](t, stream, "first event") + assert.Equal(t, "first", first.Title) + <-time.After(200 * time.Millisecond) + close(proceed) + var got []Event + for event := range stream { + got = append(got, event) + } + require.Len(t, got, 1) + last, ok := got[0].(*SessionTitleEvent) + require.True(t, ok, "got %T", got[0]) + assert.Equal(t, "last", last.Title) +} + +func TestClient_RunAgentIncompleteStreamIsAnError(t *testing.T) { + t.Parallel() + + for _, data := range []string{ + "", + `{"type":"agent_choice","content":"partial"}`, + `{"type":"stream_stopped","session_id":"child"}`, + `{"type":"stream_stopped",`, + } { + t.Run(data, func(t *testing.T) { + t.Parallel() + var runs atomic.Int32 + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + runs.Add(1) + w.Header().Set("Content-Type", "text/event-stream") + if data != "" { + fmt.Fprintf(w, "data: %s\n\n", data) + } + })) + t.Cleanup(srv.Close) + client, err := NewClient(srv.URL) + require.NoError(t, err) + stream, err := client.RunAgent(t.Context(), "s", "agent.yaml", nil, "") + require.NoError(t, err) + var got []Event + for event := range stream { + got = append(got, event) + } + require.NotEmpty(t, got) + assert.IsType(t, &ErrorEvent{}, got[len(got)-1]) + assert.Equal(t, int32(1), runs.Load(), "never resubmit a truncated run") + }) + } +} diff --git a/pkg/runtime/elicitation.go b/pkg/runtime/elicitation.go index aa939c5105..7bea6d53b3 100644 --- a/pkg/runtime/elicitation.go +++ b/pkg/runtime/elicitation.go @@ -353,6 +353,15 @@ func (r *LocalRuntime) OnElicitationRequest(handler func(Event)) { r.elicitationSinkMu.Lock() defer r.elicitationSinkMu.Unlock() r.onElicitationRequest = handler + r.onElicitationContext = nil +} + +// OnElicitationRequestWithContext preserves the originating run's ownership. +func (r *LocalRuntime) OnElicitationRequestWithContext(handler func(context.Context, Event)) { + r.elicitationSinkMu.Lock() + defer r.elicitationSinkMu.Unlock() + r.onElicitationContext = handler + r.onElicitationRequest = nil } // MirrorsElicitationOnRunStream marks LocalRuntime as a runtime whose @@ -380,9 +389,27 @@ func (r *LocalRuntime) MirrorsElicitationOnRunStream() {} // this type documents no longer holds (#3584 item 5 — dual delivery // previously required a stateful App-side dedupe to paper over). func (r *LocalRuntime) emitElicitationRequest(event Event) { + r.elicitationSinkMu.RLock() + handler, contextual := r.onElicitationRequest, r.onElicitationContext + r.elicitationSinkMu.RUnlock() + if contextual != nil { + contextual(context.Background(), event) + return + } // Test seam has no originating run. + if handler != nil { + handler(event) + } +} + +func (r *LocalRuntime) emitElicitationRequestContext(ctx context.Context, event Event) { r.elicitationSinkMu.RLock() handler := r.onElicitationRequest + contextual := r.onElicitationContext r.elicitationSinkMu.RUnlock() + if contextual != nil { + contextual(ctx, event) + return + } if handler != nil { handler(event) } @@ -407,7 +434,7 @@ func (r *LocalRuntime) EmitElicitationRequestForTesting(event Event) { func (r *LocalRuntime) hasElicitationSink() bool { r.elicitationSinkMu.RLock() defer r.elicitationSinkMu.RUnlock() - return r.onElicitationRequest != nil + return r.onElicitationRequest != nil || r.onElicitationContext != nil } // elicitationDeclineNotes accumulates model-readable notes for elicitations @@ -571,7 +598,7 @@ func (r *LocalRuntime) requestElicitation(ctx context.Context, spec elicitationS // Reliable delivery: invoked synchronously, unconditionally, and exactly // once, BEFORE anything that could block (#3584 review item 1). This // must never be gated behind the best-effort bridge below. - r.emitElicitationRequest(ev) + r.emitElicitationRequestContext(ctx, ev) // Best-effort secondary delivery on the owning stream's events channel, // kept for remote/SSE consumers that read directly off RunStream diff --git a/pkg/runtime/elicitation_test.go b/pkg/runtime/elicitation_test.go index dad58f99a8..89e6aaf9af 100644 --- a/pkg/runtime/elicitation_test.go +++ b/pkg/runtime/elicitation_test.go @@ -404,3 +404,15 @@ func TestDirectElicitationHandlerBypassesEventsAndHonorsHeadless(t *testing.T) { require.NoError(t, rt.Close()) } } + +func TestContextualElicitationPreservesProducerContext(t *testing.T) { + t.Parallel() + rt := &LocalRuntime{} + type ownerKey struct{} + origin := context.WithValue(t.Context(), ownerKey{}, "original-conversation") + var observed string + rt.OnElicitationRequestWithContext(func(ctx context.Context, event Event) { observed = ctx.Value(ownerKey{}).(string) }) + rt.emitElicitationRequestContext(origin, ElicitationRequest("request", "form", nil, "", "id", "", "child", nil, "root")) + require.Equal(t, "original-conversation", observed) + require.True(t, rt.hasElicitationSink()) +} diff --git a/pkg/runtime/event.go b/pkg/runtime/event.go index 871022aa8a..ef545f1337 100644 --- a/pkg/runtime/event.go +++ b/pkg/runtime/event.go @@ -654,6 +654,25 @@ func StreamStopped(sessionID, agentName, reason string) Event { func (e *StreamStoppedEvent) GetSessionID() string { return e.SessionID } +// SessionRecoveredEvent is an authoritative idle boundary after snapshot recovery. +// Reset stream depth and transient interactions without running stop-triggered actions. +type SessionRecoveredEvent struct { + AgentContext + + Type string `json:"type"` + SessionID string `json:"session_id"` +} + +func SessionRecovered(sessionID string) Event { + return &SessionRecoveredEvent{ + Type: "session_recovered", + SessionID: sessionID, + AgentContext: newAgentContext(""), + } +} + +func (e *SessionRecoveredEvent) GetSessionID() string { return e.SessionID } + // PausedEvent reports that the run loop has reached an iteration // boundary and is now blocked because /pause was toggled on. It is emitted // once the in-flight LLM request and its tool calls have finished — i.e. the diff --git a/pkg/runtime/recovery_event_test.go b/pkg/runtime/recovery_event_test.go new file mode 100644 index 0000000000..f1b0345fcc --- /dev/null +++ b/pkg/runtime/recovery_event_test.go @@ -0,0 +1,42 @@ +package runtime + +import ( + "encoding/json" + "fmt" + "net/http" + "net/http/httptest" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestSessionRecoveredEventContract(t *testing.T) { + t.Parallel() + event := SessionRecovered("root") + data, err := json.Marshal(event) + require.NoError(t, err) + var wire map[string]any + require.NoError(t, json.Unmarshal(data, &wire)) + assert.Equal(t, "session_recovered", wire["type"]) + assert.Equal(t, "root", wire["session_id"]) + assert.Equal(t, "root", event.(SessionScoped).GetSessionID()) + assert.NotContains(t, wire, "reason", "reset is not a stop or a completed turn") + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.Header().Set("Content-Type", "text/event-stream") + fmt.Fprintf(w, "data: %s\n\n", data) + })) + t.Cleanup(srv.Close) + client, err := NewClient(srv.URL) + require.NoError(t, err) + events, err := client.StreamSessionEvents(t.Context(), "root") + require.NoError(t, err) + var got []Event + for event := range events { + got = append(got, event) + } + require.Len(t, got, 1) + reset, ok := got[0].(*SessionRecoveredEvent) + require.True(t, ok) + assert.Equal(t, "root", reset.SessionID) +} diff --git a/pkg/runtime/remote_runtime.go b/pkg/runtime/remote_runtime.go index ceac484ec6..fb8918feeb 100644 --- a/pkg/runtime/remote_runtime.go +++ b/pkg/runtime/remote_runtime.go @@ -7,6 +7,7 @@ import ( "errors" "fmt" "log/slog" + "net/http" "strings" "sync" "time" @@ -29,12 +30,13 @@ import ( // It works with any client that implements the RemoteClient interface, // including both HTTP (Client) and Connect-RPC (ConnectRPCClient) clients. type RemoteRuntime struct { - client RemoteClient - currentAgent string - agentFilename string - sessionID string - team *team.Team - pendingOAuthElicitation *ElicitationRequestEvent + client RemoteClient + currentAgent string + agentFilename string + sessionID string + team *team.Team + pendingOAuthElicitations map[string]*ElicitationRequestEvent + oauthMu sync.Mutex // pendingModelOverride is the model ref to apply to the current agent // on the next [RemoteRuntime.RunStream] call. It is set by @@ -49,6 +51,213 @@ type RemoteRuntime struct { // field read when no specific agent has been selected. resolvedDefault string resolvedDefaultMu sync.Mutex + + stateMu sync.Mutex // sessionID and pendingOAuthElicitations + + reconcileMu sync.RWMutex // snapshots must not race foreground delivery + backgroundInit sync.Mutex + backgroundMu sync.Mutex + backgroundHandler func(Event) + background *remoteEventSubscription + closed bool +} + +// Cursor-based subscriptions are optional so other RemoteClient implementations +// do not accidentally replay foreground history through the background sink. +type remoteEventClient interface { + GetSessionSnapshot(ctx context.Context, sessionID string) (*api.SessionSnapshotResponse, error) + StreamSessionEventsSince(ctx context.Context, sessionID string, since uint64) (<-chan Event, error) +} + +type remoteEventSubscription struct { + cancel context.CancelFunc + sessionID string + history *remoteMessageHistory +} + +// Snapshots flatten sub-sessions; recovery exposes only plain assistant text. +type remoteMessageHistory struct { + mu sync.Mutex + content map[string]*strings.Builder + sessions map[string]string + complete map[string]bool + elicitations map[string]bool + streamDepth map[string]int + unidentified bool + retentionErr error + bytes int +} + +const ( + remoteHistoryMaxMessages = 4096 + remoteHistoryMaxBytes = 8 << 20 + remoteHistoryMaxElicitations = 4096 +) + +// Eviction cannot turn old answers into new answers: disable recovery on overflow. +func (h *remoteMessageHistory) limit() { + if h.retentionErr == nil && len(h.content)+len(h.complete) <= remoteHistoryMaxMessages && h.bytes <= remoteHistoryMaxBytes && len(h.elicitations) <= remoteHistoryMaxElicitations && len(h.streamDepth) <= remoteHistoryMaxMessages { + return + } + if h.retentionErr == nil { + h.retentionErr = errors.New("remote recovery history exceeded its retention limit") + } + clear(h.content) + clear(h.sessions) + clear(h.streamDepth) + h.bytes = 0 +} + +func newRemoteMessageHistory(messages []session.Message) *remoteMessageHistory { + h := &remoteMessageHistory{ + content: make(map[string]*strings.Builder), + sessions: make(map[string]string), + complete: make(map[string]bool), + elicitations: make(map[string]bool), + streamDepth: make(map[string]int), + } + for _, msg := range messages { + if !recoverableAssistantText(msg) { + continue + } + if msg.Message.MessageID == "" { + h.unidentified = true + } else { + h.complete[msg.Message.MessageID] = true + } + h.limit() + if h.retentionErr != nil { + break + } + } + return h +} + +func recoverableAssistantText(msg session.Message) bool { + return !msg.Implicit && msg.Message.Role == chat.MessageRoleAssistant && chat.VisibleAssistantContent(msg.Message.Content) != "" && + len(msg.Message.ToolCalls) == 0 && msg.Message.FunctionCall == nil +} + +func (h *remoteMessageHistory) deliver(event Event, send func(Event) bool) bool { + h.mu.Lock() + defer h.mu.Unlock() + request, isElicitation := event.(*ElicitationRequestEvent) + if isElicitation && request.ElicitationID != "" && h.elicitations[request.ElicitationID] { + return true + } + if isElicitation && request.ElicitationID != "" && len(h.elicitations) >= remoteHistoryMaxElicitations { + send(Error("remote elicitation history exceeded its retention limit; interaction was not delivered")) + return false + } + choice, ok := event.(*AgentChoiceEvent) + if ok && choice.MessageID != "" && h.complete[choice.MessageID] { + return true // already restored from a snapshot + } + if !send(event) { + return false + } + if isElicitation && request.ElicitationID != "" { + h.elicitations[request.ElicitationID] = true + } + if started, ok := event.(*StreamStartedEvent); ok && h.retentionErr == nil { + if depth := h.streamDepth[started.SessionID]; depth < remoteHistoryMaxMessages { + h.streamDepth[started.SessionID]++ + } else { + h.retentionErr = errors.New("remote recovery history exceeded its retention limit") + } + } + if stopped, ok := event.(*StreamStoppedEvent); ok { + if depth := h.streamDepth[stopped.SessionID]; depth > 1 { + h.streamDepth[stopped.SessionID]-- + return true + } + delete(h.streamDepth, stopped.SessionID) + for id, content := range h.content { + if scope := h.sessions[id]; scope != "" && scope != stopped.SessionID { + continue + } + if chat.VisibleAssistantContent(content.String()) != "" { + h.complete[id] = true + } + h.bytes -= content.Len() + delete(h.content, id) + delete(h.sessions, id) + } + } + if ok && h.retentionErr == nil { + if choice.MessageID == "" { + h.unidentified = true + } else { + content := h.content[choice.MessageID] + if content == nil { + content = new(strings.Builder) + h.content[choice.MessageID] = content + } + content.WriteString(choice.Content) + h.bytes += len(choice.Content) + h.sessions[choice.MessageID] = choice.SessionID + } + } + h.limit() + return true +} + +func (h *remoteMessageHistory) reconcile(snapshot *api.SessionSnapshotResponse, send func(Event)) error { + h.mu.Lock() + defer h.mu.Unlock() + if h.retentionErr != nil { + return h.retentionErr + } + if h.unidentified { + return errors.New("cannot safely reconcile assistant messages without message IDs") + } + // Validate the whole snapshot before emitting any append-only deltas. + seen := make(map[string]bool) + retained := len(h.content) + len(h.complete) + for _, msg := range snapshot.Messages { + if !recoverableAssistantText(msg) { + continue + } + id := msg.Message.MessageID + var delivered string + if content := h.content[id]; content != nil { + delivered = chat.VisibleAssistantContent(content.String()) + } + if id == "" || seen[id] || (!h.complete[id] && !strings.HasPrefix(chat.VisibleAssistantContent(msg.Message.Content), delivered)) { + return errors.New("cannot safely reconcile changed or ambiguous assistant messages") + } + if !h.complete[id] && h.content[id] == nil { + retained++ + } + if retained > remoteHistoryMaxMessages || len(seen) >= remoteHistoryMaxMessages { + return errors.New("remote recovery history exceeded its retention limit") + } + seen[id] = true + } + for _, msg := range snapshot.Messages { + if !recoverableAssistantText(msg) { + continue + } + id := msg.Message.MessageID + if h.complete[id] { + continue + } + var delivered string + if content := h.content[id]; content != nil { + delivered = chat.VisibleAssistantContent(content.String()) + h.bytes -= content.Len() + } + if suffix := strings.TrimPrefix(chat.VisibleAssistantContent(msg.Message.Content), delivered); suffix != "" { + sessionID := cmp.Or(h.sessions[id], snapshot.ID) + send(AgentChoice(msg.AgentName, sessionID, suffix, id)) + } + delete(h.content, id) + delete(h.sessions, id) + h.complete[id] = true + } + clear(h.streamDepth) + h.limit() + return nil } // RemoteRuntimeOption is a function for configuring the RemoteRuntime @@ -156,10 +365,11 @@ func (r *RemoteRuntime) SetCurrentAgent(ctx context.Context, agentName string) e // CurrentAgentTools returns the tools for the current agent from the session. func (r *RemoteRuntime) CurrentAgentTools(ctx context.Context) ([]tools.Tool, error) { - if r.sessionID == "" { + sessionID := r.activeSessionID() + if sessionID == "" { return nil, nil } - return r.client.GetSessionTools(ctx, r.sessionID) + return r.client.GetSessionTools(ctx, sessionID) } // CurrentAgentToolsetStatuses is not implemented for remote runtimes; the @@ -171,10 +381,11 @@ func (r *RemoteRuntime) CurrentAgentToolsetStatuses() []tools.ToolsetStatus { // RestartToolset restarts a toolset on the remote server. func (r *RemoteRuntime) RestartToolset(ctx context.Context, toolsetName string) error { - if r.sessionID == "" { + sessionID := r.activeSessionID() + if sessionID == "" { return errors.New("no active session") } - return r.client.RestartSessionToolset(ctx, r.sessionID, toolsetName) + return r.client.RestartSessionToolset(ctx, sessionID, toolsetName) } // EmitStartupInfo emits initial agent, team, and toolset information @@ -257,14 +468,26 @@ func (r *RemoteRuntime) readCurrentAgentConfig(ctx context.Context) latest.Agent // RunStream starts the agent's interaction loop and returns a channel of events func (r *RemoteRuntime) RunStream(ctx context.Context, sess *session.Session) <-chan Event { - slog.DebugContext(ctx, "Starting remote runtime stream", "agent", r.currentAgent, "session_id", r.sessionID) + slog.DebugContext(ctx, "Starting remote runtime stream", "agent", r.currentAgent, "session_id", r.activeSessionID()) events := make(chan Event, defaultEventChannelCapacity) go func() { defer close(events) + for !r.reconcileMu.TryRLock() { + if !waitEventStreamRetry(ctx, 250*time.Millisecond) { + return + } + } + defer r.reconcileMu.RUnlock() messages := r.convertSessionMessages(sess) + r.stateMu.Lock() + if r.sessionID != sess.ID { + clear(r.pendingOAuthElicitations) + } r.sessionID = sess.ID + r.stateMu.Unlock() + r.startBackgroundEvents(ctx, sess.ID) // Snapshot the queued override but do NOT clear it yet: if the // request fails before the server can persist it, clearing here @@ -278,13 +501,13 @@ func (r *RemoteRuntime) RunStream(ctx context.Context, sess *session.Session) <- var err error if r.currentAgent != "" { - streamChan, err = r.client.RunAgentWithAgentName(ctx, r.sessionID, r.agentFilename, r.currentAgent, messages, model) + streamChan, err = r.client.RunAgentWithAgentName(ctx, sess.ID, r.agentFilename, r.currentAgent, messages, model) } else { - streamChan, err = r.client.RunAgent(ctx, r.sessionID, r.agentFilename, messages, model) + streamChan, err = r.client.RunAgent(ctx, sess.ID, r.agentFilename, messages, model) } if err != nil { - events <- Error(fmt.Sprintf("failed to start remote agent: %v", err)) + sendClientEvent(ctx, events, Error(fmt.Sprintf("failed to start remote agent: %v", err))) return } @@ -299,12 +522,40 @@ func (r *RemoteRuntime) RunStream(ctx context.Context, sess *session.Session) <- r.pendingMu.Unlock() } - // Consume events from the agent stream + // Drain on cancellation too: alternate clients may finish teardown + // by emitting events after the context is cancelled. + defer func() { + for range streamChan { + } + }() + var sawRootStop, sawError bool + send := func(event Event) bool { + r.trackOAuthElicitation(event) + return sendClientEvent(ctx, events, event) + } for streamEvent := range streamChan { - if elicitationRequest, ok := streamEvent.(*ElicitationRequestEvent); ok { - r.pendingOAuthElicitation = elicitationRequest + switch event := streamEvent.(type) { + case *StreamStoppedEvent: + sawRootStop = sawRootStop || event.SessionID == "" || event.SessionID == sess.ID + case *ErrorEvent: + sawError = true + } + r.backgroundMu.Lock() + subscription := r.background + r.backgroundMu.Unlock() + if subscription != nil && subscription.sessionID == sess.ID { + if !subscription.history.deliver(streamEvent, send) { + return + } + } else if !send(streamEvent) { + return } - events <- streamEvent + } + if !sawRootStop && ctx.Err() == nil { + if !sawError { + sendClientEvent(ctx, events, Error("remote agent stream ended before completion; the response may be incomplete")) + } + sendClientEvent(ctx, events, StreamStopped(sess.ID, r.currentAgent, "error")) } }() @@ -334,20 +585,22 @@ func (r *RemoteRuntime) Run(ctx context.Context, sess *session.Session) ([]sessi // Steer enqueues a user message for mid-turn injection into the running // agent loop on the remote server. func (r *RemoteRuntime) Steer(ctx context.Context, msg QueuedMessage) error { - if r.sessionID == "" { + sessionID := r.activeSessionID() + if sessionID == "" { return errors.New("no active session") } - return r.client.SteerSession(ctx, r.sessionID, []api.Message{ + return r.client.SteerSession(ctx, sessionID, []api.Message{ {Content: msg.Content, MultiContent: msg.MultiContent}, }) } // FollowUp enqueues a message for end-of-turn processing on the remote server. func (r *RemoteRuntime) FollowUp(ctx context.Context, msg QueuedMessage) error { - if r.sessionID == "" { + sessionID := r.activeSessionID() + if sessionID == "" { return errors.New("no active session") } - return r.client.FollowUpSession(ctx, r.sessionID, []api.Message{ + return r.client.FollowUpSession(ctx, sessionID, []api.Message{ {Content: msg.Content, MultiContent: msg.MultiContent}, }) } @@ -358,25 +611,27 @@ func (r *RemoteRuntime) QueueStatus() QueueStatus { // Resume allows resuming execution after user confirmation func (r *RemoteRuntime) Resume(ctx context.Context, req ResumeRequest) { - slog.DebugContext(ctx, "Resuming remote runtime", "agent", r.currentAgent, "type", req.Type, "reason", req.Reason, "tool_name", req.ToolName, "session_id", r.sessionID) + sessionID := r.activeSessionID() + slog.DebugContext(ctx, "Resuming remote runtime", "agent", r.currentAgent, "type", req.Type, "reason", req.Reason, "tool_name", req.ToolName, "session_id", sessionID) - if r.sessionID == "" { + if sessionID == "" { slog.ErrorContext(ctx, "Cannot resume: no session ID available") return } - if err := r.client.ResumeSession(ctx, r.sessionID, string(req.Type), req.Reason, req.ToolName); err != nil { - slog.ErrorContext(ctx, "Failed to resume remote session", "error", err, "session_id", r.sessionID) + if err := r.client.ResumeSession(ctx, sessionID, string(req.Type), req.Reason, req.ToolName); err != nil { + slog.ErrorContext(ctx, "Failed to resume remote session", "error", err, "session_id", sessionID) } } // Summarize generates a summary for the session by compacting it server-side. func (r *RemoteRuntime) Summarize(ctx context.Context, sess *session.Session, _ string, sink EventSink) { - if r.sessionID == "" { + sessionID := r.activeSessionID() + if sessionID == "" { sink.Emit(SessionSummary(sess.ID, "No active session to summarize", r.currentAgent, 0, 0, "", nil)) return } - if err := r.client.CompactSession(ctx, r.sessionID); err != nil { + if err := r.client.CompactSession(ctx, sessionID); err != nil { slog.WarnContext(ctx, "Failed to compact session", "error", err) sink.Emit(SessionSummary(sess.ID, fmt.Sprintf("Compaction failed: %v", err), r.currentAgent, 0, 0, "", nil)) return @@ -400,24 +655,51 @@ func (r *RemoteRuntime) convertSessionMessages(sess *session.Session) []api.Mess return messages } -// ResumeElicitation sends an elicitation response back to a waiting elicitation request -func (r *RemoteRuntime) ResumeElicitation(ctx context.Context, action tools.ElicitationAction, content map[string]any, elicitationID ...string) error { - id := firstElicitationID(elicitationID) - slog.DebugContext(ctx, "Resuming remote runtime with elicitation response", "agent", r.currentAgent, "action", action, "session_id", r.sessionID, "elicitation_id", id) - - err := r.handleOAuthElicitation(ctx, r.pendingOAuthElicitation) - if err != nil { - return err +func (r *RemoteRuntime) trackOAuthElicitation(event Event) { + request, ok := event.(*ElicitationRequestEvent) + if !ok || request.Meta["docker-agent/type"] != "oauth_flow" || request.ElicitationID == "" { + return } - - if err := r.client.ResumeElicitation(ctx, r.sessionID, action, content, id); err != nil { - return err + r.stateMu.Lock() + defer r.stateMu.Unlock() + if r.pendingOAuthElicitations == nil { + r.pendingOAuthElicitations = make(map[string]*ElicitationRequestEvent) } + if r.pendingOAuthElicitations[request.ElicitationID] == nil && len(r.pendingOAuthElicitations) < remoteHistoryMaxElicitations { + r.pendingOAuthElicitations[request.ElicitationID] = request + } +} - return nil +// ResumeElicitation answers a request, running OAuth only for its explicit ID. +func (r *RemoteRuntime) ResumeElicitation(ctx context.Context, action tools.ElicitationAction, content map[string]any, elicitationID ...string) error { + // Serialize responses, including the browser flow, so concurrent accepts + // cannot authorize the same request twice. + r.oauthMu.Lock() + defer r.oauthMu.Unlock() + id := firstElicitationID(elicitationID) + r.stateMu.Lock() + sessionID := r.sessionID + pending := r.pendingOAuthElicitations[id] + ambiguous := id == "" && len(r.pendingOAuthElicitations) > 0 + r.stateMu.Unlock() + if ambiguous { + return errors.New("OAuth elicitation requires an explicit elicitation ID") + } + if pending != nil { + defer func() { + r.stateMu.Lock() + defer r.stateMu.Unlock() + delete(r.pendingOAuthElicitations, id) + }() + if action == tools.ElicitationActionAccept { + // handleOAuthElicitation sends the token response itself. + return r.handleOAuthElicitation(ctx, sessionID, pending) + } + } + return r.client.ResumeElicitation(ctx, sessionID, action, content, id) } -func (r *RemoteRuntime) handleOAuthElicitation(ctx context.Context, req *ElicitationRequestEvent) error { +func (r *RemoteRuntime) handleOAuthElicitation(ctx context.Context, sessionID string, req *ElicitationRequestEvent) error { if req == nil { return nil } @@ -428,7 +710,7 @@ func (r *RemoteRuntime) handleOAuthElicitation(ctx context.Context, req *Elicita if !ok { err := errors.New("server_url missing from elicitation metadata") slog.ErrorContext(ctx, "Failed to extract server_url", "error", err) - _ = r.client.ResumeElicitation(ctx, r.sessionID, "decline", nil, req.ElicitationID) + _ = r.client.ResumeElicitation(ctx, sessionID, "decline", nil, req.ElicitationID) return err } @@ -436,7 +718,7 @@ func (r *RemoteRuntime) handleOAuthElicitation(ctx context.Context, req *Elicita if !ok { err := errors.New("auth_server_metadata missing from elicitation metadata") slog.ErrorContext(ctx, "Failed to extract auth_server_metadata", "error", err) - _ = r.client.ResumeElicitation(ctx, r.sessionID, "decline", nil, req.ElicitationID) + _ = r.client.ResumeElicitation(ctx, sessionID, "decline", nil, req.ElicitationID) return err } @@ -444,12 +726,12 @@ func (r *RemoteRuntime) handleOAuthElicitation(ctx context.Context, req *Elicita metadataBytes, err := json.Marshal(authServerMetadata) if err != nil { slog.ErrorContext(ctx, "Failed to marshal auth_server_metadata", "error", err) - _ = r.client.ResumeElicitation(ctx, r.sessionID, "decline", nil, req.ElicitationID) + _ = r.client.ResumeElicitation(ctx, sessionID, "decline", nil, req.ElicitationID) return fmt.Errorf("failed to marshal auth_server_metadata: %w", err) } if err := json.Unmarshal(metadataBytes, &authMetadata); err != nil { slog.ErrorContext(ctx, "Failed to unmarshal auth_server_metadata", "error", err) - _ = r.client.ResumeElicitation(ctx, r.sessionID, "decline", nil, req.ElicitationID) + _ = r.client.ResumeElicitation(ctx, sessionID, "decline", nil, req.ElicitationID) return fmt.Errorf("failed to unmarshal auth_server_metadata: %w", err) } @@ -469,7 +751,7 @@ func (r *RemoteRuntime) handleOAuthElicitation(ctx context.Context, req *Elicita callbackServer, err := oauthflow.NewCallbackServer(ctx) if err != nil { slog.ErrorContext(ctx, "Failed to create callback server", "error", err) - _ = r.client.ResumeElicitation(ctx, r.sessionID, "decline", nil, req.ElicitationID) + _ = r.client.ResumeElicitation(ctx, sessionID, "decline", nil, req.ElicitationID) return fmt.Errorf("failed to create callback server: %w", err) } defer func() { @@ -484,7 +766,7 @@ func (r *RemoteRuntime) handleOAuthElicitation(ctx context.Context, req *Elicita if err := callbackServer.Start(); err != nil { slog.ErrorContext(ctx, "Failed to start callback server", "error", err) - _ = r.client.ResumeElicitation(ctx, r.sessionID, "decline", nil, req.ElicitationID) + _ = r.client.ResumeElicitation(ctx, sessionID, "decline", nil, req.ElicitationID) return fmt.Errorf("failed to start callback server: %w", err) } @@ -497,21 +779,21 @@ func (r *RemoteRuntime) handleOAuthElicitation(ctx context.Context, req *Elicita clientID, clientSecret, err = oauthflow.RegisterClient(oauthCtx, &authMetadata, redirectURI, nil) if err != nil { slog.ErrorContext(ctx, "Dynamic client registration failed", "error", err) - _ = r.client.ResumeElicitation(ctx, r.sessionID, "decline", nil, req.ElicitationID) + _ = r.client.ResumeElicitation(ctx, sessionID, "decline", nil, req.ElicitationID) return fmt.Errorf("failed to register client: %w", err) } slog.DebugContext(ctx, "Client registered successfully", "client_id", clientID) } else { err := errors.New("authorization server does not support dynamic client registration") slog.ErrorContext(ctx, "Client registration not supported", "error", err) - _ = r.client.ResumeElicitation(ctx, r.sessionID, "decline", nil, req.ElicitationID) + _ = r.client.ResumeElicitation(ctx, sessionID, "decline", nil, req.ElicitationID) return err } state, err := oauthflow.GenerateState() if err != nil { slog.ErrorContext(ctx, "Failed to generate state", "error", err) - _ = r.client.ResumeElicitation(ctx, r.sessionID, "decline", nil, req.ElicitationID) + _ = r.client.ResumeElicitation(ctx, sessionID, "decline", nil, req.ElicitationID) return fmt.Errorf("failed to generate state: %w", err) } @@ -534,14 +816,14 @@ func (r *RemoteRuntime) handleOAuthElicitation(ctx context.Context, req *Elicita code, receivedState, err := oauthflow.RequestAuthorizationCode(oauthCtx, authURL, callbackServer, state) if err != nil { slog.ErrorContext(ctx, "Failed to get authorization code", "error", err) - _ = r.client.ResumeElicitation(ctx, r.sessionID, "decline", nil, req.ElicitationID) + _ = r.client.ResumeElicitation(ctx, sessionID, "decline", nil, req.ElicitationID) return fmt.Errorf("failed to get authorization code: %w", err) } if receivedState != state { err := fmt.Errorf("state mismatch: expected %s, got %s", state, receivedState) slog.ErrorContext(ctx, "State mismatch in authorization response", "error", err) - _ = r.client.ResumeElicitation(ctx, r.sessionID, "decline", nil, req.ElicitationID) + _ = r.client.ResumeElicitation(ctx, sessionID, "decline", nil, req.ElicitationID) return err } @@ -559,7 +841,7 @@ func (r *RemoteRuntime) handleOAuthElicitation(ctx context.Context, req *Elicita ) if err != nil { slog.ErrorContext(ctx, "Failed to exchange code for token", "error", err) - _ = r.client.ResumeElicitation(ctx, r.sessionID, "decline", nil, req.ElicitationID) + _ = r.client.ResumeElicitation(ctx, sessionID, "decline", nil, req.ElicitationID) return fmt.Errorf("failed to exchange code for token: %w", err) } @@ -577,7 +859,7 @@ func (r *RemoteRuntime) handleOAuthElicitation(ctx context.Context, req *Elicita } slog.DebugContext(ctx, "Sending token to server") - if err := r.client.ResumeElicitation(ctx, r.sessionID, tools.ElicitationActionAccept, tokenData, req.ElicitationID); err != nil { + if err := r.client.ResumeElicitation(ctx, sessionID, tools.ElicitationActionAccept, tokenData, req.ElicitationID); err != nil { slog.ErrorContext(ctx, "Failed to send token to server", "error", err) return fmt.Errorf("failed to send token to server: %w", err) } @@ -661,19 +943,21 @@ func (r *RemoteRuntime) RunSkillFork(context.Context, *session.Session, skills.R // UpdateSessionTitle updates the title of the current session on the remote server. func (r *RemoteRuntime) UpdateSessionTitle(ctx context.Context, sess *session.Session, title string) error { + sessionID := r.activeSessionID() sess.SetTitle(title) - if r.sessionID == "" { + if sessionID == "" { return errors.New("cannot update session title: no session ID available") } - return r.client.UpdateSessionTitle(ctx, r.sessionID, title) + return r.client.UpdateSessionTitle(ctx, sessionID, title) } // CurrentMCPPrompts returns available MCP prompts from the server. func (r *RemoteRuntime) CurrentMCPPrompts(ctx context.Context) map[string]tools.PromptInfo { - if r.sessionID == "" { + sessionID := r.activeSessionID() + if sessionID == "" { return make(map[string]tools.PromptInfo) } - prompts, err := r.client.GetSessionMCPPrompts(ctx, r.sessionID) + prompts, err := r.client.GetSessionMCPPrompts(ctx, sessionID) if err != nil { slog.WarnContext(ctx, "Failed to get MCP prompts", "error", err) return make(map[string]tools.PromptInfo) @@ -700,10 +984,11 @@ func (r *RemoteRuntime) CurrentMCPPrompts(ctx context.Context) map[string]tools. // ExecuteMCPPrompt executes an MCP prompt on the server. func (r *RemoteRuntime) ExecuteMCPPrompt(ctx context.Context, promptName string, args map[string]string) (string, error) { - if r.sessionID == "" { + sessionID := r.activeSessionID() + if sessionID == "" { return "", errors.New("no active session") } - return r.client.ExecuteSessionMCPPrompt(ctx, r.sessionID, promptName, args) + return r.client.ExecuteSessionMCPPrompt(ctx, sessionID, promptName, args) } // TitleGenerator is not supported on remote runtimes (titles are generated server-side). @@ -713,10 +998,11 @@ func (r *RemoteRuntime) TitleGenerator(context.Context) *sessiontitle.Generator // TogglePause pauses/resumes a session on the server. func (r *RemoteRuntime) TogglePause(ctx context.Context) (bool, error) { - if r.sessionID == "" { + sessionID := r.activeSessionID() + if sessionID == "" { return false, errors.New("no active session") } - return false, r.client.PauseSession(ctx, r.sessionID) + return false, r.client.PauseSession(ctx, sessionID) } // OnToolsChanged is a no-op for remote runtimes; tool-list changes are @@ -724,89 +1010,324 @@ func (r *RemoteRuntime) TogglePause(ctx context.Context) (bool, error) { // than via an out-of-band callback. func (r *RemoteRuntime) OnToolsChanged(func(Event)) {} -// OnBackgroundEvent is a no-op for remote runtimes; background agent tasks -// run server-side and their events are not forwarded out-of-band. -func (r *RemoteRuntime) OnBackgroundEvent(func(Event)) {} +// OnBackgroundEvent receives server-side recalls and out-of-band events. +// The subscription outlives individual turns and stops on Close or unregister. +func (r *RemoteRuntime) OnBackgroundEvent(handler func(Event)) { + r.backgroundMu.Lock() + defer r.backgroundMu.Unlock() + r.backgroundHandler = handler + if handler == nil && r.background != nil { + r.background.cancel() + r.background = nil + } +} + +func (r *RemoteRuntime) activeSessionID() string { + r.stateMu.Lock() + defer r.stateMu.Unlock() + return r.sessionID +} + +func (r *RemoteRuntime) deliverBackgroundEvent(subscription *remoteEventSubscription, event Event) bool { + r.backgroundMu.Lock() + handler := r.backgroundHandler + active := r.background == subscription && !r.closed + r.backgroundMu.Unlock() + if !active || handler == nil { + return false + } + if recovered, ok := event.(*SessionRecoveredEvent); ok { + r.stateMu.Lock() + // Root idleness does not complete detached sub-session requests. + for id, request := range r.pendingOAuthElicitations { + if request.SessionID == "" || request.SessionID == recovered.SessionID { + delete(r.pendingOAuthElicitations, id) + } + } + r.stateMu.Unlock() + } else { + r.trackOAuthElicitation(event) + } + handler(event) + return true +} + +func (r *RemoteRuntime) emitBackgroundEvent(subscription *remoteEventSubscription, event Event) bool { + return subscription.history.deliver(event, func(event Event) bool { + return r.deliverBackgroundEvent(subscription, event) + }) +} + +func (r *RemoteRuntime) clearBackgroundSubscription(subscription *remoteEventSubscription) { + subscription.cancel() + r.backgroundMu.Lock() + defer r.backgroundMu.Unlock() + if r.background == subscription { + r.background = nil + } +} + +func (r *RemoteRuntime) startBackgroundEvents(ctx context.Context, sessionID string) { + client, ok := r.client.(remoteEventClient) + if !ok { + return + } + r.backgroundInit.Lock() + defer r.backgroundInit.Unlock() + r.backgroundMu.Lock() + if r.closed || r.backgroundHandler == nil { + r.backgroundMu.Unlock() + return + } + if r.background != nil && r.background.sessionID == sessionID { + r.backgroundMu.Unlock() + return + } + if r.background != nil { + r.background.cancel() + } + backgroundCtx, cancel := context.WithCancel(context.WithoutCancel(ctx)) + subscription := &remoteEventSubscription{cancel: cancel, sessionID: sessionID, history: newRemoteMessageHistory(nil)} + r.background = subscription + r.backgroundMu.Unlock() + + // Capture the cursor BEFORE submitting a turn; no old answers are replayed. + snapshotCtx, cancelSnapshot := context.WithCancel(ctx) + stop := context.AfterFunc(backgroundCtx, cancelSnapshot) + snapshot, err := client.GetSessionSnapshot(snapshotCtx, sessionID) + stop() + cancelSnapshot() + if err == nil && (snapshot == nil || snapshot.ID != sessionID) { + err = errors.New("snapshot session ID does not match subscription") + } + if err == nil && snapshot.Streaming { + err = errors.New("session is already streaming; no safe background baseline") + } + if err != nil { + r.emitBackgroundEvent(subscription, Warning(fmt.Sprintf("remote background events unavailable: %v", err), "")) + r.clearBackgroundSubscription(subscription) + return + } + r.backgroundMu.Lock() + if r.background != subscription || r.closed { + r.backgroundMu.Unlock() + cancel() + return + } + subscription.history = newRemoteMessageHistory(snapshot.Messages) + r.backgroundMu.Unlock() + + go func() { + defer r.clearBackgroundSubscription(subscription) + events, err := r.openBackgroundEvents(backgroundCtx, client, sessionID, snapshot.LastEventSeq) + if err != nil { + r.emitBackgroundEvent(subscription, Error(fmt.Sprintf("subscribing to remote background events: %v", err))) + return + } + for { + select { + case <-backgroundCtx.Done(): + return + case event, ok := <-events: + if !ok { + return // session ended; never blindly replay old history + } + if failure, ok := event.(*ErrorEvent); ok && failure.Error == sessionEventGapError { + r.emitBackgroundEvent(subscription, Warning("remote event gap; recovering saved assistant text when idle (tool and interaction events may be incomplete)", "")) + if !waitEventStreamRetry(backgroundCtx, 250*time.Millisecond) { + return + } + var err error + snapshot, err = r.reconcileBackgroundSnapshot(backgroundCtx, subscription, client) + if err != nil { + r.emitBackgroundEvent(subscription, Error(fmt.Sprintf("recovering remote event gap: %v", err))) + return + } + events, err = r.openBackgroundEvents(backgroundCtx, client, sessionID, snapshot.LastEventSeq) + if err != nil { + r.emitBackgroundEvent(subscription, Error(fmt.Sprintf("resuming remote background events: %v", err))) + return + } + continue + } + if !r.emitBackgroundEvent(subscription, event) { + return + } + } + } + }() +} + +// /events may not exist until the first recall or elicitation creates its log. +func (r *RemoteRuntime) openBackgroundEvents(ctx context.Context, client remoteEventClient, sessionID string, since uint64) (<-chan Event, error) { + delay := 250 * time.Millisecond + for { + events, err := client.StreamSessionEventsSince(ctx, sessionID, since) + if err == nil { + return events, nil + } + var httpErr *sessionEventHTTPError + if errors.As(err, &httpErr) && httpErr.status >= 400 && httpErr.status < 500 && httpErr.status != http.StatusNotFound && httpErr.status != http.StatusTooManyRequests { + return nil, err + } + if !waitEventStreamRetry(ctx, delay) { + return nil, ctx.Err() + } + delay = min(2*delay, 5*time.Second) + } +} + +func (r *RemoteRuntime) reconcileBackgroundSnapshot(ctx context.Context, subscription *remoteEventSubscription, client remoteEventClient) (*api.SessionSnapshotResponse, error) { + for { + if !r.reconcileMu.TryLock() { + if !waitEventStreamRetry(ctx, 250*time.Millisecond) { + return nil, ctx.Err() + } + continue + } + snapshot, err := r.reconcileIdleSnapshot(ctx, subscription, client) + r.reconcileMu.Unlock() + if err != nil || !snapshot.Streaming { + return snapshot, err + } + if !waitEventStreamRetry(ctx, 250*time.Millisecond) { + return nil, ctx.Err() + } + } +} + +func (r *RemoteRuntime) reconcileIdleSnapshot(ctx context.Context, subscription *remoteEventSubscription, client remoteEventClient) (*api.SessionSnapshotResponse, error) { + snapshot, err := client.GetSessionSnapshot(ctx, subscription.sessionID) + if err != nil { + return nil, err + } + if snapshot == nil || snapshot.ID != subscription.sessionID { + return nil, errors.New("snapshot session ID does not match subscription") + } + if snapshot.Streaming { + return snapshot, nil // partial saved messages cannot define a safe cursor + } + r.backgroundMu.Lock() + handler := r.backgroundHandler + active := r.background == subscription && !r.closed + r.backgroundMu.Unlock() + if !active || handler == nil || ctx.Err() != nil { + return nil, context.Canceled + } + if err := subscription.history.reconcile(snapshot, func(event Event) { + r.deliverBackgroundEvent(subscription, event) + }); err != nil { + return nil, err + } + if !r.deliverBackgroundEvent(subscription, SessionRecovered(snapshot.ID)) { + return nil, context.Canceled + } + return snapshot, nil +} // OnElicitationRequest is a no-op for remote runtimes; elicitation requests -// (including from server-side background jobs) arrive as ElicitationRequestEvent -// values on the RunStream channel itself (see the RunStream forwarding loop -// above), so there is no separate out-of-band sink to register. +// arrive on RunStream, or through OnBackgroundEvent when a cursor-based +// subscription is active, so there is no additional sink to register. // // RemoteRuntime deliberately does NOT implement // LocalRuntime.MirrorsElicitationOnRunStream: embedders that forward // RunStream events verbatim (e.g. pkg/app.App) rely on that capability check -// to tell that this runtime's RunStream copy is its ONLY delivery and must -// reach them unfiltered (#3584 review). +// to forward RunStream elicitations when no background subscription is active. func (r *RemoteRuntime) OnElicitationRequest(func(Event)) {} -// Close is a no-op for remote runtimes. +// RetireBackgroundEvents stops deliveries for a replaced conversation. +func (r *RemoteRuntime) RetireBackgroundEvents() { + r.backgroundMu.Lock() + defer r.backgroundMu.Unlock() + if r.background != nil { + r.background.cancel() + r.background = nil + } +} + +// Close stops the out-of-band session subscription. func (r *RemoteRuntime) Close() error { + r.backgroundMu.Lock() + defer r.backgroundMu.Unlock() + r.closed = true + r.backgroundHandler = nil + if r.background != nil { + r.background.cancel() + r.background = nil + } return nil } // GetSnapshots retrieves available snapshots for the current session. func (r *RemoteRuntime) GetSnapshots(ctx context.Context) ([]map[string]any, error) { - if r.sessionID == "" { + sessionID := r.activeSessionID() + if sessionID == "" { return nil, errors.New("no active session") } - return r.client.GetSessionSnapshots(ctx, r.sessionID) + return r.client.GetSessionSnapshots(ctx, sessionID) } // Undo reverts to the previous snapshot on the remote server. func (r *RemoteRuntime) Undo(ctx context.Context) error { - if r.sessionID == "" { + sessionID := r.activeSessionID() + if sessionID == "" { return errors.New("no active session") } - return r.client.UndoSession(ctx, r.sessionID) + return r.client.UndoSession(ctx, sessionID) } // Reset resets the session to its initial state on the remote server. func (r *RemoteRuntime) Reset(ctx context.Context) error { - if r.sessionID == "" { + sessionID := r.activeSessionID() + if sessionID == "" { return errors.New("no active session") } - return r.client.ResetSession(ctx, r.sessionID) + return r.client.ResetSession(ctx, sessionID) } // AddMessageToSession adds a message to the current session on the remote server. func (r *RemoteRuntime) AddMessageToSession(ctx context.Context, msg *session.Message) error { - if r.sessionID == "" { + sessionID := r.activeSessionID() + if sessionID == "" { return errors.New("no active session") } - return r.client.AddMessage(ctx, r.sessionID, msg) + return r.client.AddMessage(ctx, sessionID, msg) } // UpdateSessionMessage updates a message in the current session on the remote server. func (r *RemoteRuntime) UpdateSessionMessage(ctx context.Context, msgID string, msg *session.Message) error { - if r.sessionID == "" { + sessionID := r.activeSessionID() + if sessionID == "" { return errors.New("no active session") } - return r.client.UpdateMessage(ctx, r.sessionID, msgID, msg) + return r.client.UpdateMessage(ctx, sessionID, msgID, msg) } // AddSessionSummary adds a summary item to the current session on the remote server. func (r *RemoteRuntime) AddSessionSummary(ctx context.Context, item session.Item) error { - if r.sessionID == "" { + sessionID := r.activeSessionID() + if sessionID == "" { return errors.New("no active session") } - return r.client.AddSummary(ctx, r.sessionID, item) + return r.client.AddSummary(ctx, sessionID, item) } // UpdateSessionTokens updates token counts for the current session on the remote server. func (r *RemoteRuntime) UpdateSessionTokens(ctx context.Context, inputTokens, outputTokens int64, cost float64) error { - if r.sessionID == "" { + sessionID := r.activeSessionID() + if sessionID == "" { return errors.New("no active session") } - return r.client.UpdateSessionTokens(ctx, r.sessionID, inputTokens, outputTokens, cost) + return r.client.UpdateSessionTokens(ctx, sessionID, inputTokens, outputTokens, cost) } // SetSessionStarred sets the starred status for the current session on the remote server. func (r *RemoteRuntime) SetSessionStarred(ctx context.Context, starred bool) error { - if r.sessionID == "" { + sessionID := r.activeSessionID() + if sessionID == "" { return errors.New("no active session") } - return r.client.SetSessionStarred(ctx, r.sessionID, starred) + return r.client.SetSessionStarred(ctx, sessionID, starred) } var _ Runtime = (*RemoteRuntime)(nil) diff --git a/pkg/runtime/remote_runtime_test.go b/pkg/runtime/remote_runtime_test.go index 87a971350d..5ceaeb2b9a 100644 --- a/pkg/runtime/remote_runtime_test.go +++ b/pkg/runtime/remote_runtime_test.go @@ -3,14 +3,24 @@ package runtime import ( "context" "errors" + "fmt" + "net/http" + "net/http/httptest" + "strconv" + "strings" + "sync" + "sync/atomic" "testing" + "time" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" "github.com/docker/docker-agent/pkg/api" + "github.com/docker/docker-agent/pkg/chat" "github.com/docker/docker-agent/pkg/config/latest" "github.com/docker/docker-agent/pkg/session" + "github.com/docker/docker-agent/pkg/tools" ) // runStreamRecordingClient is a stubRemoteClient variant that records the @@ -146,3 +156,710 @@ func TestRemoteRuntime_CurrentMCPPrompts(t *testing.T) { assert.Equal(t, "path", prompts["review"].Arguments[0].Name) assert.True(t, prompts["review"].Arguments[0].Required) } + +func TestRemoteRuntime_BackgroundEventsSurviveTurnsWithoutReplayingHistory(t *testing.T) { + t.Parallel() + + var runs, snapshots, subscriptions atomic.Int32 + backgroundReady := make(chan struct{}) + deliverRecall := make(chan struct{}) + backgroundStopped := make(chan struct{}) + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) { + switch { + case strings.HasSuffix(req.URL.Path, "/snapshot"): + snapshots.Add(1) + fmt.Fprint(w, `{"id":"s","last_event_seq":12}`) + case strings.HasSuffix(req.URL.Path, "/events"): + subscriptions.Add(1) + assert.Equal(t, "12", req.Header.Get("Last-Event-ID")) + w.Header().Set("Content-Type", "text/event-stream") + w.WriteHeader(http.StatusOK) + w.(http.Flusher).Flush() + close(backgroundReady) + select { + case <-deliverRecall: + fmt.Fprint(w, "id: 13\ndata: {\"type\":\"agent_choice\",\"message_id\":\"recall\",\"content\":\"recall answer\"}\n\n") + w.(http.Flusher).Flush() + case <-req.Context().Done(): + } + <-req.Context().Done() + close(backgroundStopped) + case req.Method == http.MethodPost: + runs.Add(1) + assert.Equal(t, int32(1), snapshots.Load(), "cursor must be taken before RunAgent") + w.Header().Set("Content-Type", "text/event-stream") + fmt.Fprint(w, "data: {\"type\":\"agent_choice\",\"message_id\":\"foreground\",\"content\":\"foreground answer\"}\n\ndata: {\"type\":\"stream_stopped\"}\n\n") + default: + http.NotFound(w, req) + } + })) + t.Cleanup(srv.Close) + client, err := NewClient(srv.URL) + require.NoError(t, err) + rt, err := NewRemoteRuntime(client) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, rt.Close()) }) + background := make(chan Event, 4) + rt.OnBackgroundEvent(func(event Event) { background <- event }) + + ctx, cancel := context.WithCancel(t.Context()) + var foreground []Event + for event := range rt.RunStream(ctx, &session.Session{ID: "s"}) { + foreground = append(foreground, event) + } + require.Len(t, foreground, 2) + cancel() // The turn context must not end the session subscription. + select { + case <-backgroundReady: + case <-time.After(2 * time.Second): + t.Fatal("background subscription did not connect") + } + close(deliverRecall) + select { + case event := <-background: + answer, ok := event.(*AgentChoiceEvent) + require.True(t, ok, "got %T", event) + assert.Equal(t, "recall answer", answer.Content) + case <-time.After(2 * time.Second): + t.Fatal("idle recall was not delivered") + } + + for range rt.RunStream(t.Context(), &session.Session{ID: "s"}) { + } + assert.Equal(t, int32(2), runs.Load(), "subscription must not spawn extra runs") + assert.Equal(t, int32(1), snapshots.Load(), "do not re-snapshot and skip pending events between turns") + assert.Equal(t, int32(1), subscriptions.Load()) + assert.Empty(t, background, "foreground answers must not also reach the background sink") + require.NoError(t, rt.Close()) + select { + case <-backgroundStopped: + case <-time.After(2 * time.Second): + t.Fatal("Close did not stop background subscription") + } +} + +func TestRemoteRuntime_BackgroundSubscriptionWaitsForEventLogAndDeliversElicitationOnce(t *testing.T) { + t.Parallel() + + var logAvailable atomic.Bool + var attempts, runs atomic.Int32 + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) { + switch { + case strings.HasSuffix(req.URL.Path, "/snapshot"): + fmt.Fprint(w, `{"id":"s","last_event_seq":0}`) + case strings.HasSuffix(req.URL.Path, "/events"): + attempts.Add(1) + assert.Equal(t, "0", req.Header.Get("Last-Event-ID")) + if !logAvailable.Load() { + http.NotFound(w, req) + return + } + w.Header().Set("Content-Type", "text/event-stream") + fmt.Fprint(w, "id: 1\ndata: {\"type\":\"elicitation_request\",\"elicitation_id\":\"eid\"}\n\nid: 2\ndata: {\"type\":\"session_exited\"}\n\n") + case req.Method == http.MethodPost: + runs.Add(1) + assert.Eventually(t, func() bool { return attempts.Load() > 0 }, 2*time.Second, time.Millisecond) + logAvailable.Store(true) + w.Header().Set("Content-Type", "text/event-stream") + fmt.Fprint(w, "data: {\"type\":\"elicitation_request\",\"elicitation_id\":\"eid\"}\n\ndata: {\"type\":\"stream_stopped\"}\n\n") + } + })) + t.Cleanup(srv.Close) + client, err := NewClient(srv.URL) + require.NoError(t, err) + rt, err := NewRemoteRuntime(client) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, rt.Close()) }) + background := make(chan Event, 4) + rt.OnBackgroundEvent(func(event Event) { background <- event }) + var foreground []Event + for event := range rt.RunStream(t.Context(), &session.Session{ID: "s"}) { + foreground = append(foreground, event) + } + var requests []*ElicitationRequestEvent + for _, event := range foreground { + if request, ok := event.(*ElicitationRequestEvent); ok { + requests = append(requests, request) + } + } + require.Eventually(t, func() bool { return attempts.Load() >= 2 }, 3*time.Second, time.Millisecond) + select { + case event := <-background: + request, ok := event.(*ElicitationRequestEvent) + require.True(t, ok, "got %T", event) + requests = append(requests, request) + default: + } + require.Len(t, requests, 1, "foreground/background copies must be delivered once") + assert.Equal(t, "eid", requests[0].ElicitationID) + assert.Empty(t, background) + assert.Equal(t, int32(1), runs.Load()) + require.GreaterOrEqual(t, attempts.Load(), int32(2)) +} + +func TestRemoteRuntime_BackgroundGapReconcilesSavedTextWithoutReplay(t *testing.T) { + t.Parallel() + + var runs, snapshots, subscriptions atomic.Int32 + foregroundDone := make(chan struct{}) + streamReady := make(chan struct{}) + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) { + switch { + case strings.HasSuffix(req.URL.Path, "/snapshot"): + w.Header().Set("Content-Type", "application/json") + if snapshots.Add(1) == 1 { + fmt.Fprint(w, `{"id":"s","last_event_seq":10,"messages":[{"agent_name":"root","message":{"role":"assistant","message_id":"old","content":"old answer"}}]}`) + } else { + fmt.Fprint(w, `{"id":"s","last_event_seq":30,"messages":[ + {"agent_name":"root","message":{"role":"assistant","message_id":"old","content":"old answer"}}, + {"agent_name":"root","message":{"role":"assistant","message_id":"foreground","content":"foreground answer"}}, + {"agent_name":"root","message":{"role":"assistant","message_id":"recall","content":"recall final","reasoning_content":"PRIVATE reasoning"}}, + {"message":{"role":"tool","message_id":"tool","content":"PRIVATE tool output"}}, + {"message":{"role":"assistant","message_id":"tool-call","content":"PRIVATE tool content","tool_calls":[{"id":"call","type":"function"}]}}, + {"implicit":true,"message":{"role":"assistant","message_id":"implicit","content":"PRIVATE implicit"}} + ]}`) + } + case strings.HasSuffix(req.URL.Path, "/events"): + w.Header().Set("Content-Type", "text/event-stream") + if subscriptions.Add(1) == 1 { + assert.Equal(t, "10", req.Header.Get("Last-Event-ID")) + w.(http.Flusher).Flush() + select { + case <-foregroundDone: + case <-req.Context().Done(): + return + } + fmt.Fprint(w, "id: 11\ndata: {\"type\":\"agent_choice\",\"message_id\":\"recall\",\"content\":\"recall \"}\n\ndata: {\"type\":\"gap\"}\n\nid: 20\ndata: {\"type\":\"agent_choice\",\"message_id\":\"recall\",\"content\":\"blind replay\"}\n\n") + return + } + assert.Equal(t, "30", req.Header.Get("Last-Event-ID")) + // Snapshot/log overlap is possible: the completed ID must not append twice. + fmt.Fprint(w, "id: 31\ndata: {\"type\":\"agent_choice\",\"message_id\":\"recall\",\"content\":\"recall final\"}\n\nid: 32\ndata: {\"type\":\"agent_choice\",\"message_id\":\"later\",\"content\":\"later answer\"}\n\n") + w.(http.Flusher).Flush() + close(streamReady) + <-req.Context().Done() + case req.Method == http.MethodPost: + runs.Add(1) + w.Header().Set("Content-Type", "text/event-stream") + fmt.Fprint(w, "data: {\"type\":\"agent_choice\",\"message_id\":\"foreground\",\"content\":\"foreground answer\"}\n\ndata: {\"type\":\"stream_stopped\",\"session_id\":\"s\"}\n\n") + } + })) + t.Cleanup(srv.Close) + client, err := NewClient(srv.URL) + require.NoError(t, err) + rt, err := NewRemoteRuntime(client) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, rt.Close()) }) + background := make(chan Event, 16) + rt.OnBackgroundEvent(func(event Event) { background <- event }) + var foreground []Event + for event := range rt.RunStream(t.Context(), &session.Session{ID: "s"}) { + foreground = append(foreground, event) + } + require.Len(t, foreground, 2) + close(foregroundDone) + select { + case <-streamReady: + case <-time.After(3 * time.Second): + t.Fatal("gap recovery did not reconnect") + } + var choices []*AgentChoiceEvent + var recovered int + deadline := time.After(3 * time.Second) + for len(choices) < 3 { + select { + case event := <-background: + switch event := event.(type) { + case *AgentChoiceEvent: + choices = append(choices, event) + case *WarningEvent: + assert.Contains(t, event.Message, "gap") + case *SessionRecoveredEvent: + recovered++ + assert.Equal(t, "s", event.SessionID) + default: + t.Fatalf("unexpected event %T", event) + } + case <-deadline: + t.Fatal("saved final answer was not reconciled") + } + } + assert.Equal(t, 1, recovered, "idle recovery must reset lifecycle state exactly once") + assert.Equal(t, "recall ", choices[0].Content) + assert.Equal(t, "final", choices[1].Content, "append only the missing suffix") + assert.Equal(t, "recall", choices[1].MessageID) + assert.Equal(t, "s", choices[1].SessionID) + assert.Equal(t, "later answer", choices[2].Content) + assert.Empty(t, background) + assert.Equal(t, int32(1), runs.Load()) + assert.Equal(t, int32(2), snapshots.Load()) + assert.Equal(t, int32(2), subscriptions.Load()) +} + +func TestRemoteMessageHistoryRejectsUnsafeRecovery(t *testing.T) { + t.Parallel() + + for _, tc := range []struct { + name string + delivered *AgentChoiceEvent + savedID string + saved string + }{ + {name: "unidentified delivered message", delivered: AgentChoice("root", "s", "partial").(*AgentChoiceEvent), savedID: "id", saved: "partial final"}, + {name: "unidentified saved message", delivered: AgentChoice("root", "s", "partial", "id").(*AgentChoiceEvent), saved: "partial final"}, + {name: "changed prefix", delivered: AgentChoice("root", "s", "partial", "id").(*AgentChoiceEvent), savedID: "id", saved: "rewritten final"}, + } { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + history := newRemoteMessageHistory(nil) + require.True(t, history.deliver(tc.delivered, func(Event) bool { return true })) + snapshot := &api.SessionSnapshotResponse{ID: "s", Messages: []session.Message{ + {Message: chat.Message{Role: chat.MessageRoleAssistant, MessageID: "new", Content: "must not partially replay"}}, + {Message: chat.Message{Role: chat.MessageRoleAssistant, MessageID: tc.savedID, Content: tc.saved}}, + }} + var got []Event + err := history.reconcile(snapshot, func(event Event) { got = append(got, event) }) + require.Error(t, err) + assert.Empty(t, got) + }) + } +} + +func TestRemoteRuntime_ForegroundEOFHasErrorStop(t *testing.T) { + t.Parallel() + + client := &runStreamRecordingClient{} + rt, err := NewRemoteRuntime(client) + require.NoError(t, err) + var got []Event + for event := range rt.RunStream(t.Context(), &session.Session{ID: "s"}) { + got = append(got, event) + } + require.Len(t, got, 2) + assert.IsType(t, &ErrorEvent{}, got[0]) + stop, ok := got[1].(*StreamStoppedEvent) + require.True(t, ok) + assert.Equal(t, "error", stop.Reason) + assert.Equal(t, "s", stop.SessionID) + assert.Equal(t, 1, client.gotInvocation) +} + +func TestRemoteRuntime_BackgroundGapWaitsForIdleAndCloseCancelsRecovery(t *testing.T) { + t.Parallel() + + var runs, snapshots, subscriptions atomic.Int32 + recovering := make(chan struct{}) + recoveryStopped := make(chan struct{}) + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) { + switch { + case strings.HasSuffix(req.URL.Path, "/snapshot"): + switch snapshots.Add(1) { + case 1: + fmt.Fprint(w, `{"id":"s","last_event_seq":0}`) + case 2: + fmt.Fprint(w, `{"id":"s","last_event_seq":20,"streaming":true,"messages":[{"message":{"role":"assistant","message_id":"id","content":"not final"}}]}`) + default: + close(recovering) + <-req.Context().Done() + close(recoveryStopped) + } + case strings.HasSuffix(req.URL.Path, "/events"): + subscriptions.Add(1) + w.Header().Set("Content-Type", "text/event-stream") + fmt.Fprint(w, "data: {\"type\":\"gap\"}\n\n") + case req.Method == http.MethodPost: + runs.Add(1) + w.Header().Set("Content-Type", "text/event-stream") + fmt.Fprint(w, "data: {\"type\":\"stream_stopped\"}\n\n") + } + })) + t.Cleanup(srv.Close) + client, err := NewClient(srv.URL) + require.NoError(t, err) + rt, err := NewRemoteRuntime(client) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, rt.Close()) }) + background := make(chan Event, 4) + rt.OnBackgroundEvent(func(event Event) { background <- event }) + for range rt.RunStream(t.Context(), &session.Session{ID: "s"}) { + } + select { + case <-recovering: + case <-time.After(3 * time.Second): + t.Fatal("did not wait for an idle snapshot") + } + require.NoError(t, rt.Close()) + select { + case <-recoveryStopped: + case <-time.After(2 * time.Second): + t.Fatal("Close did not cancel snapshot recovery") + } + assert.Equal(t, int32(1), runs.Load()) + assert.Equal(t, int32(1), subscriptions.Load(), "never advance the cursor from a running snapshot") + for len(background) > 0 { + assert.IsType(t, &WarningEvent{}, <-background, "partial saved text must not be emitted") + } +} + +func TestRemoteMessageHistoryPreservesKnownSubSessionScope(t *testing.T) { + t.Parallel() + + history := newRemoteMessageHistory(nil) + require.True(t, history.deliver(AgentChoice("worker", "child", "partial ", "id"), func(Event) bool { return true })) + snapshot := &api.SessionSnapshotResponse{ID: "parent", Messages: []session.Message{ + {AgentName: "worker", Message: chat.Message{Role: chat.MessageRoleAssistant, MessageID: "id", Content: "partial final"}}, + }} + var got []Event + require.NoError(t, history.reconcile(snapshot, func(event Event) { got = append(got, event) })) + require.Len(t, got, 1) + choice, ok := got[0].(*AgentChoiceEvent) + require.True(t, ok) + assert.Equal(t, "child", choice.SessionID) + assert.Equal(t, "final", choice.Content) +} + +func TestRemoteMessageHistoryDeduplicatesElicitationEitherDeliveryOrder(t *testing.T) { + t.Parallel() + + history := newRemoteMessageHistory(nil) + for _, id := range []string{"foreground-first", "background-first"} { + request := &ElicitationRequestEvent{ElicitationID: id} + var deliveries int + for range 2 { + require.True(t, history.deliver(request, func(Event) bool { + deliveries++ + return true + })) + } + assert.Equal(t, 1, deliveries) + } +} + +func TestRemoteMessageHistoryRejectsDuplicateSnapshotIDs(t *testing.T) { + t.Parallel() + + history := newRemoteMessageHistory(nil) + snapshot := &api.SessionSnapshotResponse{ID: "s", Messages: []session.Message{ + {Message: chat.Message{Role: chat.MessageRoleAssistant, MessageID: "id", Content: "first"}}, + {Message: chat.Message{Role: chat.MessageRoleAssistant, MessageID: "id", Content: "second"}}, + }} + var got []Event + require.Error(t, history.reconcile(snapshot, func(event Event) { got = append(got, event) })) + assert.Empty(t, got) +} + +func TestRemoteRuntimeRetiredBackgroundSubscriptionDoesNotEmit(t *testing.T) { + t.Parallel() + rt, err := NewRemoteRuntime(&stubRemoteClient{}) + require.NoError(t, err) + canceled := make(chan struct{}) + sub := &remoteEventSubscription{cancel: func() { close(canceled) }, sessionID: "old", history: newRemoteMessageHistory(nil)} + rt.background = sub + var delivered []Event + rt.OnBackgroundEvent(func(event Event) { delivered = append(delivered, event) }) + rt.RetireBackgroundEvents() + <-canceled + rt.emitBackgroundEvent(sub, AgentChoice("root", "old", "OLD-ANSWER", "old")) + require.Empty(t, delivered) + require.Nil(t, rt.background) +} + +func TestRemoteMessageHistoryUsesVisibleContent(t *testing.T) { + t.Parallel() + msg := func(id, content string) session.Message { + return session.Message{Message: chat.Message{Role: chat.MessageRoleAssistant, MessageID: id, Content: content}} + } + history := newRemoteMessageHistory([]session.Message{msg("old", "visiblePRIVATE"), msg("hidden", "PRIVATE")}) + assert.True(t, history.complete["old"]) + assert.False(t, history.complete["hidden"], "hidden-only text must not mask later visible chunks") + assert.Empty(t, history.content, "baseline needs IDs, not copies of old answers") + require.True(t, history.deliver(AgentChoice("root", "s", "answerPRIVATE", "new"), func(Event) bool { return true })) + var got []Event + require.NoError(t, history.reconcile(&api.SessionSnapshotResponse{ID: "s", Messages: []session.Message{ + msg("old", "visiblePRIVATE"), msg("new", "answerPRIVATE MORE"), msg("recovered", "safeSECRET"), + }}, func(event Event) { got = append(got, event) })) + require.Len(t, got, 1) + assert.Equal(t, "safe", got[0].(*AgentChoiceEvent).Content) + assert.Empty(t, history.content, "finalized answers retain only dedup IDs") +} + +func TestRemoteMessageHistoryRetentionFailsClosed(t *testing.T) { + t.Parallel() + history := newRemoteMessageHistory(nil) + for i := range remoteHistoryMaxMessages + 2 { + require.True(t, history.deliver(AgentChoice("root", "s", "answer", strconv.Itoa(i)), func(Event) bool { return true })) + } + assert.Empty(t, history.content) + assert.Empty(t, history.sessions) + assert.Zero(t, history.bytes) + var got []Event + err := history.reconcile(&api.SessionSnapshotResponse{ID: "s"}, func(event Event) { got = append(got, event) }) + require.ErrorContains(t, err, "retention limit") + assert.Empty(t, got, "forgotten IDs must never cause blind replay") +} + +type recoveryRemoteClient struct { + runStreamRecordingClient + + snapshot *api.SessionSnapshotResponse + openErr error + streams chan Event + opened chan struct{} +} + +func (c *recoveryRemoteClient) GetSessionSnapshot(context.Context, string) (*api.SessionSnapshotResponse, error) { + return c.snapshot, nil +} + +func (c *recoveryRemoteClient) StreamSessionEventsSince(context.Context, string, uint64) (<-chan Event, error) { + if c.opened != nil { + close(c.opened) + } + return c.streams, c.openErr +} + +func TestRemoteRuntimeRecoveryStopsReplayAfterRetirement(t *testing.T) { + t.Parallel() + client := &recoveryRemoteClient{snapshot: &api.SessionSnapshotResponse{ID: "s", Messages: []session.Message{ + {Message: chat.Message{Role: chat.MessageRoleAssistant, MessageID: "first", Content: "first"}}, + {Message: chat.Message{Role: chat.MessageRoleAssistant, MessageID: "second", Content: "stale"}}, + }}} + rt, err := NewRemoteRuntime(client) + require.NoError(t, err) + sub := &remoteEventSubscription{sessionID: "s", history: newRemoteMessageHistory(nil), cancel: func() {}} + rt.background = sub + var got []Event + rt.OnBackgroundEvent(func(event Event) { + got = append(got, event) + rt.RetireBackgroundEvents() + }) + _, err = rt.reconcileIdleSnapshot(t.Context(), sub, client) + require.ErrorIs(t, err, context.Canceled) + require.Len(t, got, 1, "captured handler must not receive the rest of a retired snapshot") + assert.Equal(t, "first", got[0].(*AgentChoiceEvent).Content) +} + +func TestRemoteRuntimeFailedBackgroundSubscriptionCanRestart(t *testing.T) { + t.Parallel() + client := &recoveryRemoteClient{snapshot: &api.SessionSnapshotResponse{ID: "s"}, openErr: &sessionEventHTTPError{status: http.StatusForbidden}} + rt, err := NewRemoteRuntime(client) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, rt.Close()) }) + delivered := make(chan Event, 4) + rt.OnBackgroundEvent(func(event Event) { delivered <- event }) + rt.startBackgroundEvents(t.Context(), "s") + select { + case event := <-delivered: + assert.IsType(t, &ErrorEvent{}, event) + case <-time.After(2 * time.Second): + t.Fatal("subscription did not fail") + } + require.Eventually(t, func() bool { + rt.backgroundMu.Lock() + defer rt.backgroundMu.Unlock() + return rt.background == nil + }, time.Second, time.Millisecond) + client.openErr = nil + client.streams = make(chan Event, 1) + client.streams <- Warning("retry succeeded", "root") + close(client.streams) + rt.startBackgroundEvents(t.Context(), "s") + select { + case event := <-delivered: + assert.Equal(t, "retry succeeded", event.(*WarningEvent).Message) + case <-time.After(2 * time.Second): + t.Fatal("subscription did not restart") + } +} + +func TestRemoteRuntimeRecoveryIsAnIdleResetNotAStop(t *testing.T) { + t.Parallel() + client := &recoveryRemoteClient{snapshot: &api.SessionSnapshotResponse{ID: "s", LastEventSeq: 42}} + rt, err := NewRemoteRuntime(client) + require.NoError(t, err) + sub := &remoteEventSubscription{sessionID: "s", history: newRemoteMessageHistory(nil), cancel: func() {}} + rt.background = sub + depth := map[string]int{} + var stops, resets int + rt.OnBackgroundEvent(func(event Event) { + switch event := event.(type) { + case *StreamStartedEvent: + depth[event.SessionID]++ + case *StreamStoppedEvent: + stops++ + depth[event.SessionID]-- + case *SessionRecoveredEvent: + resets++ + clear(depth) + } + }) + for _, id := range []string{"s", "s", "child"} { + rt.emitBackgroundEvent(sub, StreamStarted(id, "root")) + } + assert.Equal(t, 2, depth["s"]) + _, err = rt.reconcileIdleSnapshot(t.Context(), sub, client) + require.NoError(t, err) + assert.Empty(t, depth, "one boundary resets all nested starts lost to the gap") + assert.Equal(t, 1, resets) + assert.Zero(t, stops, "recovery cannot trigger queued stop actions") +} + +type elicitationRecordingClient struct { + stubRemoteClient + + mu sync.Mutex + ids []string + actions []tools.ElicitationAction +} + +func (c *elicitationRecordingClient) ResumeElicitation(_ context.Context, _ string, action tools.ElicitationAction, _ map[string]any, id ...string) error { + c.mu.Lock() + defer c.mu.Unlock() + c.ids = append(c.ids, firstElicitationID(id)) + c.actions = append(c.actions, action) + return nil +} + +func TestRemoteRuntimeOAuthResponsesAreCorrelated(t *testing.T) { + t.Parallel() + client := &elicitationRecordingClient{} + rt, err := NewRemoteRuntime(client) + require.NoError(t, err) + rt.sessionID = "s" + for _, id := range []string{"oauth-a", "oauth-b"} { + rt.trackOAuthElicitation(&ElicitationRequestEvent{ElicitationID: id, Meta: map[string]any{"docker-agent/type": "oauth_flow"}}) + } + // A form response must not start the unrelated OAuth flow. + require.NoError(t, rt.ResumeElicitation(t.Context(), tools.ElicitationActionAccept, nil, "form")) + assert.Equal(t, []string{"form"}, client.ids) + require.ErrorContains(t, rt.ResumeElicitation(t.Context(), tools.ElicitationActionAccept, nil), "explicit") + assert.Len(t, rt.pendingOAuthElicitations, 2) + // Missing metadata fails before network/browser side effects and declines only A. + require.ErrorContains(t, rt.ResumeElicitation(t.Context(), tools.ElicitationActionAccept, nil, "oauth-a"), "server_url") + assert.Equal(t, []string{"form", "oauth-a"}, client.ids) + assert.Equal(t, tools.ElicitationActionDecline, client.actions[1]) + assert.NotContains(t, rt.pendingOAuthElicitations, "oauth-a") + assert.Contains(t, rt.pendingOAuthElicitations, "oauth-b") + require.NoError(t, rt.ResumeElicitation(t.Context(), tools.ElicitationActionDecline, nil, "oauth-b")) + assert.Empty(t, rt.pendingOAuthElicitations, "declining must not start OAuth") +} + +func TestRemoteRuntimeDuplicateElicitationDoesNotRestorePendingOAuth(t *testing.T) { + t.Parallel() + rt, err := NewRemoteRuntime(&elicitationRecordingClient{}) + require.NoError(t, err) + rt.sessionID = "s" + sub := &remoteEventSubscription{sessionID: "s", history: newRemoteMessageHistory(nil), cancel: func() {}} + rt.background = sub + var got []Event + rt.OnBackgroundEvent(func(event Event) { got = append(got, event) }) + request := &ElicitationRequestEvent{ElicitationID: "oauth", Meta: map[string]any{"docker-agent/type": "oauth_flow"}} + rt.emitBackgroundEvent(sub, request) + require.NoError(t, rt.ResumeElicitation(t.Context(), tools.ElicitationActionDecline, nil, "oauth")) + rt.emitBackgroundEvent(sub, request) + assert.Len(t, got, 1) + assert.Empty(t, rt.pendingOAuthElicitations, "suppressed duplicates must not mutate OAuth state") +} + +func TestRemoteMessageHistoryReleasesStoppedTextAndKeepsNestedScope(t *testing.T) { + t.Parallel() + history := newRemoteMessageHistory(nil) + for _, event := range []Event{ + AgentChoice("root", "s", "answer", "root"), + AgentChoice("root", "s", "PRIVATE", "hidden"), + AgentChoice("worker", "child", "partial", "child"), + StreamStopped("s", "root", ""), + } { + require.True(t, history.deliver(event, func(Event) bool { return true })) + } + assert.True(t, history.complete["root"]) + assert.False(t, history.complete["hidden"]) + assert.NotContains(t, history.content, "root") + assert.NotContains(t, history.content, "hidden") + assert.Contains(t, history.content, "child", "parent stop cannot finalize a nested stream") + assert.Equal(t, len("partial"), history.bytes) +} + +func TestRemoteMessageHistoryNestedStopDoesNotCompleteOuterMessage(t *testing.T) { + t.Parallel() + history := newRemoteMessageHistory(nil) + for _, event := range []Event{StreamStarted("s", "root"), StreamStarted("s", "worker"), AgentChoice("root", "s", "partial", "id"), StreamStopped("s", "worker", "")} { + require.True(t, history.deliver(event, func(Event) bool { return true })) + } + assert.False(t, history.complete["id"]) + var got []Event + require.True(t, history.deliver(AgentChoice("root", "s", " final", "id"), func(event Event) bool { got = append(got, event); return true })) + require.Len(t, got, 1) + require.True(t, history.deliver(StreamStopped("s", "root", ""), func(Event) bool { return true })) + assert.True(t, history.complete["id"]) + assert.Empty(t, history.content) + assert.Empty(t, history.streamDepth) +} + +func TestRemoteMessageHistoryElicitationRetentionDoesNotPermitDuplicates(t *testing.T) { + t.Parallel() + history := newRemoteMessageHistory(nil) + for i := range remoteHistoryMaxElicitations { + history.elicitations[strconv.Itoa(i)] = true + } + var got []Event + request := &ElicitationRequestEvent{ElicitationID: "overflow", Meta: map[string]any{"docker-agent/type": "oauth_flow"}} + require.False(t, history.deliver(request, func(event Event) bool { got = append(got, event); return true })) + require.Len(t, got, 1) + assert.IsType(t, &ErrorEvent{}, got[0], "unknown request cannot bypass bounded duplicate protection") + assert.Len(t, history.elicitations, remoteHistoryMaxElicitations) +} + +func TestRemoteRuntimeRootRecoveryPreservesDetachedOAuthRequests(t *testing.T) { + t.Parallel() + client := &elicitationRecordingClient{} + rt, err := NewRemoteRuntime(client) + require.NoError(t, err) + rt.sessionID = "root" + sub := &remoteEventSubscription{sessionID: "root", history: newRemoteMessageHistory(nil), cancel: func() {}} + rt.background = sub + var got []Event + rt.OnBackgroundEvent(func(event Event) { got = append(got, event) }) + requests := []*ElicitationRequestEvent{ + {SessionID: "root", ElicitationID: "foreground"}, + {ElicitationID: "legacy-foreground"}, + {SessionID: "detached-a", ElicitationID: "oauth-a", ServerElicitationID: "wire-id"}, + {SessionID: "detached-b", ElicitationID: "oauth-b", ServerElicitationID: "wire-id"}, + } + for _, request := range requests { + request.Meta = map[string]any{"docker-agent/type": "oauth_flow"} + require.True(t, rt.emitBackgroundEvent(sub, request)) + } + snapshotClient := &recoveryRemoteClient{snapshot: &api.SessionSnapshotResponse{ID: "root", LastEventSeq: 42}} + _, err = rt.reconcileIdleSnapshot(t.Context(), sub, snapshotClient) + require.NoError(t, err) + require.Len(t, got, 5) + assert.IsType(t, &SessionRecoveredEvent{}, got[4]) + assert.NotContains(t, rt.pendingOAuthElicitations, "foreground") + assert.NotContains(t, rt.pendingOAuthElicitations, "legacy-foreground") + require.Len(t, rt.pendingOAuthElicitations, 2, "root idleness does not complete detached requests") + assert.Same(t, requests[2], rt.pendingOAuthElicitations["oauth-a"]) + assert.Same(t, requests[3], rt.pendingOAuthElicitations["oauth-b"]) + + require.ErrorContains(t, rt.ResumeElicitation(t.Context(), tools.ElicitationActionAccept, nil), "explicit") + require.NoError(t, rt.ResumeElicitation(t.Context(), tools.ElicitationActionAccept, nil, "form")) + // Missing metadata fails before browser/network work, proving the OAuth path is still selected. + require.ErrorContains(t, rt.ResumeElicitation(t.Context(), tools.ElicitationActionAccept, nil, "oauth-a"), "server_url") + assert.Contains(t, rt.pendingOAuthElicitations, "oauth-b") + require.NoError(t, rt.ResumeElicitation(t.Context(), tools.ElicitationActionDecline, nil, "oauth-b")) + assert.Equal(t, []string{"form", "oauth-a", "oauth-b"}, client.ids) + assert.Equal(t, []tools.ElicitationAction{tools.ElicitationActionAccept, tools.ElicitationActionDecline, tools.ElicitationActionDecline}, client.actions) + assert.Empty(t, rt.pendingOAuthElicitations) +} + +func TestRemoteMessageHistoryBoundsUnmatchedStreamStarts(t *testing.T) { + t.Parallel() + history := newRemoteMessageHistory(nil) + for range remoteHistoryMaxMessages + 2 { + require.True(t, history.deliver(StreamStarted("s", "root"), func(Event) bool { return true })) + } + assert.Empty(t, history.streamDepth) + require.ErrorContains(t, history.reconcile(&api.SessionSnapshotResponse{ID: "s"}, func(Event) { + t.Fatal("overflowed history must not replay messages") + }), "retention limit") +} diff --git a/pkg/runtime/runtime.go b/pkg/runtime/runtime.go index fee0559a07..6429c41445 100644 --- a/pkg/runtime/runtime.go +++ b/pkg/runtime/runtime.go @@ -265,6 +265,7 @@ type LocalRuntime struct { // the sink being (re)registered. elicitationSinkMu sync.RWMutex onElicitationRequest func(Event) + onElicitationContext func(context.Context, Event) sessionStore session.Store // generatedFiles caches per-owner-session workspace roots and manifest // records for [LocalRuntime.ResolveGeneratedFile]. Seeded by @@ -378,8 +379,9 @@ type LocalRuntime struct { // background work (e.g. background agent tasks). Protected by // backgroundEventMu because background tasks read it from their own // goroutines. - backgroundEventMu sync.RWMutex - onBackgroundEvent func(Event) + backgroundEventMu sync.RWMutex + onBackgroundEvent func(Event) + onBackgroundContext func(context.Context, Event) bgAgents *agenttool.Handler @@ -1654,14 +1656,28 @@ func (r *LocalRuntime) OnBackgroundEvent(handler func(Event)) { r.backgroundEventMu.Lock() defer r.backgroundEventMu.Unlock() r.onBackgroundEvent = handler + r.onBackgroundContext = nil +} + +// OnBackgroundEventWithContext preserves detached producer ownership. +func (r *LocalRuntime) OnBackgroundEventWithContext(handler func(context.Context, Event)) { + r.backgroundEventMu.Lock() + defer r.backgroundEventMu.Unlock() + r.onBackgroundContext = handler + r.onBackgroundEvent = nil } // emitBackgroundEvent forwards an event from detached background work to the // registered handler, if any. -func (r *LocalRuntime) emitBackgroundEvent(event Event) { +func (r *LocalRuntime) emitBackgroundEvent(ctx context.Context, event Event) { r.backgroundEventMu.RLock() handler := r.onBackgroundEvent + contextual := r.onBackgroundContext r.backgroundEventMu.RUnlock() + if contextual != nil { + contextual(ctx, event) + return + } if handler != nil { handler(event) } diff --git a/pkg/runtime/synthetic_content_test.go b/pkg/runtime/synthetic_content_test.go new file mode 100644 index 0000000000..8da347ce06 --- /dev/null +++ b/pkg/runtime/synthetic_content_test.go @@ -0,0 +1,36 @@ +package runtime + +import ( + "testing" + + "github.com/stretchr/testify/require" + + "github.com/docker/docker-agent/pkg/agent" + "github.com/docker/docker-agent/pkg/chat" + "github.com/docker/docker-agent/pkg/session" +) + +func TestSyntheticAssistantAnnouncesTextBeforeCanonicalMessage(t *testing.T) { + t.Parallel() + sess := session.New() + a := agent.New("root", "test") + var events []Event + msg := &chat.Message{Role: chat.MessageRoleAssistant, Content: "FINAL-ANSWER"} + addAgentMessage(sess, a, msg, EventSinkFunc(func(e Event) { events = append(events, e) })) + require.Len(t, events, 2) + choice := events[0].(*AgentChoiceEvent) + require.Equal(t, "FINAL-ANSWER", choice.Content) + require.NotEmpty(t, choice.MessageID) + added := events[1].(*MessageAddedEvent) + require.Equal(t, choice.MessageID, added.Message.Message.MessageID) +} + +func TestAlreadyStreamedAssistantDoesNotAnnounceTextTwice(t *testing.T) { + t.Parallel() + sess := session.New() + var events []Event + msg := &chat.Message{Role: chat.MessageRoleAssistant, MessageID: "streamed", Content: "FINAL-ANSWER"} + addAgentMessage(sess, agent.New("root", "test"), msg, EventSinkFunc(func(e Event) { events = append(events, e) })) + require.Len(t, events, 1) + require.IsType(t, &MessageAddedEvent{}, events[0]) +} diff --git a/pkg/runtime/tool_dispatch.go b/pkg/runtime/tool_dispatch.go index 425cfedaaf..8120ad9788 100644 --- a/pkg/runtime/tool_dispatch.go +++ b/pkg/runtime/tool_dispatch.go @@ -2,6 +2,7 @@ package runtime import ( "context" + "uuid" "github.com/docker/docker-agent/pkg/agent" "github.com/docker/docker-agent/pkg/chat" @@ -188,7 +189,15 @@ func denySourceFor(checkerSource string) string { // and max-iteration stop messages. The dispatcher emits its own variant // directly via the [toolexec.Emitter] interface. func addAgentMessage(sess *session.Session, a *agent.Agent, msg *chat.Message, events EventSink) { + // Synthetic answers have no streaming phase; announce their text on the wire too. + synthetic := msg.Role == chat.MessageRoleAssistant && msg.MessageID == "" + if synthetic { + msg.MessageID = uuid.NewV4().String() + } agentMsg := session.NewAgentMessage(a.Name(), msg) sess.AddMessage(agentMsg) + if synthetic && msg.Content != "" { + events.Emit(AgentChoice(a.Name(), sess.ID, chat.VisibleAssistantContent(msg.Content), msg.MessageID)) + } events.Emit(MessageAdded(sess.ID, agentMsg, a.Name())) } diff --git a/pkg/server/session_manager.go b/pkg/server/session_manager.go index f13b8ef6ef..ea1e96e501 100644 --- a/pkg/server/session_manager.go +++ b/pkg/server/session_manager.go @@ -40,9 +40,10 @@ type activeRuntimes struct { session *session.Session // The actual session object used by the runtime titleGen *sessiontitle.Generator // Title generator (includes fallback models) - streaming sync.Mutex // Held while a RunStream is in progress; serialises concurrent requests - modelSwitch sync.Mutex // Serialises model changes with manager-owned session persistence - deleting atomic.Bool // Set before deletion waits for an in-flight model transaction + snapshotBoundary sync.Mutex // Separates idle snapshot ownership from live turn state. + streaming sync.Mutex // Held while a RunStream is in progress; serialises concurrent requests + modelSwitch sync.Mutex // Serialises model changes with manager-owned session persistence + deleting atomic.Bool // Set before deletion waits for an in-flight model transaction } // SessionManager manages sessions for HTTP and Connect-RPC servers. @@ -551,15 +552,18 @@ func (sm *SessionManager) GetSessionSnapshot(ctx context.Context, id string) (*a streaming := false agentName := "" if rs, ok := sm.runtimeSessions.Load(id); ok { - sess = rs.session - agentName = rs.runtime.CurrentAgentName(ctx) - // Probe streaming state without interfering: TryLock succeeds only - // when no RunStream is in progress. + rs.snapshotBoundary.Lock() + defer rs.snapshotBoundary.Unlock() + // Keep an idle turn stable until messages and cursor have been copied. if rs.streaming.TryLock() { - rs.streaming.Unlock() + defer rs.streaming.Unlock() } else { streaming = true } + rs.modelSwitch.Lock() + defer rs.modelSwitch.Unlock() + sess = rs.session + agentName = rs.runtime.CurrentAgentName(ctx) } if sess == nil { var err error @@ -570,6 +574,11 @@ func (sm *SessionManager) GetSessionSnapshot(ctx context.Context, id string) (*a } lastSeq, _ := sm.LastEventSeq(id) + sess = sess.Clone() + // Detached background events can advance the log even between turns. + if seq, _ := sm.LastEventSeq(id); seq != lastSeq { + streaming = true + } title := sess.TitleSnapshot() inputTokens, outputTokens := sess.Usage() @@ -1377,7 +1386,10 @@ func (sm *SessionManager) recallSession(ctx context.Context, sessionID string, m if !exists { return ErrSessionNotRunning } - if !rt.streaming.TryLock() { + rt.snapshotBoundary.Lock() + acquired := rt.streaming.TryLock() + rt.snapshotBoundary.Unlock() + if !acquired { return rt.runtime.Steer(ctx, msg) } @@ -1415,6 +1427,7 @@ func (sm *SessionManager) recallSession(ctx context.Context, sessionID string, m } rt.modelSwitch.Unlock() + sm.ensureEventLog(sessionID) _, skipMirroredElicitation := rt.runtime.(elicitationSinkMirror) go func() { defer rt.streaming.Unlock() @@ -1428,9 +1441,7 @@ func (sm *SessionManager) recallSession(ctx context.Context, sessionID string, m if _, isElicitation := event.(*runtime.ElicitationRequestEvent); isElicitation && skipMirroredElicitation { continue } - if pe, ok := sm.eventLogs.Load(sessionID); ok { - pe.log.append(event) - } + sm.appendSessionEvent(sessionID, event) } if err := sm.persistActiveSession(context.WithoutCancel(ctx), sessionID, rt, sess); err != nil && !errors.Is(err, ErrSessionNotRunning) { slog.WarnContext(ctx, "Failed to persist recalled session", "session_id", sessionID, "error", err) diff --git a/pkg/server/session_manager_recall_test.go b/pkg/server/session_manager_recall_test.go new file mode 100644 index 0000000000..594e25f6f0 --- /dev/null +++ b/pkg/server/session_manager_recall_test.go @@ -0,0 +1,74 @@ +package server + +import ( + "context" + "net/http/httptest" + "testing" + "testing/synctest" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/docker/docker-agent/pkg/runtime" + "github.com/docker/docker-agent/pkg/session" +) + +func TestRecallSession_CreatesReplayableEventLogBeforeFirstEvent(t *testing.T) { + t.Parallel() + + synctest.Test(t, func(t *testing.T) { + sess := session.New() + events := []runtime.Event{runtime.StreamStarted(sess.ID, "root"), runtime.Warning("recall answer", "root"), runtime.StreamStopped(sess.ID, "root", "")} + fake := &scriptedStreamRuntime{events: events} + sm := newTestSessionManager(t, sess, fake) + require.False(t, sm.HasEventSource(sess.ID)) + require.NoError(t, sm.recallSession(t.Context(), sess.ID, runtime.QueuedMessage{Content: "wake up"})) + require.True(t, sm.HasEventSource(sess.ID), "GET /events must be available as soon as recall starts") + waitSessionIdle(t, sm, sess.ID) + seq, ok := sm.LastEventSeq(sess.ID) + require.True(t, ok) + assert.Equal(t, uint64(len(events)), seq) + got := replaySessionEvents(t, sm, sess.ID, len(events)) + for i := range events { + assert.Same(t, events[i], got[i]) + } + require.NoError(t, sm.DeleteSession(t.Context(), sess.ID)) + require.False(t, sm.HasEventSource(sess.ID)) + sm.appendSessionEvent(sess.ID, runtime.Warning("late event", "root")) + assert.False(t, sm.HasEventSource(sess.ID), "late recalls must not resurrect deleted logs") + }) +} + +func TestRecallSession_RemoteRuntimeReceivesIdleRecallAfterTurnCancellation(t *testing.T) { + t.Parallel() + + sess := session.New() + fake := &scriptedStreamRuntime{events: []runtime.Event{runtime.Warning("answer", "root")}} + sm := newTestSessionManager(t, sess, fake) + srv := httptest.NewServer(NewWithManager(sm, "").e) + t.Cleanup(srv.Close) + client, err := runtime.NewClient(srv.URL) + require.NoError(t, err) + remote, err := runtime.NewRemoteRuntime(client) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, remote.Close()) }) + background := make(chan runtime.Event, 4) + remote.OnBackgroundEvent(func(event runtime.Event) { background <- event }) + + turnCtx, cancel := context.WithCancel(t.Context()) + for range remote.RunStream(turnCtx, &session.Session{ID: sess.ID}) { + } + cancel() + require.False(t, sm.HasEventSource(sess.ID), "ordinary RunSession must not duplicate its foreground events in the log") + require.NoError(t, sm.recallSession(t.Context(), sess.ID, runtime.QueuedMessage{Content: "wake up"})) + select { + case event := <-background: + warning, ok := event.(*runtime.WarningEvent) + require.True(t, ok, "got %T", event) + assert.Equal(t, "answer", warning.Message) + case <-time.After(3 * time.Second): + t.Fatal("remote runtime missed server-owned idle recall") + } + assert.Empty(t, background, "foreground answer must not be delivered again") +} diff --git a/pkg/server/session_manager_snapshot_test.go b/pkg/server/session_manager_snapshot_test.go new file mode 100644 index 0000000000..eca6463061 --- /dev/null +++ b/pkg/server/session_manager_snapshot_test.go @@ -0,0 +1,123 @@ +package server + +import ( + "context" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/docker/docker-agent/pkg/chat" + "github.com/docker/docker-agent/pkg/runtime" + "github.com/docker/docker-agent/pkg/session" +) + +type snapshotBarrierRuntime struct { + fakeRuntime + + once sync.Once + reading chan struct{} + release chan struct{} +} + +func (r *snapshotBarrierRuntime) CurrentAgentName(context.Context) string { + r.once.Do(func() { + close(r.reading) + <-r.release + }) + return "root" +} + +func TestGetSessionSnapshotHoldsIdleBoundaryThroughMessagesAndCursor(t *testing.T) { + t.Parallel() + sess := session.New() + sess.AddMessage(&session.Message{Message: chat.Message{Role: chat.MessageRoleAssistant, MessageID: "before", Content: "before"}}) + rt := &snapshotBarrierRuntime{reading: make(chan struct{}), release: make(chan struct{})} + sm := newTestSessionManager(t, sess, rt) + sm.appendSessionEvent(sess.ID, runtime.StreamStopped(sess.ID, "root", "")) + done := make(chan struct{}) + go func() { + defer close(done) + snapshot, err := sm.GetSessionSnapshot(t.Context(), sess.ID) + assert.NoError(t, err) + if assert.NotNil(t, snapshot) { + assert.False(t, snapshot.Streaming) + assert.Equal(t, uint64(1), snapshot.LastEventSeq) + if assert.Len(t, snapshot.Messages, 1) { + assert.Equal(t, "before", snapshot.Messages[0].Message.Content) + } + } + }() + <-rt.reading + rs, ok := sm.runtimeSessions.Load(sess.ID) + require.True(t, ok) + unlocked := rs.streaming.TryLock() + if unlocked { + rs.streaming.Unlock() + } + assert.False(t, unlocked, "new recall/run must not change messages or cursor during an idle snapshot") + close(rt.release) + <-done + require.True(t, rs.streaming.TryLock(), "snapshot must release ownership") + rs.streaming.Unlock() +} + +func TestGetSessionSnapshotRunningCannotSupplyIdleCursor(t *testing.T) { + t.Parallel() + sess := session.New() + sm := newTestSessionManager(t, sess, &fakeRuntime{}) + rs, ok := sm.runtimeSessions.Load(sess.ID) + require.True(t, ok) + rs.streaming.Lock() + defer rs.streaming.Unlock() + snapshot, err := sm.GetSessionSnapshot(t.Context(), sess.ID) + require.NoError(t, err) + assert.True(t, snapshot.Streaming) +} + +type snapshotRecallRuntime struct { + snapshotBarrierRuntime + + runs chan struct{} + steers atomic.Int32 +} + +func (r *snapshotRecallRuntime) RunStream(context.Context, *session.Session) <-chan runtime.Event { + r.runs <- struct{}{} + events := make(chan runtime.Event) + close(events) + return events +} + +func (r *snapshotRecallRuntime) Steer(context.Context, runtime.QueuedMessage) error { + r.steers.Add(1) + return nil +} + +func TestRecallWaitsForIdleSnapshotInsteadOfSteeringAbsentRun(t *testing.T) { + t.Parallel() + sess := session.New() + rt := &snapshotRecallRuntime{snapshotBarrierRuntime: snapshotBarrierRuntime{reading: make(chan struct{}), release: make(chan struct{})}, runs: make(chan struct{}, 1)} + sm := newTestSessionManager(t, sess, rt) + snapshotDone := make(chan struct{}) + go func() { + defer close(snapshotDone) + _, err := sm.GetSessionSnapshot(t.Context(), sess.ID) + assert.NoError(t, err) + }() + <-rt.reading + recalled := make(chan error, 1) + go func() { recalled <- sm.recallSession(t.Context(), sess.ID, runtime.QueuedMessage{Content: "wake up"}) }() + close(rt.release) + <-snapshotDone + require.NoError(t, <-recalled) + select { + case <-rt.runs: + case <-time.After(5 * time.Second): + t.Fatal("idle recall never started") + } + require.Zero(t, rt.steers.Load()) +} diff --git a/pkg/tui/components/messages/assistant_content.go b/pkg/tui/components/messages/assistant_content.go new file mode 100644 index 0000000000..bd6fc8c735 --- /dev/null +++ b/pkg/tui/components/messages/assistant_content.go @@ -0,0 +1,141 @@ +package messages + +import ( + "slices" + "strings" + + tea "charm.land/bubbletea/v2" + + "github.com/docker/docker-agent/pkg/tui/components/message" + "github.com/docker/docker-agent/pkg/tui/types" +) + +func (m *model) trackMessageIdentity(sessionID, messageID string) { + for _, last := range slices.Backward(m.messages) { + if last.Type == types.MessageTypeSpinner { + continue + } + if last.SessionID != sessionID || last.MessageID != messageID { + m.BreakMessageGroup() + } + return + } +} + +// AppendAssistantContent keeps separate attempts and interleaved sessions distinct. +func (m *model) AppendAssistantContent(sessionID, messageID, agentName, content string) tea.Cmd { + m.trackMessageIdentity(sessionID, messageID) + cmd := m.AppendToLastMessage(agentName, content) + if last := m.lastMessage(); last != nil { + last.SessionID, last.MessageID = sessionID, messageID + } + return cmd +} + +// AppendReasoningContent uses the same logical boundary as assistant content. +func (m *model) AppendReasoningContent(sessionID, messageID, agentName, content string) tea.Cmd { + m.trackMessageIdentity(sessionID, messageID) + cmd := m.AppendReasoning(agentName, content) + if last := m.lastMessage(); last != nil { + last.SessionID, last.MessageID = sessionID, messageID + } + return cmd +} + +// ReconcileAssistantContent replaces incomplete streamed text with the saved answer. +func (m *model) ReconcileAssistantContent(sessionID, messageID, agentName, content string) tea.Cmd { + if content == "" { + return nil + } + materialize := m.materializeDeferredTail() + var indices []int + var streamed strings.Builder + for i, msg := range m.messages { + if msg.Type == types.MessageTypeAssistant && msg.SessionID == sessionID && msg.MessageID == messageID && msg.Sender == agentName { + indices = append(indices, i) + streamed.WriteString(msg.Content) + } + } + content = strings.ReplaceAll(content, "\t", " ") + if len(indices) == 0 { + return tea.Batch(materialize, m.AppendAssistantContent(sessionID, messageID, agentName, content)) + } + if streamed.String() == content { + return materialize + } + // Keep segment boundaries when the canonical text extends the delivered prefix. + var cmds []tea.Cmd + cmds = append(cmds, materialize) + if strings.HasPrefix(content, streamed.String()) { + i := indices[len(indices)-1] + msg := *m.messages[i] + msg.Content += strings.TrimPrefix(content, streamed.String()) + m.messages[i] = &msg + cmds = append(cmds, m.views[i].(message.Model).SetMessage(&msg)) + m.invalidateItem(i) + } else { + for n, i := range indices { + msg := *m.messages[i] + msg.Content = "" + if n == 0 { + msg.Content = content + } + m.messages[i] = &msg + cmds = append(cmds, m.views[i].(message.Model).SetMessage(&msg)) + m.invalidateItem(i) + } + } + return tea.Batch(cmds...) +} + +// AppendAssistantMediaContent keeps media joined to its logical text message. +func (m *model) AppendAssistantMediaContent(sessionID, messageID, agentName string, media []types.AssistantMedia) tea.Cmd { + for i, msg := range slices.Backward(m.messages) { + legacyMatch := msg.MessageID == "" && strings.HasPrefix(messageID, "legacy:") && slices.ContainsFunc(msg.AssistantMedia, func(existing types.AssistantMedia) bool { + return existing.Key != "" && slices.ContainsFunc(media, func(item types.AssistantMedia) bool { return item.Key == existing.Key }) + }) + if msg.Type != types.MessageTypeAssistant || msg.SessionID != sessionID || (msg.MessageID != messageID && !legacyMatch) || msg.Sender != agentName { + continue + } + var added []types.AssistantMedia + for _, item := range media { + if item.Key != "" && slices.ContainsFunc(msg.AssistantMedia, func(existing types.AssistantMedia) bool { return existing.Key == item.Key }) { + continue + } + added = append(added, item) + } + if len(added) == 0 { + return nil + } + updated := *msg + updated.AssistantMedia = slices.Concat(msg.AssistantMedia, added) + m.messages[i] = &updated + cmd := m.views[i].(message.Model).SetMessage(&updated) + m.invalidateItem(i) + return cmd + } + m.trackMessageIdentity(sessionID, messageID) + cmd := m.AppendAssistantMedia(agentName, media) + if last := m.lastMessage(); last != nil { + last.SessionID, last.MessageID = sessionID, messageID + } + return cmd +} + +// AdoptAssistantMediaIdentity joins legacy canonical replay to restored manifest media. +func (m *model) AdoptAssistantMediaIdentity(sessionID, messageID, agentName string, media []types.AssistantMedia) { + if !strings.HasPrefix(messageID, "legacy:") { + return + } + for _, msg := range m.messages { + if msg.Type != types.MessageTypeAssistant || msg.SessionID != sessionID || msg.MessageID != "" || msg.Sender != agentName { + continue + } + if slices.ContainsFunc(msg.AssistantMedia, func(existing types.AssistantMedia) bool { + return existing.Key != "" && slices.ContainsFunc(media, func(item types.AssistantMedia) bool { return item.Key == existing.Key }) + }) { + msg.MessageID = messageID + return + } + } +} diff --git a/pkg/tui/components/messages/messages.go b/pkg/tui/components/messages/messages.go index e62e4245cc..05b57c5a78 100644 --- a/pkg/tui/components/messages/messages.go +++ b/pkg/tui/components/messages/messages.go @@ -89,6 +89,11 @@ type Model interface { AppendToolOutput(msg *runtime.ToolCallOutputEvent) tea.Cmd AddToolResult(msg *runtime.ToolCallResponseEvent, status types.ToolStatus) tea.Cmd AppendToLastMessage(agentName, content string) tea.Cmd + AppendAssistantContent(sessionID, messageID, agentName, content string) tea.Cmd + AppendReasoningContent(sessionID, messageID, agentName, content string) tea.Cmd + ReconcileAssistantContent(sessionID, messageID, agentName, content string) tea.Cmd + AppendAssistantMediaContent(sessionID, messageID, agentName string, media []types.AssistantMedia) tea.Cmd + AdoptAssistantMediaIdentity(sessionID, messageID, agentName string, media []types.AssistantMedia) // BreakMessageGroup prevents merging across streams without flushing deferred content. BreakMessageGroup() // AppendAssistantMedia attaches generated media to the agent's current @@ -332,6 +337,12 @@ func (m *model) Update(msg tea.Msg) (layout.Model, tea.Cmd) { } switch msg := msg.(type) { + case *runtime.SessionRecoveredEvent: + finalizeCmd := m.FinalizeStream() + m.removeSpinner() + m.removePendingToolCallMessages() + m.stopReasoningBlockAnimations() + return m, finalizeCmd case messages.StreamCancelledMsg: finalizeCmd := m.FinalizeStream() m.removeSpinner() @@ -1915,7 +1926,8 @@ func (m *model) LoadFromSession(sess *session.Session, generatedMedia map[int][] appendSessionMessage(msg, m.createMessageView(msg)) case chat.MessageRoleAssistant: hasReasoning := smsg.Message.ReasoningContent != "" - hasContent := smsg.Message.Content != "" + visibleContent := chat.VisibleAssistantContent(smsg.Message.Content) + hasContent := visibleContent != "" hasToolCalls := len(smsg.Message.ToolCalls) > 0 var reasoningBlock *reasoningblock.Model @@ -1937,7 +1949,8 @@ func (m *model) LoadFromSession(sess *session.Session, generatedMedia map[int][] // live behavior. restoredMedia := generatedMedia[pos] if hasContent || len(restoredMedia) > 0 { - msg := types.Agent(types.MessageTypeAssistant, smsg.AgentName, smsg.Message.Content) + msg := types.Agent(types.MessageTypeAssistant, smsg.AgentName, visibleContent) + msg.SessionID, msg.MessageID = sess.ID, smsg.Message.MessageID msg.AssistantMedia = restoredMedia appendSessionMessage(msg, m.createMessageView(msg)) } @@ -2204,6 +2217,7 @@ func (m *model) UpdateAssistantMedia(media []types.AssistantMedia) tea.Cmd { changed := false for j, item := range msg.AssistantMedia { if resolved, ok := byID[item.ID]; ok { + resolved.Key = item.Key msg.AssistantMedia[j] = resolved changed = true } diff --git a/pkg/tui/components/sidebar/sidebar.go b/pkg/tui/components/sidebar/sidebar.go index 5b57dbf0bf..ec2e134129 100644 --- a/pkg/tui/components/sidebar/sidebar.go +++ b/pkg/tui/components/sidebar/sidebar.go @@ -1362,6 +1362,14 @@ func (m *model) Update(msg tea.Msg) (layout.Model, tea.Cmd) { m.invalidateCache() cmd := m.startSpinner() return m, cmd + case *runtime.SessionRecoveredEvent: + m.workingAgent = "" + m.sessionStack = nil + m.clearTransferPresentation() + m.compacting = false + m.invalidateCache() + m.stopSpinner() + return m, nil case *runtime.StreamStoppedEvent: m.workingAgent = "" if n := len(m.sessionStack); n > 0 { diff --git a/pkg/tui/messages/tabs.go b/pkg/tui/messages/tabs.go index fe96beabce..068f43c8e7 100644 --- a/pkg/tui/messages/tabs.go +++ b/pkg/tui/messages/tabs.go @@ -12,6 +12,7 @@ type RouteScope struct{ _ byte } type RoutedMsg struct { SessionID string // The session ID this message is for Inner tea.Msg // The wrapped message + Valid func() bool // Checks producer lifetime at consumption. Scope *RouteScope // Set by runtime subscriptions; nil for page-owned messages. } diff --git a/pkg/tui/page/chat/assistant_content_test.go b/pkg/tui/page/chat/assistant_content_test.go new file mode 100644 index 0000000000..c0f06ca848 --- /dev/null +++ b/pkg/tui/page/chat/assistant_content_test.go @@ -0,0 +1,123 @@ +package chat + +import ( + "strings" + "testing" + + "github.com/charmbracelet/x/ansi" + "github.com/stretchr/testify/require" + + "github.com/docker/docker-agent/pkg/app" + chatapi "github.com/docker/docker-agent/pkg/chat" + "github.com/docker/docker-agent/pkg/runtime" + "github.com/docker/docker-agent/pkg/session" + "github.com/docker/docker-agent/pkg/tui/animation" + "github.com/docker/docker-agent/pkg/tui/service" + "github.com/docker/docker-agent/pkg/tui/types" +) + +func newContentTestPage(t *testing.T) (*chatPage, *session.Session) { + t.Helper() + ar := animation.NewRuntime() + t.Cleanup(ar.Stop) + sess := session.New() + p := New(ar, t.Context(), app.New(t.Context(), queueTestRuntime{}, sess), service.NewSessionState(sess)).(*chatPage) + t.Cleanup(func() { Cleanup(p) }) + p.messages.SetSize(100, 40) + return p, sess +} + +func TestCanonicalAssistantContentAppearsWithoutReload(t *testing.T) { + t.Parallel() + p, sess := newContentTestPage(t) + msg := session.NewAgentMessage("root", &chatapi.Message{Role: chatapi.MessageRoleAssistant, MessageID: "answer", Content: "CANONICAL-FINAL-ANSWER"}) + for _, event := range []runtime.Event{runtime.StreamStarted(sess.ID, "root"), runtime.MessageAdded(sess.ID, msg, "root"), runtime.StreamStopped(sess.ID, "root", "normal")} { + _, _ = p.handleRuntimeEvent(event) + } + require.Contains(t, ansi.Strip(p.messages.View()), "CANONICAL-FINAL-ANSWER") + require.True(t, p.hasReceivedAssistantContent) +} + +func TestCanonicalAssistantRepairsMissingContentWithoutDuplicates(t *testing.T) { + t.Parallel() + for _, partial := range []string{"FINAL-", "ANSWER", "FINAL-ANSWER"} { + t.Run(partial, func(t *testing.T) { + p, sess := newContentTestPage(t) + _, _ = p.handleRuntimeEvent(runtime.StreamStarted(sess.ID, "root")) + _, _ = p.handleRuntimeEvent(runtime.AgentChoice("root", sess.ID, partial, "answer")) + msg := session.NewAgentMessage("root", &chatapi.Message{Role: chatapi.MessageRoleAssistant, MessageID: "answer", Content: "FINAL-ANSWER"}) + for range 2 { + _, _ = p.handleRuntimeEvent(runtime.MessageAdded(sess.ID, msg, "root")) + } + _, _ = p.handleRuntimeEvent(runtime.StreamStopped(sess.ID, "root", "normal")) + frame := ansi.Strip(p.messages.View()) + require.Equal(t, 1, strings.Count(frame, "FINAL-ANSWER")) + require.Equal(t, 1, p.messages.MessageTypeCount(types.MessageTypeAssistant)) + }) + } +} + +func TestDifferentMessageIDsKeepRetryAnswerVisible(t *testing.T) { + t.Parallel() + p, sess := newContentTestPage(t) + for _, event := range []runtime.Event{ + runtime.StreamStarted(sess.ID, "root"), + runtime.AgentChoice("root", sess.ID, "[diagnostic](", "attempt-one"), + runtime.AgentChoice("root", sess.ID, "FINAL-ANSWER)", "attempt-two"), + runtime.StreamStopped(sess.ID, "root", "normal"), + } { + _, _ = p.handleRuntimeEvent(event) + } + require.Contains(t, ansi.Strip(p.messages.View()), "FINAL-ANSWER") + require.Equal(t, 2, p.messages.MessageTypeCount(types.MessageTypeAssistant)) +} + +func TestCanonicalAssistantDoesNotMergeWithPreviousTurn(t *testing.T) { + t.Parallel() + p, sess := newContentTestPage(t) + for _, event := range []runtime.Event{ + runtime.StreamStarted(sess.ID, "root"), runtime.AgentChoice("root", sess.ID, "first"), runtime.StreamStopped(sess.ID, "root", "normal"), + runtime.StreamStarted(sess.ID, "root"), runtime.MessageAdded(sess.ID, session.NewAgentMessage("root", &chatapi.Message{Role: chatapi.MessageRoleAssistant, Content: "second"}), "root"), runtime.StreamStopped(sess.ID, "root", "normal"), + } { + _, _ = p.handleRuntimeEvent(event) + } + require.Equal(t, 2, p.messages.MessageTypeCount(types.MessageTypeAssistant)) + require.Contains(t, ansi.Strip(p.messages.View()), "second") +} + +func TestPendingSpinnerDoesNotMergeDistinctMessageIDs(t *testing.T) { + t.Parallel() + p, sess := newContentTestPage(t) + _, _ = p.handleRuntimeEvent(runtime.AgentChoice("root", sess.ID, "[diagnostic](", "first")) + p.setPendingResponse(true) + _, _ = p.handleRuntimeEvent(runtime.AgentChoice("root", sess.ID, "FINAL-ANSWER)", "second")) + require.Contains(t, ansi.Strip(p.messages.View()), "FINAL-ANSWER") + require.Equal(t, 2, p.messages.MessageTypeCount(types.MessageTypeAssistant)) +} + +func TestCanonicalContentCannotRevealSuppressedToolXML(t *testing.T) { + t.Parallel() + p, sess := newContentTestPage(t) + msg := session.NewAgentMessage("root", &chatapi.Message{Role: chatapi.MessageRoleAssistant, MessageID: "answer", Content: `safe prefix{"arguments":"PRIVATE-TOOL-ARGUMENT`}) + _, _ = p.handleRuntimeEvent(runtime.AgentChoice("root", sess.ID, "safe prefix", "answer")) + _, _ = p.handleRuntimeEvent(runtime.MessageAdded(sess.ID, msg, "root")) + require.Contains(t, ansi.Strip(p.messages.View()), "safe prefix") + require.NotContains(t, ansi.Strip(p.messages.View()), "PRIVATE-TOOL-ARGUMENT") + sess.AddMessage(msg) + p.messages.LoadFromSession(sess, nil) + require.NotContains(t, ansi.Strip(p.messages.View()), "PRIVATE-TOOL-ARGUMENT") +} + +func TestRecoveredSessionClearsNestedBusyWithoutStartingQueuedRun(t *testing.T) { + t.Parallel() + p, sess := newContentTestPage(t) + for _, id := range []string{sess.ID, "child", "nested"} { + _, _ = p.handleRuntimeEvent(runtime.StreamStarted(id, "root")) + } + p.messageQueue = []queuedMessage{{content: "must not run"}} + _, _ = p.handleRuntimeEvent(runtime.SessionRecovered(sess.ID)) + require.Zero(t, p.streamDepth) + require.Empty(t, p.agentStack) + require.False(t, p.working) + require.Len(t, p.messageQueue, 1) +} diff --git a/pkg/tui/page/chat/chat.go b/pkg/tui/page/chat/chat.go index c07994f8b2..320042b507 100644 --- a/pkg/tui/page/chat/chat.go +++ b/pkg/tui/page/chat/chat.go @@ -30,6 +30,7 @@ import ( "github.com/docker/docker-agent/pkg/tui/dialog" msgtypes "github.com/docker/docker-agent/pkg/tui/messages" "github.com/docker/docker-agent/pkg/tui/service" + "github.com/docker/docker-agent/pkg/tui/streamcontent" "github.com/docker/docker-agent/pkg/tui/styles" ) @@ -232,6 +233,8 @@ type chatPage struct { agentStack []string // agent per active stream level; len(agentStack)==streamDepth streamStartTime time.Time contentSessionID string + contentIdentity streamcontent.Tracker + mediaKeys map[string]uint64 // routingID is the tab identity this page's routed UI timers are // addressed to; empty for standalone pages (timers then fire unrouted, diff --git a/pkg/tui/page/chat/generated_media.go b/pkg/tui/page/chat/generated_media.go index f15fcc0045..f8ac010001 100644 --- a/pkg/tui/page/chat/generated_media.go +++ b/pkg/tui/page/chat/generated_media.go @@ -63,24 +63,44 @@ func (p *chatPage) handleMessageAdded(msg *runtime.MessageAddedEvent) tea.Cmd { if p.streamCancelled || msg.Message.Message.Role != chat.MessageRoleAssistant { return nil } - if !p.app.CanResolveGeneratedFiles() { + if msg.Message.Implicit { return nil } - placeholders, requests := generatedImageMedia(msg.Message.Message.MultiContent) - if len(placeholders) == 0 { - return nil + sessionID := p.contentSession(msg.SessionID) + identity := p.contentIdentity.Resolve(sessionID, msg.Message.Message.MessageID) + if msg.Message.Message.Content != "" { + defer p.contentIdentity.Finish(sessionID) } - - p.trackContentSession(msg.SessionID) - p.hasReceivedAssistantContent = true - p.setPendingResponse(false) agentName := msg.Message.AgentName if agentName == "" { agentName = msg.AgentName } + var placeholders []types.AssistantMedia + var requests []generatedMediaRequest + if p.app.CanResolveGeneratedFiles() { + placeholders, requests = p.generatedImageMedia(identity.SessionID, identity.MessageID, msg.Message.Message.MultiContent) + p.messages.AdoptAssistantMediaIdentity(identity.SessionID, identity.MessageID, agentName, placeholders) + } + content := chat.VisibleAssistantContent(msg.Message.Message.Content) + var contentCmd tea.Cmd + if content != "" { + contentCmd = p.messages.ReconcileAssistantContent(identity.SessionID, identity.MessageID, agentName, content) + p.hasReceivedAssistantContent = true + p.setPendingResponse(false) + } + if !p.app.CanResolveGeneratedFiles() { + return contentCmd + } + if len(placeholders) == 0 { + return contentCmd + } + p.trackContentSession(msg.SessionID) + p.hasReceivedAssistantContent = true + p.setPendingResponse(false) return tea.Batch( + contentCmd, p.sidebar.SetAgentActivity(agentName), - p.messages.AppendAssistantMedia(agentName, placeholders), + p.messages.AppendAssistantMediaContent(identity.SessionID, identity.MessageID, agentName, placeholders), p.resolveGeneratedMediaCmd(requests), ) } @@ -99,7 +119,7 @@ func (p *chatPage) collectRestoredGeneratedMedia(sess *session.Session) (map[int if !item.IsMessage() || item.Message.Implicit || item.Message.Message.Role != chat.MessageRoleAssistant { continue } - placeholders, reqs := generatedImageMedia(item.Message.Message.MultiContent) + placeholders, reqs := p.generatedImageMedia(sess.ID, item.Message.Message.MessageID, item.Message.Message.MultiContent) if len(placeholders) == 0 { continue } @@ -119,7 +139,7 @@ func (p *chatPage) collectRestoredGeneratedMedia(sess *session.Session) (map[int // supports additionally get a resolution request. References with an // unknown (empty) root kind stay unavailable by design. User attachments // (inline sources) and ownerless references are not extracted. -func generatedImageMedia(parts []chat.MessagePart) ([]types.AssistantMedia, []generatedMediaRequest) { +func (p *chatPage) generatedImageMedia(sessionID, messageID string, parts []chat.MessagePart) ([]types.AssistantMedia, []generatedMediaRequest) { var media []types.AssistantMedia var requests []generatedMediaRequest for _, part := range parts { @@ -139,9 +159,23 @@ func generatedImageMedia(parts []chat.MessagePart) ([]types.AssistantMedia, []ge if name == "" { name = "generated media" } - item := types.AssistantMedia{Fallback: fmt.Sprintf("Generated image %q is unavailable.", name)} - if src.ArtifactRoot == chat.ArtifactRootWorkspace { - item.ID = generatedMediaIDs.Add(1) + keyMessageID := messageID + if strings.HasPrefix(messageID, "legacy:") { + keyMessageID = "" + } + key := fmt.Sprintf("%q:%q:%q:%q:%q:%q:%q", sessionID, keyMessageID, src.ArtifactOwnerSessionID, src.ArtifactRoot, src.ArtifactPath, doc.MimeType, name) + if p.mediaKeys == nil { + p.mediaKeys = make(map[string]uint64) + } + id, seen := p.mediaKeys[key] + if !seen { + if src.ArtifactRoot == chat.ArtifactRootWorkspace { + id = generatedMediaIDs.Add(1) + } + p.mediaKeys[key] = id + } + item := types.AssistantMedia{ID: id, Key: key, Fallback: fmt.Sprintf("Generated image %q is unavailable.", name)} + if src.ArtifactRoot == chat.ArtifactRootWorkspace && !seen { requests = append(requests, generatedMediaRequest{ id: item.ID, ref: runtime.GeneratedFileRef{ diff --git a/pkg/tui/page/chat/generated_media_test.go b/pkg/tui/page/chat/generated_media_test.go index 6e215a545d..9445f032e8 100644 --- a/pkg/tui/page/chat/generated_media_test.go +++ b/pkg/tui/page/chat/generated_media_test.go @@ -80,6 +80,12 @@ func (r *mediaRecordingMessages) AppendAssistantMedia(agentName string, media [] return r.Model.AppendAssistantMedia(agentName, media) } +func (r *mediaRecordingMessages) AppendAssistantMediaContent(sessionID, messageID, agentName string, media []types.AssistantMedia) tea.Cmd { + r.mediaAgents = append(r.mediaAgents, agentName) + r.mediaCalls = append(r.mediaCalls, media) + return r.Model.AppendAssistantMediaContent(sessionID, messageID, agentName, media) +} + func (r *mediaRecordingMessages) UpdateAssistantMedia(media []types.AssistantMedia) tea.Cmd { r.mediaUpdates = append(r.mediaUpdates, media) return r.Model.UpdateAssistantMedia(media) @@ -673,3 +679,56 @@ func TestLocalMediaResultCannotReplaceReloadedPlaceholder(t *testing.T) { require.Len(t, rec.mediaCalls, 1) assert.NotEqual(t, result.Inner.(generatedMediaResolvedMsg).media[0].ID, rec.mediaCalls[0][0].ID) } + +func TestCanonicalMediaReplayReusesRestoredPlaceholder(t *testing.T) { + t.Parallel() + sess := restoredMediaSession("owner") + sess.Messages[2].Message.Message.MessageID = "image-answer" + p, rec := newGeneratedMediaTestPageWithSession(t, &resolverTestRuntime{}, sess) + p.Init() + before := rec.MessageTypeCount(types.MessageTypeAssistant) + _, effects := p.UpdateEffects(runtime.MessageAdded(sess.ID, sess.Messages[2].Message, "root")) + require.Equal(t, before, rec.MessageTypeCount(types.MessageTypeAssistant)) + require.Nil(t, effects.Local, "queued canonical replay must not resolve a second placeholder") +} + +func TestCanonicalMediaJoinsInterleavedLogicalAnswer(t *testing.T) { + t.Parallel() + p, rec := newGeneratedMediaTestPage(t, &resolverTestRuntime{}) + p.messages.SetSize(100, 40) + _, _ = p.handleRuntimeEvent(runtime.AgentChoice("root", "owner", "answer A", "a")) + _, _ = p.handleRuntimeEvent(runtime.AgentChoice("root", "owner", "answer B", "b")) + added := assistantMessageAdded("owner", workspaceImagePart("a.png", "a.png", "owner")) + added.Message.Message.MessageID = "a" + added.Message.Message.Content = "answer A" + _, _ = p.handleRuntimeEvent(added) + _, _ = p.handleRuntimeEvent(added) + require.Equal(t, 2, rec.MessageTypeCount(types.MessageTypeAssistant)) + require.Contains(t, p.messages.View(), "answer A") + require.Contains(t, p.messages.View(), "answer B") +} + +func TestCanonicalMediaResolveThenReplayKeepsStableKey(t *testing.T) { + t.Parallel() + rt := &resolverTestRuntime{results: map[string]resolverResult{"a.png": {data: testPNGBytes(t)}}} + p, _ := newGeneratedMediaTestPage(t, rt) + added := assistantMessageAdded("owner", workspaceImagePart("a.png", "a.png", "owner")) + added.Message.Message.MessageID = "answer" + _, effects := p.UpdateEffects(added) + _, _ = p.UpdateEffects(resolveMedia(t, effects.Local)) + _, effects = p.UpdateEffects(added) + require.Nil(t, effects.Local) + // A replay must not append an unavailable fallback beside the resolved image. + require.NotContains(t, p.messages.View(), "unavailable") +} + +func TestCanonicalLegacyMediaReplayJoinsRestoredMessage(t *testing.T) { + t.Parallel() + sess := restoredMediaSession("owner") + p, rec := newGeneratedMediaTestPageWithSession(t, &resolverTestRuntime{}, sess) + p.Init() + before := rec.MessageTypeCount(types.MessageTypeAssistant) + _, effects := p.UpdateEffects(runtime.MessageAdded(sess.ID, sess.Messages[2].Message, "root")) + require.Equal(t, before, rec.MessageTypeCount(types.MessageTypeAssistant)) + require.Nil(t, effects.Local) +} diff --git a/pkg/tui/page/chat/runtime_events.go b/pkg/tui/page/chat/runtime_events.go index bf63c9f80c..9933b43d00 100644 --- a/pkg/tui/page/chat/runtime_events.go +++ b/pkg/tui/page/chat/runtime_events.go @@ -13,6 +13,7 @@ import ( "github.com/docker/docker-agent/pkg/sound" "github.com/docker/docker-agent/pkg/tools" builtinshell "github.com/docker/docker-agent/pkg/tools/builtin/shell" + "github.com/docker/docker-agent/pkg/tui/components/messages" "github.com/docker/docker-agent/pkg/tui/components/notification" "github.com/docker/docker-agent/pkg/tui/components/sidebar" "github.com/docker/docker-agent/pkg/tui/core" @@ -82,6 +83,19 @@ func (p *chatPage) handleRuntimeEvent(msg tea.Msg) (bool, tea.Cmd) { case *runtime.StreamStartedEvent: return true, p.handleStreamStarted(msg) + case *runtime.SessionRecoveredEvent: + if p.isSubSessionEvent(msg.SessionID) { + return true, nil + } + p.streamDepth = 0 + p.agentStack = nil + p.msgCancel = nil + p.streamCancelled = false + p.contentIdentity.Finish(p.contentSession(msg.SessionID)) + p.sidebar.ResetStreamTracking() + model, cleanup := p.messages.Update(msg) + p.messages = model.(messages.Model) + return true, tea.Batch(cleanup, p.forwardToSidebar(msg), p.setWorking(false), p.setPendingResponse(false)) case *runtime.StreamStoppedEvent: return true, p.handleStreamStopped(msg) @@ -303,6 +317,7 @@ func (p *chatPage) handleTokenUsage(msg *runtime.TokenUsageEvent) { func (p *chatPage) handleStreamStarted(msg *runtime.StreamStartedEvent) tea.Cmd { slog.Debug("handleStreamStarted called", "agent", msg.AgentName, "session_id", msg.SessionID) + p.contentIdentity.Finish(p.contentSession(msg.SessionID)) if p.contentSessionID == p.contentSession(msg.SessionID) { p.messages.BreakMessageGroup() } @@ -339,6 +354,7 @@ func (p *chatPage) handleAgentChoice(msg *runtime.AgentChoiceEvent) tea.Cmd { return nil } p.trackContentSession(msg.SessionID) + identity := p.contentIdentity.Resolve(p.contentSession(msg.SessionID), msg.MessageID) // Track that we've received assistant content p.hasReceivedAssistantContent = true // Clear pending response indicator - first chunk has arrived @@ -346,7 +362,7 @@ func (p *chatPage) handleAgentChoice(msg *runtime.AgentChoiceEvent) tea.Cmd { // Content is useful activity: acknowledge the sidebar's outbound transfer // box when this agent is a delegation target. activityCmd := p.sidebar.SetAgentActivity(msg.AgentName) - return tea.Batch(activityCmd, p.messages.AppendToLastMessage(msg.AgentName, msg.Content)) + return tea.Batch(activityCmd, p.messages.AppendAssistantContent(identity.SessionID, identity.MessageID, msg.AgentName, msg.Content)) } func (p *chatPage) handleAgentChoiceReasoning(msg *runtime.AgentChoiceReasoningEvent) tea.Cmd { @@ -354,9 +370,10 @@ func (p *chatPage) handleAgentChoiceReasoning(msg *runtime.AgentChoiceReasoningE return nil } p.trackContentSession(msg.SessionID) + identity := p.contentIdentity.Resolve(p.contentSession(msg.SessionID), msg.MessageID) p.setPendingResponse(false) activityCmd := p.sidebar.SetAgentActivity(msg.AgentName) - return tea.Batch(activityCmd, p.messages.AppendReasoning(msg.AgentName, msg.Content)) + return tea.Batch(activityCmd, p.messages.AppendReasoningContent(identity.SessionID, identity.MessageID, msg.AgentName, msg.Content)) } // handleAgentSwitching forwards transfer_task hop boundaries to the sidebar @@ -391,6 +408,7 @@ func (p *chatPage) handleStreamStopped(msg *runtime.StreamStoppedEvent) tea.Cmd "has_content", p.hasReceivedAssistantContent, "stream_depth", p.streamDepth) + p.contentIdentity.Finish(p.contentSession(msg.SessionID)) if p.contentSessionID == p.contentSession(msg.SessionID) { p.messages.BreakMessageGroup() } diff --git a/pkg/tui/service/supervisor/supervisor.go b/pkg/tui/service/supervisor/supervisor.go index e970bf6426..56b13e3b53 100644 --- a/pkg/tui/service/supervisor/supervisor.go +++ b/pkg/tui/service/supervisor/supervisor.go @@ -42,8 +42,7 @@ type Supervisor struct { program *tea.Program // programReady is closed when SetProgram is called. Subscription goroutines - // wait on this before consuming events so that startup events (welcome message, - // agent info, tool info) are not silently dropped. + // wait on this before terminal delivery; registration happens before startup. programReady chan struct{} programReadyOnce sync.Once } @@ -84,10 +83,6 @@ func (s *Supervisor) AddSession(ctx context.Context, a *app.App, sess *session.S // Create a cancellable context for this session sessionCtx, cancel := context.WithCancel(ctx) runner.cancel = cancel - if a != nil { - a.Start(sessionCtx) - } - s.runners[sess.ID] = runner s.order = append(s.order, sess.ID) @@ -97,7 +92,10 @@ func (s *Supervisor) AddSession(ctx context.Context, a *app.App, sess *session.S // Start the subscription goroutine with routing if a != nil { - go s.subscribeWithRouting(sessionCtx, a, sess.ID) + ready := make(chan struct{}) + go s.subscribeWithRouting(sessionCtx, a, sess.ID, ready) + <-ready + a.Start(sessionCtx) } return sess.ID @@ -119,30 +117,32 @@ func (s *Supervisor) SpawnSession(ctx context.Context, workingDir string) (strin } // subscribeWithRouting subscribes to app events and wraps them with session ID. -// It waits for the program to be set before consuming events so that startup -// events (welcome message, agent/team/tool info) are not dropped. -func (s *Supervisor) subscribeWithRouting(ctx context.Context, a *app.App, sessionID string) { - // Wait for the program to be available before consuming any events. - // Events are buffered in app.events, so nothing is lost during this wait. - select { - case <-s.programReady: - case <-ctx.Done(): - return +// Registration precedes startup; terminal delivery waits for the program. +func (s *Supervisor) subscribeWithRouting(ctx context.Context, a *app.App, sessionID string, ready chan struct{}) { + prepare := func(msg tea.Msg, generation uint64) tea.Msg { + s.mu.RLock() + defer s.mu.RUnlock() + runner := s.runners[sessionID] + if runner == nil || runner.App != a || ctx.Err() != nil { + return nil + } + return messages.RoutedMsg{SessionID: sessionID, Scope: runner.Scope, Inner: msg, Valid: func() bool { return a.IsEventGeneration(generation) }} } - send := func(msg tea.Msg) { - s.mu.RLock() - p, runner := s.program, s.runners[sessionID] - if p == nil || runner == nil || runner.App != a || ctx.Err() != nil { - s.mu.RUnlock() + select { + case <-s.programReady: + case <-ctx.Done(): return } - scope := runner.Scope + s.mu.RLock() + p := s.program s.mu.RUnlock() - p.Send(messages.RoutedMsg{SessionID: sessionID, Scope: scope, Inner: msg}) + if p != nil && ctx.Err() == nil { + p.Send(msg) + } } - a.SubscribeWith(ctx, send) + a.SubscribeReliable(ctx, send, app.WithGenerationEventMapper(prepare), app.WithSubscriptionReady(func() { close(ready) })) } // RetirePage invalidates runtime deliveries already queued for the previous page. @@ -150,6 +150,9 @@ func (s *Supervisor) RetirePage(tabID string) { s.mu.Lock() defer s.mu.Unlock() if runner := s.runners[tabID]; runner != nil { + if runner.App != nil { + runner.App.RetireEvents() + } runner.Scope = &messages.RouteScope{} } } @@ -274,8 +277,6 @@ func (s *Supervisor) ReplaceRunnerApp(ctx context.Context, sessionID string, new // Create a new cancellable context for the replacement. sessionCtx, cancel := context.WithCancel(ctx) runner.cancel = cancel - newApp.Start(sessionCtx) - s.notifyTabsUpdated() s.mu.Unlock() @@ -285,7 +286,10 @@ func (s *Supervisor) ReplaceRunnerApp(ctx context.Context, sessionID string, new } // Start routing events from the new app. - go s.subscribeWithRouting(sessionCtx, newApp, sessionID) + ready := make(chan struct{}) + go s.subscribeWithRouting(sessionCtx, newApp, sessionID, ready) + <-ready + newApp.Start(sessionCtx) } // ActiveID returns the ID of the currently active session. diff --git a/pkg/tui/service/tabstate/state.go b/pkg/tui/service/tabstate/state.go index cc7ad9633f..0698a08913 100644 --- a/pkg/tui/service/tabstate/state.go +++ b/pkg/tui/service/tabstate/state.go @@ -99,6 +99,12 @@ func (s *State) Apply(msg tea.Msg, active bool) (changed, bell bool) { } s.running = false s.retainDetachedElicitations() + case *runtime.SessionRecoveredEvent: + if !isTopLevelStream(s.sessionID, ev.SessionID) { + return false, false + } + s.running = false + s.retainDetachedElicitations() case messages.StreamCancelledMsg: s.running = false s.retainDetachedElicitations() @@ -150,6 +156,8 @@ func (s *State) RetiresAttention(boundary, event tea.Msg) bool { if !isTopLevelStream(s.sessionID, msg.SessionID) { return false } + case *runtime.SessionRecoveredEvent: + return isTopLevelStream(s.sessionID, msg.SessionID) && !isDetachedElicitation(s.sessionID, event) case messages.StreamCancelledMsg: default: return false diff --git a/pkg/tui/service/tabstate/state_test.go b/pkg/tui/service/tabstate/state_test.go index 777d7f07ee..d672d4a695 100644 --- a/pkg/tui/service/tabstate/state_test.go +++ b/pkg/tui/service/tabstate/state_test.go @@ -151,3 +151,13 @@ func TestReplaceSessionRetainsSameConversation(t *testing.T) { assert.False(t, running) assert.False(t, attention) } + +func TestRecoveryPreservesDetachedRequests(t *testing.T) { + t.Parallel() + state := New("root", "") + detached := runtime.ElicitationRequest("request", "form", nil, "", "child-request", "", "child", nil, "worker") + state.Apply(detached, false) + state.Apply(runtime.SessionRecovered("root"), false) + require.Same(t, detached, state.Consume()) + require.False(t, state.RetiresAttention(runtime.SessionRecovered("root"), detached)) +} diff --git a/pkg/tui/stalled_delivery_test.go b/pkg/tui/stalled_delivery_test.go new file mode 100644 index 0000000000..81b158a316 --- /dev/null +++ b/pkg/tui/stalled_delivery_test.go @@ -0,0 +1,167 @@ +package tui + +import ( + "context" + "testing" + "time" + + tea "charm.land/bubbletea/v2" + "github.com/charmbracelet/x/ansi" + "github.com/stretchr/testify/require" + + "github.com/docker/docker-agent/pkg/app" + agentruntime "github.com/docker/docker-agent/pkg/runtime" + "github.com/docker/docker-agent/pkg/tui/messages" + "github.com/docker/docker-agent/pkg/tui/page/chat" +) + +type stalledDeliveryRuntime struct { + stubRuntime + + emit func(agentruntime.Event) +} + +func (r *stalledDeliveryRuntime) OnBackgroundEvent(emit func(agentruntime.Event)) { + r.emit = emit +} + +type stallDeliveryMsg struct{} + +type stalledDeliveryModel struct { + *streamingMotionModel + + stalled chan struct{} + resume chan struct{} + subscribed chan struct{} + started chan struct{} + stopped chan struct{} +} + +func (m *stalledDeliveryModel) Update(msg tea.Msg) (tea.Model, tea.Cmd) { + if _, ok := msg.(stallDeliveryMsg); ok { + close(m.stalled) + <-m.resume + return m, nil + } + if routed, ok := msg.(messages.RoutedMsg); ok { + if _, ok := routed.Inner.(*agentruntime.SessionTitleEvent); ok { + select { + case <-m.subscribed: + default: + close(m.subscribed) + } + } + } + _, cmd := m.streamingMotionModel.Update(msg) + if !m.root.activeTab.chatPage.IsWorking() { + select { + case <-m.started: + select { + case <-m.stopped: + default: + close(m.stopped) + } + default: + } + } else { + select { + case <-m.started: + default: + close(m.started) + } + } + return m, cmd +} + +func TestActualProgramStalledDeliveryKeepsFinalResponse(t *testing.T) { + for _, hidden := range []bool{false, true} { + t.Run(map[bool]string{false: "visible", true: "hidden"}[hidden], func(t *testing.T) { + root, _, _ := frozenClockRoot(t, 120, 40) + rt := &stalledDeliveryRuntime{} + a := app.New(t.Context(), rt, root.application.Session()) + root.supervisor.ReplaceRunnerApp(t.Context(), "profile", a, "", nil) + root.application = a + root.activeTab.chatPage = chat.New(root.ar, t.Context(), a, root.activeTab.sessionState, chat.WithHideSidebar()) + root.handleWindowResize(120, 40) + model := &stalledDeliveryModel{ + streamingMotionModel: &streamingMotionModel{root: root, ready: make(chan struct{})}, + stalled: make(chan struct{}), + resume: make(chan struct{}), + subscribed: make(chan struct{}), + started: make(chan struct{}), + stopped: make(chan struct{}), + } + program := startTestProgram(t, root, model, tea.WithOutput(&wallClockCountingWriter{})) + <-model.ready + root.supervisor.SetProgram(program) + ctx, cancel := context.WithCancel(t.Context()) + defer cancel() + witness := make(chan struct{}, 1) + registered := make(chan struct{}, 1) + go a.SubscribeWith(ctx, func(msg tea.Msg) { + switch msg.(type) { + case *agentruntime.StreamStoppedEvent: + witness <- struct{}{} + case *agentruntime.SessionTitleEvent: + select { + case registered <- struct{}{}: + default: + } + } + }) + require.Eventually(t, func() bool { + rt.emit(agentruntime.SessionTitle("profile", "subscription-ready")) + select { + case <-model.subscribed: + select { + case <-registered: + return true + default: + } + default: + } + return false + }, 5*time.Second, 10*time.Millisecond) + rt.emit(agentruntime.StreamStarted("profile", "root")) + select { + case <-model.started: + case <-time.After(30 * time.Second): + t.Fatal("the TUI did not receive stream start") + } + if hidden { + program.Send(tmuxVisibilityMsg{hidden: true}) + programAck(t, program) + } + program.Send(stallDeliveryMsg{}) + <-model.stalled + t.Cleanup(func() { + select { + case <-model.resume: + default: + close(model.resume) + } + }) + for range 1100 { + rt.emit(agentruntime.NewTokenUsageEvent("profile", "root", &agentruntime.Usage{})) + } + rt.emit(agentruntime.AgentChoice("root", "profile", "FINAL-RESPONSE-AFTER-STALL", "answer")) + rt.emit(agentruntime.StreamStopped("profile", "root", "normal")) + select { + case <-witness: + case <-time.After(30 * time.Second): + t.Fatal("a stalled TUI blocked fan-out to other subscribers") + } + close(model.resume) + select { + case <-model.stopped: + case <-time.After(30 * time.Second): + t.Fatal("the TUI did not receive stream completion") + } + if hidden { + program.Send(tmuxVisibilityMsg{}) + programAck(t, program) + } + require.Contains(t, ansi.Strip(programFrame(t, program)), "FINAL-RESPONSE-AFTER-STALL") + }) + } +} diff --git a/pkg/tui/streamcontent/identity.go b/pkg/tui/streamcontent/identity.go new file mode 100644 index 0000000000..a60e6b1237 --- /dev/null +++ b/pkg/tui/streamcontent/identity.go @@ -0,0 +1,32 @@ +// Package streamcontent tracks logical assistant-message boundaries. +package streamcontent + +import "strconv" + +type Identity struct { + SessionID string + MessageID string +} + +// Tracker gives legacy events without message IDs a per-message identity. +type Tracker struct { + sequence uint64 + legacy map[string]string +} + +func (t *Tracker) Resolve(sessionID, messageID string) Identity { + if messageID == "" { + if t.legacy == nil { + t.legacy = make(map[string]string) + } + messageID = t.legacy[sessionID] + if messageID == "" { + t.sequence++ + messageID = "legacy:" + strconv.FormatUint(t.sequence, 10) + t.legacy[sessionID] = messageID + } + } + return Identity{SessionID: sessionID, MessageID: messageID} +} + +func (t *Tracker) Finish(sessionID string) { delete(t.legacy, sessionID) } diff --git a/pkg/tui/subscription_startup_test.go b/pkg/tui/subscription_startup_test.go new file mode 100644 index 0000000000..2b52cc56ca --- /dev/null +++ b/pkg/tui/subscription_startup_test.go @@ -0,0 +1,54 @@ +package tui + +import ( + "context" + "strings" + "testing" + "time" + + tea "charm.land/bubbletea/v2" + "github.com/charmbracelet/x/ansi" + "github.com/stretchr/testify/require" + + "github.com/docker/docker-agent/pkg/app" + "github.com/docker/docker-agent/pkg/runtime" + "github.com/docker/docker-agent/pkg/tui/page/chat" + "github.com/docker/docker-agent/pkg/tui/service/supervisor" +) + +func TestSupervisorRegistersBeforeProgramReadyWithOtherSubscriber(t *testing.T) { + root, _, _ := frozenClockRoot(t, 120, 40) + root.supervisor.Shutdown() + rt := &stalledDeliveryRuntime{} + a := app.New(t.Context(), rt, root.application.Session()) + ctx, cancel := context.WithCancel(t.Context()) + defer cancel() + witness := make(chan struct{}, 1) + ready := make(chan struct{}) + go a.SubscribeReliable(ctx, func(msg tea.Msg) { + if _, ok := msg.(*runtime.StreamStoppedEvent); ok { + witness <- struct{}{} + } + }, app.WithSubscriptionReady(func() { close(ready) })) + <-ready + root.supervisor = supervisor.New(nil) + root.supervisor.AddSession(t.Context(), a, a.Session(), "", nil) + root.application = a + root.activeTab.state = nil + root.activeTab.chatPage = chat.New(root.ar, t.Context(), a, root.activeTab.sessionState, chat.WithHideSidebar()) + root.handleWindowResize(120, 40) + rt.emit(runtime.StreamStarted("profile", "root")) + rt.emit(runtime.AgentChoice("root", "profile", "ANSWER-BEFORE-PROGRAM-READY", "answer")) + rt.emit(runtime.StreamStopped("profile", "root", "normal")) + select { + case <-witness: + case <-time.After(5 * time.Second): + t.Fatal("other subscriber did not drain events") + } + model := &streamingMotionModel{root: root, ready: make(chan struct{})} + program := startStreamingMotionProgram(t, model, tea.WithOutput(&wallClockCountingWriter{})) + root.supervisor.SetProgram(program) + require.Eventually(t, func() bool { + return strings.Contains(ansi.Strip(programFrame(t, program)), "ANSWER-BEFORE-PROGRAM-READY") + }, 10*time.Second, 10*time.Millisecond) +} diff --git a/pkg/tui/tui.go b/pkg/tui/tui.go index 3187ccbab2..45b7990fb1 100644 --- a/pkg/tui/tui.go +++ b/pkg/tui/tui.go @@ -1542,6 +1542,9 @@ func (m *appModel) update(msg tea.Msg) (tea.Model, tea.Cmd) { // handleRoutedMsg processes messages routed to specific sessions. func (m *appModel) handleRoutedMsg(msg messages.RoutedMsg) (tea.Model, tea.Cmd) { + if msg.Valid != nil && !msg.Valid() { + return m, nil + } runner := m.supervisor.GetRunner(msg.SessionID) if msg.Scope != nil && (runner == nil || runner.Scope != msg.Scope) { return m, nil diff --git a/pkg/tui/types/types.go b/pkg/tui/types/types.go index 4e25403662..06a1065350 100644 --- a/pkg/tui/types/types.go +++ b/pkg/tui/types/types.go @@ -73,12 +73,15 @@ type AssistantMedia struct { // replaces it (see messages.Model.UpdateAssistantMedia). Zero means // static: the item is final and never replaced. ID uint64 + Key string // Stable manifest identity; never rendered. Image *tuiimage.Inline Fallback string } // Message represents a single message in the chat type Message struct { + SessionID string + MessageID string Type MessageType Content string Sender string // Agent name for assistant messages