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:
@@ -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
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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{
|
||||
|
||||
Reference in New Issue
Block a user