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 <bryan@getprobo.com>
This commit is contained in:
Bryan Frimin
2026-03-13 19:17:17 +01:00
parent 9e842b8ef3
commit 2df9a3fe59
5 changed files with 35 additions and 15 deletions

View File

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

View File

@@ -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
}

View File

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

View File

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

View File

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