mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-10-07 15:38:22 +08:00
防止 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:
@@ -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)
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user