package service import ( "bytes" "context" "encoding/json" "errors" "fmt" "net/http" "strings" "github.com/Wei-Shaw/sub2api/internal/pkg/ctxkey" "github.com/Wei-Shaw/sub2api/internal/util/urlvalidator" "github.com/gin-gonic/gin" "github.com/google/uuid" "github.com/tidwall/gjson" "github.com/tidwall/sjson" ) func (s *OpenAIGatewayService) validateUpstreamBaseURL(raw string) (string, error) { if s.cfg != nil && !s.cfg.Security.URLAllowlist.Enabled { normalized, err := urlvalidator.ValidateURLFormat(raw, s.cfg.Security.URLAllowlist.AllowInsecureHTTP) if err != nil { return "", fmt.Errorf("invalid base_url: %w", err) } return normalized, nil } normalized, err := urlvalidator.ValidateHTTPSURL(raw, urlvalidator.ValidationOptions{ AllowedHosts: s.cfg.Security.URLAllowlist.UpstreamHosts, RequireAllowlist: true, AllowPrivate: s.cfg.Security.URLAllowlist.AllowPrivateHosts, }) if err != nil { return "", fmt.Errorf("invalid base_url: %w", err) } return normalized, nil } // buildOpenAIResponsesURL 组装 OpenAI Responses 端点。 // - base 以 /v1 结尾:追加 /responses // - base 以其他版本段结尾(如 /v4):追加 /responses // - base 已是 /responses:原样返回 // - 其他情况:追加 /v1/responses func buildOpenAIResponsesURL(base string) string { return buildOpenAIEndpointURL(base, "/v1/responses") } // buildOpenAIResponsesURLForPlatform 组装 Responses 端点(平台感知)。 // DeepSeek 官方 Responses 端点为 /responses(无 /v1 前缀,适配 Codex); // 其余平台维持 /v1/responses。 func buildOpenAIResponsesURLForPlatform(platform string, base string) string { if platform == PlatformDeepseek { return buildOpenAIEndpointURL(base, "/responses") } return buildOpenAIResponsesURL(base) } // normalizeDeepSeekResponsesRequestBody 适配 DeepSeek 无状态 Responses 端点: // 强制 store=false 并清除 previous_response_id(官方 /responses 不支持服务端 // 状态存储,携带这些字段会被拒绝)。非 deepseek responses 协议账号原样返回。 func normalizeDeepSeekResponsesRequestBody(account *Account, body []byte) []byte { if account == nil || account.Platform != PlatformDeepseek || (account.GetAPIProtocol() != APIProtocolResponses && !account.IsAdaptiveAPIProtocol()) { return body } normalized, err := sjson.SetBytes(body, "store", false) if err != nil { return body } if stripped, err := sjson.DeleteBytes(normalized, "previous_response_id"); err == nil { normalized = stripped } return normalized } func trimOpenAIEncryptedReasoningItems(reqBody map[string]any) bool { if len(reqBody) == 0 { return false } inputValue, has := reqBody["input"] if !has { return false } switch input := inputValue.(type) { case []any: filtered := input[:0] changed := false for _, item := range input { nextItem, itemChanged, keep := sanitizeEncryptedReasoningInputItem(item) if itemChanged { changed = true } if !keep { continue } filtered = append(filtered, nextItem) } if !changed { return false } if len(filtered) == 0 { delete(reqBody, "input") return true } reqBody["input"] = filtered return true case []map[string]any: filtered := input[:0] changed := false for _, item := range input { nextItem, itemChanged, keep := sanitizeEncryptedReasoningInputItem(item) if itemChanged { changed = true } if !keep { continue } nextMap, ok := nextItem.(map[string]any) if !ok { filtered = append(filtered, item) continue } filtered = append(filtered, nextMap) } if !changed { return false } if len(filtered) == 0 { delete(reqBody, "input") return true } reqBody["input"] = filtered return true case map[string]any: nextItem, changed, keep := sanitizeEncryptedReasoningInputItem(input) if !changed { return false } if !keep { delete(reqBody, "input") return true } nextMap, ok := nextItem.(map[string]any) if !ok { return false } reqBody["input"] = nextMap return true default: return false } } func sanitizeEncryptedReasoningInputItem(item any) (next any, changed bool, keep bool) { inputItem, ok := item.(map[string]any) if !ok { return item, false, true } itemType, _ := inputItem["type"].(string) switch strings.TrimSpace(itemType) { case "compaction", "compaction_summary": if _, encrypted := inputItem["encrypted_content"]; encrypted { return nil, true, false } return item, false, true case "reasoning": default: return item, false, true } if _, has := inputItem["encrypted_content"]; has { delete(inputItem, "encrypted_content") changed = true } // xAI 422: "content": null 导致 untagged enum 反序列化失败 if v, has := inputItem["content"]; has && v == nil { delete(inputItem, "content") changed = true } if !changed { return item, false, true } if len(inputItem) == 1 { return nil, true, false } return inputItem, true, true } // SanitizeOpenAICrossModeFailoverReasoning derives a failover attempt body from // the canonical request body by dropping provider-specific encrypted reasoning // input items in full (encrypted_content plus the coupled id/summary shape). // // This is the proactive counterpart to the reactive same-account // invalid_encrypted_content recovery in Forward: when a failover switches from an // OpenAI passthrough account (which forwards upstream-native encrypted reasoning, // e.g. Kiro) to a non-passthrough account (e.g. Bedrock Mantle) that rejects the // provider-specific reasoning IDs/shape, the whole reasoning item must go before // the request reaches the new upstream. Unlike trimOpenAIEncryptedReasoningItems, // which only strips the encrypted_content / null-content fields while preserving // the reasoning item's id and summary, this drops the entire item. // // The input slice is treated as immutable and is never mutated; a distinct slice // is returned only when changed is true. func SanitizeOpenAICrossModeFailoverReasoning(body []byte) (sanitized []byte, changed bool, err error) { if len(body) == 0 { return body, false, nil } if !gjson.GetBytes(body, "input").Exists() { return body, false, nil } var decoded map[string]any decoder := json.NewDecoder(bytes.NewReader(body)) decoder.UseNumber() if err := decoder.Decode(&decoded); err != nil { return body, false, fmt.Errorf("decode cross-mode failover body: %w", err) } if !dropOpenAIEncryptedReasoningInputItems(decoded) { return body, false, nil } out, marshalErr := marshalOpenAIUpstreamJSON(decoded) if marshalErr != nil { return body, false, fmt.Errorf("serialize cross-mode failover body: %w", marshalErr) } return out, true, nil } // dropOpenAIEncryptedReasoningInputItems removes reasoning input items that carry // provider-specific encrypted_content in full — including their coupled id and // summary — and reports whether anything changed. Contrast with // trimOpenAIEncryptedReasoningItems, which only strips fields while keeping the // reasoning item skeleton. func dropOpenAIEncryptedReasoningInputItems(reqBody map[string]any) bool { if len(reqBody) == 0 { return false } inputValue, has := reqBody["input"] if !has { return false } switch input := inputValue.(type) { case []any: filtered := input[:0] changed := false for _, item := range input { if isOpenAIEncryptedReasoningInputItem(item) { changed = true continue } filtered = append(filtered, item) } if !changed { return false } if len(filtered) == 0 { delete(reqBody, "input") return true } reqBody["input"] = filtered return true case []map[string]any: filtered := input[:0] changed := false for _, item := range input { if isOpenAIEncryptedReasoningInputItem(item) { changed = true continue } filtered = append(filtered, item) } if !changed { return false } if len(filtered) == 0 { delete(reqBody, "input") return true } reqBody["input"] = filtered return true case map[string]any: if isOpenAIEncryptedReasoningInputItem(input) { delete(reqBody, "input") return true } return false default: return false } } func isOpenAIEncryptedReasoningInputItem(item any) bool { inputItem, ok := item.(map[string]any) if !ok { return false } if itemType, _ := inputItem["type"].(string); strings.TrimSpace(itemType) != "reasoning" { return false } _, has := inputItem["encrypted_content"] return has } // IsOpenAIResponsesCompactPath reports whether the request targets the legacy // /responses/compact endpoint, including its forwardable subpaths. func IsOpenAIResponsesCompactPath(c *gin.Context) bool { return isOpenAIResponsesCompactPath(c) } func OpenAICompactSessionSeedKeyForTest() string { return openAICompactSessionSeedKey } func NormalizeOpenAICompactRequestBodyForTest(body []byte) ([]byte, bool, error) { return normalizeOpenAICompactRequestBody(body) } func isOpenAIResponsesCompactPath(c *gin.Context) bool { suffix := strings.TrimSpace(openAIResponsesRequestPathSuffix(c)) return suffix == "/compact" || strings.HasPrefix(suffix, "/compact/") } func normalizeOpenAICompactRequestBody(body []byte) ([]byte, bool, error) { if len(body) == 0 { return body, false, nil } normalized := []byte(`{}`) // Keep the current Codex /compact schema while still dropping request-scoped // fields such as prompt_cache_key, store, and stream. for _, field := range []string{ "model", "input", "instructions", "tools", "parallel_tool_calls", "reasoning", "service_tier", "text", "previous_response_id", } { value := gjson.GetBytes(body, field) if !value.Exists() { continue } next, err := sjson.SetRawBytes(normalized, field, []byte(value.Raw)) if err != nil { return body, false, fmt.Errorf("normalize compact body %s: %w", field, err) } normalized = next } if bytes.Equal(bytes.TrimSpace(body), bytes.TrimSpace(normalized)) { return body, false, nil } return normalized, true, nil } func normalizeOpenAICodexCompactReasoningEffortForAccount(c *gin.Context, account *Account, body []byte) ([]byte, bool, error) { if account == nil || !account.IsOpenAIOAuth() || !isOpenAIResponsesCompactPath(c) { return body, false, nil } requestedModel := strings.TrimSpace(gjson.GetBytes(body, "model").String()) effectiveModel := account.GetMappedModel(requestedModel) return normalizeOpenAICodexCompactReasoningEffort(body, effectiveModel) } func normalizeOpenAICodexCompactReasoningEffort(body []byte, effectiveModel string) ([]byte, bool, error) { if !isOpenAIGPT56Model(effectiveModel) || !strings.EqualFold(strings.TrimSpace(gjson.GetBytes(body, "reasoning.effort").String()), "max") { return body, false, nil } // Codex Ultra 在客户端编排层会下发 max;ChatGPT compact 端点目前只接受到 // xhigh。这里只降级 OpenAI OAuth 的 GPT-5.6 compact 子请求,普通 Responses、 // API Key 请求和其他平台的 OAuth 请求保留 max。 normalized, err := sjson.SetBytes(body, "reasoning.effort", "xhigh") if err != nil { return body, false, fmt.Errorf("normalize codex compact reasoning effort: %w", err) } return normalized, true, nil } func resolveOpenAICompactSessionID(c *gin.Context) string { if c != nil { if sessionID := strings.TrimSpace(c.GetHeader("session_id")); sessionID != "" { return sessionID } if conversationID := strings.TrimSpace(c.GetHeader("conversation_id")); conversationID != "" { return conversationID } if seed, ok := c.Get(openAICompactSessionSeedKey); ok { if seedStr, ok := seed.(string); ok && strings.TrimSpace(seedStr) != "" { return strings.TrimSpace(seedStr) } } } return uuid.NewString() } // openAIResponsesRequestPathSuffix 返回可拼接到上游 /responses URL 后面的子路径。 // 不可转发的子路径返回空串(退化为裸 /responses);真正的拒绝由入口守卫 // IsForwardableOpenAIResponsesRequestPath 负责。这样即便将来新增路由漏挂守卫, // 拼进上游 URL 的也只会是合规片段。 func openAIResponsesRequestPathSuffix(c *gin.Context) string { suffix, ok := sanitizedUpstreamPathSuffix(rawOpenAIResponsesRequestPathSuffix(c)) if !ok { return "" } return suffix } // IsForwardableOpenAIResponsesRequestPath 判断入站请求携带的 /responses 子路径 // 是否可以安全转发。路由层用它在鉴权后、调度前直接拒绝畸形子路径。 func IsForwardableOpenAIResponsesRequestPath(c *gin.Context) bool { _, ok := sanitizedUpstreamPathSuffix(rawOpenAIResponsesRequestPathSuffix(c)) return ok } // IsOpenAIResponsesInputTokensRequestPath reports whether the request targets // the native Responses input-token counting endpoint. func IsOpenAIResponsesInputTokensRequestPath(c *gin.Context) bool { return openAIResponsesRequestPathSuffix(c) == "/input_tokens" } // rawOpenAIResponsesRequestPathSuffix 仅做提取,不做任何安全判断。 func rawOpenAIResponsesRequestPathSuffix(c *gin.Context) string { if c == nil || c.Request == nil || c.Request.URL == nil { return "" } normalizedPath := strings.TrimRight(strings.TrimSpace(c.Request.URL.Path), "/") if normalizedPath == "" { return "" } idx := strings.LastIndex(normalizedPath, "/responses") if idx < 0 { return "" } suffix := normalizedPath[idx+len("/responses"):] if suffix == "" || suffix == "/" { return "" } if !strings.HasPrefix(suffix, "/") { return "" } return suffix } func appendOpenAIResponsesRequestPathSuffix(baseURL, suffix string) string { trimmedBase := strings.TrimRight(strings.TrimSpace(baseURL), "/") // 兜底:调用方漏了校验时,这里也不会把不合规的片段拼进上游 URL。 trimmedSuffix, ok := sanitizedUpstreamPathSuffix(suffix) if !ok || trimmedBase == "" || trimmedSuffix == "" { return trimmedBase } return trimmedBase + trimmedSuffix } func (s *OpenAIGatewayService) replaceModelInResponseBody(body []byte, fromModel, toModel string) []byte { // 使用 gjson/sjson 精确替换 model 字段,避免全量 JSON 反序列化 if m := gjson.GetBytes(body, "model"); m.Exists() && m.Str == fromModel { newBody, err := sjson.SetBytes(body, "model", toModel) if err != nil { return body } return newBody } return body } func getOpenAIReasoningEffortFromReqBody(reqBody map[string]any, requestedModel string) (value string, present bool) { if reqBody == nil { return "", false } // Primary: reasoning.effort if reasoning, ok := reqBody["reasoning"].(map[string]any); ok { if effort, ok := reasoning["effort"].(string); ok { return normalizeOpenAIReasoningEffortForModel(effort, requestedModel), true } } // Fallback: some clients may use a flat field. if effort, ok := reqBody["reasoning_effort"].(string); ok { return normalizeOpenAIReasoningEffortForModel(effort, requestedModel), true } return "", false } func deriveOpenAIReasoningEffortFromModel(model string) string { if strings.TrimSpace(model) == "" { return "" } modelID := strings.TrimSpace(model) if strings.Contains(modelID, "/") { parts := strings.Split(modelID, "/") modelID = parts[len(parts)-1] } parts := strings.FieldsFunc(strings.ToLower(modelID), func(r rune) bool { switch r { case '-', '_', ' ': return true default: return false } }) if len(parts) == 0 { return "" } return normalizeOpenAIReasoningEffortForModel(parts[len(parts)-1], modelID) } // deriveOpenAIReasoningEffortFromModelCandidates 依次对每个候选模型做后缀推导, // 返回第一个非空结果。 func deriveOpenAIReasoningEffortFromModelCandidates(models []string) string { for _, model := range models { if value := deriveOpenAIReasoningEffortFromModel(model); value != "" { return value } } return "" } type openAIRequestView struct { body []byte Model string Stream bool PromptCacheKey string PreviousResponseID string ServiceTier string ReasoningEffort string patches []openAIRequestPatch patchesDisabled bool } type openAIRequestPatch struct { path string delete bool value any } func newOpenAIRequestView(body []byte) openAIRequestView { if len(body) == 0 { return openAIRequestView{} } const ( modelField uint8 = 1 << iota streamField promptCacheKeyField previousResponseIDField serviceTierField reasoningField allRequestViewFields = modelField | streamField | promptCacheKeyField | previousResponseIDField | serviceTierField | reasoningField ) view := openAIRequestView{body: body} var seen uint8 // parseRawJSONView reads body without copying; view keeps body alive for extracted strings. parseRawJSONView(body).ForEach(func(key, value gjson.Result) bool { switch key.Str { case "model": if seen&modelField == 0 { view.Model = strings.TrimSpace(value.String()) seen |= modelField } case "stream": if seen&streamField == 0 { view.Stream = value.Bool() seen |= streamField } case "prompt_cache_key": if seen&promptCacheKeyField == 0 { view.PromptCacheKey = strings.TrimSpace(value.String()) seen |= promptCacheKeyField } case "previous_response_id": if seen&previousResponseIDField == 0 { view.PreviousResponseID = strings.TrimSpace(value.String()) seen |= previousResponseIDField } case "service_tier": if seen&serviceTierField == 0 { view.ServiceTier = strings.TrimSpace(value.String()) seen |= serviceTierField } case "reasoning": if seen&reasoningField == 0 { view.ReasoningEffort = strings.TrimSpace(value.Get("effort").String()) seen |= reasoningField } } return seen != allRequestViewFields }) return view } // Decode 保留阶段一既有 full-map 行为;后续阶段会把调用点下沉到复杂分支。 func (v openAIRequestView) Decode(c *gin.Context) (map[string]any, error) { return getOpenAIRequestBodyMap(c, v.body) } func (v *openAIRequestView) MarkPatchSet(path string, value any) { if v == nil || v.patchesDisabled { return } path = strings.TrimSpace(path) if !isSimpleOpenAIRequestPatchPath(path) { v.DisablePatches() return } v.patches = append(v.patches, openAIRequestPatch{path: path, value: value}) } func (v *openAIRequestView) MarkPatchDelete(path string) { if v == nil || v.patchesDisabled { return } path = strings.TrimSpace(path) if !isSimpleOpenAIRequestPatchPath(path) { v.DisablePatches() return } v.patches = append(v.patches, openAIRequestPatch{path: path, delete: true}) } func isSimpleOpenAIRequestPatchPath(path string) bool { if path == "" || strings.ContainsRune(path, '\\') { return false } for _, part := range strings.Split(path, ".") { if strings.TrimSpace(part) == "" { return false } } return true } func (v *openAIRequestView) DisablePatches() { if v == nil { return } v.patchesDisabled = true v.patches = nil } func (v openAIRequestView) HasPatches() bool { return !v.patchesDisabled && len(v.patches) > 0 } func (v openAIRequestView) ApplyPatches() ([]byte, error) { if v.patchesDisabled || len(v.patches) == 0 { return nil, errors.New("openai request patches disabled") } body := v.body for _, patch := range v.patches { var err error if patch.delete { body, err = sjson.DeleteBytes(body, patch.path) } else { body, err = sjson.SetBytes(body, patch.path, patch.value) } if err != nil { return nil, err } } return body, nil } func setOpenAIRequestMapPath(reqBody map[string]any, path string, value any) { path = strings.TrimSpace(path) if reqBody == nil || path == "" { return } parts := strings.Split(path, ".") current := reqBody for _, part := range parts[:len(parts)-1] { part = strings.TrimSpace(part) if part == "" { return } next, _ := current[part].(map[string]any) if next == nil { next = map[string]any{} current[part] = next } current = next } last := strings.TrimSpace(parts[len(parts)-1]) if last != "" { current[last] = value } } func deleteOpenAIRequestMapPath(reqBody map[string]any, path string) { path = strings.TrimSpace(path) if reqBody == nil || path == "" { return } parts := strings.Split(path, ".") current := reqBody for _, part := range parts[:len(parts)-1] { part = strings.TrimSpace(part) if part == "" { return } next, _ := current[part].(map[string]any) if next == nil { return } current = next } last := strings.TrimSpace(parts[len(parts)-1]) if last != "" { delete(current, last) } } func extractOpenAIRequestMetaFromBody(body []byte) (model string, stream bool, promptCacheKey string) { view := newOpenAIRequestView(body) return view.Model, view.Stream, view.PromptCacheKey } // normalizeOpenAIPassthroughOAuthBody 将透传 OAuth 请求体收敛为旧链路关键行为: // 1) 删除 ChatGPT internal API 不支持的顶层 Responses 参数 // 2) store=false 3) 非 compact 保持 stream=true;compact 强制 stream=false func normalizeOpenAIPassthroughOAuthBody(body []byte, compact bool) ([]byte, bool, error) { if len(body) == 0 { return body, false, nil } normalized := body changed := false for _, field := range openAIChatGPTInternalUnsupportedFields { if value := gjson.GetBytes(normalized, field); !value.Exists() { continue } next, err := sjson.DeleteBytes(normalized, field) if err != nil { return body, false, fmt.Errorf("normalize passthrough body delete %s: %w", field, err) } normalized = next changed = true } if inputResult := gjson.GetBytes(normalized, "input"); inputResult.Exists() { switch { case inputResult.Type == gjson.String: text := inputResult.String() var inputValue any if strings.TrimSpace(text) != "" { inputValue = []any{map[string]any{ "type": "message", "role": "user", "content": text, }} } else { inputValue = []any{} } next, err := sjson.SetBytes(normalized, "input", inputValue) if err != nil { return body, false, fmt.Errorf("normalize passthrough body input string: %w", err) } normalized = next changed = true case inputResult.Type == gjson.JSON && !inputResult.IsArray(): next, err := sjson.SetRawBytes(normalized, "input", []byte("["+inputResult.Raw+"]")) if err != nil { return body, false, fmt.Errorf("normalize passthrough body input object: %w", err) } normalized = next changed = true } } if compact { if store := gjson.GetBytes(normalized, "store"); store.Exists() { next, err := sjson.DeleteBytes(normalized, "store") if err != nil { return body, false, fmt.Errorf("normalize passthrough body delete store: %w", err) } normalized = next changed = true } if stream := gjson.GetBytes(normalized, "stream"); stream.Exists() { next, err := sjson.DeleteBytes(normalized, "stream") if err != nil { return body, false, fmt.Errorf("normalize passthrough body delete stream: %w", err) } normalized = next changed = true } } else { if store := gjson.GetBytes(normalized, "store"); !store.Exists() || store.Type != gjson.False { next, err := sjson.SetBytes(normalized, "store", false) if err != nil { return body, false, fmt.Errorf("normalize passthrough body store=false: %w", err) } normalized = next changed = true } if stream := gjson.GetBytes(normalized, "stream"); !stream.Exists() || stream.Type != gjson.True { next, err := sjson.SetBytes(normalized, "stream", true) if err != nil { return body, false, fmt.Errorf("normalize passthrough body stream=true: %w", err) } normalized = next changed = true } } return normalized, changed, nil } func detectOpenAIPassthroughInstructionsRejectReason(reqModel string, body []byte) string { if !isOpenAICodexModel(reqModel) { return "" } instructions := gjson.GetBytes(body, "instructions") if !instructions.Exists() { return "" } if instructions.Type != gjson.String { return "instructions_not_string" } if strings.TrimSpace(instructions.String()) == "" { return "instructions_empty" } return "" } func isOpenAICodexModel(model string) bool { return strings.Contains(strings.ToLower(strings.TrimSpace(model)), "codex") } // extractOpenAIReasoningEffortFromBody 按优先级传入模型候选(如 upstreamModel, // billingModel, originalModel):显式 effort 的模型归一化(max 保留判定)用第一个 // 非空候选;body 未携带 effort 时的模型后缀推导依次尝试每个候选——OAuth 的 // normalizeCodexModel 会剥掉 upstreamModel 的 effort 后缀,只有原始模型名还留着。 func extractOpenAIReasoningEffortFromBody(body []byte, modelCandidates ...string) *string { reasoningEffort := strings.TrimSpace(gjson.GetBytes(body, "reasoning.effort").String()) if reasoningEffort == "" { reasoningEffort = strings.TrimSpace(gjson.GetBytes(body, "reasoning_effort").String()) } if reasoningEffort != "" { normalized := normalizeOpenAIReasoningEffortForModel(reasoningEffort, firstNonEmpty(modelCandidates...)) if normalized == "" { return nil } return &normalized } value := deriveOpenAIReasoningEffortFromModelCandidates(modelCandidates) if value == "" { return nil } return &value } func extractOpenAIServiceTier(reqBody map[string]any) *string { if reqBody == nil { return nil } raw, ok := reqBody["service_tier"].(string) if !ok { return nil } return normalizeOpenAIServiceTier(raw) } func extractOpenAIServiceTierFromBody(body []byte) *string { if len(body) == 0 { return nil } return normalizeOpenAIServiceTier(gjson.GetBytes(body, "service_tier").String()) } func normalizeOpenAIServiceTier(raw string) *string { value := strings.ToLower(strings.TrimSpace(raw)) if value == "" { return nil } if value == "fast" { value = "priority" } // 放过 OpenAI 官方文档定义的所有合法 tier 值:priority/flex/auto/default/scale。 // 对 Codex 客户端零影响(Codex 只发 priority 或 flex,见 codex-rs/core/src/client.rs), // 但能让直连 OpenAI SDK 的用户透传 auto/default/scale 以便抓包/调试。 // 真未知值仍返回 nil,由 normalizeResponsesBodyServiceTier 从 body 中删除。 switch value { case "priority", "flex", "auto", "default", "scale": return &value default: return nil } } // OpenAIFastBlockedError indicates a request was rejected by the OpenAI fast // policy (action=block). Mirrors BetaBlockedError on the Claude side. type OpenAIFastBlockedError struct { Message string } func (e *OpenAIFastBlockedError) Error() string { return e.Message } // evaluateOpenAIFastPolicy returns the action and error message that should be // applied for a request with the given account/model/service_tier. When the // policy service is unavailable or no rule matches, it returns // (BetaPolicyActionPass, "") so callers can short-circuit safely. // // Matching rules: // - Scope filters by account type (all / oauth / apikey / bedrock) // - UserIDs, when present, filters by the trusted Sub2API user that owns the API key // - ServiceTier must be empty (= any), "all", or equal the normalized tier // - ModelWhitelist narrows the rule to specific models; FallbackAction // handles the non-matching case (default: pass) // - User-specific rules take precedence over global rules; each group keeps // the configured first-match order // // 与 Claude BetaPolicy 的差异(保留首条匹配 short-circuit): // - BetaPolicy 处理的是 anthropic-beta header 中的 token 集合,不同 // 规则可能针对不同 token,filter 需要累加成 set;block 则 first-match。 // - OpenAI fast policy 操作的是单个字段 service_tier:filter 即删字段, // 没有可累加的对象。一次请求只携带一个 service_tier,规则的 tier // 维度天然互斥;同一 (scope, tier) 下若多条规则的 model whitelist // 发生重叠,admin 可通过规则顺序明确意图。因此采用 first-match 而 // 非 BetaPolicy 那样的"block 覆盖 filter 覆盖 pass"语义。 func (s *OpenAIGatewayService) evaluateOpenAIFastPolicy(ctx context.Context, account *Account, model, serviceTier string) (action, errMsg string) { if s == nil || s.settingService == nil { return BetaPolicyActionPass, "" } tier := strings.ToLower(strings.TrimSpace(serviceTier)) if tier == "" { return BetaPolicyActionPass, "" } settings := openAIFastPolicySettingsFromContext(ctx) if settings == nil { fetched, err := s.settingService.GetOpenAIFastPolicySettings(ctx) if err != nil || fetched == nil { return BetaPolicyActionPass, "" } settings = fetched } return evaluateOpenAIFastPolicyWithSettings(settings, openAIFastPolicyUserID(ctx), account, model, tier) } // evaluateOpenAIFastPolicyWithSettings is the pure-function core extracted so // long-lived sessions (e.g. WS) can prefetch settings once and avoid hitting // the settingService on every frame. See WSSession entry and // openAIFastPolicySettingsFromContext for the caching glue. func evaluateOpenAIFastPolicyWithSettings(settings *OpenAIFastPolicySettings, userID int64, account *Account, model, tier string) (action, errMsg string) { if settings == nil { return BetaPolicyActionPass, "" } isOAuth := account != nil && account.IsOAuth() isBedrock := account != nil && account.IsBedrock() // 用户专属规则先于全局规则。规则组内仍按配置顺序首条命中,允许 // 管理员为某位用户配置例外,而不被先出现的全局规则覆盖。 for _, userScoped := range []bool{true, false} { for _, rule := range settings.Rules { if (len(rule.UserIDs) > 0) != userScoped || !openAIFastPolicyUserMatches(rule.UserIDs, userID) { continue } if !betaPolicyScopeMatches(rule.Scope, isOAuth, isBedrock) { continue } ruleTier := strings.ToLower(strings.TrimSpace(rule.ServiceTier)) if ruleTier != "" && ruleTier != OpenAIFastTierAny && ruleTier != tier { continue } eff := BetaPolicyRule{ Action: rule.Action, ErrorMessage: rule.ErrorMessage, ModelWhitelist: rule.ModelWhitelist, FallbackAction: rule.FallbackAction, FallbackErrorMessage: rule.FallbackErrorMessage, } return resolveRuleAction(eff, model) } } return BetaPolicyActionPass, "" } func openAIFastPolicyUserID(ctx context.Context) int64 { if ctx == nil { return 0 } userID, _ := ctx.Value(ctxkey.UserID).(int64) if userID <= 0 { return 0 } return userID } func openAIFastPolicyUserMatches(ruleUserIDs []int64, userID int64) bool { if len(ruleUserIDs) == 0 { return true } for _, ruleUserID := range ruleUserIDs { if ruleUserID == userID { return true } } return false } // openAIFastPolicyCtxKey 是 context 中预取的 OpenAIFastPolicySettings 缓存 // 键,仅用于 WebSocket 长会话内多帧复用同一份策略快照,避免每帧 DB 命中。 // // Trade-off:策略变更不会影响当前 WS session(只影响新 session)。这是 // 有意为之 —— 对长会话来说,"策略一致性"比"立刻生效"更重要,且 Claude // BetaPolicy 的 gin.Context 缓存也是同样取舍。需要 hot-reload 时管理员 // 可以通过踢断 session 强制刷新。 type openAIFastPolicyCtxKeyType struct{} var openAIFastPolicyCtxKey = openAIFastPolicyCtxKeyType{} // withOpenAIFastPolicyContext 将一份 settings 快照绑定到 context,供该 ctx // 衍生 goroutine 中的 evaluateOpenAIFastPolicy 复用。 func withOpenAIFastPolicyContext(ctx context.Context, settings *OpenAIFastPolicySettings) context.Context { if ctx == nil || settings == nil { return ctx } return context.WithValue(ctx, openAIFastPolicyCtxKey, settings) } func openAIFastPolicySettingsFromContext(ctx context.Context) *OpenAIFastPolicySettings { if ctx == nil { return nil } if v, ok := ctx.Value(openAIFastPolicyCtxKey).(*OpenAIFastPolicySettings); ok { return v } return nil } // applyOpenAIFastPolicyToBody applies the OpenAI fast policy to a raw request // body. When action=filter it removes the service_tier field; when // action=block it returns (body, *OpenAIFastBlockedError). On pass it // normalizes the service_tier value (e.g. client alias "fast" → "priority"). // action=force_priority rewrites any matched known tier to "priority". // // Rationale for normalize-on-pass: chat-completions / messages 入口在调用本 // 函数之前已经通过 normalizeResponsesBodyServiceTier 把 service_tier 归一化 // 到了上游可识别值;passthrough(OpenAI 自动透传) / native /responses 等 // 入口没有这一前置步骤,pass 路径下若不在此处归一化,"fast" 就会被原样 // 透传到 OpenAI 上游导致 400/拒绝。把归一化收敛到本函数,所有入口行为一致。 func (s *OpenAIGatewayService) applyOpenAIFastPolicyToBody(ctx context.Context, account *Account, model string, body []byte) ([]byte, error) { if len(body) == 0 { return body, nil } rawTier := gjson.GetBytes(body, "service_tier").String() if rawTier == "" { return body, nil } normTier := normalizedOpenAIServiceTierValue(rawTier) if normTier == "" { return body, nil } action, errMsg := s.evaluateOpenAIFastPolicy(ctx, account, model, normTier) switch action { case BetaPolicyActionBlock: msg := errMsg if msg == "" { msg = fmt.Sprintf("openai service_tier=%s is not allowed for model %s", normTier, model) } return body, &OpenAIFastBlockedError{Message: msg} case BetaPolicyActionFilter: trimmed, err := sjson.DeleteBytes(body, "service_tier") if err != nil { return body, fmt.Errorf("strip service_tier from body: %w", err) } return trimmed, nil case OpenAIFastPolicyActionForcePriority: updated, err := sjson.SetBytes(body, "service_tier", OpenAIFastTierPriority) if err != nil { return body, fmt.Errorf("force service_tier priority on body: %w", err) } return updated, nil default: // pass:把别名(如 "fast")写回为规范值("priority")。 if normTier == rawTier { return body, nil } updated, err := sjson.SetBytes(body, "service_tier", normTier) if err != nil { return body, fmt.Errorf("normalize service_tier on pass: %w", err) } return updated, nil } } // writeOpenAIFastPolicyBlockedResponse writes a 403 JSON response for a // request blocked by the OpenAI fast policy. func writeOpenAIFastPolicyBlockedResponse(c *gin.Context, err *OpenAIFastBlockedError) { if c == nil || err == nil { return } MarkOpsClientBusinessLimited(c, OpsClientBusinessLimitedReasonLocalPolicyDenied) // body-signal compact 心跳可能已把响应头提交为 200(长排队后才进入 // Forward),此时以 response.failed 终止事件回传;未提交时先停拍再写 // JSON,保持原状态码语义(#3887)。 if StopOpenAICompactSSEKeepaliveCommitted(c) { writeOpenAICompactSSEFailureMessage(c, http.StatusForbidden, "permission_error", err.Message) return } c.JSON(http.StatusForbidden, gin.H{ "error": gin.H{ "type": "permission_error", "message": err.Message, }, }) } // applyOpenAIFastPolicyToWSResponseCreate evaluates the OpenAI fast policy // against a single client→upstream WebSocket frame whose top-level // "type"=="response.create". It mirrors the HTTP-side // applyOpenAIFastPolicyToBody contract but operates on a Realtime/Responses // WS payload: // // - pass: keeps service_tier, normalizing aliases such as "fast" to "priority" // - filter: returns a copy with top-level service_tier removed // - force_priority: keeps service_tier and rewrites it to "priority" // - block: returns (frame, *OpenAIFastBlockedError) // // Only frames whose "type" field strictly equals "response.create" are // inspected/mutated. Any other frame type — including the empty string — // passes through untouched. The OpenAI Realtime client-event spec requires // "type" to be set, so an empty type is treated as a malformed frame we do // not police; the upstream is the source of truth for rejecting it. // // service_tier lives at the top level of response.create — same as the // Responses HTTP body shape (see openai_gateway_chat_completions.go:304 + // extractOpenAIServiceTierFromBody at line 5593, and the test fixture at // openai_ws_forwarder_ingress_session_test.go:402). We therefore only need // to inspect / strip the top-level field; there is no nested form in the // schema today. // // The caller is responsible for choosing the upstream model passed in — // this helper does not re-derive it. func (s *OpenAIGatewayService) applyOpenAIFastPolicyToWSResponseCreate( ctx context.Context, account *Account, model string, frame []byte, ) ([]byte, *OpenAIFastBlockedError, error) { if len(frame) == 0 { return frame, nil, nil } if !gjson.ValidBytes(frame) { return frame, nil, nil } frameType := strings.TrimSpace(gjson.GetBytes(frame, "type").String()) // Strict match: only response.create is policy-checked. Empty / other // types pass through untouched so we never accidentally strip fields // from response.cancel, conversation.item.create, or any future // client-event the spec adds. The Realtime spec requires "type" on // every client event, so an empty type is malformed input — let the // upstream reject it rather than guessing at our layer. if frameType != "response.create" { return frame, nil, nil } rawTier := gjson.GetBytes(frame, "service_tier").String() if rawTier == "" { return frame, nil, nil } normTier := normalizedOpenAIServiceTierValue(rawTier) if normTier == "" { return frame, nil, nil } action, errMsg := s.evaluateOpenAIFastPolicy(ctx, account, model, normTier) switch action { case BetaPolicyActionBlock: msg := errMsg if msg == "" { msg = fmt.Sprintf("openai service_tier=%s is not allowed for model %s", normTier, model) } return frame, &OpenAIFastBlockedError{Message: msg}, nil case BetaPolicyActionFilter: trimmed, err := sjson.DeleteBytes(frame, "service_tier") if err != nil { return frame, nil, fmt.Errorf("strip service_tier from ws frame: %w", err) } return trimmed, nil, nil case OpenAIFastPolicyActionForcePriority: updated, err := sjson.SetBytes(frame, "service_tier", OpenAIFastTierPriority) if err != nil { return frame, nil, fmt.Errorf("force service_tier priority in ws frame: %w", err) } return updated, nil, nil default: if normTier == rawTier { return frame, nil, nil } updated, err := sjson.SetBytes(frame, "service_tier", normTier) if err != nil { return frame, nil, fmt.Errorf("normalize service_tier in ws frame: %w", err) } return updated, nil, nil } } // newOpenAIFastPolicyWSEventID returns a Realtime-style event_id for a // server-emitted error event. Matches the loose "evt_" convention used // by upstream Realtime servers; the exact value is not load-bearing and is // only required for client-side log correlation. We reuse the existing // google/uuid dependency rather than pulling a new one. func newOpenAIFastPolicyWSEventID() string { id, err := uuid.NewRandom() if err != nil { // Extremely unlikely; fall back to a fixed prefix so the field is // still non-empty and the schema stays self-consistent. return "evt_openai_fast_policy" } // Strip dashes so it visually matches "evt_" rather than UUID v4 // canonical form, mirroring what real Realtime traces look like. return "evt_" + strings.ReplaceAll(id.String(), "-", "") } // buildOpenAIFastPolicyBlockedWSEvent renders an OpenAI Realtime/Responses // style "error" event payload for a request blocked by the OpenAI fast // policy. The shape mirrors Realtime error events as observed in upstream // traces and per the spec's server "error" event: // // { // "event_id": "evt_", // "type": "error", // "error": { // "type": "invalid_request_error", // "code": "policy_violation", // "message": "..." // } // } // // event_id lets clients correlate the rejection in their logs; "code" gives // programmatic clients a stable identifier (HTTP-side equivalent is the // 403 permission_error JSON body). func buildOpenAIFastPolicyBlockedWSEvent(err *OpenAIFastBlockedError) []byte { if err == nil { return nil } eventID := newOpenAIFastPolicyWSEventID() payload, mErr := json.Marshal(map[string]any{ "event_id": eventID, "type": "error", "error": map[string]any{ "type": "invalid_request_error", "code": "policy_violation", "message": err.Message, }, }) if mErr != nil { // Fallback to a minimal hand-rolled payload; Marshal of the literal // shape above should never fail in practice. return []byte(`{"event_id":"` + eventID + `","type":"error","error":{"type":"invalid_request_error","code":"policy_violation","message":"openai fast policy blocked this request"}}`) } return payload } func openAIRequestBodyMayContainImageInput(body []byte) bool { if len(body) == 0 { return false } input := gjson.GetBytes(body, "input") messages := gjson.GetBytes(body, "messages.#-1") return openAIJSONValueMayContainImageInput(input) || openAIJSONValueMayContainImageInput(messages) } func openAIJSONValueMayContainImageInput(value gjson.Result) bool { if !value.Exists() { return false } if value.IsArray() { found := false value.ForEach(func(_, item gjson.Result) bool { if openAIJSONValueMayContainImageInput(item) { found = true return false } return true }) return found } if value.IsObject() { if strings.TrimSpace(value.Get("type").String()) == "input_image" || value.Get("image_url").Exists() { return true } return openAIJSONValueMayContainImageInput(value.Get("content")) } return false } func openAIRequestBodyMayContainEmptyBase64InputImage(body []byte) bool { if len(body) == 0 || !openAIRequestBodyMayContainInputImageToken(body) { return false } input := gjson.GetBytes(body, "input") if !input.Exists() { return false } return openAIJSONValueMayContainEmptyBase64InputImage(input) } func openAIRequestBodyMayContainInputImageToken(body []byte) bool { if bytes.Contains(body, []byte("input_image")) { return true } // JSON 字符串任意字符都可能被 unicode escape,遇到 \u 时交给 gjson 解码后的结构扫描兜底。 return bytes.Contains(body, []byte("\\u")) } func openAIJSONValueMayContainEmptyBase64InputImage(value gjson.Result) bool { if !value.Exists() { return false } if value.IsArray() { found := false value.ForEach(func(_, item gjson.Result) bool { if openAIJSONValueMayContainEmptyBase64InputImage(item) { found = true return false } return true }) return found } if value.IsObject() { if strings.TrimSpace(value.Get("type").String()) == "input_image" && isEmptyBase64DataURI(value.Get("image_url").String()) { return true } return openAIJSONValueMayContainEmptyBase64InputImage(value.Get("content")) } return false } func sanitizeEmptyBase64InputImagesInOpenAIBody(body []byte) ([]byte, bool, error) { if !openAIRequestBodyMayContainEmptyBase64InputImage(body) { return body, false, nil } var reqBody map[string]any if err := json.Unmarshal(body, &reqBody); err != nil { return body, false, fmt.Errorf("sanitize request body: %w", err) } if !sanitizeEmptyBase64InputImagesInOpenAIRequestBodyMap(reqBody) { return body, false, nil } normalized, err := marshalOpenAIUpstreamJSON(reqBody) if err != nil { return body, false, fmt.Errorf("serialize sanitized request body: %w", err) } return normalized, true, nil } func sanitizeEmptyBase64InputImagesInOpenAIRequestBodyMap(reqBody map[string]any) bool { if reqBody == nil { return false } input, ok := reqBody["input"] if !ok { return false } normalizedInput, changed := sanitizeEmptyBase64InputImagesInOpenAIInput(input) if !changed { return false } reqBody["input"] = normalizedInput return true } func sanitizeEmptyBase64InputImagesInOpenAIInput(input any) (any, bool) { items, ok := input.([]any) if !ok { return input, false } normalizedItems := make([]any, 0, len(items)) changed := false for _, item := range items { itemMap, ok := item.(map[string]any) if !ok { normalizedItems = append(normalizedItems, item) continue } if shouldDropEmptyBase64InputImagePart(itemMap) { changed = true continue } content, ok := itemMap["content"] if !ok { normalizedItems = append(normalizedItems, itemMap) continue } parts, ok := content.([]any) if !ok { normalizedItems = append(normalizedItems, itemMap) continue } normalizedParts := make([]any, 0, len(parts)) itemChanged := false for _, part := range parts { if shouldDropEmptyBase64InputImagePart(part) { changed = true itemChanged = true continue } normalizedParts = append(normalizedParts, part) } if itemChanged { if len(normalizedParts) == 0 { continue } itemMap["content"] = normalizedParts } normalizedItems = append(normalizedItems, itemMap) } if !changed { return input, false } return normalizedItems, true } func shouldDropEmptyBase64InputImagePart(part any) bool { partMap, ok := part.(map[string]any) if !ok { return false } typeValue, _ := partMap["type"].(string) if strings.TrimSpace(typeValue) != "input_image" { return false } imageURL, _ := partMap["image_url"].(string) return isEmptyBase64DataURI(imageURL) } func isEmptyBase64DataURI(raw string) bool { if !strings.HasPrefix(raw, "data:") { return false } rest := strings.TrimPrefix(raw, "data:") semicolonIdx := strings.Index(rest, ";") if semicolonIdx < 0 { return false } rest = rest[semicolonIdx+1:] if !strings.HasPrefix(rest, "base64,") { return false } return strings.TrimSpace(strings.TrimPrefix(rest, "base64,")) == "" } func getOpenAIRequestBodyMap(_ *gin.Context, body []byte) (map[string]any, error) { var reqBody map[string]any if err := json.Unmarshal(body, &reqBody); err != nil { return nil, fmt.Errorf("parse request: %w", err) } return reqBody, nil } // extractOpenAIReasoningEffort 的模型候选语义同 extractOpenAIReasoningEffortFromBody。 func extractOpenAIReasoningEffort(reqBody map[string]any, modelCandidates ...string) *string { if value, present := getOpenAIReasoningEffortFromReqBody(reqBody, firstNonEmpty(modelCandidates...)); present { if value == "" { return nil } return &value } value := deriveOpenAIReasoningEffortFromModelCandidates(modelCandidates) if value == "" { return nil } return &value } func normalizeOpenAIReasoningEffort(raw string) string { value := strings.ToLower(strings.TrimSpace(raw)) if value == "" { return "" } // Normalize separators for "x-high"/"x_high" variants. value = strings.NewReplacer("-", "", "_", "", " ", "").Replace(value) switch value { case "none", "minimal": return "" case "low", "medium", "high": return value case "xhigh", "extrahigh", "max": return "xhigh" default: // Only store known effort levels for now to keep UI consistent. return "" } } func normalizeOpenAIReasoningEffortForModel(raw, model string) string { if strings.EqualFold(strings.TrimSpace(raw), "max") && isOpenAIGPT56Model(model) { return "max" } return normalizeOpenAIReasoningEffort(raw) }