From da9a594f7945b08874837b52bd4aa83cbee4ad4d Mon Sep 17 00:00:00 2001 From: Bryan Frimin Date: Fri, 13 Mar 2026 19:22:46 +0100 Subject: [PATCH] Here's a commit message following the repo's style: MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Fix Wait draining events from concurrent consumers Wait() was ranging over the public Events channel, competing with any concurrent reader for events. Callers that streamed events in one goroutine and called Wait() in another would lose an arbitrary subset of events. Wait now only blocks on the done channel; the result fields are already visible thanks to the close ordering (set fields → close events → close done). Signed-off-by: Bryan Frimin --- pkg/agent/agent_test.go | 68 +++++++++++++++++++++++++++++++++++++++++ pkg/agent/stream.go | 3 +- 2 files changed, 69 insertions(+), 2 deletions(-) diff --git a/pkg/agent/agent_test.go b/pkg/agent/agent_test.go index 2c2df7dec..2842595b8 100644 --- a/pkg/agent/agent_test.go +++ b/pkg/agent/agent_test.go @@ -2270,6 +2270,74 @@ func TestRunStreamed(t *testing.T) { assert.True(t, gotError, "StreamEventError should be emitted when session save fails") }, ) + + t.Run( + "concurrent consumer receives all events before Wait returns", + func(t *testing.T) { + t.Parallel() + + mockStream := &mockChatStream{ + events: []llm.ChatCompletionStreamEvent{ + {Delta: llm.MessageDelta{Content: "Hello"}}, + {Delta: llm.MessageDelta{Content: " world"}}, + { + Delta: llm.MessageDelta{Content: "!"}, + Usage: &llm.Usage{InputTokens: 10, OutputTokens: 3}, + FinishReason: finishReasonPtr(llm.FinishReasonStop), + }, + }, + } + + streamProvider := &mockStreamProvider{stream: mockStream} + client := llm.NewClient(streamProvider, "test") + + ag := agent.New( + "assistant", + client, + agent.WithModel("test-model"), + ) + + sr := ag.RunStreamed( + context.Background(), + []llm.Message{userMessage("Hi")}, + ) + + var collected []agent.StreamEvent + done := make(chan struct{}) + go func() { + defer close(done) + for ev := range sr.Events { + collected = append(collected, ev) + } + }() + + result, err := sr.Wait() + <-done + + require.NoError(t, err) + assert.Equal(t, "Hello world!", result.FinalMessage().Text()) + + var deltaCount int + var gotAgentStart, gotAgentEnd, gotComplete bool + for _, ev := range collected { + switch ev.Type { + case agent.StreamEventLLMDelta: + deltaCount++ + case agent.StreamEventAgentStart: + gotAgentStart = true + case agent.StreamEventAgentEnd: + gotAgentEnd = true + case agent.StreamEventComplete: + gotComplete = true + } + } + + assert.Equal(t, 3, deltaCount, "all delta events should reach the consumer") + assert.True(t, gotAgentStart, "agent_start event should reach the consumer") + assert.True(t, gotAgentEnd, "agent_end event should reach the consumer") + assert.True(t, gotComplete, "complete event should reach the consumer") + }, + ) } func TestClone(t *testing.T) { diff --git a/pkg/agent/stream.go b/pkg/agent/stream.go index 2899bb041..b77171bf6 100644 --- a/pkg/agent/stream.go +++ b/pkg/agent/stream.go @@ -53,9 +53,8 @@ const ( ) func (sr *StreamedRun) Wait() (*Result, error) { - for range sr.Events { - } <-sr.done + return sr.result, sr.err }