Here's a commit message following the repo's style:

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 <bryan@getprobo.com>
This commit is contained in:
Bryan Frimin
2026-03-13 19:22:46 +01:00
parent 2df9a3fe59
commit da9a594f79
2 changed files with 69 additions and 2 deletions

View File

@@ -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) {

View File

@@ -53,9 +53,8 @@ const (
)
func (sr *StreamedRun) Wait() (*Result, error) {
for range sr.Events {
}
<-sr.done
return sr.result, sr.err
}