Files
sub2api/backend/internal/service/openai_ws_forwarder_support.go
李建琦 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

759 lines
25 KiB
Go

package service
import (
"bytes"
"context"
"encoding/json"
"errors"
"net/http"
"strings"
"time"
"github.com/gin-gonic/gin"
"github.com/tidwall/gjson"
"github.com/tidwall/sjson"
)
func (s *OpenAIGatewayService) isOpenAIWSGeneratePrewarmEnabled() bool {
return s != nil && s.cfg != nil && s.cfg.Gateway.OpenAIWS.PrewarmGenerateEnabled
}
// performOpenAIWSGeneratePrewarm 在 WSv2 下执行可选的 generate=false 预热。
// 预热默认关闭,仅在配置开启后生效;失败时按可恢复错误回退到 HTTP。
func (s *OpenAIGatewayService) performOpenAIWSGeneratePrewarm(
ctx context.Context,
lease *openAIWSConnLease,
decision OpenAIWSProtocolDecision,
payload map[string]any,
previousResponseID string,
reqBody map[string]any,
account *Account,
stateStore OpenAIWSStateStore,
groupID int64,
) error {
if s == nil {
return nil
}
if lease == nil || account == nil {
logOpenAIWSModeInfo("prewarm_skip reason=invalid_state has_lease=%v has_account=%v", lease != nil, account != nil)
return nil
}
connID := strings.TrimSpace(lease.ConnID())
if !s.isOpenAIWSGeneratePrewarmEnabled() {
return nil
}
if decision.Transport != OpenAIUpstreamTransportResponsesWebsocketV2 {
logOpenAIWSModeInfo(
"prewarm_skip account_id=%d conn_id=%s reason=transport_not_v2 transport=%s",
account.ID,
connID,
normalizeOpenAIWSLogValue(string(decision.Transport)),
)
return nil
}
if strings.TrimSpace(previousResponseID) != "" {
logOpenAIWSModeInfo(
"prewarm_skip account_id=%d conn_id=%s reason=has_previous_response_id previous_response_id=%s",
account.ID,
connID,
truncateOpenAIWSLogValue(previousResponseID, openAIWSIDValueMaxLen),
)
return nil
}
if lease.IsPrewarmed() {
logOpenAIWSModeInfo("prewarm_skip account_id=%d conn_id=%s reason=already_prewarmed", account.ID, connID)
return nil
}
if NeedsToolContinuation(reqBody) {
logOpenAIWSModeInfo("prewarm_skip account_id=%d conn_id=%s reason=tool_continuation", account.ID, connID)
return nil
}
prewarmStart := time.Now()
logOpenAIWSModeInfo("prewarm_start account_id=%d conn_id=%s", account.ID, connID)
prewarmPayload := make(map[string]any, len(payload)+1)
for k, v := range payload {
prewarmPayload[k] = v
}
prewarmPayload["generate"] = false
prewarmPayloadJSON := payloadAsJSONBytes(prewarmPayload)
if err := lease.WriteJSONWithContextTimeout(ctx, prewarmPayload, s.openAIWSWriteTimeout()); err != nil {
lease.MarkBroken()
logOpenAIWSModeInfo(
"prewarm_write_fail account_id=%d conn_id=%s cause=%s",
account.ID,
connID,
truncateOpenAIWSLogValue(err.Error(), openAIWSLogValueMaxLen),
)
return wrapOpenAIWSFallback("prewarm_write", err)
}
logOpenAIWSModeInfo("prewarm_write_sent account_id=%d conn_id=%s payload_bytes=%d", account.ID, connID, len(prewarmPayloadJSON))
prewarmResponseID := ""
prewarmEventCount := 0
prewarmTerminalCount := 0
for {
message, readErr := lease.ReadMessageWithContextTimeout(ctx, s.openAIWSReadTimeout())
if readErr != nil {
lease.MarkBroken()
closeStatus, closeReason := summarizeOpenAIWSReadCloseError(readErr)
logOpenAIWSModeInfo(
"prewarm_read_fail account_id=%d conn_id=%s close_status=%s close_reason=%s cause=%s events=%d",
account.ID,
connID,
closeStatus,
closeReason,
truncateOpenAIWSLogValue(readErr.Error(), openAIWSLogValueMaxLen),
prewarmEventCount,
)
return wrapOpenAIWSFallback("prewarm_"+classifyOpenAIWSReadFallbackReason(readErr), readErr)
}
eventType, eventResponseID, _ := parseOpenAIWSEventEnvelope(message)
if eventType == "" {
continue
}
prewarmEventCount++
if prewarmResponseID == "" && eventResponseID != "" {
prewarmResponseID = eventResponseID
}
if prewarmEventCount <= openAIWSPrewarmEventLogHead || eventType == "error" || isOpenAIWSTerminalEvent(eventType) {
logOpenAIWSModeInfo(
"prewarm_event account_id=%d conn_id=%s idx=%d type=%s bytes=%d",
account.ID,
connID,
prewarmEventCount,
truncateOpenAIWSLogValue(eventType, openAIWSLogValueMaxLen),
len(message),
)
}
if eventType == "error" {
errCodeRaw, errTypeRaw, errMsgRaw := parseOpenAIWSErrorEventFields(message)
s.persistOpenAIWSRateLimitSignal(ctx, account, lease.HandshakeHeaders(), message, errCodeRaw, errTypeRaw, errMsgRaw)
errMsg := strings.TrimSpace(errMsgRaw)
if errMsg == "" {
errMsg = "OpenAI websocket prewarm error"
}
fallbackReason, canFallback := classifyOpenAIWSErrorEventFromRaw(errCodeRaw, errTypeRaw, errMsgRaw)
errCode, errType, errMessage := summarizeOpenAIWSErrorEventFieldsFromRaw(errCodeRaw, errTypeRaw, errMsgRaw)
logOpenAIWSModeInfo(
"prewarm_error_event account_id=%d conn_id=%s idx=%d fallback_reason=%s can_fallback=%v err_code=%s err_type=%s err_message=%s",
account.ID,
connID,
prewarmEventCount,
truncateOpenAIWSLogValue(fallbackReason, openAIWSLogValueMaxLen),
canFallback,
errCode,
errType,
errMessage,
)
lease.MarkBroken()
if canFallback {
return wrapOpenAIWSFallback("prewarm_"+fallbackReason, errors.New(errMsg))
}
return wrapOpenAIWSFallback("prewarm_error_event", errors.New(errMsg))
}
if isOpenAIWSTerminalEvent(eventType) {
prewarmTerminalCount++
break
}
}
lease.MarkPrewarmed()
if prewarmResponseID != "" && stateStore != nil {
ttl := s.openAIWSResponseStickyTTL()
logOpenAIWSBindResponseAccountWarn(groupID, account.ID, prewarmResponseID, stateStore.BindResponseAccount(ctx, groupID, prewarmResponseID, account.ID, ttl))
stateStore.BindResponseConn(prewarmResponseID, lease.ConnID(), ttl)
}
logOpenAIWSModeInfo(
"prewarm_done account_id=%d conn_id=%s response_id=%s events=%d terminal_events=%d duration_ms=%d",
account.ID,
connID,
truncateOpenAIWSLogValue(prewarmResponseID, openAIWSIDValueMaxLen),
prewarmEventCount,
prewarmTerminalCount,
time.Since(prewarmStart).Milliseconds(),
)
return nil
}
func payloadAsJSON(payload map[string]any) string {
return string(payloadAsJSONBytes(payload))
}
func payloadAsJSONBytes(payload map[string]any) []byte {
if len(payload) == 0 {
return []byte("{}")
}
body, err := json.Marshal(payload)
if err != nil {
return []byte("{}")
}
return body
}
func isOpenAIWSTerminalEvent(eventType string) bool {
switch strings.TrimSpace(eventType) {
case "response.completed", "response.done", "response.failed", "response.incomplete", "response.cancelled", "response.canceled":
return true
default:
return false
}
}
func normalizeOpenAIWSTerminalEvent(eventType string) string {
switch strings.TrimSpace(eventType) {
case "response.completed":
return "response.completed"
case "response.done":
return "response.done"
case "response.failed":
return "response.failed"
case "response.incomplete":
return "response.incomplete"
case "response.cancelled", "response.canceled":
return "response.cancelled"
default:
return ""
}
}
func openAIWSPayloadTransientStatus(payload []byte) int {
if len(payload) == 0 {
return 0
}
status := int(gjson.GetBytes(payload, "response.error.status_code").Int())
if status == 0 {
status = int(gjson.GetBytes(payload, "response.error.status").Int())
}
if status == 0 {
status = int(gjson.GetBytes(payload, "error.status_code").Int())
}
if status == 0 {
status = int(gjson.GetBytes(payload, "error.status").Int())
}
if shouldCooldownOpenAITransientUpstreamError(status, payload) {
return status
}
if status != 0 {
return 0
}
code := strings.ToLower(strings.TrimSpace(gjson.GetBytes(payload, "response.error.code").String()))
errType := strings.ToLower(strings.TrimSpace(gjson.GetBytes(payload, "response.error.type").String()))
if code == "" {
code = strings.ToLower(strings.TrimSpace(gjson.GetBytes(payload, "error.code").String()))
}
if errType == "" {
errType = strings.ToLower(strings.TrimSpace(gjson.GetBytes(payload, "error.type").String()))
}
switch {
case code == "server_is_overloaded", code == "slow_down":
return http.StatusServiceUnavailable
case strings.Contains(code, "server_error"),
strings.Contains(code, "internal_error"),
strings.Contains(code, "upstream_error"),
strings.Contains(errType, "server_error"),
strings.Contains(errType, "internal_error"),
strings.Contains(errType, "upstream_error"):
return http.StatusInternalServerError
default:
return 0
}
}
func (s *OpenAIGatewayService) handleOpenAIWSTerminalTransientFailure(ctx context.Context, account *Account, canonicalModel string, headers http.Header, payload []byte) string {
eventType, _, _ := parseOpenAIWSEventEnvelope(payload)
terminalEvent := normalizeOpenAIWSTerminalEvent(eventType)
if terminalEvent != "response.failed" {
return terminalEvent
}
status := openAIWSPayloadTransientStatus(payload)
if status != 0 {
s.handleOpenAIAccountUpstreamError(ctx, account, status, headers, payload, canonicalModel)
}
return terminalEvent
}
func (s *OpenAIGatewayService) handleOpenAIWSErrorEventTransientFailure(ctx context.Context, account *Account, canonicalModel string, headers http.Header, payload []byte) {
eventType, _, _ := parseOpenAIWSEventEnvelope(payload)
if eventType != "error" {
return
}
status := openAIWSPayloadTransientStatus(payload)
if status != 0 {
s.handleOpenAIAccountUpstreamError(ctx, account, status, headers, payload, canonicalModel)
}
}
func (s *OpenAIGatewayService) handleOpenAIWSDialTransientFailure(ctx context.Context, account *Account, canonicalModel string, err error) {
var dialErr *openAIWSDialError
if !errors.As(err, &dialErr) || dialErr == nil || !shouldCooldownOpenAITransientUpstreamError(dialErr.StatusCode, dialErr.ResponseBody) {
return
}
s.handleOpenAIAccountUpstreamError(ctx, account, dialErr.StatusCode, dialErr.ResponseHeaders, dialErr.ResponseBody, canonicalModel)
}
func isOpenAIWSTokenEvent(eventType string) bool {
eventType = strings.TrimSpace(eventType)
if eventType == "" {
return false
}
switch eventType {
case "response.created", "response.in_progress", "response.output_item.added", "response.output_item.done":
return false
}
if strings.Contains(eventType, ".delta") {
return true
}
if strings.HasPrefix(eventType, "response.output_text") {
return true
}
if strings.HasPrefix(eventType, "response.output") {
return true
}
// 终止事件(response.completed/done/failed/...)由 isOpenAIWSTerminalEvent 单独处理。
// 不能把它们当作 token event,否则当上游没有可识别的 delta 时,
// firstTokenMs 会被填到终止时刻,等于把"总耗时"误报为"首 token 延迟"。
return false
}
func replaceOpenAIWSMessageModel(message []byte, fromModel, toModel string) []byte {
if len(message) == 0 {
return message
}
if strings.TrimSpace(fromModel) == "" || strings.TrimSpace(toModel) == "" || fromModel == toModel {
return message
}
if !bytes.Contains(message, []byte(`"model"`)) || !bytes.Contains(message, []byte(fromModel)) {
return message
}
modelValues := gjson.GetManyBytes(message, "model", "response.model")
replaceModel := modelValues[0].Exists() && modelValues[0].Str == fromModel
replaceResponseModel := modelValues[1].Exists() && modelValues[1].Str == fromModel
if !replaceModel && !replaceResponseModel {
return message
}
updated := message
if replaceModel {
if next, err := sjson.SetBytes(updated, "model", toModel); err == nil {
updated = next
}
}
if replaceResponseModel {
if next, err := sjson.SetBytes(updated, "response.model", toModel); err == nil {
updated = next
}
}
return updated
}
func populateOpenAIUsageFromResponseJSON(body []byte, usage *OpenAIUsage) {
if usage == nil || len(body) == 0 {
return
}
if parsed, ok := extractOpenAIUsageFromJSONBytes(body); ok {
*usage = parsed
}
}
func getOpenAIGroupIDFromContext(c *gin.Context) int64 {
if c == nil {
return 0
}
value, exists := c.Get("api_key")
if !exists {
return 0
}
apiKey, ok := value.(*APIKey)
if !ok || apiKey == nil || apiKey.GroupID == nil {
return 0
}
return *apiKey.GroupID
}
// SelectAccountByPreviousResponseID 按 previous_response_id 命中账号粘连。
// 未命中或账号不可用时返回 (nil, nil),由调用方继续走常规调度。
func (s *OpenAIGatewayService) SelectAccountByPreviousResponseID(
ctx context.Context,
groupID *int64,
previousResponseID string,
requestedModel string,
excludedIDs map[int64]struct{},
requireCompact bool,
) (*AccountSelectionResult, error) {
// 分组利润控制:公共入口装门,保证不经 selectAccountWithScheduler
// 的调用方也无法绕过利润准入(scheduler 内部路径已在唯一调度入口装门)。
ctx = s.withOpenAIProfitControlGate(ctx, groupID)
return s.selectAccountByPreviousResponseIDForCapability(ctx, groupID, previousResponseID, requestedModel, excludedIDs, "", requireCompact)
}
func (s *OpenAIGatewayService) selectAccountByPreviousResponseIDForCapability(
ctx context.Context,
groupID *int64,
previousResponseID string,
requestedModel string,
excludedIDs map[int64]struct{},
requiredCapability OpenAIEndpointCapability,
requireCompact bool,
) (*AccountSelectionResult, error) {
if s == nil {
return nil, nil
}
accountID, account, responseID, store := s.resolveAccountByPreviousResponseIDForCapability(ctx, groupID, previousResponseID, requestedModel, excludedIDs, requiredCapability, requireCompact)
if accountID <= 0 || account == nil || store == nil {
return nil, nil
}
result, acquireErr := s.tryAcquireAccountSlot(ctx, accountID, account.Concurrency)
if acquireErr == nil && result.Acquired {
logOpenAIWSBindResponseAccountWarn(
derefGroupID(groupID),
accountID,
responseID,
store.BindResponseAccount(ctx, derefGroupID(groupID), responseID, accountID, s.openAIWSResponseStickyTTL()),
)
return attachSelectionProfitGate(ctx, &AccountSelectionResult{
Account: account,
Acquired: true,
ReleaseFunc: result.ReleaseFunc,
}), nil
}
cfg := s.schedulingConfig()
if s.concurrencyService != nil {
return attachSelectionProfitGate(ctx, &AccountSelectionResult{
Account: account,
WaitPlan: &AccountWaitPlan{
AccountID: accountID,
MaxConcurrency: account.Concurrency,
Timeout: cfg.StickySessionWaitTimeout,
MaxWaiting: cfg.StickySessionMaxWaiting,
},
}), nil
}
return nil, nil
}
func (s *OpenAIGatewayService) ResolveAccountIDByPreviousResponseIDForScheduler(
ctx context.Context,
groupID *int64,
previousResponseID string,
requestedModel string,
excludedIDs map[int64]struct{},
requiredCapability OpenAIEndpointCapability,
requireCompact bool,
) int64 {
accountID, _, _, _ := s.resolveAccountByPreviousResponseIDForCapability(ctx, groupID, previousResponseID, requestedModel, excludedIDs, requiredCapability, requireCompact)
return accountID
}
func (s *OpenAIGatewayService) resolveAccountByPreviousResponseIDForCapability(
ctx context.Context,
groupID *int64,
previousResponseID string,
requestedModel string,
excludedIDs map[int64]struct{},
requiredCapability OpenAIEndpointCapability,
requireCompact bool,
) (int64, *Account, string, OpenAIWSStateStore) {
if s == nil {
return 0, nil, "", nil
}
responseID := strings.TrimSpace(previousResponseID)
if responseID == "" {
return 0, nil, "", nil
}
store := s.getOpenAIWSStateStore()
if store == nil {
return 0, nil, "", nil
}
accountID, err := store.GetResponseAccount(ctx, derefGroupID(groupID), responseID)
if err != nil || accountID <= 0 {
return 0, nil, "", nil
}
if excludedIDs != nil {
if _, excluded := excludedIDs[accountID]; excluded {
return 0, nil, "", nil
}
}
account, err := s.getSchedulableAccount(ctx, accountID)
if err != nil || account == nil {
_ = store.DeleteResponseAccount(ctx, derefGroupID(groupID), responseID)
return 0, nil, "", nil
}
// 非 WSv2 场景(如 force_http/全局关闭)不应使用 previous_response_id 粘连,
// 以保持“回滚到 HTTP”后的历史行为一致性。
if s.getOpenAIWSProtocolResolver().Resolve(account).Transport != OpenAIUpstreamTransportResponsesWebsocketV2 {
return 0, nil, "", nil
}
if shouldClearStickySession(account, requestedModel) || !account.IsOpenAI() || !account.IsSchedulable() {
_ = store.DeleteResponseAccount(ctx, derefGroupID(groupID), responseID)
return 0, nil, "", nil
}
if !parentHealthyForShadow(account, s.parentAccountLookup(ctx)) {
_ = store.DeleteResponseAccount(ctx, derefGroupID(groupID), responseID)
return 0, nil, "", nil
}
if requestedModel != "" && !account.IsModelSupported(requestedModel) {
return 0, nil, "", nil
}
if !account.SupportsOpenAIEndpointCapability(requiredCapability) {
return 0, nil, "", nil
}
// Quota auto-pause must also gate the previous_response_id sticky path; otherwise an
// account over its 5h/7d threshold keeps serving the same response chain even though
// normal scheduling skips it. Pause is transient, so fall through to normal scheduling
// without deleting the binding (the window may reset before the next turn).
if paused, _ := shouldAutoPauseOpenAIAccountByQuota(ctx, account); paused {
return 0, nil, "", nil
}
// 分组利润控制:与 quota auto-pause 同语义——利润不合格是暂时
// 状态(上游倍率/高峰随时间变化),只跳过本次复用、落回普通调度,不删除
// 绑定(倍率恢复后可继续按 previous_response_id 粘连)。
if vetoed, _ := openAIProfitControlVetoReason(ctx, account); vetoed {
return 0, nil, "", nil
}
if s.schedulerSnapshot != nil && s.accountRepo != nil {
latest, latestErr := s.accountRepo.GetByID(ctx, account.ID)
if latestErr != nil || latest == nil {
_ = store.DeleteResponseAccount(ctx, derefGroupID(groupID), responseID)
return 0, nil, "", nil
}
if shouldClearStickySession(latest, requestedModel) || !latest.IsOpenAI() || !latest.IsSchedulable() {
_ = store.DeleteResponseAccount(ctx, derefGroupID(groupID), responseID)
return 0, nil, "", nil
}
if !parentHealthyForShadow(latest, s.parentAccountLookup(ctx)) {
_ = store.DeleteResponseAccount(ctx, derefGroupID(groupID), responseID)
return 0, nil, "", nil
}
if requestedModel != "" && !latest.IsModelSupported(requestedModel) {
return 0, nil, "", nil
}
if !latest.SupportsOpenAIEndpointCapability(requiredCapability) {
return 0, nil, "", nil
}
if paused, _ := shouldAutoPauseOpenAIAccountByQuota(ctx, latest); paused {
return 0, nil, "", nil
}
// 利润门对最新账号状态复检一次,语义同上:跳过复用、不删绑定。
if vetoed, _ := openAIProfitControlVetoReason(ctx, latest); vetoed {
return 0, nil, "", nil
}
if s.isOpenAIAccountRequestRuntimeBlocked(latest, requestedModel) {
_ = store.DeleteResponseAccount(ctx, derefGroupID(groupID), responseID)
return 0, nil, "", nil
}
account = latest
}
if requireCompact && openAICompactSupportTier(account) == 0 {
_ = store.DeleteResponseAccount(ctx, derefGroupID(groupID), responseID)
return 0, nil, "", nil
}
return accountID, account, responseID, store
}
func classifyOpenAIWSAcquireError(err error) string {
if err == nil {
return "acquire_conn"
}
var dialErr *openAIWSDialError
if errors.As(err, &dialErr) {
switch dialErr.StatusCode {
case 426:
return "upgrade_required"
case 401, 403:
return "auth_failed"
case 429:
return "upstream_rate_limited"
}
if dialErr.StatusCode >= 500 {
return "upstream_5xx"
}
return "dial_failed"
}
if errors.Is(err, errOpenAIWSConnQueueFull) {
return "conn_queue_full"
}
if errors.Is(err, errOpenAIWSPreferredConnUnavailable) {
return "preferred_conn_unavailable"
}
if errors.Is(err, context.DeadlineExceeded) {
return "acquire_timeout"
}
return "acquire_conn"
}
func isOpenAIWSRateLimitError(codeRaw, errTypeRaw, msgRaw string) bool {
code := strings.ToLower(strings.TrimSpace(codeRaw))
errType := strings.ToLower(strings.TrimSpace(errTypeRaw))
msg := strings.ToLower(strings.TrimSpace(msgRaw))
if strings.Contains(errType, "rate_limit") || strings.Contains(errType, "usage_limit") {
return true
}
if strings.Contains(code, "rate_limit") || strings.Contains(code, "usage_limit") || strings.Contains(code, "insufficient_quota") {
return true
}
if strings.Contains(msg, "usage limit") && strings.Contains(msg, "reached") {
return true
}
if strings.Contains(msg, "rate limit") && (strings.Contains(msg, "reached") || strings.Contains(msg, "exceeded")) {
return true
}
return false
}
func (s *OpenAIGatewayService) persistOpenAIWSRateLimitSignal(ctx context.Context, account *Account, headers http.Header, responseBody []byte, codeRaw, errTypeRaw, msgRaw string) {
if s == nil || s.rateLimitService == nil || account == nil || account.Platform != PlatformOpenAI {
return
}
if !isOpenAIWSRateLimitError(codeRaw, errTypeRaw, msgRaw) {
return
}
s.handleOpenAIAccountUpstreamError(ctx, account, http.StatusTooManyRequests, headers, responseBody)
}
func classifyOpenAIWSErrorEventFromRaw(codeRaw, errTypeRaw, msgRaw string) (string, bool) {
code := strings.ToLower(strings.TrimSpace(codeRaw))
errType := strings.ToLower(strings.TrimSpace(errTypeRaw))
msg := strings.ToLower(strings.TrimSpace(msgRaw))
switch code {
case "upgrade_required":
return "upgrade_required", true
case "websocket_not_supported", "websocket_unsupported":
return "ws_unsupported", true
case "websocket_connection_limit_reached":
return "ws_connection_limit_reached", true
case "invalid_encrypted_content":
return "invalid_encrypted_content", true
case "previous_response_not_found":
return "previous_response_not_found", true
}
if isOpenAIWSRateLimitError(codeRaw, errTypeRaw, msgRaw) {
return "upstream_rate_limited", false
}
if strings.Contains(msg, "upgrade required") || strings.Contains(msg, "status 426") {
return "upgrade_required", true
}
if strings.Contains(errType, "upgrade") {
return "upgrade_required", true
}
if strings.Contains(msg, "websocket") && strings.Contains(msg, "unsupported") {
return "ws_unsupported", true
}
if strings.Contains(msg, "connection limit") && strings.Contains(msg, "websocket") {
return "ws_connection_limit_reached", true
}
if strings.Contains(msg, "invalid_encrypted_content") ||
(strings.Contains(msg, "encrypted content") && strings.Contains(msg, "could not be verified")) {
return "invalid_encrypted_content", true
}
if strings.Contains(msg, "previous_response_not_found") ||
(strings.Contains(msg, "previous response") && strings.Contains(msg, "not found")) {
return "previous_response_not_found", true
}
if strings.Contains(errType, "server_error") || strings.Contains(code, "server_error") {
return "upstream_error_event", true
}
return "event_error", false
}
func classifyOpenAIWSErrorEvent(message []byte) (string, bool) {
if len(message) == 0 {
return "event_error", false
}
return classifyOpenAIWSErrorEventFromRaw(parseOpenAIWSErrorEventFields(message))
}
func openAIWSErrorHTTPStatusFromRaw(codeRaw, errTypeRaw string) int {
code := strings.ToLower(strings.TrimSpace(codeRaw))
errType := strings.ToLower(strings.TrimSpace(errTypeRaw))
switch {
case strings.Contains(errType, "invalid_request"),
strings.Contains(code, "invalid_request"),
strings.Contains(code, "bad_request"),
code == "invalid_encrypted_content",
code == "previous_response_not_found":
return http.StatusBadRequest
case strings.Contains(errType, "authentication"),
strings.Contains(code, "invalid_api_key"),
strings.Contains(code, "unauthorized"):
return http.StatusUnauthorized
case strings.Contains(errType, "permission"),
strings.Contains(code, "forbidden"):
return http.StatusForbidden
case isOpenAIWSRateLimitError(codeRaw, errTypeRaw, ""):
return http.StatusTooManyRequests
default:
return http.StatusBadGateway
}
}
func openAIWSErrorHTTPStatus(message []byte) int {
if len(message) == 0 {
return http.StatusBadGateway
}
codeRaw, errTypeRaw, _ := parseOpenAIWSErrorEventFields(message)
return openAIWSErrorHTTPStatusFromRaw(codeRaw, errTypeRaw)
}
func (s *OpenAIGatewayService) openAIWSFallbackCooldown() time.Duration {
if s == nil || s.cfg == nil {
return 30 * time.Second
}
seconds := s.cfg.Gateway.OpenAIWS.FallbackCooldownSeconds
if seconds <= 0 {
return 0
}
return time.Duration(seconds) * time.Second
}
func (s *OpenAIGatewayService) isOpenAIWSFallbackCooling(accountID int64) bool {
if s == nil || accountID <= 0 {
return false
}
cooldown := s.openAIWSFallbackCooldown()
if cooldown <= 0 {
return false
}
rawUntil, ok := s.openaiWSFallbackUntil.Load(accountID)
if !ok || rawUntil == nil {
return false
}
until, ok := rawUntil.(time.Time)
if !ok || until.IsZero() {
s.openaiWSFallbackUntil.Delete(accountID)
return false
}
if time.Now().Before(until) {
return true
}
s.openaiWSFallbackUntil.Delete(accountID)
return false
}
func (s *OpenAIGatewayService) markOpenAIWSFallbackCooling(accountID int64, _ string) {
if s == nil || accountID <= 0 {
return
}
cooldown := s.openAIWSFallbackCooldown()
if cooldown <= 0 {
return
}
s.openaiWSFallbackUntil.Store(accountID, time.Now().Add(cooldown))
}
func (s *OpenAIGatewayService) clearOpenAIWSFallbackCooling(accountID int64) {
if s == nil || accountID <= 0 {
return
}
s.openaiWSFallbackUntil.Delete(accountID)
}