diff --git a/backend/internal/service/openai_gateway_response_handling.go b/backend/internal/service/openai_gateway_response_handling.go index 0b26a580cf..a60e37feb5 100644 --- a/backend/internal/service/openai_gateway_response_handling.go +++ b/backend/internal/service/openai_gateway_response_handling.go @@ -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) diff --git a/backend/internal/service/openai_responses_stream_output_items_test.go b/backend/internal/service/openai_responses_stream_output_items_test.go new file mode 100644 index 0000000000..9cfe3f3d78 --- /dev/null +++ b/backend/internal/service/openai_responses_stream_output_items_test.go @@ -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)) +}