From 2df9a3fe59e3a1de6bb1592e0e5bafbc451bcbf2 Mon Sep 17 00:00:00 2001 From: Bryan Frimin Date: Fri, 13 Mar 2026 19:17:17 +0100 Subject: [PATCH] Propagate model through streamed responses StreamAccumulator never set its model field, so Response().Model was always empty for streamed completions. Add a Model field to stream events and populate it in all three providers (OpenAI, Anthropic, Bedrock). Signed-off-by: Bryan Frimin --- pkg/llm/anthropic/provider.go | 17 +++++++++-------- pkg/llm/bedrock/provider.go | 11 +++++++++-- pkg/llm/chat.go | 5 +++++ pkg/llm/llm_test.go | 3 ++- pkg/llm/openai/provider.go | 14 ++++++++++---- 5 files changed, 35 insertions(+), 15 deletions(-) diff --git a/pkg/llm/anthropic/provider.go b/pkg/llm/anthropic/provider.go index 7ae9b11dc..32f62dfda 100644 --- a/pkg/llm/anthropic/provider.go +++ b/pkg/llm/anthropic/provider.go @@ -431,15 +431,16 @@ func (s *anthropicStream) mapStreamEvent(event *anthropic.MessageStreamEventUnio return evt, true case "message_start": - if event.Message.Usage.InputTokens > 0 { - return llm.ChatCompletionStreamEvent{ - Usage: &llm.Usage{ - InputTokens: int(event.Message.Usage.InputTokens), - OutputTokens: int(event.Message.Usage.OutputTokens), - }, - }, true + evt := llm.ChatCompletionStreamEvent{ + Model: string(event.Message.Model), } - return llm.ChatCompletionStreamEvent{}, false + if event.Message.Usage.InputTokens > 0 || event.Message.Usage.OutputTokens > 0 { + evt.Usage = &llm.Usage{ + InputTokens: int(event.Message.Usage.InputTokens), + OutputTokens: int(event.Message.Usage.OutputTokens), + } + } + return evt, true default: return llm.ChatCompletionStreamEvent{}, false diff --git a/pkg/llm/bedrock/provider.go b/pkg/llm/bedrock/provider.go index 0ae9aee18..2b14d9137 100644 --- a/pkg/llm/bedrock/provider.go +++ b/pkg/llm/bedrock/provider.go @@ -83,7 +83,7 @@ func (p *Provider) ChatCompletionStream(ctx context.Context, req *llm.ChatComple return nil, mapError(err) } - return newBedrockStream(output.GetStream()), nil + return newBedrockStream(output.GetStream(), req.Model), nil } func buildInput(req *llm.ChatCompletionRequest) *bedrockruntime.ConverseInput { @@ -346,14 +346,17 @@ type bedrockStream struct { events <-chan types.ConverseStreamOutput current llm.ChatCompletionStreamEvent err error + model string + modelSent bool toolIndex int inToolUse bool } -func newBedrockStream(eventStream *bedrockruntime.ConverseStreamEventStream) *bedrockStream { +func newBedrockStream(eventStream *bedrockruntime.ConverseStreamEventStream, model string) *bedrockStream { return &bedrockStream{ eventStream: eventStream, events: eventStream.Events(), + model: model, } } @@ -361,6 +364,10 @@ func (s *bedrockStream) Next() bool { for event := range s.events { mapped, ok := s.mapEvent(event) if ok { + if !s.modelSent { + mapped.Model = s.model + s.modelSent = true + } s.current = mapped return true } diff --git a/pkg/llm/chat.go b/pkg/llm/chat.go index 1eb85c178..931ff98b0 100644 --- a/pkg/llm/chat.go +++ b/pkg/llm/chat.go @@ -90,6 +90,7 @@ type ( } ChatCompletionStreamEvent struct { + Model string // present on first event if provider supports it Delta MessageDelta Usage *Usage // present on final event if provider supports it FinishReason *FinishReason // present on final event @@ -206,6 +207,10 @@ func (a *StreamAccumulator) Response() *ChatCompletionResponse { } func (a *StreamAccumulator) accumulate(event ChatCompletionStreamEvent) { + if a.model == "" && event.Model != "" { + a.model = event.Model + } + a.content.WriteString(event.Delta.Content) for _, tcd := range event.Delta.ToolCalls { diff --git a/pkg/llm/llm_test.go b/pkg/llm/llm_test.go index 5ab961aca..d2f64b7ba 100644 --- a/pkg/llm/llm_test.go +++ b/pkg/llm/llm_test.go @@ -508,7 +508,7 @@ func TestStreamAccumulator(t *testing.T) { t.Parallel() events := []llm.ChatCompletionStreamEvent{ - {Delta: llm.MessageDelta{Content: "Hello"}}, + {Model: "gpt-4o", Delta: llm.MessageDelta{Content: "Hello"}}, {Delta: llm.MessageDelta{Content: " world"}}, {Delta: llm.MessageDelta{ ToolCalls: []llm.ToolCallDelta{ @@ -537,6 +537,7 @@ func TestStreamAccumulator(t *testing.T) { require.NoError(t, acc.Err()) resp := acc.Response() + assert.Equal(t, "gpt-4o", resp.Model) assert.Equal(t, "Hello world", resp.Message.Text()) assert.Equal(t, llm.RoleAssistant, resp.Message.Role) assert.Equal(t, llm.FinishReasonToolCalls, resp.FinishReason) diff --git a/pkg/llm/openai/provider.go b/pkg/llm/openai/provider.go index fbb26e5bc..41c488920 100644 --- a/pkg/llm/openai/provider.go +++ b/pkg/llm/openai/provider.go @@ -181,9 +181,13 @@ func buildMessages(messages []llm.Message) []openai.ChatCompletionMessageParamUn case llm.TextPart: parts = append(parts, openai.TextContentPart(p.Text)) case llm.ImagePart: - parts = append(parts, openai.ImageContentPart(openai.ChatCompletionContentPartImageImageURLParam{ - URL: p.URL, - })) + parts = append( + parts, + openai.ImageContentPart(openai.ChatCompletionContentPartImageImageURLParam{ + URL: p.URL, + }, + ), + ) } } out = append(out, openai.UserMessage(parts)) @@ -408,7 +412,9 @@ func (s *openaiStream) Close() error { } func mapChunkToEvent(chunk *openai.ChatCompletionChunk) llm.ChatCompletionStreamEvent { - event := llm.ChatCompletionStreamEvent{} + event := llm.ChatCompletionStreamEvent{ + Model: chunk.Model, + } if chunk.Usage.PromptTokens > 0 || chunk.Usage.CompletionTokens > 0 { usage := llm.Usage{