From 1c3a08b3815eb20213e3bd1b887de13dd053beb3 Mon Sep 17 00:00:00 2001 From: David Gageot Date: Thu, 1 Oct 2026 10:29:48 +0200 Subject: [PATCH 01/10] fix: stop dropping assistant messages across TUIs and remote sessions Make event delivery reliable end to end: app subscriptions track generation lifetimes so stale events can't leak into a new session, both TUIs reconcile messages by ID instead of blindly appending so startup and reconnects can't duplicate or lose text, and the remote runtime recovers from streaming timeouts and event-log gaps by replaying from a safe cursor instead of guessing. Synthetic answers now get a message ID so they can be tracked the same way. --- cmd/root/run.go | 1 - cmd/root/run_listen.go | 2 +- docs/features/api-server/index.md | 16 +- pkg/app/app.go | 171 +++++- pkg/app/app_test.go | 16 + pkg/app/event_generation.go | 43 ++ pkg/app/event_generation_test.go | 91 +++ pkg/app/fanout_test.go | 6 +- pkg/app/reliable_subscriber_test.go | 164 ++++++ pkg/app/session_lifecycle_test.go | 4 +- pkg/app/subscriber.go | 103 ++++ pkg/app/subscriber_test.go | 40 ++ pkg/leantui/assistant_content_test.go | 62 ++ pkg/leantui/events.go | 28 +- pkg/leantui/leantui.go | 25 +- pkg/leantui/ui/renderer.go | 26 + pkg/leantui/ui/renderer_test.go | 43 ++ pkg/leantui/ui/transcript.go | 107 +++- pkg/leantui/update.go | 8 +- pkg/runtime/client.go | 288 +++++++--- pkg/runtime/client_test.go | 244 ++++++++ pkg/runtime/remote_runtime.go | 539 +++++++++++++++--- pkg/runtime/remote_runtime_test.go | 409 +++++++++++++ pkg/runtime/synthetic_content_test.go | 36 ++ pkg/runtime/tool_dispatch.go | 9 + pkg/server/session_manager.go | 5 +- pkg/server/session_manager_recall_test.go | 74 +++ .../components/messages/assistant_content.go | 99 ++++ pkg/tui/components/messages/messages.go | 5 + pkg/tui/messages/tabs.go | 1 + pkg/tui/page/chat/assistant_content_test.go | 96 ++++ pkg/tui/page/chat/chat.go | 2 + pkg/tui/page/chat/generated_media.go | 31 +- pkg/tui/page/chat/generated_media_test.go | 6 + pkg/tui/page/chat/runtime_events.go | 8 +- pkg/tui/service/supervisor/supervisor.go | 58 +- pkg/tui/stalled_delivery_test.go | 167 ++++++ pkg/tui/streamcontent/identity.go | 32 ++ pkg/tui/subscription_startup_test.go | 54 ++ pkg/tui/tui.go | 3 + pkg/tui/types/types.go | 2 + 41 files changed, 2868 insertions(+), 256 deletions(-) create mode 100644 pkg/app/event_generation.go create mode 100644 pkg/app/event_generation_test.go create mode 100644 pkg/app/reliable_subscriber_test.go create mode 100644 pkg/app/subscriber.go create mode 100644 pkg/app/subscriber_test.go create mode 100644 pkg/leantui/assistant_content_test.go create mode 100644 pkg/runtime/synthetic_content_test.go create mode 100644 pkg/server/session_manager_recall_test.go create mode 100644 pkg/tui/components/messages/assistant_content.go create mode 100644 pkg/tui/page/chat/assistant_content_test.go create mode 100644 pkg/tui/stalled_delivery_test.go create mode 100644 pkg/tui/streamcontent/identity.go create mode 100644 pkg/tui/subscription_startup_test.go 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..5f57dbf8af 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,18 @@ $ 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 +296,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..94a3be717b 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,8 +198,12 @@ 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. + backgroundCtx := ctx + if _, ok := a.runtime.(interface{ RetireBackgroundEvents() }); ok { + backgroundCtx = a.eventContext(ctx) + } a.runtime.OnBackgroundEvent(func(event runtime.Event) { - a.sendEvent(ctx, event) + a.sendEvent(backgroundCtx, event) }) // Forward elicitation requests raised anywhere in the runtime — @@ -419,6 +425,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 +570,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 +762,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 +817,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 +967,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 +998,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 +1033,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 +1102,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 +1120,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.push(delivery) + continue + } + ch := sub.ch select { case ch <- msg: default: @@ -1155,8 +1242,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 +1367,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 +1407,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 +1434,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 +1529,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 +1812,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 +1867,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 +1966,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 +1997,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 +2262,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..39098194ad --- /dev/null +++ b/pkg/app/event_generation.go @@ -0,0 +1,43 @@ +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() { + a.eventGeneration.Add(1) + 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..94a6e2a952 --- /dev/null +++ b/pkg/app/event_generation_test.go @@ -0,0 +1,91 @@ +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") + }) +} 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..c68261c454 --- /dev/null +++ b/pkg/app/subscriber.go @@ -0,0 +1,103 @@ +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 eventQueue struct { + mu sync.Mutex + pending []tea.Msg + ready chan struct{} + closed bool +} + +func newEventQueue() *eventQueue { + return &eventQueue{ready: make(chan struct{}, 1)} +} + +func (q *eventQueue) push(msg tea.Msg) { + q.mu.Lock() + defer q.mu.Unlock() + if q.closed { + return + } + q.pending = append(q.pending, 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] + q.pending[0] = nil + 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) 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/leantui/assistant_content_test.go b/pkg/leantui/assistant_content_test.go new file mode 100644 index 0000000000..8875d3f05b --- /dev/null +++ b/pkg/leantui/assistant_content_test.go @@ -0,0 +1,62 @@ +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/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") +} diff --git a/pkg/leantui/events.go b/pkg/leantui/events.go index 0f09509c54..d7ca33d06c 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() } @@ -37,6 +45,7 @@ func (m *model) handleEvent(ctx context.Context, ev any) { case *runtime.UserMessageEvent: m.handleUserMessageEvent(e) 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 +54,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, 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 +262,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..73003d0525 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,27 @@ func (r *Renderer) repaintVisible(newLines []string, cursorLine, cursorCol int) r.prev = newLines } +// redrawSuffix emits changed offscreen rows before they become immutable scrollback. +func (r *Renderer) redrawSuffix(newLines []string, first, cursorLine, cursorCol int) { + var b strings.Builder + b.WriteString(seqSyncStart) + b.WriteString(seqHideCursor) + b.WriteString("\x1b[2J\x1b[H") + for i, line := range newLines[first:] { + if i > 0 { + b.WriteString("\r\n") + } + b.WriteString(seqEraseLine) + b.WriteString(line) + } + r.viewportTop = max(0, len(newLines)-r.height) + 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_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/update.go b/pkg/leantui/update.go index 721f96b0d4..55abf90019 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() @@ -538,8 +540,10 @@ func (m *model) loadSessionTranscript(sess *session.Session) { 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/client.go b/pkg/runtime/client.go index fcc3b75a2b..78f483c1bb 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 { @@ -336,7 +339,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 +379,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 +409,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 +428,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 +441,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 +461,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 +511,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/remote_runtime.go b/pkg/runtime/remote_runtime.go index ceac484ec6..7a4b535827 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" @@ -49,6 +50,140 @@ type RemoteRuntime struct { // field read when no specific agent has been selected. resolvedDefault string resolvedDefaultMu sync.Mutex + + stateMu sync.Mutex // sessionID and pendingOAuthElicitation + + 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 + unidentified bool +} + +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), + } + for _, msg := range messages { + if msg.Message.Role != chat.MessageRoleAssistant { + continue + } + if msg.Message.MessageID == "" && recoverableAssistantText(msg) { + h.unidentified = true + } + content := new(strings.Builder) + content.WriteString(msg.Message.Content) + h.content[msg.Message.MessageID] = content + h.complete[msg.Message.MessageID] = true + } + return h +} + +func recoverableAssistantText(msg session.Message) bool { + return !msg.Implicit && msg.Message.Role == chat.MessageRoleAssistant && 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 + } + 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 ok { + 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.sessions[choice.MessageID] = choice.SessionID + } + } + return true +} + +func (h *remoteMessageHistory) reconcile(snapshot *api.SessionSnapshotResponse, send func(Event)) error { + h.mu.Lock() + defer h.mu.Unlock() + 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) + for _, msg := range snapshot.Messages { + if !recoverableAssistantText(msg) { + continue + } + id := msg.Message.MessageID + var delivered string + if content := h.content[id]; content != nil { + delivered = content.String() + } + if id == "" || seen[id] || !strings.HasPrefix(msg.Message.Content, delivered) { + return errors.New("cannot safely reconcile changed or ambiguous assistant messages") + } + seen[id] = true + } + for _, msg := range snapshot.Messages { + if !recoverableAssistantText(msg) { + continue + } + id := msg.Message.MessageID + content := h.content[id] + if content == nil { + content = new(strings.Builder) + h.content[id] = content + } + if suffix := strings.TrimPrefix(msg.Message.Content, content.String()); suffix != "" { + sessionID := cmp.Or(h.sessions[id], snapshot.ID) + send(AgentChoice(msg.AgentName, sessionID, suffix, id)) + content.WriteString(suffix) + } + h.complete[id] = true + } + return nil } // RemoteRuntimeOption is a function for configuring the RemoteRuntime @@ -156,10 +291,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 +307,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 +394,23 @@ 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() 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 +424,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 +445,42 @@ 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 { return sendClientEvent(ctx, events, event) } for streamEvent := range streamChan { + switch event := streamEvent.(type) { + case *StreamStoppedEvent: + sawRootStop = sawRootStop || event.SessionID == "" || event.SessionID == sess.ID + case *ErrorEvent: + sawError = true + } if elicitationRequest, ok := streamEvent.(*ElicitationRequestEvent); ok { + r.stateMu.Lock() r.pendingOAuthElicitation = elicitationRequest + r.stateMu.Unlock() + } + 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 + } + } + if !sawRootStop && ctx.Err() == nil { + if !sawError { + sendClientEvent(ctx, events, Error("remote agent stream ended before completion; the response may be incomplete")) } - events <- streamEvent + sendClientEvent(ctx, events, StreamStopped(sess.ID, r.currentAgent, "error")) } }() @@ -334,20 +510,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 +536,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 @@ -402,15 +582,19 @@ func (r *RemoteRuntime) convertSessionMessages(sess *session.Session) []api.Mess // 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 { + sessionID := r.activeSessionID() id := firstElicitationID(elicitationID) - slog.DebugContext(ctx, "Resuming remote runtime with elicitation response", "agent", r.currentAgent, "action", action, "session_id", r.sessionID, "elicitation_id", id) + slog.DebugContext(ctx, "Resuming remote runtime with elicitation response", "agent", r.currentAgent, "action", action, "session_id", sessionID, "elicitation_id", id) - err := r.handleOAuthElicitation(ctx, r.pendingOAuthElicitation) + r.stateMu.Lock() + pending := r.pendingOAuthElicitation + r.stateMu.Unlock() + err := r.handleOAuthElicitation(ctx, pending) if err != nil { return err } - if err := r.client.ResumeElicitation(ctx, r.sessionID, action, content, id); err != nil { + if err := r.client.ResumeElicitation(ctx, sessionID, action, content, id); err != nil { return err } @@ -418,6 +602,7 @@ func (r *RemoteRuntime) ResumeElicitation(ctx context.Context, action tools.Elic } func (r *RemoteRuntime) handleOAuthElicitation(ctx context.Context, req *ElicitationRequestEvent) error { + sessionID := r.activeSessionID() if req == nil { return nil } @@ -428,7 +613,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 +621,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 +629,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 +654,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 +669,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 +682,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 +719,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 +744,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 +762,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 +846,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 +887,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 +901,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 +913,301 @@ 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) emitBackgroundEvent(subscription *remoteEventSubscription, event Event) { + r.backgroundMu.Lock() + handler := r.backgroundHandler + active := r.background == subscription && !r.closed + r.backgroundMu.Unlock() + if active && handler != nil { + if request, ok := event.(*ElicitationRequestEvent); ok { + r.stateMu.Lock() + r.pendingOAuthElicitation = request + r.stateMu.Unlock() + } + subscription.history.deliver(event, func(event Event) bool { + handler(event) + return true + }) + } +} + +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), "")) + cancel() + r.backgroundMu.Lock() + if r.background == subscription { + r.background = nil + } + r.backgroundMu.Unlock() + 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 cancel() + 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 + } + r.emitBackgroundEvent(subscription, event) + } + } + }() +} + +// /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, handler); err != nil { + return nil, err + } + 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..80a1c537a0 100644 --- a/pkg/runtime/remote_runtime_test.go +++ b/pkg/runtime/remote_runtime_test.go @@ -3,12 +3,19 @@ package runtime import ( "context" "errors" + "fmt" + "net/http" + "net/http/httptest" + "strings" + "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" ) @@ -146,3 +153,405 @@ 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 + 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") + default: + t.Fatalf("unexpected event %T", event) + } + case <-deadline: + t.Fatal("saved final answer was not reconciled") + } + } + 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) +} 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..ef0b45c2f7 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, 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..392684bf7d 100644 --- a/pkg/server/session_manager.go +++ b/pkg/server/session_manager.go @@ -1415,6 +1415,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 +1429,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/tui/components/messages/assistant_content.go b/pkg/tui/components/messages/assistant_content.go new file mode 100644 index 0000000000..31b971bda9 --- /dev/null +++ b/pkg/tui/components/messages/assistant_content.go @@ -0,0 +1,99 @@ +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 { + m.trackMessageIdentity(sessionID, messageID) + cmd := m.AppendAssistantMedia(agentName, media) + if last := m.lastMessage(); last != nil { + last.SessionID, last.MessageID = sessionID, messageID + } + return cmd +} diff --git a/pkg/tui/components/messages/messages.go b/pkg/tui/components/messages/messages.go index e62e4245cc..52585fb06b 100644 --- a/pkg/tui/components/messages/messages.go +++ b/pkg/tui/components/messages/messages.go @@ -89,6 +89,10 @@ 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 // BreakMessageGroup prevents merging across streams without flushing deferred content. BreakMessageGroup() // AppendAssistantMedia attaches generated media to the agent's current @@ -1938,6 +1942,7 @@ func (m *model) LoadFromSession(sess *session.Session, generatedMedia map[int][] restoredMedia := generatedMedia[pos] if hasContent || len(restoredMedia) > 0 { msg := types.Agent(types.MessageTypeAssistant, smsg.AgentName, smsg.Message.Content) + msg.SessionID, msg.MessageID = sess.ID, smsg.Message.MessageID msg.AssistantMedia = restoredMedia appendSessionMessage(msg, m.createMessageView(msg)) } 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..86d16b101a --- /dev/null +++ b/pkg/tui/page/chat/assistant_content_test.go @@ -0,0 +1,96 @@ +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)) +} diff --git a/pkg/tui/page/chat/chat.go b/pkg/tui/page/chat/chat.go index c07994f8b2..9abfe8d7cf 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,7 @@ type chatPage struct { agentStack []string // agent per active stream level; len(agentStack)==streamDepth streamStartTime time.Time contentSessionID string + contentIdentity streamcontent.Tracker // 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..fed925cdc9 100644 --- a/pkg/tui/page/chat/generated_media.go +++ b/pkg/tui/page/chat/generated_media.go @@ -63,24 +63,39 @@ 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 } + sessionID := p.contentSession(msg.SessionID) + identity := p.contentIdentity.Resolve(sessionID, msg.Message.Message.MessageID) + if msg.Message.Message.Content != "" { + defer p.contentIdentity.Finish(sessionID) + } + agentName := msg.Message.AgentName + if agentName == "" { + agentName = msg.AgentName + } + content := 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 + } placeholders, requests := generatedImageMedia(msg.Message.Message.MultiContent) if len(placeholders) == 0 { - return nil + return contentCmd } - p.trackContentSession(msg.SessionID) p.hasReceivedAssistantContent = true p.setPendingResponse(false) - agentName := msg.Message.AgentName - if agentName == "" { - agentName = msg.AgentName - } 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), ) } diff --git a/pkg/tui/page/chat/generated_media_test.go b/pkg/tui/page/chat/generated_media_test.go index 6e215a545d..0cfdaf4cda 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) diff --git a/pkg/tui/page/chat/runtime_events.go b/pkg/tui/page/chat/runtime_events.go index bf63c9f80c..42f48a01a4 100644 --- a/pkg/tui/page/chat/runtime_events.go +++ b/pkg/tui/page/chat/runtime_events.go @@ -303,6 +303,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 +340,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 +348,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 +356,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 +394,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/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..33011e4f83 100644 --- a/pkg/tui/types/types.go +++ b/pkg/tui/types/types.go @@ -79,6 +79,8 @@ type AssistantMedia struct { // 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 From ea540e9f0126a46617807ded43e6fc63a461c628 Mon Sep 17 00:00:00 2001 From: David Gageot Date: Thu, 1 Oct 2026 11:35:23 +0200 Subject: [PATCH 02/10] fix: suppress tool-call XML from canonical and reloaded assistant text Add chat.VisibleAssistantContent so synthetic answers, lean TUI session reload, and remote recovery never surface raw payloads that were already hidden from the live stream. Assisted-By: docker-agent --- pkg/chat/visible_content.go | 9 +++++++++ pkg/chat/visible_content_test.go | 19 +++++++++++++++++++ pkg/leantui/update.go | 1 + pkg/runtime/tool_dispatch.go | 2 +- 4 files changed, 30 insertions(+), 1 deletion(-) create mode 100644 pkg/chat/visible_content.go create mode 100644 pkg/chat/visible_content_test.go 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/update.go b/pkg/leantui/update.go index 55abf90019..61881da6d7 100644 --- a/pkg/leantui/update.go +++ b/pkg/leantui/update.go @@ -535,6 +535,7 @@ 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) }) diff --git a/pkg/runtime/tool_dispatch.go b/pkg/runtime/tool_dispatch.go index ef0b45c2f7..8120ad9788 100644 --- a/pkg/runtime/tool_dispatch.go +++ b/pkg/runtime/tool_dispatch.go @@ -197,7 +197,7 @@ func addAgentMessage(sess *session.Session, a *agent.Agent, msg *chat.Message, e agentMsg := session.NewAgentMessage(a.Name(), msg) sess.AddMessage(agentMsg) if synthetic && msg.Content != "" { - events.Emit(AgentChoice(a.Name(), sess.ID, msg.Content, msg.MessageID)) + events.Emit(AgentChoice(a.Name(), sess.ID, chat.VisibleAssistantContent(msg.Content), msg.MessageID)) } events.Emit(MessageAdded(sess.ID, agentMsg, a.Name())) } From b938331de64c269c0768db7f57df0071310e42e2 Mon Sep 17 00:00:00 2001 From: David Gageot Date: Thu, 1 Oct 2026 11:35:31 +0200 Subject: [PATCH 03/10] feat(runtime): add SessionRecovered event for idle recovery boundaries A dedicated event lets consumers distinguish an authoritative idle reset from a stream stop, so recovery doesn't fire stop-triggered actions. Wires it into the SSE client decoder. Assisted-By: docker-agent --- pkg/runtime/client.go | 1 + pkg/runtime/event.go | 19 ++++++++++++++ pkg/runtime/recovery_event_test.go | 42 ++++++++++++++++++++++++++++++ 3 files changed, 62 insertions(+) create mode 100644 pkg/runtime/recovery_event_test.go diff --git a/pkg/runtime/client.go b/pkg/runtime/client.go index 78f483c1bb..0e5f88ec3d 100644 --- a/pkg/runtime/client.go +++ b/pkg/runtime/client.go @@ -88,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{} }, 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) +} From e588f069e5a2caa36140071f591075dba38c281d Mon Sep 17 00:00:00 2001 From: David Gageot Date: Thu, 1 Oct 2026 11:35:43 +0200 Subject: [PATCH 04/10] fix(runtime,app): keep detached background/elicitation events on their origin context LocalRuntime can now register context-aware background and elicitation sinks (OnBackgroundEventWithContext, OnElicitationRequestWithContext), and App uses them when available so a detached background task or elicitation keeps the conversation it started in, even across session replacement. Per-subscriber queues now carry a generation so stale queued deliveries are discarded on retirement without racing newly accepted ones. Assisted-By: docker-agent --- pkg/app/app.go | 32 ++++++++++++---- pkg/app/event_generation.go | 9 ++++- pkg/app/event_generation_test.go | 65 ++++++++++++++++++++++++++++++++ pkg/app/subscriber.go | 33 +++++++++++++--- pkg/runtime/agent_delegation.go | 10 ++--- pkg/runtime/elicitation.go | 31 ++++++++++++++- pkg/runtime/elicitation_test.go | 12 ++++++ pkg/runtime/runtime.go | 22 +++++++++-- 8 files changed, 191 insertions(+), 23 deletions(-) diff --git a/pkg/app/app.go b/pkg/app/app.go index 94a3be717b..3420bb61c4 100644 --- a/pkg/app/app.go +++ b/pkg/app/app.go @@ -202,9 +202,18 @@ func (a *App) Start(ctx context.Context) { if _, ok := a.runtime.(interface{ RetireBackgroundEvents() }); ok { backgroundCtx = a.eventContext(ctx) } - a.runtime.OnBackgroundEvent(func(event runtime.Event) { - a.sendEvent(backgroundCtx, event) - }) + 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 @@ -218,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) }) + } }) } @@ -1210,7 +1228,7 @@ func (a *App) startFanOut() { } } if sub.queue != nil { - sub.queue.push(delivery) + sub.queue.pushGeneration(delivery, generation) continue } ch := sub.ch diff --git a/pkg/app/event_generation.go b/pkg/app/event_generation.go index 39098194ad..9de097df34 100644 --- a/pkg/app/event_generation.go +++ b/pkg/app/event_generation.go @@ -17,7 +17,14 @@ type generationEvent struct { // RetireEvents discards deliveries belonging to a replaced conversation. func (a *App) RetireEvents() { - a.eventGeneration.Add(1) + 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()) diff --git a/pkg/app/event_generation_test.go b/pkg/app/event_generation_test.go index 94a6e2a952..87f26ac66e 100644 --- a/pkg/app/event_generation_test.go +++ b/pkg/app/event_generation_test.go @@ -89,3 +89,68 @@ func TestRoutingRetainsOriginWhenRetirementRacesMapping(t *testing.T) { 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/subscriber.go b/pkg/app/subscriber.go index c68261c454..f66a7dfda5 100644 --- a/pkg/app/subscriber.go +++ b/pkg/app/subscriber.go @@ -38,9 +38,14 @@ 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 []tea.Msg + pending []queuedEvent ready chan struct{} closed bool } @@ -49,13 +54,15 @@ func newEventQueue() *eventQueue { return &eventQueue{ready: make(chan struct{}, 1)} } -func (q *eventQueue) push(msg tea.Msg) { +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, msg) + q.pending = append(q.pending, queuedEvent{generation: generation, msg: msg}) select { case q.ready <- struct{}{}: default: @@ -74,8 +81,8 @@ func (q *eventQueue) next(ctx context.Context, done <-chan struct{}) (tea.Msg, b q.mu.Lock() if len(q.pending) > 0 { - msg := q.pending[0] - q.pending[0] = nil + msg := q.pending[0].msg + q.pending[0] = queuedEvent{} q.pending = q.pending[1:] if len(q.pending) == 0 { q.pending = nil @@ -95,6 +102,22 @@ func (q *eventQueue) next(ctx context.Context, done <-chan struct{}) (tea.Msg, b } } +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() 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/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/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) } From 82d514968d6b819e0f785c56a159dc550d0d11d0 Mon Sep 17 00:00:00 2001 From: David Gageot Date: Thu, 1 Oct 2026 11:35:59 +0200 Subject: [PATCH 05/10] fix(runtime): bound remote recovery history and correlate OAuth elicitations remoteMessageHistory now tracks nested stream depth and byte/message counts and fails recovery closed instead of replaying once any limit is exceeded, release message buffers as soon as a stream completes, and reconcile against suppressed (non-tool) content. OAuth elicitations are tracked per elicitation ID instead of a single shared pointer, serialized so concurrent accepts can't double-authorize, and cleared on root SessionRecovered without touching detached sub-session requests. A failed background subscription can now restart on the next call instead of sticking around as a dead placeholder. Assisted-By: docker-agent --- pkg/runtime/remote_runtime.go | 246 +++++++++++++++++------ pkg/runtime/remote_runtime_test.go | 308 +++++++++++++++++++++++++++++ 2 files changed, 491 insertions(+), 63 deletions(-) diff --git a/pkg/runtime/remote_runtime.go b/pkg/runtime/remote_runtime.go index 7a4b535827..fb8918feeb 100644 --- a/pkg/runtime/remote_runtime.go +++ b/pkg/runtime/remote_runtime.go @@ -30,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 @@ -51,7 +52,7 @@ type RemoteRuntime struct { resolvedDefault string resolvedDefaultMu sync.Mutex - stateMu sync.Mutex // sessionID and pendingOAuthElicitation + stateMu sync.Mutex // sessionID and pendingOAuthElicitations reconcileMu sync.RWMutex // snapshots must not race foreground delivery backgroundInit sync.Mutex @@ -81,7 +82,30 @@ type remoteMessageHistory struct { 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 { @@ -90,24 +114,27 @@ func newRemoteMessageHistory(messages []session.Message) *remoteMessageHistory { 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 msg.Message.Role != chat.MessageRoleAssistant { + if !recoverableAssistantText(msg) { continue } - if msg.Message.MessageID == "" && recoverableAssistantText(msg) { + if msg.Message.MessageID == "" { h.unidentified = true + } else { + h.complete[msg.Message.MessageID] = true + } + h.limit() + if h.retentionErr != nil { + break } - content := new(strings.Builder) - content.WriteString(msg.Message.Content) - h.content[msg.Message.MessageID] = content - h.complete[msg.Message.MessageID] = true } return h } func recoverableAssistantText(msg session.Message) bool { - return !msg.Implicit && msg.Message.Role == chat.MessageRoleAssistant && msg.Message.Content != "" && + return !msg.Implicit && msg.Message.Role == chat.MessageRoleAssistant && chat.VisibleAssistantContent(msg.Message.Content) != "" && len(msg.Message.ToolCalls) == 0 && msg.Message.FunctionCall == nil } @@ -118,6 +145,10 @@ func (h *remoteMessageHistory) deliver(event Event, send func(Event) bool) bool 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 @@ -128,7 +159,32 @@ func (h *remoteMessageHistory) deliver(event Event, send func(Event) bool) bool if isElicitation && request.ElicitationID != "" { h.elicitations[request.ElicitationID] = true } - if ok { + 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 { @@ -138,20 +194,26 @@ func (h *remoteMessageHistory) deliver(event Event, send func(Event) bool) bool 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 @@ -159,11 +221,17 @@ func (h *remoteMessageHistory) reconcile(snapshot *api.SessionSnapshotResponse, id := msg.Message.MessageID var delivered string if content := h.content[id]; content != nil { - delivered = content.String() + delivered = chat.VisibleAssistantContent(content.String()) } - if id == "" || seen[id] || !strings.HasPrefix(msg.Message.Content, delivered) { + 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 { @@ -171,18 +239,24 @@ func (h *remoteMessageHistory) reconcile(snapshot *api.SessionSnapshotResponse, continue } id := msg.Message.MessageID - content := h.content[id] - if content == nil { - content = new(strings.Builder) - h.content[id] = content + if h.complete[id] { + continue } - if suffix := strings.TrimPrefix(msg.Message.Content, content.String()); suffix != "" { + 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)) - content.WriteString(suffix) } + delete(h.content, id) + delete(h.sessions, id) h.complete[id] = true } + clear(h.streamDepth) + h.limit() return nil } @@ -408,6 +482,9 @@ func (r *RemoteRuntime) RunStream(ctx context.Context, sess *session.Session) <- 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) @@ -452,7 +529,10 @@ func (r *RemoteRuntime) RunStream(ctx context.Context, sess *session.Session) <- } }() var sawRootStop, sawError bool - send := func(event Event) bool { return sendClientEvent(ctx, events, event) } + send := func(event Event) bool { + r.trackOAuthElicitation(event) + return sendClientEvent(ctx, events, event) + } for streamEvent := range streamChan { switch event := streamEvent.(type) { case *StreamStoppedEvent: @@ -460,11 +540,6 @@ func (r *RemoteRuntime) RunStream(ctx context.Context, sess *session.Session) <- case *ErrorEvent: sawError = true } - if elicitationRequest, ok := streamEvent.(*ElicitationRequestEvent); ok { - r.stateMu.Lock() - r.pendingOAuthElicitation = elicitationRequest - r.stateMu.Unlock() - } r.backgroundMu.Lock() subscription := r.background r.backgroundMu.Unlock() @@ -580,29 +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) trackOAuthElicitation(event Event) { + request, ok := event.(*ElicitationRequestEvent) + if !ok || request.Meta["docker-agent/type"] != "oauth_flow" || request.ElicitationID == "" { + return + } + 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 + } +} + +// 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 { - sessionID := r.activeSessionID() + // 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) - slog.DebugContext(ctx, "Resuming remote runtime with elicitation response", "agent", r.currentAgent, "action", action, "session_id", sessionID, "elicitation_id", id) - r.stateMu.Lock() - pending := r.pendingOAuthElicitation + sessionID := r.sessionID + pending := r.pendingOAuthElicitations[id] + ambiguous := id == "" && len(r.pendingOAuthElicitations) > 0 r.stateMu.Unlock() - err := r.handleOAuthElicitation(ctx, pending) - if err != nil { - return err + if ambiguous { + return errors.New("OAuth elicitation requires an explicit elicitation ID") } - - if err := r.client.ResumeElicitation(ctx, sessionID, action, content, id); err != nil { - return err + 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 nil + return r.client.ResumeElicitation(ctx, sessionID, action, content, id) } -func (r *RemoteRuntime) handleOAuthElicitation(ctx context.Context, req *ElicitationRequestEvent) error { - sessionID := r.activeSessionID() +func (r *RemoteRuntime) handleOAuthElicitation(ctx context.Context, sessionID string, req *ElicitationRequestEvent) error { if req == nil { return nil } @@ -931,21 +1028,42 @@ func (r *RemoteRuntime) activeSessionID() string { return r.sessionID } -func (r *RemoteRuntime) emitBackgroundEvent(subscription *remoteEventSubscription, event Event) { +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 { - if request, ok := event.(*ElicitationRequestEvent); ok { - r.stateMu.Lock() - r.pendingOAuthElicitation = request - r.stateMu.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) + } } - subscription.history.deliver(event, func(event Event) bool { - handler(event) - return true - }) + 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 } } @@ -987,12 +1105,7 @@ func (r *RemoteRuntime) startBackgroundEvents(ctx context.Context, sessionID str } if err != nil { r.emitBackgroundEvent(subscription, Warning(fmt.Sprintf("remote background events unavailable: %v", err), "")) - cancel() - r.backgroundMu.Lock() - if r.background == subscription { - r.background = nil - } - r.backgroundMu.Unlock() + r.clearBackgroundSubscription(subscription) return } r.backgroundMu.Lock() @@ -1005,7 +1118,7 @@ func (r *RemoteRuntime) startBackgroundEvents(ctx context.Context, sessionID str r.backgroundMu.Unlock() go func() { - defer cancel() + 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))) @@ -1037,7 +1150,9 @@ func (r *RemoteRuntime) startBackgroundEvents(ctx context.Context, sessionID str } continue } - r.emitBackgroundEvent(subscription, event) + if !r.emitBackgroundEvent(subscription, event) { + return + } } } }() @@ -1099,9 +1214,14 @@ func (r *RemoteRuntime) reconcileIdleSnapshot(ctx context.Context, subscription if !active || handler == nil || ctx.Err() != nil { return nil, context.Canceled } - if err := subscription.history.reconcile(snapshot, handler); err != nil { + 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 } diff --git a/pkg/runtime/remote_runtime_test.go b/pkg/runtime/remote_runtime_test.go index 80a1c537a0..5ceaeb2b9a 100644 --- a/pkg/runtime/remote_runtime_test.go +++ b/pkg/runtime/remote_runtime_test.go @@ -6,7 +6,9 @@ import ( "fmt" "net/http" "net/http/httptest" + "strconv" "strings" + "sync" "sync/atomic" "testing" "time" @@ -18,6 +20,7 @@ import ( "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 @@ -361,6 +364,7 @@ func TestRemoteRuntime_BackgroundGapReconcilesSavedTextWithoutReplay(t *testing. t.Fatal("gap recovery did not reconnect") } var choices []*AgentChoiceEvent + var recovered int deadline := time.After(3 * time.Second) for len(choices) < 3 { select { @@ -370,6 +374,9 @@ func TestRemoteRuntime_BackgroundGapReconcilesSavedTextWithoutReplay(t *testing. 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) } @@ -377,6 +384,7 @@ func TestRemoteRuntime_BackgroundGapReconcilesSavedTextWithoutReplay(t *testing. 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) @@ -555,3 +563,303 @@ func TestRemoteRuntimeRetiredBackgroundSubscriptionDoesNotEmit(t *testing.T) { 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") +} From a5a67add9be48045eeeb65cec89e291f33355458 Mon Sep 17 00:00:00 2001 From: David Gageot Date: Thu, 1 Oct 2026 11:36:08 +0200 Subject: [PATCH 06/10] fix(server): hold an idle snapshot boundary across GetSessionSnapshot A new snapshotBoundary mutex is held while copying session messages and event cursor so a concurrent recall can't start a run (and move both) mid-copy; recall now waits on the same boundary instead of racing it. Snapshots clone the session and recheck the event cursor so background events landing between the streaming check and the copy are reflected as still-streaming rather than silently dropped. Assisted-By: docker-agent --- pkg/server/session_manager.go | 30 +++-- pkg/server/session_manager_snapshot_test.go | 123 ++++++++++++++++++++ 2 files changed, 144 insertions(+), 9 deletions(-) create mode 100644 pkg/server/session_manager_snapshot_test.go diff --git a/pkg/server/session_manager.go b/pkg/server/session_manager.go index 392684bf7d..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) } 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()) +} From a886ca9274f0f92322a9bd95a30753e8f1ef218c Mon Sep 17 00:00:00 2001 From: David Gageot Date: Thu, 1 Oct 2026 11:36:20 +0200 Subject: [PATCH 07/10] fix(tui,leantui): treat SessionRecovered as an idle reset, not a stop Both TUIs and tabstate now clear nested stream depth, spinners, and transfer/compaction presentation on SessionRecovered without running stop-triggered actions (draining the queued-message buffer, firing attention) and without discarding elicitations detached to a sub-session, since root idleness doesn't mean they completed. Assisted-By: docker-agent --- pkg/leantui/assistant_content_test.go | 37 +++++++++++++++++++++ pkg/leantui/events.go | 17 +++++++++- pkg/leantui/ui/usage.go | 8 +++++ pkg/tui/components/sidebar/sidebar.go | 8 +++++ pkg/tui/page/chat/assistant_content_test.go | 27 +++++++++++++++ pkg/tui/page/chat/runtime_events.go | 14 ++++++++ pkg/tui/service/tabstate/state.go | 8 +++++ pkg/tui/service/tabstate/state_test.go | 10 ++++++ 8 files changed, 128 insertions(+), 1 deletion(-) diff --git a/pkg/leantui/assistant_content_test.go b/pkg/leantui/assistant_content_test.go index 8875d3f05b..a8f9caf5f1 100644 --- a/pkg/leantui/assistant_content_test.go +++ b/pkg/leantui/assistant_content_test.go @@ -8,6 +8,7 @@ import ( "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" ) @@ -60,3 +61,39 @@ func TestRetiredLeanEventsCannotAffectReplacement(t *testing.T) { 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 d7ca33d06c..854336d2da 100644 --- a/pkg/leantui/events.go +++ b/pkg/leantui/events.go @@ -44,6 +44,21 @@ 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() @@ -63,7 +78,7 @@ func (m *model) handleEvent(ctx context.Context, ev any) { } sessionID := m.contentSession(e.SessionID) identity := m.contentIdentity.Resolve(sessionID, e.Message.Message.MessageID) - m.screen.Transcript.ReconcileAssistantContent(identity, e.Message.Message.Content) + m.screen.Transcript.ReconcileAssistantContent(identity, chat.VisibleAssistantContent(e.Message.Message.Content)) m.contentIdentity.Finish(sessionID) case *runtime.PartialToolCallEvent: m.screen.Transcript.FlushPending() 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/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/page/chat/assistant_content_test.go b/pkg/tui/page/chat/assistant_content_test.go index 86d16b101a..c0f06ca848 100644 --- a/pkg/tui/page/chat/assistant_content_test.go +++ b/pkg/tui/page/chat/assistant_content_test.go @@ -94,3 +94,30 @@ func TestPendingSpinnerDoesNotMergeDistinctMessageIDs(t *testing.T) { 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/runtime_events.go b/pkg/tui/page/chat/runtime_events.go index 42f48a01a4..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) 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)) +} From 540261aa05aacd6f31022250588a5a3ba60fdee3 Mon Sep 17 00:00:00 2001 From: David Gageot Date: Thu, 1 Oct 2026 11:36:31 +0200 Subject: [PATCH 08/10] fix(tui): give canonical generated media a stable identity Derive a dedup key from session/message/artifact instead of resolution order, so re-resolving, restoring from a session reload, and replaying canonical content that interleaves with other messages all join the same placeholder instead of appending duplicates. Legacy (ID-less) canonical replay adopts the identity of the already-restored manifest entry it matches. Assisted-By: docker-agent --- .../components/messages/assistant_content.go | 42 +++++++++++++++ pkg/tui/components/messages/messages.go | 13 ++++- pkg/tui/page/chat/chat.go | 1 + pkg/tui/page/chat/generated_media.go | 33 +++++++++--- pkg/tui/page/chat/generated_media_test.go | 53 +++++++++++++++++++ pkg/tui/types/types.go | 1 + 6 files changed, 134 insertions(+), 9 deletions(-) diff --git a/pkg/tui/components/messages/assistant_content.go b/pkg/tui/components/messages/assistant_content.go index 31b971bda9..bd6fc8c735 100644 --- a/pkg/tui/components/messages/assistant_content.go +++ b/pkg/tui/components/messages/assistant_content.go @@ -90,6 +90,30 @@ func (m *model) ReconcileAssistantContent(sessionID, messageID, agentName, conte // 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 { @@ -97,3 +121,21 @@ func (m *model) AppendAssistantMediaContent(sessionID, messageID, agentName stri } 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 52585fb06b..05b57c5a78 100644 --- a/pkg/tui/components/messages/messages.go +++ b/pkg/tui/components/messages/messages.go @@ -93,6 +93,7 @@ type Model interface { 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 @@ -336,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() @@ -1919,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 @@ -1941,7 +1949,7 @@ 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)) @@ -2209,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/page/chat/chat.go b/pkg/tui/page/chat/chat.go index 9abfe8d7cf..320042b507 100644 --- a/pkg/tui/page/chat/chat.go +++ b/pkg/tui/page/chat/chat.go @@ -234,6 +234,7 @@ type chatPage struct { 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 fed925cdc9..f8ac010001 100644 --- a/pkg/tui/page/chat/generated_media.go +++ b/pkg/tui/page/chat/generated_media.go @@ -75,7 +75,13 @@ func (p *chatPage) handleMessageAdded(msg *runtime.MessageAddedEvent) tea.Cmd { if agentName == "" { agentName = msg.AgentName } - content := msg.Message.Message.Content + 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) @@ -85,7 +91,6 @@ func (p *chatPage) handleMessageAdded(msg *runtime.MessageAddedEvent) tea.Cmd { if !p.app.CanResolveGeneratedFiles() { return contentCmd } - placeholders, requests := generatedImageMedia(msg.Message.Message.MultiContent) if len(placeholders) == 0 { return contentCmd } @@ -114,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 } @@ -134,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 { @@ -154,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 0cfdaf4cda..9445f032e8 100644 --- a/pkg/tui/page/chat/generated_media_test.go +++ b/pkg/tui/page/chat/generated_media_test.go @@ -679,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/types/types.go b/pkg/tui/types/types.go index 33011e4f83..06a1065350 100644 --- a/pkg/tui/types/types.go +++ b/pkg/tui/types/types.go @@ -73,6 +73,7 @@ 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 } From 7e82e0e1eee6b031cd86a81396a2c5b25ad160a5 Mon Sep 17 00:00:00 2001 From: David Gageot Date: Thu, 1 Oct 2026 11:36:40 +0200 Subject: [PATCH 09/10] fix(leantui): stop replaying the full offscreen suffix on every redraw redrawSuffix now diffs against the previous frame and only re-emits offscreen rows that actually changed before restoring the visible tail, so an answer pushed into scrollback by unrelated tool-timer churn is archived once instead of duplicated on each tick. Assisted-By: docker-agent --- pkg/leantui/ui/renderer.go | 21 +++- pkg/leantui/ui/renderer_scrollback_test.go | 109 +++++++++++++++++++++ 2 files changed, 127 insertions(+), 3 deletions(-) create mode 100644 pkg/leantui/ui/renderer_scrollback_test.go diff --git a/pkg/leantui/ui/renderer.go b/pkg/leantui/ui/renderer.go index 73003d0525..c8f6b2043d 100644 --- a/pkg/leantui/ui/renderer.go +++ b/pkg/leantui/ui/renderer.go @@ -179,20 +179,35 @@ func (r *Renderer) repaintVisible(newLines []string, cursorLine, cursorCol int) r.prev = newLines } -// redrawSuffix emits changed offscreen rows before they become immutable scrollback. +// 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") - for i, line := range newLines[first:] { + 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 = max(0, len(newLines)-r.height) + r.viewportTop = top r.cursorRow = r.moveCursor(&b, len(newLines)-1, cursorLine, cursorCol) b.WriteString(seqShowCursor) b.WriteString(seqSyncEnd) 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) +} From 9e5d60f61249975d90b64e7863b9e59cfb27176c Mon Sep 17 00:00:00 2001 From: David Gageot Date: Thu, 1 Oct 2026 12:05:16 +0200 Subject: [PATCH 10/10] docs: fix API server markdown lint Signed-off-by: David Gageot --- docs/features/api-server/index.md | 1 - 1 file changed, 1 deletion(-) diff --git a/docs/features/api-server/index.md b/docs/features/api-server/index.md index 5f57dbf8af..2d0a13dbc9 100644 --- a/docs/features/api-server/index.md +++ b/docs/features/api-server/index.md @@ -288,7 +288,6 @@ 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