Sub2API v1.0 - AI API 网关(二开初始版本,基于上游 Wei-Shaw/sub2api)
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
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
This commit is contained in:
@@ -0,0 +1,758 @@
|
||||
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)
|
||||
}
|
||||
Reference in New Issue
Block a user