防止 WebSocket 生图结果停留在进行中状态

HTTP/SSE 已会在终态事件中修正带结果的图片 item,但 WebSocket 下行路径缺少同一处理。现在四条文本帧出口复用现有归一化,并原样保留 result 数据。

Constraint: 仅修改终态事件中带非空 result 的 image_generation_call
Rejected: 解码并改写二进制帧 | 二进制帧不属于 JSON 事件契约
Confidence: high
Scope-risk: narrow
Reversibility: clean
Directive: 保持 HTTP、SSE 与 WebSocket 的图片终态规则一致
Tested: 后端 unit、integration、go vet
Not-tested: 真实 Codex WebSocket 生图会话

Co-authored-by: OmX <omx@oh-my-codex.dev>
This commit is contained in:
viccy
2026-07-16 17:11:00 +08:00
co-authored by OmX
parent 43e95dfa45
commit 2b5944fd6c
7 changed files with 33 additions and 3 deletions
@@ -649,6 +649,12 @@ func TestNormalizeCompletedImageGenerationStatus(t *testing.T) {
want: `{"type":"response.output_item.added","item":{"type":"image_generation_call","status":"generating","result":"image-data"}}`,
wantChanged: false,
},
{
name: "done preserves base64 result",
input: `{"type":"response.done","response":{"output":[{"type":"image_generation_call","status":"generating","result":"iVBORw0KGgoAAAANSUhEUg/+=="}]}}`,
want: `{"type":"response.done","response":{"output":[{"type":"image_generation_call","status":"completed","result":"iVBORw0KGgoAAAANSUhEUg/+=="}]}}`,
wantChanged: true,
},
}
for _, tt := range tests {
@@ -803,6 +803,9 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
wroteDownstream,
)
}
if normalized, changed := normalizeCompletedImageGenerationStatus(upstreamMessage); changed {
upstreamMessage = normalized
}
eventType, eventResponseID, _ := parseOpenAIWSEventEnvelope(upstreamMessage)
if responseID == "" && eventResponseID != "" {
@@ -38,6 +38,7 @@ func TestOpenAIGatewayService_ProxyResponsesWebSocketFromClient_KeepLeaseAcrossT
captureConn := &openAIWSCaptureConn{
events: [][]byte{
[]byte(`{"type":"response.output_item.done","item":{"id":"ig_ingress_1","type":"image_generation_call","status":"generating","result":"iVBORw0KGgoAAAANSUhEUg/+=="}}`),
[]byte(`{"type":"response.completed","response":{"id":"resp_ingress_turn_1","model":"gpt-5.1","usage":{"input_tokens":1,"output_tokens":1}}}`),
[]byte(`{"type":"response.completed","response":{"id":"resp_ingress_turn_2","model":"gpt-5.1","usage":{"input_tokens":1,"output_tokens":1}}}`),
},
@@ -138,6 +139,10 @@ func TestOpenAIGatewayService_ProxyResponsesWebSocketFromClient_KeepLeaseAcrossT
}
writeMessage(`{"type":"response.create","model":"gpt-5.1","stream":false}`)
firstTurnImageEvent := readMessage()
require.Equal(t, "response.output_item.done", gjson.GetBytes(firstTurnImageEvent, "type").String())
require.Equal(t, "completed", gjson.GetBytes(firstTurnImageEvent, "item.status").String())
require.Equal(t, "iVBORw0KGgoAAAANSUhEUg/+==", gjson.GetBytes(firstTurnImageEvent, "item.result").String())
firstTurnEvent := readMessage()
require.Equal(t, "response.completed", gjson.GetBytes(firstTurnEvent, "type").String())
require.Equal(t, "resp_ingress_turn_1", gjson.GetBytes(firstTurnEvent, "response.id").String())
@@ -765,7 +770,7 @@ func TestOpenAIGatewayService_ProxyResponsesWebSocketFromClient_PassthroughModeR
upstreamConn := &openAIWSCaptureConn{
events: [][]byte{
[]byte(`{"type":"response.completed","response":{"id":"resp_passthrough_turn_1","model":"gpt-5.1","usage":{"input_tokens":2,"output_tokens":3}}}`),
[]byte(`{"type":"response.completed","response":{"id":"resp_passthrough_turn_1","model":"gpt-5.1","output":[{"id":"ig_passthrough_1","type":"image_generation_call","status":"generating","result":"final-image"}],"usage":{"input_tokens":2,"output_tokens":3}}}`),
},
}
captureDialer := &openAIWSCaptureDialer{conn: upstreamConn}
@@ -858,6 +863,7 @@ func TestOpenAIGatewayService_ProxyResponsesWebSocketFromClient_PassthroughModeR
require.NoError(t, readErr)
require.Equal(t, "response.completed", gjson.GetBytes(event, "type").String())
require.Equal(t, "resp_passthrough_turn_1", gjson.GetBytes(event, "response.id").String())
require.Equal(t, "completed", gjson.GetBytes(event, "response.output.0.status").String())
_ = clientConn.Close(coderws.StatusNormalClosure, "done")
select {
@@ -1047,7 +1053,7 @@ func TestOpenAIGatewayService_ProxyResponsesWebSocketFromClient_HTTPBridgeModeRe
Header: http.Header{"Content-Type": []string{"text/event-stream"}, "x-request-id": []string{"rid_bridge_1"}},
Body: io.NopCloser(strings.NewReader(
"data: {\"type\":\"response.output_text.delta\",\"delta\":\"hello\"}\n\n" +
"data: {\"type\":\"response.completed\",\"response\":{\"id\":\"resp_http_bridge_1\",\"usage\":{\"input_tokens\":2,\"output_tokens\":1,\"input_tokens_details\":{\"cached_tokens\":1}}}}\n\n" +
"data: {\"type\":\"response.done\",\"response\":{\"id\":\"resp_http_bridge_1\",\"output\":[{\"id\":\"ig_bridge_1\",\"type\":\"image_generation_call\",\"status\":\"in_progress\",\"result\":\"final-image\"}],\"usage\":{\"input_tokens\":2,\"output_tokens\":1,\"input_tokens_details\":{\"cached_tokens\":1}}}}\n\n" +
"data: [DONE]\n\n",
)),
},
@@ -1145,8 +1151,9 @@ func TestOpenAIGatewayService_ProxyResponsesWebSocketFromClient_HTTPBridgeModeRe
_, event2, readErr2 := clientConn.Read(readCtx2)
cancelRead2()
require.NoError(t, readErr2)
require.Equal(t, "response.completed", gjson.GetBytes(event2, "type").String())
require.Equal(t, "response.done", gjson.GetBytes(event2, "type").String())
require.Equal(t, "resp_http_bridge_1", gjson.GetBytes(event2, "response.id").String())
require.Equal(t, "completed", gjson.GetBytes(event2, "response.output.0.status").String())
_ = clientConn.Close(coderws.StatusNormalClosure, "done")
@@ -281,6 +281,7 @@ func TestOpenAIGatewayService_Forward_WSv2_ImageGenerationCountsOutputs(t *testi
"item": map[string]any{
"id": "ig_ws_1",
"type": "image_generation_call",
"status": "generating",
"result": "final-image",
},
}); err != nil {
@@ -296,6 +297,7 @@ func TestOpenAIGatewayService_Forward_WSv2_ImageGenerationCountsOutputs(t *testi
map[string]any{
"id": "ig_ws_1",
"type": "image_generation_call",
"status": "in_progress",
"result": "final-image",
},
},
@@ -375,6 +377,7 @@ func TestOpenAIGatewayService_Forward_WSv2_ImageGenerationCountsOutputs(t *testi
require.Equal(t, 4, result.Usage.OutputTokens)
require.True(t, result.OpenAIWSMode)
require.Equal(t, "resp_ws_image_1", gjson.GetBytes(rec.Body.Bytes(), "id").String())
require.Equal(t, "completed", gjson.GetBytes(rec.Body.Bytes(), "output.0.status").String())
}
func requestToJSONString(payload map[string]any) string {
@@ -504,6 +504,9 @@ func (s *OpenAIGatewayService) forwardOpenAIWSV2(
setOpsUpstreamError(c, 0, sanitizeUpstreamErrorMessage(readErr.Error()), "")
return nil, fmt.Errorf("openai ws read event: %w", readErr)
}
if normalized, changed := normalizeCompletedImageGenerationStatus(message); changed {
message = normalized
}
eventType, eventResponseID, responseField := parseOpenAIWSEventEnvelope(message)
if eventType == "" {
@@ -331,6 +331,9 @@ func (s *OpenAIGatewayService) proxyOpenAIWSHTTPBridgeTurn(
}
upstreamMessage := []byte(trimmedData)
if normalized, changed := normalizeCompletedImageGenerationStatus(upstreamMessage); changed {
upstreamMessage = normalized
}
eventType, eventResponseID, _ := parseOpenAIWSEventEnvelope(upstreamMessage)
if responseID == "" && eventResponseID != "" {
responseID = eventResponseID
@@ -212,6 +212,11 @@ func (c *openAIWSClientFrameConn) WriteFrame(ctx context.Context, msgType coderw
if ctx == nil {
ctx = context.Background()
}
if msgType == coderws.MessageText {
if normalized, changed := normalizeCompletedImageGenerationStatus(payload); changed {
payload = normalized
}
}
return c.conn.Write(ctx, msgType, payload)
}