Skip to content

Commit d6b2035

Browse files
committed
fix(terminal): prevent duplicate messages after agent navigation
1 parent aca2032 commit d6b2035

11 files changed

Lines changed: 390 additions & 78 deletions

internal/terminal/agent_tasks.go

Lines changed: 24 additions & 20 deletions
Original file line numberDiff line numberDiff line change
@@ -2061,29 +2061,11 @@ func (app *App) appendMissingSessionMessages(messages []database.SessionMessageE
20612061

20622062
for index := range messages {
20632063
message := &messages[index]
2064-
role := transcript.FromDatabaseRole(message.Role)
2065-
found := false
2066-
2067-
for historyIndex := range app.transcript.History {
2068-
history := &app.transcript.History[historyIndex]
2069-
if history.CreatedAt.Equal(message.CreatedAt) &&
2070-
history.Role == role && history.Content == message.Content {
2071-
found = true
2072-
2073-
break
2074-
}
2075-
}
2076-
2077-
if found {
2064+
if app.hasSessionMessage(message) {
20782065
continue
20792066
}
20802067

2081-
app.appendMessage(chatMessage{
2082-
CreatedAt: message.CreatedAt,
2083-
Role: role,
2084-
Content: message.Content,
2085-
Attachments: databaseAttachmentSummaries(message.Parts),
2086-
})
2068+
app.appendMessage(chatMessageFromSessionMessage(message))
20872069

20882070
appended = true
20892071

@@ -2102,6 +2084,28 @@ func (app *App) appendMissingSessionMessages(messages []database.SessionMessageE
21022084
}
21032085
}
21042086

2087+
func (app *App) hasSessionMessage(message *database.SessionMessageEntity) bool {
2088+
role := transcript.FromDatabaseRole(message.Role)
2089+
2090+
for index := range app.transcript.History {
2091+
history := &app.transcript.History[index]
2092+
if message.EntryID != "" && history.EntryID != nil && *history.EntryID == message.EntryID {
2093+
return true
2094+
}
2095+
2096+
if history.EntryID == nil && history.CreatedAt.Equal(message.CreatedAt) &&
2097+
history.Role == role && history.Content == message.Content {
2098+
if message.EntryID != "" {
2099+
history.EntryID = cloneStringPtr(&message.EntryID)
2100+
}
2101+
2102+
return true
2103+
}
2104+
}
2105+
2106+
return false
2107+
}
2108+
21052109
func (app *App) resumeInspectedAgentTask(ctx context.Context, childSessionID string) {
21062110
ownerSessionID := app.agentTaskSessionStack[len(app.agentTaskSessionStack)-1]
21072111

internal/terminal/agent_tasks_behavior_internal_test.go

Lines changed: 120 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -764,6 +764,126 @@ func TestLeaveAgentTaskSessionRefreshesDurableParentTranscript(t *testing.T) {
764764
assert.Contains(t, app.transcript.History[0].Content, "new durable parent message")
765765
}
766766

767+
func TestAgentBackCommandReconcilesOptimisticPromptByEntryID(t *testing.T) {
768+
t.Parallel()
769+
770+
fixture, _, app := newAgentTaskSessionTestApp(t, database.TaskSucceeded)
771+
772+
const prompt = "optimistic parent prompt"
773+
774+
entry, err := fixture.sessions.AppendMessage(t.Context(), fixture.parent.ID, nil, &database.MessageEntity{
775+
Timestamp: time.Now().UTC().Add(time.Second),
776+
Role: database.RoleUser,
777+
Content: prompt,
778+
Provider: "",
779+
Model: "", Parts: nil,
780+
})
781+
require.NoError(t, err)
782+
783+
message := newChatMessage(transcript.RoleUser, prompt)
784+
message.EntryID = cloneStringPtr(&entry.ID)
785+
app.appendMessage(message)
786+
787+
require.NoError(t, app.inspectAgentTask(t.Context(), behaviorTaskID))
788+
quit, err := app.runSessionCommand(t.Context(), "agents", "back", "/agents back")
789+
require.NoError(t, err)
790+
assert.False(t, quit)
791+
assert.Equal(t, fixture.parent.ID, app.sessionID)
792+
assert.Empty(t, app.agentTaskSessionStack)
793+
794+
matches := 0
795+
796+
for index := range app.transcript.History {
797+
if app.transcript.History[index].Content == prompt {
798+
matches++
799+
}
800+
}
801+
802+
assert.Equal(t, 1, matches)
803+
}
804+
805+
func TestAgentBackAfterParentPromptCompletionDoesNotDuplicateDurableMessages(t *testing.T) {
806+
t.Parallel()
807+
808+
fixture, _, app := newActivePromptInspectionTestApp(t, database.TaskSucceeded, nil)
809+
810+
const (
811+
prompt = "parent prompt completed during inspection"
812+
response = "parent response completed during inspection"
813+
)
814+
815+
userMessage := newChatMessage(transcript.RoleUser, prompt)
816+
app.activePrompt.Prompt = prompt
817+
app.activePrompt.UserMessageTimestamp = userMessage.CreatedAt.UnixNano()
818+
app.appendMessage(userMessage)
819+
promptID := app.activePrompt.ID
820+
821+
require.NoError(t, app.inspectAgentTask(t.Context(), behaviorTaskID))
822+
823+
userEntry, err := fixture.sessions.AppendMessage(t.Context(), fixture.parent.ID, nil, &database.MessageEntity{
824+
Timestamp: userMessage.CreatedAt.Add(time.Second),
825+
Role: database.RoleUser,
826+
Content: prompt,
827+
Provider: "",
828+
Model: "", Parts: nil,
829+
})
830+
require.NoError(t, err)
831+
assistantEntry, err := fixture.sessions.AppendMessage(
832+
t.Context(),
833+
fixture.parent.ID,
834+
&userEntry.ID,
835+
&database.MessageEntity{
836+
Timestamp: userMessage.CreatedAt.Add(2 * time.Second),
837+
Role: database.RoleAssistant,
838+
Content: response,
839+
Provider: "",
840+
Model: "", Parts: nil,
841+
},
842+
)
843+
require.NoError(t, err)
844+
845+
app.handlePromptAsyncEvent(t.Context(), asyncTestEvent(
846+
asyncEventPromptUserEntry,
847+
fixture.parent.ID,
848+
userEntry.ID,
849+
promptID,
850+
))
851+
852+
promptResponse := newTestPromptResponse(response)
853+
promptResponse.SessionID = fixture.parent.ID
854+
promptResponse.UserEntryID = userEntry.ID
855+
promptResponse.AssistantEntryID = assistantEntry.ID
856+
app.handlePromptAsyncEvent(t.Context(), &asyncEvent{
857+
Response: promptResponse, ToolCallEvent: nil, ToolEvent: nil, Usage: nil,
858+
Kind: asyncEventPromptDone, Provider: "", Text: "", PromptID: promptID,
859+
})
860+
861+
assert.Equal(t, fixture.child.ID, app.sessionID)
862+
assert.Nil(t, app.activePrompt)
863+
864+
quit, err := app.runSessionCommand(t.Context(), "agents", "back", "/agents back")
865+
require.NoError(t, err)
866+
assert.False(t, quit)
867+
assert.Equal(t, fixture.parent.ID, app.sessionID)
868+
869+
messageCounts := make(map[string]int)
870+
entryIDsByContent := make(map[string][]string)
871+
872+
for index := range app.transcript.History {
873+
message := &app.transcript.History[index]
874+
875+
messageCounts[message.Content]++
876+
if message.EntryID != nil {
877+
entryIDsByContent[message.Content] = append(entryIDsByContent[message.Content], *message.EntryID)
878+
}
879+
}
880+
881+
assert.Equal(t, 1, messageCounts[prompt])
882+
assert.Equal(t, 1, messageCounts[response])
883+
assert.Equal(t, []string{userEntry.ID}, entryIDsByContent[prompt])
884+
assert.Equal(t, []string{assistantEntry.ID}, entryIDsByContent[response])
885+
}
886+
767887
func TestRevisitAgentTaskSessionRefreshesDurableTranscript(t *testing.T) {
768888
t.Parallel()
769889

internal/terminal/app.go

Lines changed: 32 additions & 15 deletions
Original file line numberDiff line numberDiff line change
@@ -53,19 +53,20 @@ const (
5353
type chatMessage struct {
5454
Attachments *attachmentSummaries
5555
CreatedAt time.Time
56+
EntryID *string
5657
Role transcript.Role
5758
Content string
5859
}
5960

6061
type activePromptState struct {
61-
Cancel context.CancelFunc
62-
ParentEntryID *string
63-
SessionID string
64-
UserEntryID string
65-
Prompt string
66-
Images []imageAttachment
67-
ID uint64
68-
Canceled bool
62+
Cancel context.CancelFunc
63+
SessionID string
64+
UserEntryID string
65+
Prompt string
66+
Images []imageAttachment
67+
UserMessageTimestamp int64
68+
ID uint64
69+
Canceled bool
6970
}
7071

7172
type resizeCoalescedEvent struct {
@@ -739,12 +740,7 @@ func (app *App) sessionMessages(ctx context.Context, sessionID string) ([]databa
739740
func (app *App) appendSessionMessages(messages []database.SessionMessageEntity) {
740741
for index := range messages {
741742
message := &messages[index]
742-
app.appendMessage(chatMessage{
743-
CreatedAt: message.CreatedAt,
744-
Role: transcript.FromDatabaseRole(message.Role),
745-
Content: message.Content,
746-
Attachments: databaseAttachmentSummaries(message.Parts),
747-
})
743+
app.appendMessage(chatMessageFromSessionMessage(message))
748744

749745
if message.Role == database.RoleUser {
750746
app.recordPromptDraftHistory(promptDraft{
@@ -754,6 +750,21 @@ func (app *App) appendSessionMessages(messages []database.SessionMessageEntity)
754750
}
755751
}
756752

753+
func chatMessageFromSessionMessage(message *database.SessionMessageEntity) chatMessage {
754+
chat := chatMessage{
755+
CreatedAt: message.CreatedAt,
756+
EntryID: nil,
757+
Role: transcript.FromDatabaseRole(message.Role),
758+
Content: message.Content,
759+
Attachments: databaseAttachmentSummaries(message.Parts),
760+
}
761+
if message.EntryID != "" {
762+
chat.EntryID = cloneStringPtr(&message.EntryID)
763+
}
764+
765+
return chat
766+
}
767+
757768
func (app *App) addSystemMessage(content string) {
758769
app.addMessage(transcript.RoleCustom, content)
759770
}
@@ -763,7 +774,13 @@ func (app *App) addMessage(role transcript.Role, content string) {
763774
}
764775

765776
func newChatMessage(role transcript.Role, content string) chatMessage {
766-
return chatMessage{CreatedAt: time.Now().UTC(), Role: role, Content: content, Attachments: nil}
777+
return chatMessage{
778+
Attachments: nil,
779+
CreatedAt: time.Now().UTC(),
780+
EntryID: nil,
781+
Role: role,
782+
Content: content,
783+
}
767784
}
768785

769786
func emptyCachedRenderedMessage() cachedRenderedMessage {

internal/terminal/async_events.go

Lines changed: 19 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -624,12 +624,31 @@ func (app *App) applyPromptUserEntry(_ context.Context, sessionID, entryID strin
624624
previousSessionID := app.activePrompt.SessionID
625625
app.activePrompt.SessionID = sessionID
626626
app.activePrompt.UserEntryID = entryID
627+
app.bindPromptUserMessageEntryID(entryID)
627628

628629
if app.sessionID == previousSessionID {
629630
app.sessionID = sessionID
630631
}
631632
}
632633

634+
func (app *App) bindPromptUserMessageEntryID(entryID string) {
635+
if entryID == "" || app.activePrompt.UserMessageTimestamp == 0 {
636+
return
637+
}
638+
639+
for index := range app.transcript.History {
640+
message := &app.transcript.History[index]
641+
642+
isPromptUserMessage := message.CreatedAt.UnixNano() == app.activePrompt.UserMessageTimestamp &&
643+
message.Role == transcript.RoleUser
644+
if isPromptUserMessage {
645+
message.EntryID = &entryID
646+
647+
return
648+
}
649+
}
650+
}
651+
633652
func (app *App) applyPromptError(ctx context.Context, message string, promptID uint64) {
634653
streamingBlocks := append([]chatMessage(nil), app.transcript.Streaming.Blocks...)
635654
canceled := app.activePrompt != nil && app.activePrompt.ID == promptID && app.activePrompt.Canceled

internal/terminal/async_events_internal_test.go

Lines changed: 28 additions & 17 deletions
Original file line numberDiff line numberDiff line change
@@ -516,6 +516,30 @@ func promptCompactionLifecycleEventCases() []promptLifecycleCase {
516516
}
517517
}
518518

519+
func promptUserEntryLifecycleCase() promptLifecycleCase {
520+
return promptLifecycleCase{
521+
name: "prompt user entry",
522+
payload: asyncTestEvent(asyncEventPromptUserEntry, asyncTestSessionID, asyncTestEntryID, 3),
523+
setup: func(app *App) {
524+
app.activePrompt = newTestActivePrompt(nil)
525+
app.activePrompt.ID = 3
526+
message := newChatMessage(transcript.RoleUser, app.activePrompt.Prompt)
527+
app.activePrompt.UserMessageTimestamp = message.CreatedAt.UnixNano()
528+
app.appendMessage(message)
529+
},
530+
assert: func(t *testing.T, app *App) {
531+
t.Helper()
532+
require.NotNil(t, app.activePrompt)
533+
assert.Equal(t, asyncTestSessionID, app.activePrompt.SessionID)
534+
assert.Equal(t, asyncTestEntryID, app.activePrompt.UserEntryID)
535+
require.Len(t, app.transcript.History, 1)
536+
require.NotNil(t, app.transcript.History[0].EntryID)
537+
assert.Equal(t, asyncTestEntryID, *app.transcript.History[0].EntryID)
538+
},
539+
wantHandled: true,
540+
}
541+
}
542+
519543
func promptLifecycleEventCases() []promptLifecycleCase {
520544
return append(promptCompactionLifecycleEventCases(), []promptLifecycleCase{
521545
{
@@ -525,10 +549,11 @@ func promptLifecycleEventCases() []promptLifecycleCase {
525549
app.streamingText = asyncTestPartial
526550
app.streamingThinkingText = "thought"
527551
app.transcript.Streaming.Blocks = []chatMessage{{
552+
Attachments: nil,
553+
CreatedAt: time.Time{},
554+
EntryID: nil,
528555
Role: transcript.RoleAssistant,
529556
Content: asyncTestPartial,
530-
CreatedAt: time.Time{},
531-
Attachments: nil,
532557
}}
533558
app.runningToolBlocks = []runningToolBlock{{
534559
StartedAt: time.Time{},
@@ -554,21 +579,7 @@ func promptLifecycleEventCases() []promptLifecycleCase {
554579
},
555580
wantHandled: true,
556581
},
557-
{
558-
name: "prompt user entry",
559-
payload: asyncTestEvent(asyncEventPromptUserEntry, asyncTestSessionID, asyncTestEntryID, 3),
560-
setup: func(app *App) {
561-
app.activePrompt = newTestActivePrompt(nil)
562-
app.activePrompt.ID = 3
563-
},
564-
assert: func(t *testing.T, app *App) {
565-
t.Helper()
566-
require.NotNil(t, app.activePrompt)
567-
assert.Equal(t, asyncTestSessionID, app.activePrompt.SessionID)
568-
assert.Equal(t, asyncTestEntryID, app.activePrompt.UserEntryID)
569-
},
570-
wantHandled: true,
571-
},
582+
promptUserEntryLifecycleCase(),
572583
{
573584
name: "prompt error",
574585
payload: asyncTestEvent(asyncEventPromptError, "", "provider failed", 4),

internal/terminal/interrupt_internal_test.go

Lines changed: 8 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -183,13 +183,13 @@ func newInterruptTestApp(t *testing.T, cancel context.CancelFunc) *App {
183183

184184
func newTestActivePrompt(cancel context.CancelFunc) *activePromptState {
185185
return &activePromptState{
186-
Cancel: cancel,
187-
ParentEntryID: nil,
188-
SessionID: "",
189-
UserEntryID: "",
190-
Images: nil,
191-
Prompt: interruptTestPrompt,
192-
ID: 1,
193-
Canceled: false,
186+
Cancel: cancel,
187+
SessionID: "",
188+
UserEntryID: "",
189+
Images: nil,
190+
Prompt: interruptTestPrompt,
191+
ID: 1,
192+
UserMessageTimestamp: 0,
193+
Canceled: false,
194194
}
195195
}

0 commit comments

Comments
 (0)