Files
sub2api/backend/internal/service/openai_ws_forwarder_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

209 lines
9.6 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 (
"context"
"net/http"
"net/http/httptest"
"testing"
"time"
"github.com/Wei-Shaw/sub2api/internal/config"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/require"
)
// TestIsOpenAIWSTokenEvent_TerminalEventsExcluded 覆盖 isOpenAIWSTokenEvent 的回归用例。
// 重点验证终止事件(response.completed / response.done)不再被当作 token event
// 否则当上游没有可识别的 delta 时,firstTokenMs 会被填到终止时刻,
// 等于把"总耗时"误报为"首 token 延迟"issue #2651)。
func TestIsOpenAIWSTokenEvent_TerminalEventsExcluded(t *testing.T) {
cases := []struct {
name string
eventType string
want bool
}{
{name: "empty", eventType: "", want: false},
{name: "whitespace_trimmed_empty", eventType: " ", want: false},
{name: "response.created", eventType: "response.created", want: false},
{name: "response.in_progress", eventType: "response.in_progress", want: false},
{name: "response.output_item.added", eventType: "response.output_item.added", want: false},
{name: "response.output_item.done", eventType: "response.output_item.done", want: false},
{name: "terminal_response.completed", eventType: "response.completed", want: false},
{name: "terminal_response.done", eventType: "response.done", want: false},
{name: "terminal_response.completed_padded", eventType: " response.completed ", want: false},
{name: "terminal_response.done_padded", eventType: " response.done ", want: false},
{name: "delta_text", eventType: "response.output_text.delta", want: true},
{name: "delta_audio_transcript", eventType: "response.audio_transcript.delta", want: true},
{name: "delta_function_call_arguments", eventType: "response.function_call_arguments.delta", want: true},
{name: "output_text_done", eventType: "response.output_text.done", want: true},
{name: "output_text_annotation_added", eventType: "response.output_text.annotation.added", want: true},
{name: "output_audio_done", eventType: "response.output_audio.done", want: true},
{name: "reasoning_summary_delta", eventType: "response.reasoning_summary_text.delta", want: true},
{name: "unrelated_event_error", eventType: "error", want: false},
{name: "unknown_event_without_match", eventType: "response.reasoning_summary_part.added", want: false},
}
for _, tc := range cases {
tc := tc
t.Run(tc.name, func(t *testing.T) {
got := isOpenAIWSTokenEvent(tc.eventType)
require.Equal(t, tc.want, got, "isOpenAIWSTokenEvent(%q)", tc.eventType)
})
}
}
// TestOpenAIWSCyberPolicyMark_ResponseFailed 验证 WS 路径 response.failed cyber_policy 标记逻辑。
//
// 全量转发循环(forwardOpenAIWSV2 / sendAndRelay)依赖真实 WebSocket 连接,
// 无法在单元测试中驱动。本测试通过直接调用转发循环内使用的两个函数
// detectOpenAICyberPolicy + MarkOpsCyberPolicy,覆盖「从 response.failed 帧
// 到 gin context 写入」的完整调用序列,等同于循环体内对应代码段的逻辑验证。
// 全量 WS 端到端覆盖由后续集成测试(Task 12 handler 编排)承担。
func TestOpenAIWSCyberPolicyMark_ResponseFailed(t *testing.T) {
// 构造一个真实的 response.failed 帧(cyber_policy 命中路径)。
payload := []byte(`{"type":"response.failed","response":{"id":"resp_abc","status":"failed","error":{"code":"cyber_policy","message":"Request blocked by content policy."}}}`)
// 验证 detectOpenAICyberPolicy 能从 response.error.code 路径识别。
hit, code, msg := detectOpenAICyberPolicy(payload)
require.True(t, hit, "detectOpenAICyberPolicy should return true for cyber_policy payload")
require.Equal(t, "cyber_policy", code)
require.Equal(t, "Request blocked by content policy.", msg)
// 构造 gin test context,模拟转发循环调用 MarkOpsCyberPolicy。
rec := httptest.NewRecorder()
c, _ := gin.CreateTestContext(rec)
usage := OpenAIUsage{InputTokens: 42, OutputTokens: 7}
MarkOpsCyberPolicy(c, CyberPolicyMark{
Code: code,
Message: msg,
Body: truncateString(string(payload), 4096),
UpstreamStatus: 200,
UpstreamInTok: usage.InputTokens,
UpstreamOutTok: usage.OutputTokens,
})
mark := GetOpsCyberPolicy(c)
require.NotNil(t, mark, "GetOpsCyberPolicy should return non-nil after MarkOpsCyberPolicy")
require.Equal(t, "cyber_policy", mark.Code)
require.Equal(t, "Request blocked by content policy.", mark.Message)
require.Equal(t, 200, mark.UpstreamStatus)
require.Equal(t, 42, mark.UpstreamInTok)
require.Equal(t, 7, mark.UpstreamOutTok)
// 验证幂等性:再次标记不覆盖首个。
MarkOpsCyberPolicy(c, CyberPolicyMark{Code: "cyber_policy", Message: "second call"})
require.Equal(t, "Request blocked by content policy.", GetOpsCyberPolicy(c).Message, "second MarkOpsCyberPolicy call must not overwrite first")
}
// TestOpenAIWSCyberPolicyMark_NonCyberPayload 验证非 cyber_policy 的 response.failed 不触发标记。
func TestOpenAIWSCyberPolicyMark_NonCyberPayload(t *testing.T) {
payload := []byte(`{"type":"response.failed","response":{"id":"resp_xyz","status":"failed","error":{"code":"server_error","message":"Internal error"}}}`)
hit, _, _ := detectOpenAICyberPolicy(payload)
require.False(t, hit, "detectOpenAICyberPolicy should return false for non-cyber_policy error code")
}
func TestOpenAIForwardResultSucceededForScheduling_TerminalEvents(t *testing.T) {
tests := []struct {
name string
result *OpenAIForwardResult
expected bool
}{
{name: "nil legacy result", result: nil, expected: true},
{name: "non websocket zero value", result: &OpenAIForwardResult{}, expected: true},
{name: "websocket legacy empty terminal", result: &OpenAIForwardResult{OpenAIWSMode: true}, expected: true},
{name: "completed", result: &OpenAIForwardResult{OpenAIWSMode: true, UpstreamTerminalEvent: "response.completed"}, expected: true},
{name: "done", result: &OpenAIForwardResult{OpenAIWSMode: true, UpstreamTerminalEvent: "response.done"}, expected: true},
{name: "failed", result: &OpenAIForwardResult{OpenAIWSMode: true, UpstreamTerminalEvent: "response.failed"}, expected: false},
{name: "incomplete", result: &OpenAIForwardResult{OpenAIWSMode: true, UpstreamTerminalEvent: "response.incomplete"}, expected: false},
{name: "cancelled", result: &OpenAIForwardResult{OpenAIWSMode: true, UpstreamTerminalEvent: "response.cancelled"}, expected: false},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
require.Equal(t, tt.expected, tt.result.SucceededForScheduling())
})
}
}
func TestOpenAIWSTerminalEvent_ResponseFailedRecordsModelTransient(t *testing.T) {
svc := &OpenAIGatewayService{}
svc.rateLimitService = NewRateLimitService(transientCooldownAccountRepo{}, nil, &config.Config{}, nil, nil)
account := &Account{ID: 5201, Platform: PlatformOpenAI, Type: AccountTypeAPIKey}
payload := []byte(`{"type":"response.failed","response":{"error":{"code":"server_error","message":"Internal error"}}}`)
for range 2 {
terminalEvent := svc.handleOpenAIWSTerminalTransientFailure(context.Background(), account, "gpt-5.5", http.Header{}, payload)
require.Equal(t, "response.failed", terminalEvent)
}
require.True(t, svc.isOpenAIAccountModelRuntimeBlocked(account, "gpt-5.5"))
}
func TestOpenAIWSErrorEvent_ServerErrorRecordsModelTransient(t *testing.T) {
svc := &OpenAIGatewayService{}
svc.rateLimitService = NewRateLimitService(transientCooldownAccountRepo{}, nil, &config.Config{}, nil, nil)
account := &Account{ID: 5203, Platform: PlatformOpenAI, Type: AccountTypeAPIKey}
payload := []byte(`{"type":"error","error":{"code":"server_error","type":"server_error","message":"Internal error"}}`)
for range 2 {
svc.handleOpenAIWSErrorEventTransientFailure(context.Background(), account, "gpt-5.5", http.Header{}, payload)
}
require.True(t, svc.isOpenAIAccountModelRuntimeBlocked(account, "gpt-5.5"))
}
func TestOpenAIWSPayloadTransientStatus_Explicit529IsNotModelTransient(t *testing.T) {
payload := []byte(`{"type":"response.failed","response":{"error":{"status_code":529,"code":"server_error","message":"overloaded"}}}`)
require.Zero(t, openAIWSPayloadTransientStatus(payload))
}
func TestOpenAIWSDial5xxRecordsModelTransient(t *testing.T) {
svc := &OpenAIGatewayService{}
svc.rateLimitService = NewRateLimitService(transientCooldownAccountRepo{}, nil, &config.Config{}, nil, nil)
account := &Account{ID: 5202, Platform: PlatformOpenAI, Type: AccountTypeAPIKey}
dialErr := &openAIWSDialError{
StatusCode: http.StatusBadGateway,
ResponseHeaders: http.Header{"X-Request-Id": []string{"req-ws-502"}},
ResponseBody: []byte(`{"error":{"message":"bad gateway"}}`),
}
for range 2 {
svc.handleOpenAIWSDialTransientFailure(context.Background(), account, "gpt-5.5", dialErr)
}
require.Eventually(t, func() bool {
return svc.isOpenAIAccountModelRuntimeBlocked(account, "gpt-5.5")
}, time.Second, 10*time.Millisecond)
}
// TestIsOpenAIWSTokenEvent_DisjointWithTerminal 守护「token 事件集合与终止事件集合互斥」的不变量。
// firstTokenMs 的计算依赖于 isTokenEvent && !isTerminalEvent
// 若两者再次出现交集,则 issue #2651 描述的 latency 误报会重现。
func TestIsOpenAIWSTokenEvent_DisjointWithTerminal(t *testing.T) {
terminalEvents := []string{
"response.completed",
"response.done",
"response.failed",
"response.incomplete",
"response.cancelled",
"response.canceled",
}
for _, ev := range terminalEvents {
ev := ev
t.Run(ev, func(t *testing.T) {
require.True(t, isOpenAIWSTerminalEvent(ev), "expected terminal event %q to be classified as terminal", ev)
require.False(t, isOpenAIWSTokenEvent(ev), "terminal event %q must NOT be classified as token event (issue #2651)", ev)
})
}
}