Files
sub2api/backend/internal/service/openai_images_upstream_context_test.go
T
李建琦 6d655c9903
Release / update-version (push) Has been cancelled
Release / build-frontend (push) Has been cancelled
Release / release (push) Has been cancelled
Release / sync-version-file (push) Has been cancelled
CI / shell (push) Canceled after 0s
CI / test (push) Canceled after 0s
CI / frontend (push) Canceled after 0s
CI / golangci-lint (push) Canceled after 0s
Security Scan / backend-security (push) Canceled after 0s
Security Scan / frontend-security (push) Canceled after 0s
Sub2API v1.0 - AI API 网关(二开初始版本,基于上游 Wei-Shaw/sub2api)
2026-08-21 18:30:13 +08:00

155 lines
5.3 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
package service
import (
"bytes"
"context"
"io"
"net/http"
"net/http/httptest"
"strings"
"testing"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/require"
"github.com/Wei-Shaw/sub2api/internal/config"
)
func newOpenAIImagesTestContext(t *testing.T, body []byte) (*gin.Context, *httptest.ResponseRecorder) {
t.Helper()
gin.SetMode(gin.TestMode)
req := httptest.NewRequest(http.MethodPost, "/v1/images/generations", bytes.NewReader(body))
req.Header.Set("Content-Type", "application/json")
rec := httptest.NewRecorder()
c, _ := gin.CreateTestContext(rec)
c.Request = req
return c, rec
}
func newOpenAIImagesTestService(upstream HTTPUpstream) *OpenAIGatewayService {
return &OpenAIGatewayService{
httpUpstream: upstream,
cfg: &config.Config{
Security: config.SecurityConfig{
URLAllowlist: config.URLAllowlistConfig{Enabled: false},
},
},
}
}
func newOpenAIImagesAPIKeyAccount() *Account {
return &Account{
ID: 31,
Name: "openai-apikey-images",
Platform: PlatformOpenAI,
Type: AccountTypeAPIKey,
Credentials: map[string]any{
"api_key": "sk-test",
"base_url": "https://api.openai.com/v1",
},
}
}
func openAIImagesJSONResponse() *http.Response {
return &http.Response{
StatusCode: http.StatusOK,
Header: http.Header{
"Content-Type": []string{"application/json"},
"X-Request-Id": []string{"req_img_ctx"},
},
Body: io.NopCloser(strings.NewReader(
`{"created":1710000000,"data":[{"b64_json":"aGVsbG8="}],"usage":{"input_tokens":10,"output_tokens":20,"total_tokens":30}}`,
)),
}
}
// issue #5411:生图是长耗时、上游侧已经产生实际成本的操作。客户端中途断开时,
// 如果连带取消上游请求,就会出现「上游已出图并计费、网关记 502 context canceled、
// 用户不扣费」。非流式路径以前走 detachStreamUpstreamContext(ctx, false)
// 该函数在非流式时原样返回请求 context,因此不脱钩。
func TestForwardOpenAIImagesAPIKey_NonStreamDetachesUpstreamContext(t *testing.T) {
body := []byte(`{"model":"gpt-image-2","prompt":"draw a cat","response_format":"b64_json"}`)
c, _ := newOpenAIImagesTestContext(t, body)
recorder := &httpUpstreamRecorder{resp: openAIImagesJSONResponse()}
svc := newOpenAIImagesTestService(recorder)
parsed, err := svc.ParseOpenAIImagesRequest(c, body)
require.NoError(t, err)
require.False(t, parsed.Stream, "本用例覆盖非流式生图")
ctx, cancel := context.WithCancel(context.Background())
cancel() // 客户端已断开
result, err := svc.ForwardImages(ctx, c, newOpenAIImagesAPIKeyAccount(), body, parsed, "")
require.NoError(t, err, "客户端断开不应把已在出图的上游调用打断成 context canceled")
require.NotNil(t, result)
require.Equal(t, 1, result.ImageCount, "图片已产出,必须带回结果供计费")
require.NotNil(t, recorder.lastReq)
require.NoError(t, recorder.lastReq.Context().Err(),
"交给上游的请求 context 必须已脱钩,不随客户端断开取消")
}
// 流式路径本来就脱钩,这条守卫防止对齐时把它改坏。
func TestForwardOpenAIImagesAPIKey_StreamKeepsDetachedUpstreamContext(t *testing.T) {
body := []byte(`{"model":"gpt-image-2","prompt":"draw a cat","stream":true,"response_format":"b64_json"}`)
c, _ := newOpenAIImagesTestContext(t, body)
recorder := &httpUpstreamRecorder{resp: &http.Response{
StatusCode: http.StatusOK,
Header: http.Header{
"Content-Type": []string{"text/event-stream"},
"X-Request-Id": []string{"req_img_ctx_stream"},
},
Body: io.NopCloser(strings.NewReader(
"data: {\"type\":\"response.created\",\"response\":{\"created_at\":1710000000}}\n\n" +
"data: {\"type\":\"response.image_generation_call.completed\",\"result\":\"aGVsbG8=\"}\n\n" +
"data: {\"type\":\"response.completed\",\"response\":{\"usage\":{\"input_tokens\":10,\"output_tokens\":20}}}\n\n",
)),
}}
svc := newOpenAIImagesTestService(recorder)
parsed, err := svc.ParseOpenAIImagesRequest(c, body)
require.NoError(t, err)
require.True(t, parsed.Stream)
ctx, cancel := context.WithCancel(context.Background())
cancel()
_, _ = svc.ForwardImages(ctx, c, newOpenAIImagesAPIKeyAccount(), body, parsed, "")
require.NotNil(t, recorder.lastReq)
require.NoError(t, recorder.lastReq.Context().Err(),
"流式路径原本就脱钩,不能被改回随客户端取消")
}
// 两个 detach 辅助函数的语义差异是本次修复的根据,锁死它们防止被悄悄改动。
func TestDetachUpstreamContextSemantics(t *testing.T) {
t.Run("detachUpstreamContext_always_detaches", func(t *testing.T) {
ctx, cancel := context.WithCancel(context.Background())
cancel()
detached, release := detachUpstreamContext(ctx)
defer release()
require.NoError(t, detached.Err())
})
t.Run("detachStreamUpstreamContext_keeps_cancel_when_not_streaming", func(t *testing.T) {
ctx, cancel := context.WithCancel(context.Background())
cancel()
same, release := detachStreamUpstreamContext(ctx, false)
defer release()
require.ErrorIs(t, same.Err(), context.Canceled,
"非流式时该函数原样返回请求 context —— 生图路径不能用它")
})
t.Run("detachStreamUpstreamContext_detaches_when_streaming", func(t *testing.T) {
ctx, cancel := context.WithCancel(context.Background())
cancel()
detached, release := detachStreamUpstreamContext(ctx, true)
defer release()
require.NoError(t, detached.Err())
})
}