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
|
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{
|
|
||||||
InputTokens: int(event.Message.Usage.InputTokens),
|
|
||||||
OutputTokens: int(event.Message.Usage.OutputTokens),
|
|
||||||
},
|
|
||||||
}, true
|
|
||||||
}
|
}
|
||||||
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:
|
default:
|
||||||
return llm.ChatCompletionStreamEvent{}, false
|
return llm.ChatCompletionStreamEvent{}, false
|
||||||
|
|||||||
@@ -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
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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 {
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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(
|
||||||
URL: p.URL,
|
parts,
|
||||||
}))
|
openai.ImageContentPart(openai.ChatCompletionContentPartImageImageURLParam{
|
||||||
|
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{
|
||||||
|
|||||||
Reference in New Issue
Block a user