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
1515 lines
46 KiB
Go
1515 lines
46 KiB
Go
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_<rand>" 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_<hex>" 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_<random>",
|
||
// "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)
|
||
}
|