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 return evt, true
case "message_start": case "message_start":
if event.Message.Usage.InputTokens > 0 { evt := llm.ChatCompletionStreamEvent{
return llm.ChatCompletionStreamEvent{ Model: string(event.Message.Model),
Usage: &llm.Usage{ }
if event.Message.Usage.InputTokens > 0 || event.Message.Usage.OutputTokens > 0 {
evt.Usage = &llm.Usage{
InputTokens: int(event.Message.Usage.InputTokens), InputTokens: int(event.Message.Usage.InputTokens),
OutputTokens: int(event.Message.Usage.OutputTokens), OutputTokens: int(event.Message.Usage.OutputTokens),
},
}, true
} }
return llm.ChatCompletionStreamEvent{}, false }
return evt, true
default: default:
return llm.ChatCompletionStreamEvent{}, false 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 nil, mapError(err)
} }
return newBedrockStream(output.GetStream()), nil return newBedrockStream(output.GetStream(), req.Model), nil
} }
func buildInput(req *llm.ChatCompletionRequest) *bedrockruntime.ConverseInput { func buildInput(req *llm.ChatCompletionRequest) *bedrockruntime.ConverseInput {
@@ -346,14 +346,17 @@ type bedrockStream struct {
events <-chan types.ConverseStreamOutput events <-chan types.ConverseStreamOutput
current llm.ChatCompletionStreamEvent current llm.ChatCompletionStreamEvent
err error err error
model string
modelSent bool
toolIndex int toolIndex int
inToolUse bool inToolUse bool
} }
func newBedrockStream(eventStream *bedrockruntime.ConverseStreamEventStream) *bedrockStream { func newBedrockStream(eventStream *bedrockruntime.ConverseStreamEventStream, model string) *bedrockStream {
return &bedrockStream{ return &bedrockStream{
eventStream: eventStream, eventStream: eventStream,
events: eventStream.Events(), events: eventStream.Events(),
model: model,
} }
} }
@@ -361,6 +364,10 @@ func (s *bedrockStream) Next() bool {
for event := range s.events { for event := range s.events {
mapped, ok := s.mapEvent(event) mapped, ok := s.mapEvent(event)
if ok { if ok {
if !s.modelSent {
mapped.Model = s.model
s.modelSent = true
}
s.current = mapped s.current = mapped
return true return true
} }

View File

@@ -90,6 +90,7 @@ type (
} }
ChatCompletionStreamEvent struct { ChatCompletionStreamEvent struct {
Model string // present on first event if provider supports it
Delta MessageDelta Delta MessageDelta
Usage *Usage // present on final event if provider supports it Usage *Usage // present on final event if provider supports it
FinishReason *FinishReason // present on final event FinishReason *FinishReason // present on final event
@@ -206,6 +207,10 @@ func (a *StreamAccumulator) Response() *ChatCompletionResponse {
} }
func (a *StreamAccumulator) accumulate(event ChatCompletionStreamEvent) { func (a *StreamAccumulator) accumulate(event ChatCompletionStreamEvent) {
if a.model == "" && event.Model != "" {
a.model = event.Model
}
a.content.WriteString(event.Delta.Content) a.content.WriteString(event.Delta.Content)
for _, tcd := range event.Delta.ToolCalls { for _, tcd := range event.Delta.ToolCalls {

View File

@@ -508,7 +508,7 @@ func TestStreamAccumulator(t *testing.T) {
t.Parallel() t.Parallel()
events := []llm.ChatCompletionStreamEvent{ events := []llm.ChatCompletionStreamEvent{
{Delta: llm.MessageDelta{Content: "Hello"}}, {Model: "gpt-4o", Delta: llm.MessageDelta{Content: "Hello"}},
{Delta: llm.MessageDelta{Content: " world"}}, {Delta: llm.MessageDelta{Content: " world"}},
{Delta: llm.MessageDelta{ {Delta: llm.MessageDelta{
ToolCalls: []llm.ToolCallDelta{ ToolCalls: []llm.ToolCallDelta{
@@ -537,6 +537,7 @@ func TestStreamAccumulator(t *testing.T) {
require.NoError(t, acc.Err()) require.NoError(t, acc.Err())
resp := acc.Response() resp := acc.Response()
assert.Equal(t, "gpt-4o", resp.Model)
assert.Equal(t, "Hello world", resp.Message.Text()) assert.Equal(t, "Hello world", resp.Message.Text())
assert.Equal(t, llm.RoleAssistant, resp.Message.Role) assert.Equal(t, llm.RoleAssistant, resp.Message.Role)
assert.Equal(t, llm.FinishReasonToolCalls, resp.FinishReason) assert.Equal(t, llm.FinishReasonToolCalls, resp.FinishReason)

View File

@@ -181,9 +181,13 @@ func buildMessages(messages []llm.Message) []openai.ChatCompletionMessageParamUn
case llm.TextPart: case llm.TextPart:
parts = append(parts, openai.TextContentPart(p.Text)) parts = append(parts, openai.TextContentPart(p.Text))
case llm.ImagePart: case llm.ImagePart:
parts = append(parts, openai.ImageContentPart(openai.ChatCompletionContentPartImageImageURLParam{ parts = append(
parts,
openai.ImageContentPart(openai.ChatCompletionContentPartImageImageURLParam{
URL: p.URL, URL: p.URL,
})) },
),
)
} }
} }
out = append(out, openai.UserMessage(parts)) out = append(out, openai.UserMessage(parts))
@@ -408,7 +412,9 @@ func (s *openaiStream) Close() error {
} }
func mapChunkToEvent(chunk *openai.ChatCompletionChunk) llm.ChatCompletionStreamEvent { 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 { if chunk.Usage.PromptTokens > 0 || chunk.Usage.CompletionTokens > 0 {
usage := llm.Usage{ usage := llm.Usage{