mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-10-06 15:03:54 +08:00
Merge pull request #6118 from akihitohyh/fix/terminal-output-item-preservation
fix(openai): rebuild streaming terminal output from the reported items
This commit is contained in:
@@ -8,6 +8,7 @@ import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"sort"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync/atomic"
|
||||
@@ -336,6 +337,7 @@ func (s *OpenAIGatewayService) handleStreamingResponseWithReasoning(ctx context.
|
||||
|
||||
needModelReplace := originalModel != mappedModel
|
||||
streamOutputAccumulator := apicompat.NewBufferedResponseAccumulator()
|
||||
streamDoneItems := newResponsesStreamOutputItems()
|
||||
streamImageOutputs := make([]json.RawMessage, 0, 1)
|
||||
streamSeenImages := make(map[string]struct{})
|
||||
searchCounter := 0
|
||||
@@ -614,13 +616,14 @@ func (s *OpenAIGatewayService) handleStreamingResponseWithReasoning(ctx context.
|
||||
if imageOutput, ok := extractImageGenerationOutputFromSSEData(dataBytes, streamSeenImages); ok {
|
||||
streamImageOutputs = append(streamImageOutputs, imageOutput)
|
||||
}
|
||||
streamDoneItems.Observe(dataBytes)
|
||||
if responsesStreamEventMayContributeToOutput(eventType) {
|
||||
var streamEvent apicompat.ResponsesStreamEvent
|
||||
if err := json.Unmarshal(dataBytes, &streamEvent); err == nil {
|
||||
streamOutputAccumulator.ProcessEvent(&streamEvent)
|
||||
}
|
||||
}
|
||||
if normalizedData, normalized := normalizeResponsesStreamingTerminalOutput(dataBytes, streamOutputAccumulator, streamImageOutputs); normalized {
|
||||
if normalizedData, normalized := normalizeResponsesStreamingTerminalOutput(dataBytes, streamOutputAccumulator, streamDoneItems, streamImageOutputs); normalized {
|
||||
dataBytes = normalizedData
|
||||
data = string(normalizedData)
|
||||
line = "data: " + data
|
||||
@@ -1989,7 +1992,75 @@ func normalizeCompletedImageGenerationStatus(data []byte) ([]byte, bool) {
|
||||
}
|
||||
}
|
||||
|
||||
func normalizeResponsesStreamingTerminalOutput(data []byte, acc *apicompat.BufferedResponseAccumulator, imageOutputs []json.RawMessage) ([]byte, bool) {
|
||||
// responsesStreamOutputItems remembers the raw item carried by each
|
||||
// response.output_item.done event, keyed by output_index.
|
||||
//
|
||||
// reconstructResponseOutputFromSSE already prefers the raw done items over
|
||||
// delta accumulation when it rebuilds a buffered response, because the
|
||||
// accumulator models only "one reasoning, one message, N function calls" and
|
||||
// therefore cannot preserve item identity, per-item status/phase, ordering, or
|
||||
// item types it does not know about. The streaming path had no equivalent
|
||||
// because it never sees the whole body at once; this collector gives it one.
|
||||
type responsesStreamOutputItems struct {
|
||||
items map[int]json.RawMessage
|
||||
}
|
||||
|
||||
func newResponsesStreamOutputItems() *responsesStreamOutputItems {
|
||||
return &responsesStreamOutputItems{items: make(map[int]json.RawMessage)}
|
||||
}
|
||||
|
||||
// Observe records the item of a response.output_item.done event verbatim. The
|
||||
// raw JSON is kept byte for byte so vendor extensions and future fields survive
|
||||
// the rebuild.
|
||||
func (r *responsesStreamOutputItems) Observe(data []byte) {
|
||||
if r == nil || len(data) == 0 || !gjson.ValidBytes(data) {
|
||||
return
|
||||
}
|
||||
if strings.TrimSpace(gjson.GetBytes(data, "type").String()) != "response.output_item.done" {
|
||||
return
|
||||
}
|
||||
item := gjson.GetBytes(data, "item")
|
||||
if !item.Exists() || !item.IsObject() {
|
||||
return
|
||||
}
|
||||
index := int(gjson.GetBytes(data, "output_index").Int())
|
||||
r.items[index] = json.RawMessage(append([]byte(nil), item.Raw...))
|
||||
}
|
||||
|
||||
func (r *responsesStreamOutputItems) HasItems() bool {
|
||||
return r != nil && len(r.items) > 0
|
||||
}
|
||||
|
||||
// Count reports how many distinct output items the stream reported as done.
|
||||
func (r *responsesStreamOutputItems) Count() int {
|
||||
if r == nil {
|
||||
return 0
|
||||
}
|
||||
return len(r.items)
|
||||
}
|
||||
|
||||
// BuildOutput returns the remembered items ordered by output_index.
|
||||
func (r *responsesStreamOutputItems) BuildOutput() ([]byte, bool) {
|
||||
if !r.HasItems() {
|
||||
return nil, false
|
||||
}
|
||||
indexes := make([]int, 0, len(r.items))
|
||||
for index := range r.items {
|
||||
indexes = append(indexes, index)
|
||||
}
|
||||
sort.Ints(indexes)
|
||||
ordered := make([]json.RawMessage, 0, len(indexes))
|
||||
for _, index := range indexes {
|
||||
ordered = append(ordered, r.items[index])
|
||||
}
|
||||
encoded, err := json.Marshal(ordered)
|
||||
if err != nil {
|
||||
return nil, false
|
||||
}
|
||||
return encoded, true
|
||||
}
|
||||
|
||||
func normalizeResponsesStreamingTerminalOutput(data []byte, acc *apicompat.BufferedResponseAccumulator, doneItems *responsesStreamOutputItems, imageOutputs []json.RawMessage) ([]byte, bool) {
|
||||
eventType := strings.TrimSpace(gjson.GetBytes(data, "type").String())
|
||||
switch eventType {
|
||||
case "response.completed", "response.done", "response.incomplete", "response.cancelled", "response.canceled":
|
||||
@@ -1998,15 +2069,28 @@ func normalizeResponsesStreamingTerminalOutput(data []byte, acc *apicompat.Buffe
|
||||
}
|
||||
|
||||
output := gjson.GetBytes(data, "response.output")
|
||||
hasAccumulatedOutput := (acc != nil && acc.HasContent()) || len(imageOutputs) > 0
|
||||
hasAccumulatedOutput := (acc != nil && acc.HasContent()) || len(imageOutputs) > 0 || doneItems.HasItems()
|
||||
if output.Exists() && output.IsArray() {
|
||||
if len(output.Array()) > 0 || !hasAccumulatedOutput {
|
||||
terminalCount := len(output.Array())
|
||||
// A terminal output carrying at least as many items as the stream
|
||||
// reported is left untouched. Carrying fewer means the terminal
|
||||
// dropped items the stream already reported as done, and those
|
||||
// reported items are the authoritative record of the turn.
|
||||
if terminalCount > 0 && terminalCount >= doneItems.Count() {
|
||||
return data, false
|
||||
}
|
||||
if terminalCount == 0 && !hasAccumulatedOutput {
|
||||
return data, false
|
||||
}
|
||||
}
|
||||
|
||||
outputJSON := []byte("[]")
|
||||
if reconstructed, ok := buildResponsesOutputJSON(acc, imageOutputs); ok {
|
||||
// Same precedence as reconstructResponseOutputFromSSE: the items the stream
|
||||
// actually reported win over anything rebuilt from deltas. Image generation
|
||||
// items arrive as done events too, so imageOutputs would duplicate them here.
|
||||
if reconstructed, ok := doneItems.BuildOutput(); ok {
|
||||
outputJSON = reconstructed
|
||||
} else if reconstructed, ok := buildResponsesOutputJSON(acc, imageOutputs); ok {
|
||||
outputJSON = reconstructed
|
||||
}
|
||||
updated, err := sjson.SetRawBytes(data, "response.output", outputJSON)
|
||||
|
||||
@@ -0,0 +1,117 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
"github.com/tidwall/gjson"
|
||||
)
|
||||
|
||||
// A terminal event that arrives with an empty output must be rebuilt from the
|
||||
// items the stream reported, not from delta accumulation. The accumulator
|
||||
// models only one reasoning and one message, so rebuilding through it collapses
|
||||
// a multi-item turn into a single fabricated message.
|
||||
func TestNormalizeResponsesStreamingTerminalOutputPreservesReportedItems(t *testing.T) {
|
||||
doneItems := newResponsesStreamOutputItems()
|
||||
|
||||
doneItems.Observe([]byte(`{
|
||||
"type":"response.output_item.done",
|
||||
"output_index":0,
|
||||
"item":{"id":"rs_1","type":"reasoning","summary":[{"type":"summary_text","text":"thinking"}],"encrypted_content":"opaque"}
|
||||
}`))
|
||||
doneItems.Observe([]byte(`{
|
||||
"type":"response.output_item.done",
|
||||
"output_index":1,
|
||||
"item":{"id":"msg_1","type":"message","status":"completed","phase":"final_answer","role":"assistant","content":[{"type":"output_text","text":"shipped","annotations":[],"logprobs":[]}]}
|
||||
}`))
|
||||
|
||||
normalized, changed := normalizeResponsesStreamingTerminalOutput(
|
||||
[]byte(`{"type":"response.completed","response":{"status":"completed","output":[]}}`),
|
||||
nil,
|
||||
doneItems,
|
||||
nil,
|
||||
)
|
||||
require.True(t, changed)
|
||||
|
||||
output := gjson.GetBytes(normalized, "response.output")
|
||||
require.True(t, output.IsArray())
|
||||
require.Len(t, output.Array(), 2, "both reported items must survive")
|
||||
|
||||
require.Equal(t, "reasoning", gjson.GetBytes(normalized, "response.output.0.type").String())
|
||||
require.Equal(t, "rs_1", gjson.GetBytes(normalized, "response.output.0.id").String())
|
||||
require.Equal(t, "opaque", gjson.GetBytes(normalized, "response.output.0.encrypted_content").String(),
|
||||
"fields the gateway does not model must survive verbatim")
|
||||
|
||||
require.Equal(t, "message", gjson.GetBytes(normalized, "response.output.1.type").String())
|
||||
require.Equal(t, "msg_1", gjson.GetBytes(normalized, "response.output.1.id").String(),
|
||||
"the reported id must be reused, not regenerated")
|
||||
require.Equal(t, "completed", gjson.GetBytes(normalized, "response.output.1.status").String())
|
||||
require.Equal(t, "final_answer", gjson.GetBytes(normalized, "response.output.1.phase").String())
|
||||
require.Equal(t, "shipped", gjson.GetBytes(normalized, "response.output.1.content.0.text").String())
|
||||
}
|
||||
|
||||
// Items are ordered by output_index, not by arrival order.
|
||||
func TestResponsesStreamOutputItemsOrderByOutputIndex(t *testing.T) {
|
||||
doneItems := newResponsesStreamOutputItems()
|
||||
doneItems.Observe([]byte(`{"type":"response.output_item.done","output_index":2,"item":{"id":"c","type":"message"}}`))
|
||||
doneItems.Observe([]byte(`{"type":"response.output_item.done","output_index":0,"item":{"id":"a","type":"reasoning"}}`))
|
||||
|
||||
built, ok := doneItems.BuildOutput()
|
||||
require.True(t, ok)
|
||||
require.Equal(t, "a", gjson.GetBytes(built, "0.id").String())
|
||||
require.Equal(t, "c", gjson.GetBytes(built, "1.id").String())
|
||||
}
|
||||
|
||||
// A stream that never reports a done item keeps the previous rebuild path.
|
||||
func TestNormalizeResponsesStreamingTerminalOutputIgnoresNonDoneEvents(t *testing.T) {
|
||||
doneItems := newResponsesStreamOutputItems()
|
||||
doneItems.Observe([]byte(`{"type":"response.output_item.added","output_index":0,"item":{"id":"msg_1","type":"message"}}`))
|
||||
doneItems.Observe([]byte(`{"type":"response.output_text.delta","output_index":0,"delta":"hi"}`))
|
||||
require.False(t, doneItems.HasItems())
|
||||
|
||||
raw := []byte(`{"type":"response.completed","response":{"status":"completed","output":[]}}`)
|
||||
normalized, changed := normalizeResponsesStreamingTerminalOutput(raw, nil, doneItems, nil)
|
||||
require.False(t, changed)
|
||||
require.Equal(t, string(raw), string(normalized))
|
||||
}
|
||||
|
||||
// The terminal event can arrive with a non-empty but truncated output: the
|
||||
// stream reported two items, the terminal carries one, and its id was not the
|
||||
// one the stream reported. The reported items win.
|
||||
func TestNormalizeResponsesStreamingTerminalOutputRepairsTruncatedOutput(t *testing.T) {
|
||||
doneItems := newResponsesStreamOutputItems()
|
||||
doneItems.Observe([]byte(`{
|
||||
"type":"response.output_item.done","output_index":0,
|
||||
"item":{"id":"rs_real","type":"reasoning","status":"in_progress","summary":[]}
|
||||
}`))
|
||||
doneItems.Observe([]byte(`{
|
||||
"type":"response.output_item.done","output_index":1,
|
||||
"item":{"id":"msg_real","type":"message","status":"completed","role":"assistant","content":[{"type":"output_text","text":"shipped","annotations":[],"logprobs":[]}]}
|
||||
}`))
|
||||
|
||||
normalized, changed := normalizeResponsesStreamingTerminalOutput([]byte(`{
|
||||
"type":"response.completed",
|
||||
"response":{"status":"completed","output":[{"type":"message","role":"assistant","id":"msg_fabricated","status":"completed","content":[{"type":"output_text","text":"shipped","annotations":[],"logprobs":[]}]}]}
|
||||
}`), nil, doneItems, nil)
|
||||
require.True(t, changed)
|
||||
|
||||
require.Len(t, gjson.GetBytes(normalized, "response.output").Array(), 2)
|
||||
require.Equal(t, "reasoning", gjson.GetBytes(normalized, "response.output.0.type").String())
|
||||
require.Equal(t, "rs_real", gjson.GetBytes(normalized, "response.output.0.id").String())
|
||||
require.Equal(t, "msg_real", gjson.GetBytes(normalized, "response.output.1.id").String(),
|
||||
"the id the stream reported must replace the fabricated one")
|
||||
}
|
||||
|
||||
// A terminal output that is already complete is never rewritten.
|
||||
func TestNormalizeResponsesStreamingTerminalOutputLeavesCompleteOutputAlone(t *testing.T) {
|
||||
doneItems := newResponsesStreamOutputItems()
|
||||
doneItems.Observe([]byte(`{
|
||||
"type":"response.output_item.done","output_index":0,
|
||||
"item":{"id":"msg_real","type":"message","status":"completed"}
|
||||
}`))
|
||||
|
||||
raw := []byte(`{"type":"response.completed","response":{"status":"completed","output":[{"type":"message","id":"msg_upstream","status":"completed","vendor":"keep"}]}}`)
|
||||
normalized, changed := normalizeResponsesStreamingTerminalOutput(raw, nil, doneItems, nil)
|
||||
require.False(t, changed)
|
||||
require.Equal(t, string(raw), string(normalized))
|
||||
}
|
||||
Reference in New Issue
Block a user