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

1218 lines
42 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"
"crypto/sha256"
"encoding/hex"
"encoding/json"
"errors"
"fmt"
"log/slog"
"math/rand"
"net/http"
"strings"
"sync"
"sync/atomic"
"time"
"github.com/Wei-Shaw/sub2api/internal/config"
"github.com/Wei-Shaw/sub2api/internal/pkg/ip"
"github.com/Wei-Shaw/sub2api/internal/pkg/logger"
"github.com/Wei-Shaw/sub2api/internal/pkg/openai"
"github.com/Wei-Shaw/sub2api/internal/platform/liveattestation"
"github.com/Wei-Shaw/sub2api/internal/util/responseheaders"
"github.com/cespare/xxhash/v2"
"github.com/gin-gonic/gin"
"go.uber.org/zap"
)
const (
// ChatGPT internal API for OAuth accounts
chatgptCodexURL = "https://chatgpt.com/backend-api/codex/responses"
// OpenAI Platform API for API Key accounts (fallback)
openaiPlatformAPIURL = "https://api.openai.com/v1/responses"
openaiPlatformAPIInputTokensURL = "https://api.openai.com/v1/responses/input_tokens"
openaiStickySessionTTL = time.Hour // 粘性会话TTL
// 与真实 Codex TUI 的 User-Agent 结构对齐:
// {originator}/{version} ({OS} {OS_version}; {arch}) {terminal}
// 缺少 OS/架构/终端后缀的形态易被上游指纹识别为非官方客户端。
// 该后缀是 UA 形态的唯一定义处,buildCodexCLIUserAgent 按运行时版本号复用它。
codexCLIUserAgentSuffix = " (Ubuntu 22.4.0; x86_64) xterm-256color"
// codexCLIUserAgent 是编译期兜底 UA;运行时优先使用由后台版本号拼出的规范 UA。
// 版本段必须来自 codexCLIVersionUA 与 version 头是同一个版本声明的两个出口,
// 各自硬编码会漂移成互相矛盾的身份。
codexCLIUserAgent = openai.CodexDefaultOriginator + "/" + codexCLIVersion + codexCLIUserAgentSuffix
// codex_cli_only 拒绝时单个请求头日志长度上限(字符)
codexCLIOnlyHeaderValueMaxBytes = 256
// OpenAI WS Mode 失败后的重连次数上限(不含首次尝试)。
// 与 Codex 客户端保持一致:失败后最多重连 5 次。
openAIWSReconnectRetryLimit = 5
// 上游错误体只需要提取错误 JSON/日志摘要,默认 512KiB 避免错误风暴叠加大请求体。
openAIUpstreamErrorBodyReadLimit int64 = 512 << 10
// OpenAI WS Mode 重连退避默认值(可由配置覆盖)。
openAIWSRetryBackoffInitialDefault = 120 * time.Millisecond
openAIWSRetryBackoffMaxDefault = 2 * time.Second
openAIWSRetryJitterRatioDefault = 0.2
openAICompactSessionSeedKey = "openai_compact_session_seed"
openAIUpstreamEndpointContextKey = "openai_actual_upstream_endpoint"
// codexCLIVersion 是网关对上游声明的 Codex 客户端版本,同时供 codexCLIUserAgent
// 与 version 头使用。上游 /backend-api/codex 在容量紧张时按客户端身份分优先级降载,
// 陈旧版本会被优先丢弃(HTTP 200 + 流内 server_is_overloaded);非官方客户端配不出
// 官方身份时整体回退到本常量,因此它必须跟随官方 CLI 的当前发布版本,
// 落后多个版本会让这些请求稳定落在被优先丢弃的一侧。
codexCLIVersion = "0.146.0"
// Codex 限额快照仅用于后台展示/诊断,不需要每个成功请求都立即落库。
openAICodexSnapshotPersistMinInterval = 30 * time.Second
// 配额自动暂停时,超过该时长仍未刷新的 used% 快照视为陈旧,不再据此暂停账号。
// 被暂停的账号收不到流量,其快照永远不会从上游响应头刷新;该兜底让账号在快照
// 陈旧时放行一次请求,从而通过正常响应头自愈,而无需等待整个窗口(5h/7d)重置。
openAICodexAutoPauseStaleAfter = 2 * time.Hour
)
// OpenAI allowed headers whitelist (for non-passthrough).
var openaiAllowedHeaders = map[string]bool{
"accept-language": true,
"content-type": true,
"conversation_id": true,
"user-agent": true,
"originator": true,
"session_id": true,
"x-codex-beta-features": true,
"x-codex-installation-id": true,
"x-codex-turn-state": true,
"x-codex-turn-metadata": true,
"x-codex-window-id": true,
responsesLiteHeaderKey: true,
}
// OpenAI passthrough allowed headers whitelist.
// 透传模式下仅放行这些低风险请求头,避免将非标准/环境噪声头传给上游触发风控。
var openaiPassthroughAllowedHeaders = map[string]bool{
"accept": true,
"accept-language": true,
"content-type": true,
"conversation_id": true,
"openai-beta": true,
"user-agent": true,
"originator": true,
"session_id": true,
"x-codex-beta-features": true,
"x-codex-installation-id": true,
"x-codex-turn-state": true,
"x-codex-turn-metadata": true,
"x-codex-window-id": true,
responsesLiteHeaderKey: true,
}
// codex_cli_only 拒绝时记录的请求头白名单(仅用于诊断日志,不参与上游透传)
var codexCLIOnlyDebugHeaderWhitelist = []string{
"User-Agent",
"Content-Type",
"Accept",
"Accept-Language",
"OpenAI-Beta",
"Originator",
"Session_ID",
"Conversation_ID",
"X-Request-ID",
"X-Client-Request-ID",
"X-Forwarded-For",
"X-Real-IP",
}
// OpenAICodexUsageSnapshot represents Codex API usage limits from response headers
type OpenAICodexUsageSnapshot struct {
PrimaryUsedPercent *float64 `json:"primary_used_percent,omitempty"`
PrimaryResetAfterSeconds *int `json:"primary_reset_after_seconds,omitempty"`
PrimaryWindowMinutes *int `json:"primary_window_minutes,omitempty"`
SecondaryUsedPercent *float64 `json:"secondary_used_percent,omitempty"`
SecondaryResetAfterSeconds *int `json:"secondary_reset_after_seconds,omitempty"`
SecondaryWindowMinutes *int `json:"secondary_window_minutes,omitempty"`
PrimaryOverSecondaryPercent *float64 `json:"primary_over_secondary_percent,omitempty"`
UpdatedAt string `json:"updated_at,omitempty"`
}
// NormalizedCodexLimits contains normalized 5h/7d rate limit data
type NormalizedCodexLimits struct {
Used5hPercent *float64
Reset5hSeconds *int
Window5hMinutes *int
Used7dPercent *float64
Reset7dSeconds *int
Window7dMinutes *int
}
// Normalize converts primary/secondary fields to canonical 5h/7d fields.
// Strategy: Compare window_minutes to determine which is 5h vs 7d.
// Returns nil if snapshot is nil or has no useful data.
func (s *OpenAICodexUsageSnapshot) Normalize() *NormalizedCodexLimits {
if s == nil {
return nil
}
result := &NormalizedCodexLimits{}
primaryMins := 0
secondaryMins := 0
hasPrimaryWindow := false
hasSecondaryWindow := false
if s.PrimaryWindowMinutes != nil {
primaryMins = *s.PrimaryWindowMinutes
hasPrimaryWindow = true
}
if s.SecondaryWindowMinutes != nil {
secondaryMins = *s.SecondaryWindowMinutes
hasSecondaryWindow = true
}
// Determine mapping based on window_minutes
use5hFromPrimary := false
use7dFromPrimary := false
if hasPrimaryWindow && hasSecondaryWindow {
// Both known: smaller window is 5h, larger is 7d
if primaryMins < secondaryMins {
use5hFromPrimary = true
} else {
use7dFromPrimary = true
}
} else if hasPrimaryWindow {
// Only primary known: classify by threshold (<=360 min = 6h -> 5h window)
if primaryMins <= 360 {
use5hFromPrimary = true
} else {
use7dFromPrimary = true
}
} else if hasSecondaryWindow {
// Only secondary known: classify by threshold
if secondaryMins <= 360 {
// 5h from secondary, so primary (if any data) is 7d
use7dFromPrimary = true
} else {
// 7d from secondary, so primary (if any data) is 5h
use5hFromPrimary = true
}
} else {
// No window_minutes: fall back to legacy assumption (primary=7d, secondary=5h)
use7dFromPrimary = true
}
// Assign values
if use5hFromPrimary {
result.Used5hPercent = s.PrimaryUsedPercent
result.Reset5hSeconds = s.PrimaryResetAfterSeconds
result.Window5hMinutes = s.PrimaryWindowMinutes
result.Used7dPercent = s.SecondaryUsedPercent
result.Reset7dSeconds = s.SecondaryResetAfterSeconds
result.Window7dMinutes = s.SecondaryWindowMinutes
} else if use7dFromPrimary {
result.Used7dPercent = s.PrimaryUsedPercent
result.Reset7dSeconds = s.PrimaryResetAfterSeconds
result.Window7dMinutes = s.PrimaryWindowMinutes
result.Used5hPercent = s.SecondaryUsedPercent
result.Reset5hSeconds = s.SecondaryResetAfterSeconds
result.Window5hMinutes = s.SecondaryWindowMinutes
}
return result
}
// OpenAIUsage represents OpenAI API response usage
type OpenAIUsage struct {
InputTokens int `json:"input_tokens"`
ImageInputTokens int `json:"image_input_tokens,omitempty"`
OutputTokens int `json:"output_tokens"`
CacheCreationInputTokens int `json:"cache_creation_input_tokens,omitempty"`
CacheReadInputTokens int `json:"cache_read_input_tokens,omitempty"`
ImageOutputTokens int `json:"image_output_tokens,omitempty"`
}
// OpenAIForwardResult represents the result of forwarding
type OpenAIForwardResult struct {
RequestID string
ResponseID string
Usage OpenAIUsage
Model string // 原始模型(用于响应和日志显示)
// BillingModel is the model used for cost calculation.
// When non-empty, CalculateCost uses this instead of Model.
// This is set by the Anthropic Messages conversion path where
// the mapped upstream model differs from the client-facing model.
BillingModel string
// UpstreamModel is the actual model sent to the upstream provider after mapping.
// Empty when no mapping was applied (requested model was used as-is).
UpstreamModel string
// UpstreamResponseModel is captured from the raw successful upstream
// response before any client-facing rewrite or protocol conversion.
UpstreamResponseModel string
UpstreamResponseModelConflict bool
// UpstreamEndpoint is the actual upstream API path used for this request.
// It avoids guessing when one downstream protocol can use multiple upstream endpoints.
UpstreamEndpoint string
// ServiceTier records the OpenAI Responses API service tier, e.g. "priority" / "flex".
// Nil means the request did not specify a recognized tier.
ServiceTier *string
// ReasoningEffort is extracted from request body (reasoning.effort) or derived from model suffix.
// Stored for usage records display; nil means not provided / not applicable.
ReasoningEffort *string
Stream bool
OpenAIWSMode bool
// UpstreamTerminalEvent is the normalized terminal event observed on an
// upstream Responses WebSocket turn. Empty preserves legacy/non-WS success.
UpstreamTerminalEvent string
ResponseHeaders http.Header
Duration time.Duration
FirstTokenMs *int
ClientDisconnect bool
ImageCount int
ImageSize string
ImageInputSize string
ImageOutputSize string
ImageOutputSizes []string
ImageSizeSource string
ImageSizeBreakdown map[string]int
VideoCount int
VideoResolution string
// VideoDurationSeconds 是提交时请求的生成时长(xAI 按输出秒数计费),已归一化到 1-15 秒。
VideoDurationSeconds int
// WebSearchCalls 是 Codex alpha/search 网页搜索调用次数(每次成功请求为 1)。
// 上游不返回 usage 字段,>0 时走按次计费(分组单价 × 次数 × 倍率)。
WebSearchCalls int
// SearchCount is Grok-native web_search / tool search call count (per 1k pricing).
SearchCount int
// AudioUsage carries Voice billing units when present.
AudioUsage *AudioUsage
wsReplayInput []json.RawMessage
wsReplayInputExists bool
wsAccountFailoverReplayInput []json.RawMessage
}
// SucceededForScheduling reports whether this result is an upstream success
// that may clear model-scoped transient state. The zero value remains a success
// for existing non-WS callers.
func (r *OpenAIForwardResult) SucceededForScheduling() bool {
if r == nil || !r.OpenAIWSMode || r.UpstreamTerminalEvent == "" {
return true
}
switch r.UpstreamTerminalEvent {
case "response.completed", "response.done":
return true
default:
return false
}
}
// SetActualOpenAIUpstreamEndpoint records the endpoint selected by the current
// forwarding attempt. It covers error paths where no OpenAIForwardResult is
// available for usage and operations logging.
func SetActualOpenAIUpstreamEndpoint(c *gin.Context, endpoint string) {
if c == nil {
return
}
if endpoint = strings.TrimSpace(endpoint); endpoint != "" {
c.Set(openAIUpstreamEndpointContextKey, endpoint)
}
}
// GetActualOpenAIUpstreamEndpoint returns the endpoint recorded by the latest
// forwarding attempt in this request.
func GetActualOpenAIUpstreamEndpoint(c *gin.Context) string {
if c == nil {
return ""
}
value, exists := c.Get(openAIUpstreamEndpointContextKey)
if !exists {
return ""
}
endpoint, _ := value.(string)
return strings.TrimSpace(endpoint)
}
type OpenAIWSRetryMetricsSnapshot struct {
RetryAttemptsTotal int64 `json:"retry_attempts_total"`
RetryBackoffMsTotal int64 `json:"retry_backoff_ms_total"`
RetryExhaustedTotal int64 `json:"retry_exhausted_total"`
NonRetryableFastFallbackTotal int64 `json:"non_retryable_fast_fallback_total"`
}
type OpenAICompatibilityFallbackMetricsSnapshot struct {
SessionHashLegacyReadFallbackTotal int64 `json:"session_hash_legacy_read_fallback_total"`
SessionHashLegacyReadFallbackHit int64 `json:"session_hash_legacy_read_fallback_hit"`
SessionHashLegacyDualWriteTotal int64 `json:"session_hash_legacy_dual_write_total"`
SessionHashLegacyReadHitRate float64 `json:"session_hash_legacy_read_hit_rate"`
MetadataLegacyFallbackIsMaxTokensOneHaikuTotal int64 `json:"metadata_legacy_fallback_is_max_tokens_one_haiku_total"`
MetadataLegacyFallbackThinkingEnabledTotal int64 `json:"metadata_legacy_fallback_thinking_enabled_total"`
MetadataLegacyFallbackPrefetchedStickyAccount int64 `json:"metadata_legacy_fallback_prefetched_sticky_account_total"`
MetadataLegacyFallbackPrefetchedStickyGroup int64 `json:"metadata_legacy_fallback_prefetched_sticky_group_total"`
MetadataLegacyFallbackSingleAccountRetryTotal int64 `json:"metadata_legacy_fallback_single_account_retry_total"`
MetadataLegacyFallbackAccountSwitchCountTotal int64 `json:"metadata_legacy_fallback_account_switch_count_total"`
MetadataLegacyFallbackTotal int64 `json:"metadata_legacy_fallback_total"`
}
type openAIWSRetryMetrics struct {
retryAttempts atomic.Int64
retryBackoffMs atomic.Int64
retryExhausted atomic.Int64
nonRetryableFastFallback atomic.Int64
}
type accountWriteThrottle struct {
minInterval time.Duration
mu sync.Mutex
lastByID map[int64]time.Time
}
func newAccountWriteThrottle(minInterval time.Duration) *accountWriteThrottle {
return &accountWriteThrottle{
minInterval: minInterval,
lastByID: make(map[int64]time.Time),
}
}
func (t *accountWriteThrottle) Allow(id int64, now time.Time) bool {
if t == nil || id <= 0 || t.minInterval <= 0 {
return true
}
t.mu.Lock()
defer t.mu.Unlock()
if last, ok := t.lastByID[id]; ok && now.Sub(last) < t.minInterval {
return false
}
t.lastByID[id] = now
if len(t.lastByID) > 4096 {
cutoff := now.Add(-4 * t.minInterval)
for accountID, writtenAt := range t.lastByID {
if writtenAt.Before(cutoff) {
delete(t.lastByID, accountID)
}
}
}
return true
}
var defaultOpenAICodexSnapshotPersistThrottle = newAccountWriteThrottle(openAICodexSnapshotPersistMinInterval)
// ErrNoAvailableCompactAccounts indicates a legacy /responses/compact request
// needs compact support but no compatible account is available.
var ErrNoAvailableCompactAccounts = errors.New("no available accounts support /responses/compact")
// OpenAIGatewayService handles OpenAI API gateway operations
type OpenAIGatewayService struct {
accountRepo AccountRepository
usageLogRepo UsageLogRepository
usageBillingRepo UsageBillingRepository
userRepo UserRepository
userSubRepo UserSubscriptionRepository
cache GatewayCache
cfg *config.Config
codexDetector CodexClientRestrictionDetector
schedulerSnapshot *SchedulerSnapshotService
concurrencyService *ConcurrencyService
billingService *BillingService
rateLimitService *RateLimitService
billingCacheService *BillingCacheService
userGroupRateResolver *userGroupRateResolver
httpUpstream HTTPUpstream
deferredService *DeferredService
openAITokenProvider *OpenAITokenProvider
grokTokenProvider *GrokTokenProvider
toolCorrector *CodexToolCorrector
openaiWSResolver OpenAIWSProtocolResolver
resolver *ModelPricingResolver
channelService *ChannelService
balanceNotifyService *BalanceNotifyService
settingService *SettingService
userPlatformQuotaRepo UserPlatformQuotaRepository
liveAttestation liveattestation.Provider
liveAttestationCipher SecretEncryptor
openaiWSPoolOnce sync.Once
openaiWSStateStoreOnce sync.Once
openaiSchedulerOnce sync.Once
openaiProxyStreamCircuitOnce sync.Once
openaiWSPassthroughDialerOnce sync.Once
openaiModelTransientOnce sync.Once
agentIdentityTaskMu sync.Mutex
openaiWSPool *openAIWSConnPool
openaiWSStateStore OpenAIWSStateStore
openaiScheduler OpenAIAccountScheduler
openaiWSPassthroughDialer openAIWSClientDialer
openaiAccountStats *openAIAccountRuntimeStats
openaiModelTransient *openAIAccountModelTransientState
openaiProxyStreamCircuit *openAIProxyStreamCircuit
openaiProxyStreamFailOpenLogAt atomic.Int64
openaiWSFallbackUntil sync.Map // key: int64(accountID), value: time.Time
openaiAccountRuntimeBlockUntil sync.Map // key: int64(accountID), value: time.Time
openaiAccountRuntimeBlockLocks sync.Map // key: int64(accountID), value: *sync.Mutex
openaiAccountRuntimeBlockGeneration sync.Map // key: int64(accountID), value: uint64
openaiAccountRuntimeBlockSequence atomic.Uint64
grokCredentialMutationLocks sync.Map // key: int64(accountID), value: *sync.Mutex
openaiOAuth429WindowStartUnixNano atomic.Int64
openaiOAuth429WindowCount atomic.Int64
openaiWSRetryMetrics openAIWSRetryMetrics
responseHeaderFilter *responseheaders.CompiledHeaderFilter
codexSnapshotThrottle *accountWriteThrottle
codexModelsManifestCache codexModelsManifestCache
openaiCompatSessionResponses sync.Map
openaiCompatAnthropicDigestSessions sync.Map
// openaiCodexTurnStateOrigins: 下游会话 seed → openAICodexTurnStateOrigin
// 记录最近一次向该会话下发 x-codex-turn-state 的铸造账号,供出站守卫
// 剥离跨账号回带(openai_codex_turn_state.go)。
openaiCodexTurnStateOrigins sync.Map
openaiCodexTurnStateWrites atomic.Uint64
}
// NewOpenAIGatewayService creates a new OpenAIGatewayService
func NewOpenAIGatewayService(
accountRepo AccountRepository,
usageLogRepo UsageLogRepository,
usageBillingRepo UsageBillingRepository,
userRepo UserRepository,
userSubRepo UserSubscriptionRepository,
userGroupRateRepo UserGroupRateRepository,
cache GatewayCache,
cfg *config.Config,
schedulerSnapshot *SchedulerSnapshotService,
concurrencyService *ConcurrencyService,
billingService *BillingService,
rateLimitService *RateLimitService,
billingCacheService *BillingCacheService,
httpUpstream HTTPUpstream,
deferredService *DeferredService,
openAITokenProvider *OpenAITokenProvider,
grokTokenProvider *GrokTokenProvider,
resolver *ModelPricingResolver,
channelService *ChannelService,
balanceNotifyService *BalanceNotifyService,
settingService *SettingService,
userPlatformQuotaRepo UserPlatformQuotaRepository,
) *OpenAIGatewayService {
// enforceCodexIdentityHeaders 是 HTTP / 透传 / WS / 探针 等出站路径共用的纯函数收口点,
// 拿不到配置,故在此发布进程级开关快照。配置取反义,零值即「强制统一出口开启」。
if cfg != nil {
SetCodexIdentityEnforcementEnabled(!cfg.Gateway.DisableCodexIdentityEnforcement)
}
svc := &OpenAIGatewayService{
accountRepo: accountRepo,
usageLogRepo: usageLogRepo,
usageBillingRepo: usageBillingRepo,
userRepo: userRepo,
userSubRepo: userSubRepo,
cache: cache,
cfg: cfg,
codexDetector: NewOpenAICodexClientRestrictionDetector(cfg),
schedulerSnapshot: schedulerSnapshot,
concurrencyService: concurrencyService,
billingService: billingService,
rateLimitService: rateLimitService,
billingCacheService: billingCacheService,
userGroupRateResolver: newUserGroupRateResolver(
userGroupRateRepo,
nil,
resolveUserGroupRateCacheTTL(cfg),
nil,
"service.openai_gateway",
),
httpUpstream: httpUpstream,
deferredService: deferredService,
openAITokenProvider: openAITokenProvider,
grokTokenProvider: grokTokenProvider,
toolCorrector: NewCodexToolCorrector(),
openaiWSResolver: NewOpenAIWSProtocolResolver(cfg),
resolver: resolver,
channelService: channelService,
balanceNotifyService: balanceNotifyService,
settingService: settingService,
userPlatformQuotaRepo: userPlatformQuotaRepo,
liveAttestation: liveattestation.NewProvider(),
liveAttestationCipher: newLiveAttestationCipher(cfg),
responseHeaderFilter: compileResponseHeaderFilter(cfg),
codexSnapshotThrottle: newAccountWriteThrottle(openAICodexSnapshotPersistMinInterval),
openaiModelTransient: newOpenAIAccountModelTransientState(openAIModelTransientDefaultMax),
}
if rateLimitService != nil {
rateLimitService.SetAccountRuntimeBlocker(svc)
}
if openAITokenProvider != nil {
openAITokenProvider.SetAccountRuntimeBlocker(svc)
}
svc.logOpenAIWSModeBootstrap()
return svc
}
// ResolveChannelMapping 解析渠道级模型映射(代理到 ChannelService
func (s *OpenAIGatewayService) ResolveChannelMapping(ctx context.Context, groupID int64, model string) ChannelMappingResult {
if s.channelService == nil {
return ChannelMappingResult{MappedModel: model}
}
return s.channelService.ResolveChannelMapping(ctx, groupID, model)
}
// IsModelRestricted 检查模型是否被渠道限制(代理到 ChannelService
func (s *OpenAIGatewayService) IsModelRestricted(ctx context.Context, groupID int64, model string) bool {
if s.channelService == nil {
return false
}
return s.channelService.IsModelRestricted(ctx, groupID, model)
}
// ResolveChannelMappingAndRestrict 解析渠道映射。
// 模型限制检查已移至调度阶段,restricted 始终返回 false。
func (s *OpenAIGatewayService) ResolveChannelMappingAndRestrict(ctx context.Context, groupID *int64, model string) (ChannelMappingResult, bool) {
if s.channelService == nil {
return ChannelMappingResult{MappedModel: model}, false
}
return s.channelService.ResolveChannelMappingAndRestrict(ctx, groupID, model)
}
func (s *OpenAIGatewayService) isCodexImageGenerationBridgeEnabled(ctx context.Context, account *Account, apiKey *APIKey) bool {
if override := account.CodexImageGenerationBridgeOverride(); override != nil {
return *override
}
if s != nil && s.channelService != nil && apiKey != nil && apiKey.GroupID != nil {
ch, err := s.channelService.GetChannelForGroup(ctx, *apiKey.GroupID)
if err != nil {
slog.Warn("failed to resolve codex image generation bridge channel override", "group_id", *apiKey.GroupID, "error", err)
} else if override := ch.CodexImageGenerationBridgeOverride(PlatformOpenAI); override != nil {
return *override
}
}
return s != nil && s.cfg != nil && s.cfg.Gateway.CodexImageGenerationBridgeEnabled
}
func (s *OpenAIGatewayService) checkChannelPricingRestriction(ctx context.Context, groupID *int64, requestedModel string) bool {
if groupID == nil || s.channelService == nil || requestedModel == "" {
return false
}
mapping := s.channelService.ResolveChannelMapping(ctx, *groupID, requestedModel)
billingModel := billingModelForRestriction(mapping.BillingModelSource, requestedModel, mapping.MappedModel)
if billingModel == "" {
return false
}
return s.channelService.IsModelRestricted(ctx, *groupID, billingModel)
}
func (s *OpenAIGatewayService) isUpstreamModelRestrictedByChannel(ctx context.Context, groupID int64, account *Account, requestedModel string, requireCompact bool) bool {
if s.channelService == nil {
return false
}
if compactForwardModel, ok := openAIForwardModelFromContext(ctx); ok {
requestedModel = compactForwardModel.model
requireCompact = compactForwardModel.useCompactModelMapping
}
upstreamModel := resolveOpenAIAccountUpstreamModelForRequest(account, requestedModel, requireCompact)
if upstreamModel == "" {
return false
}
return s.channelService.IsModelRestricted(ctx, groupID, upstreamModel)
}
func (s *OpenAIGatewayService) needsUpstreamChannelRestrictionCheck(ctx context.Context, groupID *int64) bool {
if groupID == nil || s.channelService == nil {
return false
}
ch, err := s.channelService.GetChannelForGroup(ctx, *groupID)
if err != nil {
slog.Warn("failed to check openai channel upstream restriction", "group_id", *groupID, "error", err)
return false
}
if ch == nil || !ch.RestrictModels {
return false
}
return ch.BillingModelSource == BillingModelSourceUpstream
}
// ReplaceModelInBody 替换请求体中的 JSON model 字段(通用 gjson/sjson 实现)。
func (s *OpenAIGatewayService) ReplaceModelInBody(body []byte, newModel string) []byte {
return ReplaceModelInBody(body, newModel)
}
func (s *OpenAIGatewayService) getCodexSnapshotThrottle() *accountWriteThrottle {
if s != nil && s.codexSnapshotThrottle != nil {
return s.codexSnapshotThrottle
}
return defaultOpenAICodexSnapshotPersistThrottle
}
func (s *OpenAIGatewayService) billingDeps() *billingDeps {
return &billingDeps{
accountRepo: s.accountRepo,
userRepo: s.userRepo,
userSubRepo: s.userSubRepo,
billingCacheService: s.billingCacheService,
deferredService: s.deferredService,
balanceNotifyService: s.balanceNotifyService,
userPlatformQuotaRepo: s.userPlatformQuotaRepo,
}
}
// CloseOpenAIWSPool 关闭 OpenAI WebSocket 连接池的后台 worker 和空闲连接。
// 应在应用优雅关闭时调用。
func (s *OpenAIGatewayService) CloseOpenAIWSPool() {
if s != nil && s.openaiWSPool != nil {
s.openaiWSPool.Close()
}
}
func (s *OpenAIGatewayService) InvalidateAgentIdentityWSConnections(accountID int64) {
if pool := s.getOpenAIWSConnPool(); pool != nil {
pool.ClearAccount(accountID)
}
}
func (s *OpenAIGatewayService) logOpenAIWSModeBootstrap() {
if s == nil || s.cfg == nil {
return
}
wsCfg := s.cfg.Gateway.OpenAIWS
logOpenAIWSModeInfo(
"bootstrap enabled=%v oauth_enabled=%v apikey_enabled=%v force_http=%v responses_websockets_v2=%v responses_websockets=%v payload_log_sample_rate=%.3f event_flush_batch_size=%d event_flush_interval_ms=%d prewarm_cooldown_ms=%d retry_backoff_initial_ms=%d retry_backoff_max_ms=%d retry_jitter_ratio=%.3f retry_total_budget_ms=%d ws_read_limit_bytes=%d",
wsCfg.Enabled,
wsCfg.OAuthEnabled,
wsCfg.APIKeyEnabled,
wsCfg.ForceHTTP,
wsCfg.ResponsesWebsocketsV2,
wsCfg.ResponsesWebsockets,
wsCfg.PayloadLogSampleRate,
wsCfg.EventFlushBatchSize,
wsCfg.EventFlushIntervalMS,
wsCfg.PrewarmCooldownMS,
wsCfg.RetryBackoffInitialMS,
wsCfg.RetryBackoffMaxMS,
wsCfg.RetryJitterRatio,
wsCfg.RetryTotalBudgetMS,
openAIWSMessageReadLimitBytes,
)
}
func (s *OpenAIGatewayService) getCodexClientRestrictionDetector() CodexClientRestrictionDetector {
if s != nil && s.codexDetector != nil {
return s.codexDetector
}
var cfg *config.Config
if s != nil {
cfg = s.cfg
}
return NewOpenAICodexClientRestrictionDetector(cfg)
}
func (s *OpenAIGatewayService) getOpenAIWSProtocolResolver() OpenAIWSProtocolResolver {
if s != nil && s.openaiWSResolver != nil {
return s.openaiWSResolver
}
var cfg *config.Config
if s != nil {
cfg = s.cfg
}
return NewOpenAIWSProtocolResolver(cfg)
}
func classifyOpenAIWSReconnectReason(err error) (string, bool) {
if err == nil {
return "", false
}
var fallbackErr *openAIWSFallbackError
if !errors.As(err, &fallbackErr) || fallbackErr == nil {
return "", false
}
reason := strings.TrimSpace(fallbackErr.Reason)
if reason == "" {
return "", false
}
baseReason := strings.TrimPrefix(reason, "prewarm_")
switch baseReason {
case "policy_violation",
"message_too_big",
"upgrade_required",
"ws_unsupported",
"auth_failed",
"invalid_encrypted_content",
"previous_response_not_found":
return reason, false
}
switch baseReason {
case "read_event",
"write_request",
"write",
"acquire_timeout",
"acquire_conn",
"conn_queue_full",
"dial_failed",
"upstream_5xx",
"event_error",
"error_event",
"upstream_error_event",
"ws_connection_limit_reached",
"missing_final_response":
return reason, true
default:
return reason, false
}
}
func resolveOpenAIWSFallbackErrorResponse(err error) (statusCode int, errType string, clientMessage string, upstreamMessage string, ok bool) {
if err == nil {
return 0, "", "", "", false
}
var fallbackErr *openAIWSFallbackError
if !errors.As(err, &fallbackErr) || fallbackErr == nil {
return 0, "", "", "", false
}
reason := strings.TrimSpace(fallbackErr.Reason)
reason = strings.TrimPrefix(reason, "prewarm_")
if reason == "" {
return 0, "", "", "", false
}
var dialErr *openAIWSDialError
if fallbackErr.Err != nil && errors.As(fallbackErr.Err, &dialErr) && dialErr != nil {
if dialErr.StatusCode > 0 {
statusCode = dialErr.StatusCode
}
if dialErr.Err != nil {
upstreamMessage = sanitizeUpstreamErrorMessage(strings.TrimSpace(dialErr.Err.Error()))
}
}
switch reason {
case "invalid_encrypted_content":
if statusCode == 0 {
statusCode = http.StatusBadRequest
}
errType = "invalid_request_error"
if upstreamMessage == "" {
upstreamMessage = "encrypted content could not be verified"
}
case "previous_response_not_found":
if statusCode == 0 {
statusCode = http.StatusBadRequest
}
errType = "invalid_request_error"
if upstreamMessage == "" {
upstreamMessage = "previous response not found"
}
case "upgrade_required":
if statusCode == 0 {
statusCode = http.StatusUpgradeRequired
}
case "ws_unsupported":
if statusCode == 0 {
statusCode = http.StatusBadRequest
}
case "auth_failed":
if statusCode == 0 {
statusCode = http.StatusUnauthorized
}
case "upstream_rate_limited":
if statusCode == 0 {
statusCode = http.StatusTooManyRequests
}
default:
if statusCode == 0 {
return 0, "", "", "", false
}
}
if upstreamMessage == "" && fallbackErr.Err != nil {
upstreamMessage = sanitizeUpstreamErrorMessage(strings.TrimSpace(fallbackErr.Err.Error()))
}
if upstreamMessage == "" {
switch reason {
case "upgrade_required":
upstreamMessage = "upstream websocket upgrade required"
case "ws_unsupported":
upstreamMessage = "upstream websocket not supported"
case "auth_failed":
upstreamMessage = "upstream authentication failed"
case "upstream_rate_limited":
upstreamMessage = "upstream rate limit exceeded, please retry later"
default:
upstreamMessage = "Upstream request failed"
}
}
if errType == "" {
if statusCode == http.StatusTooManyRequests {
errType = "rate_limit_error"
} else {
errType = "upstream_error"
}
}
clientMessage = upstreamMessage
return statusCode, errType, clientMessage, upstreamMessage, true
}
func (s *OpenAIGatewayService) writeOpenAIWSFallbackErrorResponse(c *gin.Context, account *Account, wsErr error) bool {
if c == nil || c.Writer == nil || c.Writer.Written() {
return false
}
statusCode, errType, clientMessage, upstreamMessage, ok := resolveOpenAIWSFallbackErrorResponse(wsErr)
if !ok {
return false
}
if strings.TrimSpace(clientMessage) == "" {
clientMessage = "Upstream request failed"
}
if strings.TrimSpace(upstreamMessage) == "" {
upstreamMessage = clientMessage
}
setOpsUpstreamError(c, statusCode, upstreamMessage, "")
if account != nil {
appendOpsUpstreamError(c, OpsUpstreamErrorEvent{
Platform: account.Platform,
AccountID: account.ID,
AccountName: account.Name,
UpstreamStatusCode: statusCode,
Kind: "ws_error",
Message: upstreamMessage,
})
}
c.JSON(statusCode, gin.H{
"error": gin.H{
"type": errType,
"message": clientMessage,
},
})
return true
}
func (s *OpenAIGatewayService) openAIWSRetryBackoff(attempt int) time.Duration {
if attempt <= 0 {
return 0
}
initial := openAIWSRetryBackoffInitialDefault
maxBackoff := openAIWSRetryBackoffMaxDefault
jitterRatio := openAIWSRetryJitterRatioDefault
if s != nil && s.cfg != nil {
wsCfg := s.cfg.Gateway.OpenAIWS
if wsCfg.RetryBackoffInitialMS > 0 {
initial = time.Duration(wsCfg.RetryBackoffInitialMS) * time.Millisecond
}
if wsCfg.RetryBackoffMaxMS > 0 {
maxBackoff = time.Duration(wsCfg.RetryBackoffMaxMS) * time.Millisecond
}
if wsCfg.RetryJitterRatio >= 0 {
jitterRatio = wsCfg.RetryJitterRatio
}
}
if initial <= 0 {
return 0
}
if maxBackoff <= 0 {
maxBackoff = initial
}
if maxBackoff < initial {
maxBackoff = initial
}
if jitterRatio < 0 {
jitterRatio = 0
}
if jitterRatio > 1 {
jitterRatio = 1
}
shift := attempt - 1
if shift < 0 {
shift = 0
}
backoff := initial
if shift > 0 {
backoff = initial * time.Duration(1<<shift)
}
if backoff > maxBackoff {
backoff = maxBackoff
}
if jitterRatio <= 0 {
return backoff
}
jitter := time.Duration(float64(backoff) * jitterRatio)
if jitter <= 0 {
return backoff
}
delta := time.Duration(rand.Int63n(int64(jitter)*2+1)) - jitter
withJitter := backoff + delta
if withJitter < 0 {
return 0
}
return withJitter
}
func (s *OpenAIGatewayService) openAIWSRetryTotalBudget() time.Duration {
if s != nil && s.cfg != nil {
ms := s.cfg.Gateway.OpenAIWS.RetryTotalBudgetMS
if ms <= 0 {
return 0
}
return time.Duration(ms) * time.Millisecond
}
return 0
}
func (s *OpenAIGatewayService) recordOpenAIWSRetryAttempt(backoff time.Duration) {
if s == nil {
return
}
s.openaiWSRetryMetrics.retryAttempts.Add(1)
if backoff > 0 {
s.openaiWSRetryMetrics.retryBackoffMs.Add(backoff.Milliseconds())
}
}
func (s *OpenAIGatewayService) recordOpenAIWSRetryExhausted() {
if s == nil {
return
}
s.openaiWSRetryMetrics.retryExhausted.Add(1)
}
func (s *OpenAIGatewayService) recordOpenAIWSNonRetryableFastFallback() {
if s == nil {
return
}
s.openaiWSRetryMetrics.nonRetryableFastFallback.Add(1)
}
func (s *OpenAIGatewayService) SnapshotOpenAIWSRetryMetrics() OpenAIWSRetryMetricsSnapshot {
if s == nil {
return OpenAIWSRetryMetricsSnapshot{}
}
return OpenAIWSRetryMetricsSnapshot{
RetryAttemptsTotal: s.openaiWSRetryMetrics.retryAttempts.Load(),
RetryBackoffMsTotal: s.openaiWSRetryMetrics.retryBackoffMs.Load(),
RetryExhaustedTotal: s.openaiWSRetryMetrics.retryExhausted.Load(),
NonRetryableFastFallbackTotal: s.openaiWSRetryMetrics.nonRetryableFastFallback.Load(),
}
}
func SnapshotOpenAICompatibilityFallbackMetrics() OpenAICompatibilityFallbackMetricsSnapshot {
legacyReadFallbackTotal, legacyReadFallbackHit, legacyDualWriteTotal := openAIStickyCompatStats()
isMaxTokensOneHaiku, thinkingEnabled, prefetchedStickyAccount, prefetchedStickyGroup, singleAccountRetry, accountSwitchCount := RequestMetadataFallbackStats()
readHitRate := float64(0)
if legacyReadFallbackTotal > 0 {
readHitRate = float64(legacyReadFallbackHit) / float64(legacyReadFallbackTotal)
}
metadataFallbackTotal := isMaxTokensOneHaiku + thinkingEnabled + prefetchedStickyAccount + prefetchedStickyGroup + singleAccountRetry + accountSwitchCount
return OpenAICompatibilityFallbackMetricsSnapshot{
SessionHashLegacyReadFallbackTotal: legacyReadFallbackTotal,
SessionHashLegacyReadFallbackHit: legacyReadFallbackHit,
SessionHashLegacyDualWriteTotal: legacyDualWriteTotal,
SessionHashLegacyReadHitRate: readHitRate,
MetadataLegacyFallbackIsMaxTokensOneHaikuTotal: isMaxTokensOneHaiku,
MetadataLegacyFallbackThinkingEnabledTotal: thinkingEnabled,
MetadataLegacyFallbackPrefetchedStickyAccount: prefetchedStickyAccount,
MetadataLegacyFallbackPrefetchedStickyGroup: prefetchedStickyGroup,
MetadataLegacyFallbackSingleAccountRetryTotal: singleAccountRetry,
MetadataLegacyFallbackAccountSwitchCountTotal: accountSwitchCount,
MetadataLegacyFallbackTotal: metadataFallbackTotal,
}
}
func (s *OpenAIGatewayService) detectCodexClientRestriction(c *gin.Context, account *Account, body []byte) CodexClientRestrictionDetectionResult {
// 安全默认:即便缺 settingService(仅测试/误配可达)也保持指纹门为默认种子,
// 避免零值 policy(nil 信号)让指纹门失败开放。有 settingService 时整体覆盖为全局策略。
policy := CodexRestrictionPolicy{EngineFingerprintSignals: openai.DefaultEngineFingerprintSignals}
if account != nil && account.IsCodexCLIOnlyEnabled() && s != nil && s.settingService != nil {
ctx := context.Background()
if c != nil && c.Request != nil {
ctx = c.Request.Context()
}
policy = s.settingService.GetCodexRestrictionPolicy(ctx)
}
return s.getCodexClientRestrictionDetector().Detect(c, account, policy, body)
}
func getAPIKeyIDFromContext(c *gin.Context) int64 {
if c == nil {
return 0
}
v, exists := c.Get("api_key")
if !exists {
return 0
}
apiKey, ok := v.(*APIKey)
if !ok || apiKey == nil {
return 0
}
return apiKey.ID
}
// isolateOpenAISessionID 将 apiKeyID 混入 session 标识符,
// 确保不同 API Key 的用户即使使用相同的原始 session_id/conversation_id
// 到达上游的标识符也不同,防止跨用户会话碰撞。
func isolateOpenAISessionID(apiKeyID int64, raw string) string {
raw = strings.TrimSpace(raw)
if raw == "" {
return ""
}
h := xxhash.New()
_, _ = fmt.Fprintf(h, "k%d:", apiKeyID)
_, _ = h.WriteString(raw)
return fmt.Sprintf("%016x", h.Sum64())
}
func logCodexCLIOnlyDetection(ctx context.Context, c *gin.Context, account *Account, apiKeyID int64, result CodexClientRestrictionDetectionResult, body []byte) {
if !result.Enabled {
return
}
if ctx == nil {
ctx = context.Background()
}
accountID := int64(0)
if account != nil {
accountID = account.ID
}
fields := []zap.Field{
zap.String("component", "service.openai_gateway"),
zap.Int64("account_id", accountID),
zap.Bool("codex_cli_only_enabled", result.Enabled),
zap.Bool("codex_official_client_match", result.Matched),
zap.String("reject_reason", result.Reason),
}
if apiKeyID > 0 {
fields = append(fields, zap.Int64("api_key_id", apiKeyID))
}
if !result.Matched {
fields = appendCodexCLIOnlyRejectedRequestFields(fields, c, body)
}
log := logger.FromContext(ctx).With(fields...)
if result.Matched {
log.Info("OpenAI codex_cli_only 放行请求")
return
}
log.Warn("OpenAI codex_cli_only 拒绝非官方客户端请求")
}
func appendCodexCLIOnlyRejectedRequestFields(fields []zap.Field, c *gin.Context, body []byte) []zap.Field {
if c == nil || c.Request == nil {
return fields
}
req := c.Request
requestModel, requestStream, promptCacheKey := extractOpenAIRequestMetaFromBody(body)
fields = append(fields,
zap.String("request_method", strings.TrimSpace(req.Method)),
zap.String("request_path", strings.TrimSpace(req.URL.Path)),
zap.String("request_query", strings.TrimSpace(req.URL.RawQuery)),
zap.String("request_host", strings.TrimSpace(req.Host)),
zap.String("request_client_ip", strings.TrimSpace(ip.GetClientIP(c))),
zap.String("request_remote_addr", strings.TrimSpace(req.RemoteAddr)),
zap.String("request_user_agent", strings.TrimSpace(req.Header.Get("User-Agent"))),
zap.String("request_content_type", strings.TrimSpace(req.Header.Get("Content-Type"))),
zap.Int64("request_content_length", req.ContentLength),
zap.Bool("request_stream", requestStream),
)
if requestModel != "" {
fields = append(fields, zap.String("request_model", requestModel))
}
if promptCacheKey != "" {
fields = append(fields, zap.String("request_prompt_cache_key_sha256", hashSensitiveValueForLog(promptCacheKey)))
}
if headers := snapshotCodexCLIOnlyHeaders(req.Header); len(headers) > 0 {
fields = append(fields, zap.Any("request_headers", headers))
}
fields = append(fields, zap.Int("request_body_size", len(body)))
return fields
}
func snapshotCodexCLIOnlyHeaders(header http.Header) map[string]string {
if len(header) == 0 {
return nil
}
result := make(map[string]string, len(codexCLIOnlyDebugHeaderWhitelist))
for _, key := range codexCLIOnlyDebugHeaderWhitelist {
value := strings.TrimSpace(header.Get(key))
if value == "" {
continue
}
result[strings.ToLower(key)] = truncateString(value, codexCLIOnlyHeaderValueMaxBytes)
}
return result
}
func hashSensitiveValueForLog(raw string) string {
value := strings.TrimSpace(raw)
if value == "" {
return ""
}
sum := sha256.Sum256([]byte(value))
return hex.EncodeToString(sum[:8])
}
// GetAccessToken gets the access token for an OpenAI account
func (s *OpenAIGatewayService) GetAccessToken(ctx context.Context, account *Account) (string, string, error) {
if account.IsShadow() {
credAccount, err := resolveCredentialAccount(ctx, s.accountRepo, account)
if err != nil {
return "", "", err
}
account = credAccount
}
switch account.Type {
case AccountTypeOAuth:
if account.IsOpenAIAgentIdentity() {
return "", OpenAIAuthModeAgentIdentity, nil
}
if account.Platform == PlatformGrok {
if s.grokTokenProvider != nil {
accessToken, err := s.grokTokenProvider.GetAccessToken(ctx, account)
if err != nil {
return "", "", err
}
return accessToken, "oauth", nil
}
accessToken := account.GetGrokAccessToken()
if accessToken == "" {
return "", "", errors.New("access_token not found in credentials")
}
return accessToken, "oauth", nil
}
// 使用 TokenProvider 获取缓存的 token
if s.openAITokenProvider != nil {
accessToken, err := s.openAITokenProvider.GetAccessToken(ctx, account)
if err != nil {
return "", "", err
}
return accessToken, "oauth", nil
}
// 降级:TokenProvider 未配置时直接从账号读取
accessToken := account.GetOpenAIAccessToken()
if accessToken == "" {
return "", "", errors.New("access_token not found in credentials")
}
return accessToken, "oauth", nil
case AccountTypeAPIKey:
if account.Platform == PlatformGrok {
apiKey := strings.TrimSpace(account.GetCredential("api_key"))
if apiKey == "" {
return "", "", errors.New("api_key not found in credentials")
}
return apiKey, "apikey", nil
}
apiKey := strings.TrimSpace(account.GetOpenAIProtocolAPIKey())
if apiKey == "" {
return "", "", errors.New("api_key not found in credentials")
}
return apiKey, "apikey", nil
default:
return "", "", fmt.Errorf("unsupported account type: %s", account.Type)
}
}