diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 5c8e599..1deffa9 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -7,7 +7,7 @@ on: branches: [master] env: - GOPRIVATE: github.com/DouDOU-start/airgate-sdk + GOPRIVATE: github.com/DevilGenius/airgate-sdk jobs: ci: diff --git a/.github/workflows/release.yml b/.github/workflows/release.yml index f5b6152..482faa3 100644 --- a/.github/workflows/release.yml +++ b/.github/workflows/release.yml @@ -9,7 +9,7 @@ permissions: contents: write env: - GOPRIVATE: github.com/DouDOU-start/airgate-sdk + GOPRIVATE: github.com/DevilGenius/airgate-sdk jobs: build: @@ -64,7 +64,7 @@ jobs: echo "Injecting PluginVersion=${VERSION}" mkdir -p bin cd backend && go build -buildvcs=false -trimpath \ - -ldflags "-X 'github.com/DouDOU-start/airgate-openai/backend/internal/gateway.PluginVersion=${VERSION}'" \ + -ldflags "-X 'github.com/DevilGenius/airgate-openai/backend/internal/gateway.PluginVersion=${VERSION}'" \ -o ../bin/gateway-openai-${{ matrix.goos }}-${{ matrix.goarch }} . - name: Generate SHA256 diff --git a/LICENSE b/LICENSE index 30c35bf..18f224f 100644 --- a/LICENSE +++ b/LICENSE @@ -1,6 +1,6 @@ MIT License -Copyright (c) 2026 DouDOU-start +Copyright (c) 2026 DevilGenius Permission is hereby granted, free of charge, to any person obtaining a copy of this software and associated documentation files (the "Software"), to deal diff --git a/Makefile b/Makefile index 06b0bd4..c5415c0 100644 --- a/Makefile +++ b/Makefile @@ -75,7 +75,7 @@ lint: ## 代码检查(需要安装 golangci-lint) fmt: ## 格式化代码 @cd backend && \ if command -v goimports > /dev/null 2>&1; then \ - goimports -w -local github.com/DouDOU-start .; \ + goimports -w -local github.com/DevilGenius .; \ else \ $(GO) fmt ./...; \ fi diff --git a/README.md b/README.md index 86619d4..085b31d 100644 --- a/README.md +++ b/README.md @@ -4,9 +4,9 @@
OpenAI / ChatGPT / Anthropic 协议三合一网关插件
@@ -14,7 +14,7 @@ --- -AirGate OpenAI 不是又一个"OpenAI 转发服务",而是 [airgate-core](https://github.com/DouDOU-start/airgate-core) 的旗舰网关插件,也是 [airgate-sdk](https://github.com/DouDOU-start/airgate-sdk) 的官方参考实现。它在一个 gRPC 子进程里同时承载: +AirGate OpenAI 不是又一个"OpenAI 转发服务",而是 [airgate-core](https://github.com/DevilGenius/airgate-core) 的旗舰网关插件,也是 [airgate-sdk](https://github.com/DevilGenius/airgate-sdk) 的官方参考实现。它在一个 gRPC 子进程里同时承载: - **OpenAI Responses / Chat Completions API** 转发(Codex 核心端点) - **ChatGPT OAuth 浏览器授权账号** 接入(PKCE + WebSocket 桥接) @@ -128,19 +128,19 @@ AirGate OpenAI 不是又一个"OpenAI 转发服务",而是 [airgate-core](http ```text 1. 插件市场 → 点击「安装」 (从 GitHub Release 自动拉取,匹配当前架构) 2. 上传安装 → 拖入二进制文件 (适合内部环境 / 自建二进制) -3. GitHub 安装 → 输入 DouDOU-start/airgate-openai +3. GitHub 安装 → 输入 DevilGenius/airgate-openai ``` market 会**定时从 GitHub API 同步**最新 release(默认 6h,使用 ETag 不消耗 API 配额),新 tag push 后通常几分钟内即可在市场看到。 ### 方式 2:源码运行(开发) -需要 Go 1.25+、Node 22+,以及兄弟目录 [`airgate-sdk`](https://github.com/DouDOU-start/airgate-sdk) 与 [`airgate-core`](https://github.com/DouDOU-start/airgate-core): +需要 Go 1.25+、Node 22+,以及兄弟目录 [`airgate-sdk`](https://github.com/DevilGenius/airgate-sdk) 与 [`airgate-core`](https://github.com/DevilGenius/airgate-core): ```bash -git clone https://github.com/DouDOU-start/airgate-sdk.git -git clone https://github.com/DouDOU-start/airgate-core.git -git clone https://github.com/DouDOU-start/airgate-openai.git +git clone https://github.com/DevilGenius/airgate-sdk.git +git clone https://github.com/DevilGenius/airgate-core.git +git clone https://github.com/DevilGenius/airgate-openai.git cd airgate-openai make install # 装 web 依赖与 Go 模块 @@ -248,9 +248,9 @@ git push origin v0.2.0 ## 🤝 贡献 / 反馈 -- Bug / Feature: [Issues](https://github.com/DouDOU-start/airgate-openai/issues) -- 主仓库: [airgate-core](https://github.com/DouDOU-start/airgate-core) -- 插件 SDK: [airgate-sdk](https://github.com/DouDOU-start/airgate-sdk) +- Bug / Feature: [Issues](https://github.com/DevilGenius/airgate-openai/issues) +- 主仓库: [airgate-core](https://github.com/DevilGenius/airgate-core) +- 插件 SDK: [airgate-sdk](https://github.com/DevilGenius/airgate-sdk) ## 📜 License diff --git a/backend/.golangci.yml b/backend/.golangci.yml index 078a931..cbcee2c 100644 --- a/backend/.golangci.yml +++ b/backend/.golangci.yml @@ -23,7 +23,7 @@ formatters: settings: goimports: local-prefixes: - - github.com/DouDOU-start + - github.com/DevilGenius issues: max-issues-per-linter: 0 diff --git a/backend/cmd/chat/main.go b/backend/cmd/chat/main.go index c1895a3..21238dd 100644 --- a/backend/cmd/chat/main.go +++ b/backend/cmd/chat/main.go @@ -10,7 +10,7 @@ import ( "os" "strings" - "github.com/DouDOU-start/airgate-openai/backend/internal/gateway" + "github.com/DevilGenius/airgate-openai/backend/internal/gateway" ) func main() { diff --git a/backend/cmd/chat/session.go b/backend/cmd/chat/session.go index 711f5d6..f5fe4fc 100644 --- a/backend/cmd/chat/session.go +++ b/backend/cmd/chat/session.go @@ -12,8 +12,8 @@ import ( "github.com/gorilla/websocket" - "github.com/DouDOU-start/airgate-openai/backend/internal/gateway" - "github.com/DouDOU-start/airgate-openai/backend/resources" + "github.com/DevilGenius/airgate-openai/backend/internal/gateway" + "github.com/DevilGenius/airgate-openai/backend/resources" ) // ────────────────────────────────────────────────────── diff --git a/backend/cmd/devserver/main.go b/backend/cmd/devserver/main.go index 2661903..b7d67df 100644 --- a/backend/cmd/devserver/main.go +++ b/backend/cmd/devserver/main.go @@ -6,9 +6,9 @@ import ( "log" "net/http" - "github.com/DouDOU-start/airgate-sdk/devkit/devserver" + "github.com/DevilGenius/airgate-sdk/devkit/devserver" - "github.com/DouDOU-start/airgate-openai/backend/internal/gateway" + "github.com/DevilGenius/airgate-openai/backend/internal/gateway" ) func main() { diff --git a/backend/cmd/genmanifest/main.go b/backend/cmd/genmanifest/main.go index ca894c2..b35fca3 100644 --- a/backend/cmd/genmanifest/main.go +++ b/backend/cmd/genmanifest/main.go @@ -9,9 +9,9 @@ import ( "gopkg.in/yaml.v3" - "github.com/DouDOU-start/airgate-openai/backend/internal/gateway" - "github.com/DouDOU-start/airgate-openai/backend/internal/model" - sdk "github.com/DouDOU-start/airgate-sdk/sdkgo" + "github.com/DevilGenius/airgate-openai/backend/internal/gateway" + "github.com/DevilGenius/airgate-openai/backend/internal/model" + sdk "github.com/DevilGenius/airgate-sdk/sdkgo" ) const generatedComment = "# 本文件由 backend/cmd/genmanifest 自动生成,请勿手工修改。\n\n" diff --git a/backend/go.mod b/backend/go.mod index 05d2fb3..37c012a 100644 --- a/backend/go.mod +++ b/backend/go.mod @@ -1,9 +1,9 @@ -module github.com/DouDOU-start/airgate-openai/backend +module github.com/DevilGenius/airgate-openai/backend go 1.25.7 require ( - github.com/DouDOU-start/airgate-sdk v0.2.1 + github.com/DevilGenius/airgate-sdk v0.2.1 github.com/google/uuid v1.6.0 github.com/gorilla/websocket v1.5.3 github.com/lib/pq v1.10.9 diff --git a/backend/go.sum b/backend/go.sum index 2081728..37d37f6 100644 --- a/backend/go.sum +++ b/backend/go.sum @@ -1,5 +1,3 @@ -github.com/DouDOU-start/airgate-sdk v0.2.1 h1:MylrFMifIJy5mT1wSFdXcSvAXjIcB6itp4dw8pP4pes= -github.com/DouDOU-start/airgate-sdk v0.2.1/go.mod h1:784vC4lIfCnUdOroWuvn55Dv/b9TucVMpFwuu40gc0Q= github.com/andybalholm/brotli v1.0.6 h1:Yf9fFpf49Zrxb9NlQaluyE92/+X7UVHlhMNJN2sxfOI= github.com/andybalholm/brotli v1.0.6/go.mod h1:fO7iG3H7G2nSZ7m0zPUDn85XEX2GTukHGRSepvi9Eig= github.com/bufbuild/protocompile v0.14.1 h1:iA73zAf/fyljNjQKwYzUHD6AD4R8KMasmwa/FBatYVw= diff --git a/backend/internal/gateway/anthropic_compat_test.go b/backend/internal/gateway/anthropic_compat_test.go index 0bf6b46..e8ed0f6 100644 --- a/backend/internal/gateway/anthropic_compat_test.go +++ b/backend/internal/gateway/anthropic_compat_test.go @@ -12,7 +12,7 @@ import ( "github.com/tidwall/gjson" - sdk "github.com/DouDOU-start/airgate-sdk/sdkgo" + sdk "github.com/DevilGenius/airgate-sdk/sdkgo" ) func TestEstimateAnthropicInputTokensCountsSystemMessagesAndTools(t *testing.T) { diff --git a/backend/internal/gateway/anthropic_count_tokens.go b/backend/internal/gateway/anthropic_count_tokens.go index ea78bb7..8736532 100644 --- a/backend/internal/gateway/anthropic_count_tokens.go +++ b/backend/internal/gateway/anthropic_count_tokens.go @@ -5,7 +5,7 @@ import ( "encoding/json" "net/http" - sdk "github.com/DouDOU-start/airgate-sdk/sdkgo" + sdk "github.com/DevilGenius/airgate-sdk/sdkgo" ) // forwardAnthropicCountTokens 返回本地估算,避免客户端把 404 当作不可计数并继续发送过长请求。 diff --git a/backend/internal/gateway/anthropic_forward.go b/backend/internal/gateway/anthropic_forward.go index 6a9fa07..21eb5a6 100644 --- a/backend/internal/gateway/anthropic_forward.go +++ b/backend/internal/gateway/anthropic_forward.go @@ -13,7 +13,7 @@ import ( "github.com/tidwall/gjson" "github.com/tidwall/sjson" - sdk "github.com/DouDOU-start/airgate-sdk/sdkgo" + sdk "github.com/DevilGenius/airgate-sdk/sdkgo" ) // ────────────────────────────────────────────────────── @@ -26,7 +26,7 @@ func (g *OpenAIGateway) forwardAnthropicMessage(ctx context.Context, req *sdk.Fo start := time.Now() body := req.Body strategy := resolveAnthropicUpstreamStrategy(req.Account) - session := resolveOpenAISession(req.Headers, req.Body) + session := resolveOpenAISession(req.Headers, req.Body, req.Account.ID) session.DigestChain = buildAnthropicDigestChain(body) if session.SessionKey == "" { if reusedSessionID, matchedChain, ok := findAnthropicDigestSession(req.Account.ID, session.DigestChain); ok { @@ -41,7 +41,7 @@ func (g *OpenAIGateway) forwardAnthropicMessage(ctx context.Context, req *sdk.Fo session.SessionSource = "anthropic_digest_new" } } - updateSessionStateFromRequest(session) + updateSessionStateFromRequest(session, req.Account.ID) logger := sdk.LoggerFromContext(ctx) logger.Debug("anthropic_request_received", @@ -516,7 +516,7 @@ func (g *OpenAIGateway) handleAnthropicNonStreamFromResponses( return transientOutcome(reason), fmt.Errorf("%s", reason) } if session.SessionKey != "" && wsResult.ResponseID != "" { - updateSessionStateResponseID(session.SessionKey, wsResult.ResponseID) + updateSessionStateResponseID(session.SessionKey, wsResult.ResponseID, accountID) } if session.SessionID != "" && session.DigestChain != "" { saveAnthropicDigestSession(accountID, session.DigestChain, session.SessionID, session.MatchedDigest) @@ -554,6 +554,7 @@ func (g *OpenAIGateway) handleAnthropicNonStreamFromResponses( elapsed.Milliseconds(), ) fillUsageCost(usage) + setUsageResponseID(usage, wsResult.ResponseID) return sdk.ForwardOutcome{ Kind: sdk.OutcomeSuccess, Upstream: sdk.UpstreamResponse{StatusCode: http.StatusOK, Headers: upstreamHeaders, Body: anthropicBody}, diff --git a/backend/internal/gateway/anthropic_response.go b/backend/internal/gateway/anthropic_response.go index 6557db5..38261c3 100644 --- a/backend/internal/gateway/anthropic_response.go +++ b/backend/internal/gateway/anthropic_response.go @@ -13,7 +13,7 @@ import ( "github.com/tidwall/gjson" "github.com/tidwall/sjson" - sdk "github.com/DouDOU-start/airgate-sdk/sdkgo" + sdk "github.com/DevilGenius/airgate-sdk/sdkgo" ) // ────────────────────────────────────────────────────── @@ -649,6 +649,7 @@ func translateResponsesSSEToAnthropicSSE( serviceTier := firstNonEmptyTier(requestServiceTier) skipCurrentOutput := false firstTokenRecorded := false + responseID := "" for scanner.Scan() { skipCurrentOutput = false @@ -687,9 +688,12 @@ func translateResponsesSSEToAnthropicSSE( } if session.SessionKey != "" { if responseID := gjson.Get(data, "response.id").String(); strings.TrimSpace(responseID) != "" { - updateSessionStateResponseID(session.SessionKey, responseID) + updateSessionStateResponseID(session.SessionKey, responseID, session.AccountID) } } + if id := strings.TrimSpace(gjson.Get(data, "response.id").String()); id != "" { + responseID = id + } if eventType == "response.completed" || eventType == "response.done" { if serviceTier == "" { serviceTier = firstNonEmptyTier(gjson.Get(data, "response.service_tier").String(), defaultServiceTier) @@ -802,6 +806,7 @@ done: } fillUsageCost(usage) + setUsageResponseID(usage, responseID) return sdk.ForwardOutcome{ Kind: sdk.OutcomeSuccess, Upstream: sdk.UpstreamResponse{StatusCode: http.StatusOK}, diff --git a/backend/internal/gateway/anthropic_strategy.go b/backend/internal/gateway/anthropic_strategy.go index 8799e6b..07138a0 100644 --- a/backend/internal/gateway/anthropic_strategy.go +++ b/backend/internal/gateway/anthropic_strategy.go @@ -1,7 +1,7 @@ package gateway import ( - sdk "github.com/DouDOU-start/airgate-sdk/sdkgo" + sdk "github.com/DevilGenius/airgate-sdk/sdkgo" ) type anthropicUpstreamStrategy string diff --git a/backend/internal/gateway/chat_completions_oauth.go b/backend/internal/gateway/chat_completions_oauth.go index 7303338..9cdbc10 100644 --- a/backend/internal/gateway/chat_completions_oauth.go +++ b/backend/internal/gateway/chat_completions_oauth.go @@ -12,7 +12,7 @@ import ( "github.com/tidwall/gjson" - sdk "github.com/DouDOU-start/airgate-sdk/sdkgo" + sdk "github.com/DevilGenius/airgate-sdk/sdkgo" ) // OAuth/Codex 上游永远返回 Responses API 的 SSE 流,但客户端走 @@ -157,7 +157,7 @@ func (s *chatCompletionsStreamWriter) OnRawEvent(eventType string, data []byte) case "response.created", "response.completed", "response.done": if s.sessionKey != "" { if responseID := gjson.GetBytes(data, "response.id").String(); strings.TrimSpace(responseID) != "" { - updateSessionStateResponseID(s.sessionKey, responseID) + updateSessionStateResponseID(s.sessionKey, responseID, s.accountID) } } } @@ -506,7 +506,7 @@ func (h *chatCompletionsSilentHandler) OnRawEvent(eventType string, data []byte) case "response.created", "response.completed", "response.done": if h.sessionKey != "" { if responseID := gjson.GetBytes(data, "response.id").String(); strings.TrimSpace(responseID) != "" { - updateSessionStateResponseID(h.sessionKey, responseID) + updateSessionStateResponseID(h.sessionKey, responseID, h.accountID) } } } @@ -557,7 +557,7 @@ func (h *responsesSilentHandler) OnRawEvent(eventType string, data []byte) { case "response.created", "response.completed", "response.done": if h.sessionKey != "" { if responseID := gjson.GetBytes(data, "response.id").String(); strings.TrimSpace(responseID) != "" { - updateSessionStateResponseID(h.sessionKey, responseID) + updateSessionStateResponseID(h.sessionKey, responseID, h.accountID) } } } diff --git a/backend/internal/gateway/chat_completions_oauth_test.go b/backend/internal/gateway/chat_completions_oauth_test.go index d33d6df..b8b024a 100644 --- a/backend/internal/gateway/chat_completions_oauth_test.go +++ b/backend/internal/gateway/chat_completions_oauth_test.go @@ -8,7 +8,7 @@ import ( "testing" "time" - sdk "github.com/DouDOU-start/airgate-sdk/sdkgo" + sdk "github.com/DevilGenius/airgate-sdk/sdkgo" ) // fakeResponseWriter 用于单测:捕获写回的 SSE / JSON 响应体。 diff --git a/backend/internal/gateway/errors.go b/backend/internal/gateway/errors.go index f448141..49d9d51 100644 --- a/backend/internal/gateway/errors.go +++ b/backend/internal/gateway/errors.go @@ -8,7 +8,7 @@ import ( "github.com/tidwall/gjson" "github.com/tidwall/sjson" - sdk "github.com/DouDOU-start/airgate-sdk/sdkgo" + sdk "github.com/DevilGenius/airgate-sdk/sdkgo" ) // 统一错误分类工具(跨 OpenAI / Anthropic 协议共用)。 @@ -20,7 +20,8 @@ import ( // 返回的 Kind 决定 Core 如何处置账号: // // 429 → AccountRateLimited -// 401 / 403 → AccountDead(附加消息关键词检查"usage limit" / "rate limit" 等会降级为 RateLimited) +// 401 → AccountDead +// 403 → AccountUnavailable(明确 disabled/deactivated/suspended 仍升级为 AccountDead) // 400 + 消息含限流关键词 → AccountRateLimited(部分上游用 400 返回 usage_limit_reached) // 400 + 消息含 disabled/deactivated → AccountDead // 5xx → UpstreamTransient @@ -35,8 +36,10 @@ func classifyHTTPFailure(statusCode int, message string) sdk.OutcomeKind { switch statusCode { case 429: return sdk.OutcomeAccountRateLimited - case 401, 403: + case 401: return sdk.OutcomeAccountDead + case 403: + return sdk.OutcomeAccountUnavailable } if statusCode >= 500 { return sdk.OutcomeUpstreamTransient diff --git a/backend/internal/gateway/forward.go b/backend/internal/gateway/forward.go index bced3f5..ee20133 100644 --- a/backend/internal/gateway/forward.go +++ b/backend/internal/gateway/forward.go @@ -17,8 +17,8 @@ import ( "github.com/tidwall/gjson" "github.com/tidwall/sjson" - "github.com/DouDOU-start/airgate-openai/backend/internal/model" - sdk "github.com/DouDOU-start/airgate-sdk/sdkgo" + "github.com/DevilGenius/airgate-openai/backend/internal/model" + sdk "github.com/DevilGenius/airgate-sdk/sdkgo" ) // redactURL 去掉 query string,仅保留 host+path(避免敏感参数泄漏到日志) @@ -228,14 +228,16 @@ func (g *OpenAIGateway) forwardAPIKey(ctx context.Context, req *sdk.ForwardReque reqMethod, reqPath := resolveAPIKeyRoute(req) targetURL := buildAPIKeyURL(account, reqPath) - imagesBillingSize := "" - if isImagesRequest(reqPath) && len(req.Body) > 0 { - if parsed, err := parseImagesRequest(req.Body, req.Headers.Get("Content-Type"), isImagesEditRequest(reqPath)); err == nil { - imagesBillingSize = parsed.Size - } - } - if isImagesEditRequest(reqPath) && len(req.Body) > 0 && !strings.HasPrefix(strings.ToLower(req.Headers.Get("Content-Type")), "multipart/") { - body, contentType, err := buildAPIKeyImagesEditMultipartBody(req.Body, req.Headers.Get("Content-Type")) + isImageReq := isImagesRequest(reqPath) + isImageEdit := isImagesEditRequest(reqPath) + reqContentType := req.Headers.Get("Content-Type") + isMultipart := isMultipartContentType(reqContentType) + imagesRespOpts := imagesResponseOptions{} + if isImageReq && len(req.Body) > 0 && (!isImageEdit || isMultipart) { + imagesRespOpts = imagesResponseOptionsFromRequestBody(req.Body, reqContentType, isImageEdit) + } + if isImageEdit && len(req.Body) > 0 && !isMultipart { + body, contentType, _, err := buildAPIKeyImagesEditMultipartBodyWithRequest(req.Body, reqContentType) if err != nil { errBody := jsonError(err.Error()) return sdk.ForwardOutcome{ @@ -251,7 +253,8 @@ func (g *OpenAIGateway) forwardAPIKey(ctx context.Context, req *sdk.ForwardReque } req.Body = body req.Headers.Set("Content-Type", contentType) - } else if isImagesRequest(reqPath) && len(req.Body) > 0 && !strings.HasPrefix(req.Headers.Get("Content-Type"), "multipart/") { + imagesRespOpts = imagesResponseOptionsFromRequestBody(body, contentType, true) + } else if isImageReq && len(req.Body) > 0 && !isMultipart { if patched, err := sjson.DeleteBytes(req.Body, "stream"); err == nil { req.Body = patched } @@ -301,7 +304,7 @@ func (g *OpenAIGateway) forwardAPIKey(ctx context.Context, req *sdk.ForwardReque Header: http.Header{"Content-Type": []string{"application/json"}}, Body: io.NopCloser(bytes.NewReader(finalBody)), } - return g.handleImagesResponse(mockResp, req.Writer, nil, start, req.Model, imagesBillingSize) + return g.handleImagesResponse(mockResp, req.Writer, nil, start, req.Model, imagesRespOpts) } reason := fmt.Sprintf("上游异步任务恢复失败: %v", pollErr) logger.Warn("images_async_task_recovery_failed", @@ -446,7 +449,7 @@ func (g *OpenAIGateway) forwardAPIKey(ctx context.Context, req *sdk.ForwardReque body = finalBody } resp.Body = io.NopCloser(bytes.NewReader(body)) - return g.handleImagesResponse(resp, req.Writer, sseKA, start, req.Model, imagesBillingSize) + return g.handleImagesResponse(resp, req.Writer, sseKA, start, req.Model, imagesRespOpts) } if req.Stream && req.Writer != nil { @@ -525,8 +528,8 @@ func (g *OpenAIGateway) forwardOAuth(ctx context.Context, req *sdk.ForwardReques start := time.Now() account := req.Account logger := sdk.LoggerFromContext(ctx) - session := resolveOpenAISession(req.Headers, req.Body) - updateSessionStateFromRequest(session) + session := resolveOpenAISession(req.Headers, req.Body, account.ID) + updateSessionStateFromRequest(session, account.ID) cfg := WSConfig{ Token: account.Credentials["access_token"], @@ -690,7 +693,7 @@ func (g *OpenAIGateway) forwardOAuth(ctx context.Context, req *sdk.ForwardReques } if session.SessionKey != "" { if result.ResponseID != "" { - updateSessionStateResponseID(session.SessionKey, result.ResponseID) + updateSessionStateResponseID(session.SessionKey, result.ResponseID, account.ID) } } @@ -809,6 +812,7 @@ func (g *OpenAIGateway) forwardOAuth(ctx context.Context, req *sdk.ForwardReques "stream", req.Stream, ) fillUsageCostWithImageTool(usage, numImages, imageToolSize) + setUsageResponseID(usage, result.ResponseID) return sdk.ForwardOutcome{ Kind: sdk.OutcomeSuccess, Upstream: sdk.UpstreamResponse{StatusCode: http.StatusOK}, @@ -864,7 +868,7 @@ func (s *sseEventWriter) OnRawEvent(eventType string, data []byte) { case "response.created", "response.completed", "response.done": if s.sessionKey != "" { if responseID := gjson.GetBytes(data, "response.id").String(); strings.TrimSpace(responseID) != "" { - updateSessionStateResponseID(s.sessionKey, responseID) + updateSessionStateResponseID(s.sessionKey, responseID, s.accountID) } } } diff --git a/backend/internal/gateway/gateway.go b/backend/internal/gateway/gateway.go index 87d4906..7dd6aaa 100644 --- a/backend/internal/gateway/gateway.go +++ b/backend/internal/gateway/gateway.go @@ -14,9 +14,9 @@ import ( "strings" "time" - sdk "github.com/DouDOU-start/airgate-sdk/sdkgo" + sdk "github.com/DevilGenius/airgate-sdk/sdkgo" - "github.com/DouDOU-start/airgate-openai/backend/internal/model" + "github.com/DevilGenius/airgate-openai/backend/internal/model" ) // OpenAIGateway OpenAI 网关插件(SimpleGatewayPlugin 实现) diff --git a/backend/internal/gateway/headers.go b/backend/internal/gateway/headers.go index 918f594..0dad9af 100644 --- a/backend/internal/gateway/headers.go +++ b/backend/internal/gateway/headers.go @@ -10,7 +10,7 @@ import ( "sync" "time" - sdk "github.com/DouDOU-start/airgate-sdk/sdkgo" + sdk "github.com/DevilGenius/airgate-sdk/sdkgo" ) // setAuthHeaders 设置认证头 diff --git a/backend/internal/gateway/headers_test.go b/backend/internal/gateway/headers_test.go index 0e0c75c..7f4ae57 100644 --- a/backend/internal/gateway/headers_test.go +++ b/backend/internal/gateway/headers_test.go @@ -5,7 +5,7 @@ import ( "testing" "time" - sdk "github.com/DouDOU-start/airgate-sdk/sdkgo" + sdk "github.com/DevilGenius/airgate-sdk/sdkgo" ) func TestPassHeadersForAccount_Sub2APIStripsClientIdentityHeaders(t *testing.T) { diff --git a/backend/internal/gateway/host_invoke.go b/backend/internal/gateway/host_invoke.go index 91cdfc7..8b2751b 100644 --- a/backend/internal/gateway/host_invoke.go +++ b/backend/internal/gateway/host_invoke.go @@ -12,7 +12,7 @@ import ( "github.com/google/uuid" - sdk "github.com/DouDOU-start/airgate-sdk/sdkgo" + sdk "github.com/DevilGenius/airgate-sdk/sdkgo" ) const ( diff --git a/backend/internal/gateway/image_token_calculator.go b/backend/internal/gateway/image_token_calculator.go new file mode 100644 index 0000000..4551040 --- /dev/null +++ b/backend/internal/gateway/image_token_calculator.go @@ -0,0 +1,81 @@ +package gateway + +import ( + "fmt" + "strings" +) + +const ( + gptImageTokenFormulaBias = int64(2_000_000) + gptImageTokenFormulaScale = int64(4_000_000) + gptImageDefaultQuality = "high" +) + +// GPTImageTokenCalculator mirrors OpenAI's GPT Image 2 token calculator: +// size + low/medium/high quality -> estimated image tokens. +type GPTImageTokenCalculator struct{} + +// Calculate parses size as WIDTHxHEIGHT and returns GPT Image 2 image tokens for one image. +func (GPTImageTokenCalculator) Calculate(size, quality string) (int, error) { + width, height, ok := parseImageSize(size) + if !ok { + return 0, fmt.Errorf("size 格式无效,应为 WIDTHxHEIGHT") + } + return calculateGPTImageTokens(width, height, quality) +} + +// CalculateDefaultQuality parses size and calculates tokens with the default high quality. +func (c GPTImageTokenCalculator) CalculateDefaultQuality(size string) (int, error) { + return c.Calculate(size, gptImageDefaultQuality) +} + +// CalculateDimensions returns GPT Image 2 image tokens for one image. +func (GPTImageTokenCalculator) CalculateDimensions(width, height int, quality string) (int, error) { + return calculateGPTImageTokens(width, height, quality) +} + +// CalculateDimensionsDefaultQuality calculates tokens with the default high quality. +func (c GPTImageTokenCalculator) CalculateDimensionsDefaultQuality(width, height int) (int, error) { + return c.CalculateDimensions(width, height, gptImageDefaultQuality) +} + +func calculateGPTImageTokens(width, height int, quality string) (int, error) { + if width <= 0 || height <= 0 { + return 0, fmt.Errorf("size 宽高必须大于 0") + } + base, err := gptImageQualityBase(quality) + if err != nil { + return 0, err + } + + longEdge, shortEdge := width, height + if shortEdge > longEdge { + longEdge, shortEdge = shortEdge, longEdge + } + + scaledShort := roundPositiveRatio(int64(base*shortEdge), int64(longEdge)) + patches := int64(base) * scaledShort + area := int64(width) * int64(height) + return int(ceilPositiveRatio(patches*(gptImageTokenFormulaBias+area), gptImageTokenFormulaScale)), nil +} + +func gptImageQualityBase(quality string) (int, error) { + switch strings.ToLower(strings.TrimSpace(quality)) { + case "low": + return 16, nil + case "medium": + return 48, nil + case "high": + return 96, nil + default: + return 0, fmt.Errorf("quality 必须是 low、medium 或 high") + } +} + +func roundPositiveRatio(numerator, denominator int64) int64 { + return (numerator*2 + denominator) / (denominator * 2) +} + +func ceilPositiveRatio(numerator, denominator int64) int64 { + return (numerator + denominator - 1) / denominator +} diff --git a/backend/internal/gateway/image_token_calculator_test.go b/backend/internal/gateway/image_token_calculator_test.go new file mode 100644 index 0000000..04c0127 --- /dev/null +++ b/backend/internal/gateway/image_token_calculator_test.go @@ -0,0 +1,104 @@ +package gateway + +import ( + "strings" + "testing" +) + +func TestGPTImageTokenCalculator(t *testing.T) { + calc := GPTImageTokenCalculator{} + cases := []struct { + size string + quality string + want int + }{ + {"1024x1024", "low", 196}, + {"1024x1024", "medium", 1756}, + {"1024x1024", "high", 7024}, + {"1024x1536", "low", 158}, + {"1536x1024", "medium", 1372}, + {"3840x2160", "low", 371}, + {"3840x2160", "medium", 3336}, + {"3840x2160", "high", 13342}, + {"3840X2160", "HIGH", 13342}, + } + + for _, tc := range cases { + t.Run(tc.size+"_"+tc.quality, func(t *testing.T) { + got, err := calc.Calculate(tc.size, tc.quality) + if err != nil { + t.Fatalf("Calculate(%q, %q) returned err: %v", tc.size, tc.quality, err) + } + if got != tc.want { + t.Fatalf("Calculate(%q, %q) = %d, want %d", tc.size, tc.quality, got, tc.want) + } + }) + } +} + +func TestGPTImageTokenCalculatorRejectsUnparseableInput(t *testing.T) { + calc := GPTImageTokenCalculator{} + cases := []struct { + name string + size string + quality string + wantSubstr string + }{ + {"auto size", "auto", "low", "WIDTHxHEIGHT"}, + {"bad size", "1024", "low", "WIDTHxHEIGHT"}, + {"unknown quality", "1024x1024", "standard", "quality"}, + {"zero dimensions", "0x1024", "low", "WIDTHxHEIGHT"}, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + _, err := calc.Calculate(tc.size, tc.quality) + if err == nil { + t.Fatalf("Calculate(%q, %q) = nil err, want invalid input error", tc.size, tc.quality) + } + if !strings.Contains(err.Error(), tc.wantSubstr) { + t.Fatalf("error = %q, want substring %q", err.Error(), tc.wantSubstr) + } + }) + } +} + +func TestGPTImageTokenCalculatorDefaultQualityUsesHigh(t *testing.T) { + calc := GPTImageTokenCalculator{} + + got, err := calc.CalculateDefaultQuality("3840x2160") + if err != nil { + t.Fatalf("CalculateDefaultQuality returned err: %v", err) + } + if got != 13342 { + t.Fatalf("CalculateDefaultQuality = %d, want high-quality 13342", got) + } + + got, err = calc.CalculateDimensionsDefaultQuality(1024, 1024) + if err != nil { + t.Fatalf("CalculateDimensionsDefaultQuality returned err: %v", err) + } + if got != 7024 { + t.Fatalf("CalculateDimensionsDefaultQuality = %d, want high-quality 7024", got) + } +} + +func TestGPTImageTokenCalculatorDoesNotValidateSizeRules(t *testing.T) { + got, err := GPTImageTokenCalculator{}.Calculate("512x512", "low") + if err != nil { + t.Fatalf("Calculate returned err: %v", err) + } + if got != 145 { + t.Fatalf("Calculate = %d, want 145", got) + } +} + +func TestGPTImageTokenCalculatorDimensions(t *testing.T) { + got, err := GPTImageTokenCalculator{}.CalculateDimensions(3840, 2160, "medium") + if err != nil { + t.Fatalf("CalculateDimensions returned err: %v", err) + } + if got != 3336 { + t.Fatalf("CalculateDimensions = %d, want 3336", got) + } +} diff --git a/backend/internal/gateway/images.go b/backend/internal/gateway/images.go index a63de60..abff9eb 100644 --- a/backend/internal/gateway/images.go +++ b/backend/internal/gateway/images.go @@ -25,8 +25,9 @@ import ( "unicode" "github.com/tidwall/gjson" + "github.com/tidwall/sjson" - sdk "github.com/DouDOU-start/airgate-sdk/sdkgo" + sdk "github.com/DevilGenius/airgate-sdk/sdkgo" ) // imagesOAuthChatModel OAuth 下 REST→tools 翻译时使用的主 chat 模型。 @@ -90,6 +91,94 @@ func lookupImageGenOutputTokens(size, quality string) int { return 1056 } +func normalizeImageQualityDefaultMedium(quality string) string { + q := strings.ToLower(strings.TrimSpace(quality)) + if q == "" || strings.EqualFold(q, "auto") { + return "medium" + } + return q +} + +func calculateGPTImageOutputTokensForImages(modelName, size, quality string, numImages int) int { + if numImages <= 0 || !isGPTImageTwoModel(modelName) { + return 0 + } + if _, _, ok := parseImageSize(size); !ok { + size = "1024x1024" + } + tokens, err := GPTImageTokenCalculator{}.Calculate(size, normalizeImageQualityDefaultMedium(quality)) + if err != nil { + return 0 + } + return tokens * numImages +} + +func estimateGPTImageInputTokensForImages(refs []string, fallbackSize string) int { + if len(refs) == 0 { + return 0 + } + calc := GPTImageTokenCalculator{} + fallbackTokens := -1 + total := 0 + for _, ref := range refs { + if width, height, ok := imageDimensionsFromRef(ref); ok { + if tokens, err := calc.CalculateDimensionsDefaultQuality(width, height); err == nil { + total += tokens + continue + } + } + if fallbackTokens < 0 { + fallbackTokens = gptImageInputFallbackTokens(calc, fallbackSize) + } + total += fallbackTokens + } + return total +} + +func imageDimensionsFromRef(ref string) (int, int, bool) { + if !strings.HasPrefix(ref, "data:") { + return 0, 0, false + } + return imageDimensionsFromDataURL(ref) +} + +func imageDimensionsFromDataURL(ref string) (int, int, bool) { + comma := strings.IndexByte(ref, ',') + if comma < 0 || !strings.HasPrefix(strings.ToLower(ref[:comma]), "data:image/") { + return 0, 0, false + } + payload := ref[comma+1:] + if width, height, ok := imageDimensionsFromBase64Reader(base64.StdEncoding, payload); ok { + return width, height, true + } + return imageDimensionsFromBase64Reader(base64.RawStdEncoding, payload) +} + +func imageDimensionsFromBase64Reader(enc *base64.Encoding, payload string) (int, int, bool) { + cfg, _, err := image.DecodeConfig(base64.NewDecoder(enc, strings.NewReader(payload))) + if err != nil || cfg.Width <= 0 || cfg.Height <= 0 { + return 0, 0, false + } + return cfg.Width, cfg.Height, true +} + +func imageDimensionsFromBytes(data []byte) (int, int, bool) { + cfg, _, err := image.DecodeConfig(bytes.NewReader(data)) + if err != nil || cfg.Width <= 0 || cfg.Height <= 0 { + return 0, 0, false + } + return cfg.Width, cfg.Height, true +} + +func gptImageInputFallbackTokens(calc GPTImageTokenCalculator, fallbackSize string) int { + size := normalizeImagesResponseSize(fallbackSize) + if size == "" { + size = "1024x1024" + } + tokens, _ := calc.CalculateDefaultQuality(size) + return tokens +} + // estimateImageGenOutputTokens 汇总所有 image_generation_call 的估算 token 数。 func estimateImageGenOutputTokens(calls []ImageGenCall) int { total := 0 @@ -195,9 +284,14 @@ type imagesRequest struct { // /generations 只支持 application/json;/edits 同时支持 JSON(image 字段是 data URL/http(s) URL/数组) // 与 multipart/form-data(OpenAI SDK 标准)。 func buildAPIKeyImagesEditMultipartBody(body []byte, contentType string) ([]byte, string, error) { + multipartBody, multipartContentType, _, err := buildAPIKeyImagesEditMultipartBodyWithRequest(body, contentType) + return multipartBody, multipartContentType, err +} + +func buildAPIKeyImagesEditMultipartBodyWithRequest(body []byte, contentType string) ([]byte, string, *imagesRequest, error) { req, err := parseImagesRequest(body, contentType, true) if err != nil { - return nil, "", err + return nil, "", nil, err } var buf bytes.Buffer @@ -238,7 +332,7 @@ func buildAPIKeyImagesEditMultipartBody(body []byte, contentType string) ([]byte } if err := writeMultipartImageBytes(mw, fieldName, fmt.Sprintf("image-%d", i+1), mimeType, data); err != nil { _ = mw.Close() - return nil, "", err + return nil, "", nil, err } } if req.Mask != "" { @@ -257,13 +351,13 @@ func buildAPIKeyImagesEditMultipartBody(body []byte, contentType string) ([]byte } if err := writeMultipartImageBytes(mw, "mask", "mask", mimeType, data); err != nil { _ = mw.Close() - return nil, "", err + return nil, "", nil, err } } if err := mw.Close(); err != nil { - return nil, "", err + return nil, "", nil, err } - return buf.Bytes(), mw.FormDataContentType(), nil + return buf.Bytes(), mw.FormDataContentType(), req, nil } func writeMultipartImageBytes(mw *multipart.Writer, fieldName, baseName, mimeType string, data []byte) error { @@ -399,6 +493,164 @@ func parseImagesRequest(body []byte, contentType string, isEdit bool) (*imagesRe return parseImagesJSON(body, isEdit) } +func isMultipartContentType(contentType string) bool { + return strings.HasPrefix(strings.ToLower(strings.TrimSpace(contentType)), "multipart/") +} + +func imagesResponseOptionsFromRequestBody(body []byte, contentType string, isEdit bool) imagesResponseOptions { + if len(body) == 0 { + return imagesResponseOptions{} + } + if isEdit && isMultipartContentType(contentType) { + return imagesResponseOptionsFromMultipartBody(body, contentType) + } + size := strings.TrimSpace(gjson.GetBytes(body, "size").String()) + opts := imagesResponseOptions{ + BillingSize: size, + RequestSize: size, + RequestQuality: normalizeImageQualityDefaultMedium(gjson.GetBytes(body, "quality").String()), + RequestOutputFormat: strings.TrimSpace(gjson.GetBytes(body, "output_format").String()), + } + if req, err := parseImagesRequest(body, contentType, isEdit); err == nil { + applyImagesRequestOptions(&opts, req) + } else if prompt := strings.TrimSpace(gjson.GetBytes(body, "prompt").String()); prompt != "" { + opts.RequestTextInputTokens = estimatePromptTokens(prompt) + } + return opts +} + +func imagesResponseOptionsFromMultipartBody(body []byte, contentType string) imagesResponseOptions { + opts := imagesResponseOptions{} + _, params, err := mime.ParseMediaType(contentType) + if err != nil || params["boundary"] == "" { + return opts + } + + calc := GPTImageTokenCalculator{} + unknownImageRefs := 0 + fields := map[string]string{} + marker := []byte("--" + params["boundary"]) + pos := 0 + for { + startRel := bytes.Index(body[pos:], marker) + if startRel < 0 { + break + } + partStart := pos + startRel + len(marker) + if partStart+2 <= len(body) && body[partStart] == '-' && body[partStart+1] == '-' { + break + } + if partStart+2 <= len(body) && body[partStart] == '\r' && body[partStart+1] == '\n' { + partStart += 2 + } else if partStart < len(body) && body[partStart] == '\n' { + partStart++ + } + + headerEndRel, sepLen := multipartHeaderEnd(body[partStart:]) + if headerEndRel < 0 { + break + } + contentStart := partStart + headerEndRel + sepLen + nextRel := bytes.Index(body[contentStart:], marker) + if nextRel < 0 { + break + } + contentEnd := contentStart + nextRel + headers := body[partStart : partStart+headerEndRel] + content := body[contentStart:contentEnd] + name := multipartPartName(headers) + switch name { + case "prompt", "size", "quality", "output_format": + fields[name] = strings.TrimSpace(string(content)) + case "image", "image[]": + contentType := strings.ToLower(strings.TrimSpace(strings.Split(multipartPartContentType(headers), ";")[0])) + switch { + case strings.HasPrefix(contentType, "image/"): + if width, height, ok := imageDimensionsFromBytes(content); ok { + if tokens, err := calc.CalculateDimensionsDefaultQuality(width, height); err == nil { + opts.RequestImageInputTokens += tokens + } + } + case multipartPartFilename(headers) == "" && len(bytes.TrimSpace(content)) > 0: + unknownImageRefs++ + } + } + pos = contentEnd + } + + size := strings.TrimSpace(fields["size"]) + opts.BillingSize = size + opts.RequestSize = size + opts.RequestQuality = normalizeImageQualityDefaultMedium(fields["quality"]) + opts.RequestOutputFormat = strings.TrimSpace(fields["output_format"]) + opts.RequestTextInputTokens = estimatePromptTokens(fields["prompt"]) + if unknownImageRefs > 0 { + opts.RequestImageInputTokens += unknownImageRefs * gptImageInputFallbackTokens(calc, size) + } + return opts +} + +func applyImagesRequestOptions(opts *imagesResponseOptions, req *imagesRequest) { + if opts == nil || req == nil { + return + } + size := strings.TrimSpace(req.Size) + opts.BillingSize = size + opts.RequestSize = size + opts.RequestQuality = normalizeImageQualityDefaultMedium(req.Quality) + opts.RequestOutputFormat = strings.TrimSpace(req.OutputFormat) + opts.RequestTextInputTokens = estimatePromptTokens(req.Prompt) + if req.IsEdit { + opts.RequestImageInputTokens = estimateGPTImageInputTokensForImages(req.Images, req.Size) + } +} + +func multipartHeaderEnd(body []byte) (int, int) { + if idx := bytes.Index(body, []byte("\r\n\r\n")); idx >= 0 { + return idx, 4 + } + if idx := bytes.Index(body, []byte("\n\n")); idx >= 0 { + return idx, 2 + } + return -1, 0 +} + +func multipartPartName(headers []byte) string { + return multipartPartDispositionParam(headers, "name") +} + +func multipartPartFilename(headers []byte) string { + return multipartPartDispositionParam(headers, "filename") +} + +func multipartPartDispositionParam(headers []byte, paramName string) string { + for _, line := range strings.Split(string(headers), "\n") { + line = strings.TrimRight(line, "\r") + colon := strings.IndexByte(line, ':') + if colon < 0 || !strings.EqualFold(strings.TrimSpace(line[:colon]), "Content-Disposition") { + continue + } + _, params, err := mime.ParseMediaType(strings.TrimSpace(line[colon+1:])) + if err != nil { + return "" + } + return strings.TrimSpace(params[paramName]) + } + return "" +} + +func multipartPartContentType(headers []byte) string { + for _, line := range strings.Split(string(headers), "\n") { + line = strings.TrimRight(line, "\r") + colon := strings.IndexByte(line, ':') + if colon < 0 || !strings.EqualFold(strings.TrimSpace(line[:colon]), "Content-Type") { + continue + } + return strings.TrimSpace(line[colon+1:]) + } + return "" +} + func parseImagesJSON(body []byte, isEdit bool) (*imagesRequest, error) { prompt := strings.TrimSpace(gjson.GetBytes(body, "prompt").String()) if prompt == "" { @@ -621,7 +873,7 @@ func imageAPIConstraintLines(req *imagesRequest, isEdit, hasRegionAnnotation boo // 写进 constraints 反而是噪声,可能让 chat 模型输出无效约束词。 // SKILL.md: "gpt-image-2 always uses high fidelity for image inputs; // do not set input_fidelity with this model." - if !isGPTImage2Model(req.Model) { + if !isGPTImageTwoModel(req.Model) { if fidelity := cleanImageConstraintValue(req.InputFidelity); fidelity != "" { lines = append(lines, "Preserve the input image with "+fidelity+" fidelity.") } @@ -633,9 +885,9 @@ func imageAPIConstraintLines(req *imagesRequest, isEdit, hasRegionAnnotation boo return lines } -// isGPTImage2Model 判断 model 是否走 gpt-image-2 链路。 +// isGPTImageTwoModel 判断 model 是否走 gpt-image-2 链路。 // 空 model 视为 gpt-image-2,因为客户端不指定时上游默认升到 gpt-image-2。 -func isGPTImage2Model(model string) bool { +func isGPTImageTwoModel(model string) bool { m := strings.ToLower(strings.TrimSpace(model)) if m == "" { return true @@ -652,24 +904,18 @@ func isGPTImage2Model(model string) bool { // 失败(base64 异常 / 非 PNG/JPEG / WebP 没注册解码器)返回 ok=false, // 调用方继续用 fallback 链。 func imageActualSizeFromBase64(b64 string) (string, bool) { - if b64 == "" { + payload := strings.TrimSpace(b64) + if payload == "" { return "", false } - data, err := base64.StdEncoding.DecodeString(strings.TrimSpace(b64)) - if err != nil { - data, err = base64.RawStdEncoding.DecodeString(strings.TrimSpace(b64)) - if err != nil { - return "", false - } - } - cfg, _, err := image.DecodeConfig(bytes.NewReader(data)) - if err != nil { - return "", false + width, height, ok := imageDimensionsFromBase64Reader(base64.StdEncoding, payload) + if !ok { + width, height, ok = imageDimensionsFromBase64Reader(base64.RawStdEncoding, payload) } - if cfg.Width <= 0 || cfg.Height <= 0 { + if !ok { return "", false } - return fmt.Sprintf("%dx%d", cfg.Width, cfg.Height), true + return fmt.Sprintf("%dx%d", width, height), true } // imagePriceForSize 把生成的 size 映射成 USD 单价(1K/2K/4K 三档)。 @@ -719,7 +965,7 @@ func validateImageSize(size, model string) error { if s == "" || s == "auto" { return nil } - if !isGPTImage2Model(model) { + if !isGPTImageTwoModel(model) { return nil } width, height, ok := parseImageSize(s) @@ -1118,7 +1364,7 @@ func classifyImageGenCallFailures(failures []ImageGenCallFailure, fallbackDetail // buildImagesToolCreateMsg 把 Images REST 请求体翻译成 Codex HTTP SSE // /backend-api/codex/responses 接受的 Responses body(tools 数组带一个 // image_generation 项)。 -// 返回:上游消息 bytes;n(当前固定 1);prompt 估算的 token 数(用于计费)。 +// 返回:上游消息 bytes;n(当前固定 1);input 估算的 token 数(用于计费)。 // // contentType 仅在 isEdit=true 时需要(可能是 multipart/form-data)。 func buildImagesToolCreateMsg( @@ -1127,30 +1373,44 @@ func buildImagesToolCreateMsg( isEdit bool, session openAISessionResolution, ) ([]byte, int, int, error) { - req, err := parseImagesRequest(body, contentType, isEdit) + msg, n, estimate, err := buildImagesToolCreateMsgWithUsage(body, contentType, isEdit, session) if err != nil { return nil, 0, 0, err } + return msg, n, estimate.Total(), nil +} + +func buildImagesToolCreateMsgWithUsage( + body []byte, + contentType string, + isEdit bool, + session openAISessionResolution, +) ([]byte, int, imagesInputTokenEstimate, error) { + req, err := parseImagesRequest(body, contentType, isEdit) + if err != nil { + return nil, 0, imagesInputTokenEstimate{}, err + } + req.Quality = normalizeImageQualityDefaultMedium(req.Quality) // Responses API 的 image_generation tool 每次仅生成 1 张;n>1 在 REST 侧的语义 // 需要多轮工具调用才能模拟,暂不支持 —— V1 限定 n=1。 if req.N > 1 { - return nil, 0, 0, fmt.Errorf("OAuth 模式下 n 只能为 1(REST→tools 翻译路径暂不支持多图)") + return nil, 0, imagesInputTokenEstimate{}, fmt.Errorf("OAuth 模式下 n 只能为 1(REST→tools 翻译路径暂不支持多图)") } if err := shrinkResponsesInputImages(req); err != nil { - return nil, 0, 0, err + return nil, 0, imagesInputTokenEstimate{}, err } if isEdit && req.Mask != "" { if err := normalizeResponsesEditTargetImage(req); err != nil { - return nil, 0, 0, err + return nil, 0, imagesInputTokenEstimate{}, err } } regionAnnotation, err := buildEditRegionAnnotation(req) if err != nil { - return nil, 0, 0, err + return nil, 0, imagesInputTokenEstimate{}, err } if regionAnnotation != "" { if err := normalizeResponsesEditTargetAnnotationPair(req, ®ionAnnotation); err != nil { - return nil, 0, 0, err + return nil, 0, imagesInputTokenEstimate{}, err } } outputFormat := normalizedImageOutputFormat(req.OutputFormat) @@ -1173,9 +1433,7 @@ func buildImagesToolCreateMsg( if size := normalizeImageSizeForUpstream(req.Size); size != "" && !strings.EqualFold(size, "auto") { tool["size"] = size } - if quality := strings.TrimSpace(req.Quality); quality != "" { - tool["quality"] = quality - } + tool["quality"] = req.Quality if background := strings.TrimSpace(req.Background); background != "" { tool["background"] = background } @@ -1219,32 +1477,31 @@ func buildImagesToolCreateMsg( payload = applySessionFields(payload, session) msg, err := json.Marshal(payload) if err != nil { - return nil, 0, 0, err + return nil, 0, imagesInputTokenEstimate{}, err } - // input token 估算:文本 prompt + 每张参考图按 size 低质档估算(~272 tokens/1024²), - // 与 image_generation tool 输入图像的 token 级别一致。 - imageInputCount := len(req.Images) + inputImageRefs := append([]string{}, req.Images...) if regionAnnotation != "" { - imageInputCount++ + inputImageRefs = append(inputImageRefs, regionAnnotation) + } + estimate := imagesInputTokenEstimate{ + TextTokens: estimatePromptTokens(req.Prompt), + ImageTokens: estimateGPTImageInputTokensForImages(inputImageRefs, req.Size), } - inputTokens := estimatePromptTokens(req.Prompt) + estimateImageInputTokens(imageInputCount, req.Size) - return msg, req.N, inputTokens, nil + return msg, req.N, estimate, nil } -// estimateImageInputTokens 估算参考图输入的 token 总数。 -// 策略:套用 OpenAI 的 size→low-quality output token 表,近似代表"单张参考图"的 token 体量。 -// OpenAI 对图像输入另有单价(gpt-image-1.5 约 $10/1M),但当前注册表里只记了文本 $5/1M, -// 精度损失不到 2×,对总价(输出图像 token 主导)影响 < 5%,V1 可接受。 -func estimateImageInputTokens(count int, size string) int { - if count <= 0 { - return 0 - } - return count * lookupImageGenOutputTokens(size, "low") +type imagesInputTokenEstimate struct { + TextTokens int + ImageTokens int +} + +func (e imagesInputTokenEstimate) Total() int { + return e.TextTokens + e.ImageTokens } -// forwardImagesViaResponsesTool 把 OpenAI Images REST 请求翻译成 Responses API -// 的 image_generation tool 调用,跑 OAuth HTTP/SSE 通道,最后把生成的 base64 图像 -// 重新包装成 Images REST 响应返回给客户端。 +// forwardImagesViaResponsesTool 把 OpenAI Images REST 请求翻译成 Codex HTTP SSE +// /backend-api/codex/responses 的 image_generation tool 调用,最后把生成的 base64 +// 图像重新包装成 Images REST 响应返回给客户端。 // // 只在 OAuth 账号处理 /v1/images/generations 时使用;API Key 账号继续走原生 // REST 通道(见 handleImagesResponse)。 @@ -1256,8 +1513,8 @@ func (g *OpenAIGateway) forwardImagesViaResponsesToolWithURL(ctx context.Context start := time.Now() account := req.Account - session := resolveOpenAISession(req.Headers, req.Body) - updateSessionStateFromRequest(session) + session := resolveOpenAISession(req.Headers, req.Body, account.ID) + updateSessionStateFromRequest(session, account.ID) _, reqPath := resolveAPIKeyRoute(req) isEdit := isImagesEditRequest(reqPath) @@ -1286,7 +1543,7 @@ func (g *OpenAIGateway) forwardImagesViaResponsesToolWithURL(ctx context.Context "size", imgReq.Size, "n", imgReq.N, ) - createMsg, n, promptTokens, err := buildImagesToolCreateMsg(req.Body, contentType, isEdit, session) + createMsg, n, inputEstimate, err := buildImagesToolCreateMsgWithUsage(req.Body, contentType, isEdit, session) if err != nil { body := jsonError(err.Error()) return sdk.ForwardOutcome{ @@ -1386,7 +1643,7 @@ func (g *OpenAIGateway) forwardImagesViaResponsesToolWithURL(ctx context.Context wsResult = ParseSSEStream(resp.Body, handler) } if wsResult.ResponseID != "" && session.SessionKey != "" { - updateSessionStateResponseID(session.SessionKey, wsResult.ResponseID) + updateSessionStateResponseID(session.SessionKey, wsResult.ResponseID, account.ID) } elapsed := time.Since(start) @@ -1501,12 +1758,13 @@ func (g *OpenAIGateway) forwardImagesViaResponsesToolWithURL(ctx context.Context } numImages := len(wsResult.ImageGenCalls) + inputTokens := inputEstimate.Total() // 对外响应仍沿用 Images API 的 usage 口径; // 对内账单则拆成两段: // 1. Responses 主模型的上下文 token; // 2. 生图产出的按张费用。 billingModel := imageGenerationBillingModel(wsResult.ToolImageModel, imgReq.Model) - usage := newTokenUsage(billingModel, "", promptTokens, 0, 0, 0, handler.firstTokenMs) + usage := newTokenUsage(billingModel, "", inputTokens, 0, 0, 0, handler.firstTokenMs) contextModel := strings.TrimSpace(wsResult.Model) if contextModel == "" { contextModel = imagesOAuthChatModel @@ -1519,8 +1777,6 @@ func (g *OpenAIGateway) forwardImagesViaResponsesToolWithURL(ctx context.Context wsResult.OutputTokens, wsResult.CachedInputTokens, wsResult.ReasoningOutputTokens, - "responses_context", - "上下文", ) g.logger.Debug("Images OAuth result", "path", reqPath, @@ -1531,7 +1787,31 @@ func (g *OpenAIGateway) forwardImagesViaResponsesToolWithURL(ctx context.Context "num_images", numImages, ) - respBody := buildImagesRESTResponse(wsResult, promptTokens, 0, billingModel) + // 计费 size 优先级(高 → 低): + // 1. 上游 image_generation_call event 的 size 字段(telemetry 反馈) + // 2. 直接解码生成的 base64 图 header 拿真实宽高(auto 时最准的来源—— + // 上游有时不返 size 字段,但图本身永远是诚实的) + // 3. 客户端请求里的 size(最可能是 "auto" 兜底) + // + // 这里得到的 billingSize 同时用于 response.usage.output_tokens 估算, + // 避免为了写 usage 再额外解析一次图片。 + billingSize := imgReq.Size + if len(wsResult.ImageGenCalls) > 0 { + first := wsResult.ImageGenCalls[0] + if first.Size != "" { + billingSize = first.Size + } else if sz, ok := imageActualSizeFromBase64(first.Result); ok { + billingSize = sz + } + } + imageOutputTokens := calculateGPTImageOutputTokensForImages(billingModel, billingSize, imgReq.Quality, numImages) + respBody := buildImagesRESTResponse(wsResult, inputEstimate.TextTokens, imageOutputTokens, billingModel, imagesResponseOptions{ + BillingSize: billingSize, + RequestSize: imgReq.Size, + RequestQuality: imgReq.Quality, + RequestOutputFormat: imgReq.OutputFormat, + RequestImageInputTokens: inputEstimate.ImageTokens, + }) outcome := sdk.ForwardOutcome{ Kind: sdk.OutcomeSuccess, Upstream: sdk.UpstreamResponse{StatusCode: http.StatusOK}, @@ -1547,64 +1827,77 @@ func (g *OpenAIGateway) forwardImagesViaResponsesToolWithURL(ctx context.Context outcome.Upstream.Headers = http.Header{"Content-Type": []string{"application/json"}} } - // 计费 size 优先级(高 → 低): - // 1. 上游 image_generation_call event 的 size 字段(telemetry 反馈) - // 2. 直接解码生成的 base64 图 header 拿真实宽高(auto 时最准的来源—— - // 上游有时不返 size 字段,但图本身永远是诚实的) - // 3. 客户端请求里的 size(最可能是 "auto" 兜底) - billingSize := imgReq.Size - if len(wsResult.ImageGenCalls) > 0 { - first := wsResult.ImageGenCalls[0] - if first.Size != "" { - billingSize = first.Size - } else if sz, ok := imageActualSizeFromBase64(first.Result); ok { - billingSize = sz - } - } - // 图片尺寸作为通用 UsageAttribute 入库,后台费用明细可用它解释 1K/2K/4K 分档。 + // 图片尺寸写入 Usage 标准字段和 metadata,后台费用明细可用它解释 1K/2K/4K 分档。 + setUsageTokens(usage, inputTokens, imageOutputTokens, 0, 0) + setUsageInputTokenDetails(usage, inputEstimate.TextTokens, inputEstimate.ImageTokens) fillUsageCostPerImageBySize(usage, numImages, billingSize) + setUsageResponseID(usage, wsResult.ResponseID) return outcome, nil } // buildImagesRESTResponse 按 OpenAI Images API 官方契约打包响应。 -// 计费口径:prompt tokens + image output tokens,按实际响应的图像模型记录。 +// 计费口径:text input tokens + image input tokens + image output tokens,按实际响应的图像模型记录。 // instructions / 工具调用包装产生的额外 chat tokens 由内层吸收,不出现在对外 usage。 // 这样: // 1. 客户端拿到的 usage 数字语义与 OpenAI 原生 Images API 完全一致 // 2. 外层再套一层 AirGate 时,两级按同一口径独立计算,金额零偏差 -func buildImagesRESTResponse(wsResult WSResult, promptTokens, imageOutputTokens int, responseModel string) []byte { +func buildImagesRESTResponse(wsResult WSResult, textInputTokens, imageOutputTokens int, responseModel string, options ...imagesResponseOptions) []byte { if responseModel == "" { responseModel = imageToolCostModel } + opts := imagesResponseOptions{} + if len(options) > 0 { + opts = options[0] + } data := make([]map[string]any, 0, len(wsResult.ImageGenCalls)) for _, call := range wsResult.ImageGenCalls { item := map[string]any{ "b64_json": call.Result, } - if call.RevisedPrompt != "" { - item["revised_prompt"] = call.RevisedPrompt - } - // 透传上游实际生效的 image_generation 工具模型(从 response.tools[].model 提取)。 - // 客户端可据此判断"请求的 model"是否被上游静默降级。 - if wsResult.ToolImageModel != "" { - item["model"] = wsResult.ToolImageModel - } data = append(data, item) } + size := normalizeImagesResponseSize(opts.BillingSize) + if size == "" { + size = normalizeImagesResponseSize(opts.RequestSize) + } + if size == "" && len(wsResult.ImageGenCalls) > 0 { + size = normalizeImagesResponseSize(wsResult.ImageGenCalls[0].Size) + } + background := "" + if len(wsResult.ImageGenCalls) > 0 { + background = normalizeImagesResponseBackground(wsResult.ImageGenCalls[0].Background) + } + outputFormat := normalizeImagesResponseOutputFormat(opts.RequestOutputFormat) + if outputFormat == "" && len(wsResult.ImageGenCalls) > 0 { + outputFormat = normalizeImagesResponseOutputFormat(wsResult.ImageGenCalls[0].OutputFormat) + } + if outputFormat == "" { + outputFormat = "png" + } payload := map[string]any{ - "created": time.Now().Unix(), - "data": data, + "created": time.Now().Unix(), + "data": data, + "quality": normalizeImageQualityDefaultMedium(opts.RequestQuality), + "output_format": outputFormat, // root 级 model,供下一级 handleImagesResponse 做 fillCost 查价和 usage 记录。 "model": responseModel, } - if promptTokens+imageOutputTokens > 0 { + if size != "" { + payload["size"] = size + } + if background != "" { + payload["background"] = background + } + imageInputTokens := opts.RequestImageInputTokens + inputTokens := textInputTokens + imageInputTokens + if inputTokens+imageOutputTokens > 0 { payload["usage"] = map[string]any{ - "input_tokens": promptTokens, + "input_tokens": inputTokens, "output_tokens": imageOutputTokens, - "total_tokens": promptTokens + imageOutputTokens, + "total_tokens": inputTokens + imageOutputTokens, "input_tokens_details": map[string]any{ - "text_tokens": promptTokens, - "image_tokens": 0, + "text_tokens": textInputTokens, + "image_tokens": imageInputTokens, }, } } @@ -1612,6 +1905,246 @@ func buildImagesRESTResponse(wsResult WSResult, promptTokens, imageOutputTokens return b } +func normalizeImagesResponseSize(size string) string { + width, height, ok := parseImageSize(size) + if !ok { + return "" + } + return fmt.Sprintf("%dx%d", width, height) +} + +func normalizeImagesResponseBackground(background string) string { + switch strings.ToLower(strings.TrimSpace(background)) { + case "transparent": + return "transparent" + case "opaque": + return "opaque" + default: + return "" + } +} + +func normalizeImagesResponseOutputFormat(format string) string { + switch strings.ToLower(strings.TrimSpace(format)) { + case "png": + return "png" + case "webp": + return "webp" + case "jpeg": + return "jpeg" + default: + return "" + } +} + +func applyImagesResponseMetadata(body []byte, opts imagesResponseOptions, summary imagesResponseSummary) []byte { + if len(body) == 0 { + return body + } + + updated := body + data := gjson.GetBytes(body, "data") + if data.Exists() && data.IsArray() { + for idx, item := range data.Array() { + if !item.IsObject() { + continue + } + for _, field := range []string{"quality", "size", "background", "output_format"} { + if !item.Get(field).Exists() { + continue + } + next, err := sjson.DeleteBytes(updated, fmt.Sprintf("data.%d.%s", idx, field)) + if err != nil { + return body + } + updated = next + } + } + } + + rootSize := strings.TrimSpace(gjson.GetBytes(body, "size").String()) + rootBackground := strings.TrimSpace(gjson.GetBytes(body, "background").String()) + rootOutputFormat := strings.TrimSpace(gjson.GetBytes(body, "output_format").String()) + firstDataBackground := "" + firstDataOutputFormat := "" + if data.Exists() && data.IsArray() { + for _, item := range data.Array() { + if !item.IsObject() { + continue + } + if firstDataBackground == "" { + firstDataBackground = item.Get("background").String() + } + if firstDataOutputFormat == "" { + firstDataOutputFormat = item.Get("output_format").String() + } + if firstDataBackground != "" && firstDataOutputFormat != "" { + break + } + } + } + + quality := normalizeImageQualityDefaultMedium(opts.RequestQuality) + size := "" + if candidate := normalizeImagesResponseSize(summary.BillingSize); candidate != "" { + size = candidate + } else if candidate := normalizeImagesResponseSize(rootSize); candidate != "" { + size = candidate + } else if candidate := normalizeImagesResponseSize(opts.RequestSize); candidate != "" { + size = candidate + } + background := "" + if candidate := normalizeImagesResponseBackground(rootBackground); candidate != "" { + background = candidate + } else if candidate := normalizeImagesResponseBackground(firstDataBackground); candidate != "" { + background = candidate + } + outputFormat := "" + if candidate := normalizeImagesResponseOutputFormat(opts.RequestOutputFormat); candidate != "" { + outputFormat = candidate + } else if candidate := normalizeImagesResponseOutputFormat(rootOutputFormat); candidate != "" { + outputFormat = candidate + } else if candidate := normalizeImagesResponseOutputFormat(firstDataOutputFormat); candidate != "" { + outputFormat = candidate + } else { + outputFormat = "png" + } + + for _, field := range []string{"size", "background"} { + if (field == "size" && size != "") || (field == "background" && background != "") || !gjson.GetBytes(body, field).Exists() { + continue + } + next, err := sjson.DeleteBytes(updated, field) + if err != nil { + return body + } + updated = next + } + for _, field := range []struct { + name string + value string + }{ + {"background", background}, + {"output_format", outputFormat}, + {"quality", quality}, + {"size", size}, + } { + if field.value == "" { + continue + } + next, err := sjson.SetBytes(updated, field.name, field.value) + if err != nil { + return body + } + updated = next + } + + return updated +} + +func applyImagesResponseUsage(body []byte, opts imagesResponseOptions, outputTokens int) []byte { + if len(body) == 0 { + return body + } + updated := body + usageNode := gjson.GetBytes(updated, "usage") + existingInputNode := usageNode.Get("input_tokens") + existingOutputNode := usageNode.Get("output_tokens") + existingTextNode := usageNode.Get("input_tokens_details.text_tokens") + existingImageNode := usageNode.Get("input_tokens_details.image_tokens") + + textTokens := 0 + switch { + case existingTextNode.Exists(): + textTokens = int(existingTextNode.Int()) + case existingInputNode.Exists(): + textTokens = int(existingInputNode.Int()) + if existingImageNode.Exists() { + imageTokens := int(existingImageNode.Int()) + if imageTokens > 0 && textTokens >= imageTokens { + textTokens -= imageTokens + } + } + case opts.RequestTextInputTokens > 0: + textTokens = opts.RequestTextInputTokens + } + + imageInputTokens := int(existingImageNode.Int()) + if opts.RequestImageInputTokens > 0 { + imageInputTokens = opts.RequestImageInputTokens + } + if outputTokens <= 0 && existingOutputNode.Exists() { + outputTokens = int(existingOutputNode.Int()) + } + inputTokens := textTokens + imageInputTokens + if inputTokens+outputTokens == 0 && !usageNode.Exists() { + return body + } + var err error + updated, err = sjson.SetBytes(updated, "usage.input_tokens", inputTokens) + if err != nil { + return body + } + updated, err = sjson.SetBytes(updated, "usage.output_tokens", outputTokens) + if err != nil { + return body + } + updated, err = sjson.SetBytes(updated, "usage.total_tokens", inputTokens+outputTokens) + if err != nil { + return body + } + updated, err = sjson.SetBytes(updated, "usage.input_tokens_details.text_tokens", textTokens) + if err != nil { + return body + } + updated, err = sjson.SetBytes(updated, "usage.input_tokens_details.image_tokens", imageInputTokens) + if err != nil { + return body + } + return updated +} + +type imagesResponseSummary struct { + NumImages int + BillingSize string +} + +func summarizeImagesResponseForBilling(body []byte, fallbackSize string) imagesResponseSummary { + summary := imagesResponseSummary{BillingSize: strings.TrimSpace(fallbackSize)} + sizeResolved := false + if rootSize := normalizeImagesResponseSize(gjson.GetBytes(body, "size").String()); rootSize != "" { + summary.BillingSize = rootSize + sizeResolved = true + } + dataArr := gjson.GetBytes(body, "data") + if !dataArr.Exists() || !dataArr.IsArray() { + return summary + } + for _, item := range dataArr.Array() { + b64 := item.Get("b64_json").String() + if b64 != "" { + summary.NumImages++ + } else if u := item.Get("url").String(); strings.HasPrefix(u, "http://") || strings.HasPrefix(u, "https://") { + summary.NumImages++ + } + if sizeResolved { + continue + } + if sz := normalizeImagesResponseSize(item.Get("size").String()); sz != "" { + summary.BillingSize = sz + sizeResolved = true + continue + } + if b64 != "" { + if sz, ok := imageActualSizeFromBase64(b64); ok { + summary.BillingSize = sz + sizeResolved = true + } + } + } + return summary +} + // buildImagesErrorBody 返回 OpenAI 风格错误 body。 func buildImagesErrorBody(status int, message string) []byte { return buildImagesErrorBodyWithCode(status, "", message) @@ -1645,15 +2178,24 @@ func buildImagesErrorBodyWithCode(status int, code, message string) []byte { // 计费字段复用 parseUsage:gpt-image-1 / gpt-image-1.5 返回的 // usage.input_tokens / usage.output_tokens / usage.input_tokens_details.cached_tokens // 与 Responses API 字段同构,parseUsage 已经处理了 cached token 扣减。 -func handleImagesResponse(resp *http.Response, w http.ResponseWriter, sseKA *ssePingKeepAlive, start time.Time, fallbackModel string, billingSize ...string) (sdk.ForwardOutcome, error) { - return handleImagesResponseWithLogger(nil, resp, w, sseKA, start, fallbackModel, billingSize...) +type imagesResponseOptions struct { + BillingSize string + RequestSize string + RequestQuality string + RequestOutputFormat string + RequestTextInputTokens int + RequestImageInputTokens int +} + +func handleImagesResponse(resp *http.Response, w http.ResponseWriter, sseKA *ssePingKeepAlive, start time.Time, fallbackModel string, options ...imagesResponseOptions) (sdk.ForwardOutcome, error) { + return handleImagesResponseWithLogger(nil, resp, w, sseKA, start, fallbackModel, options...) } -func (g *OpenAIGateway) handleImagesResponse(resp *http.Response, w http.ResponseWriter, sseKA *ssePingKeepAlive, start time.Time, fallbackModel string, billingSize ...string) (sdk.ForwardOutcome, error) { - return handleImagesResponseWithLogger(g.logger, resp, w, sseKA, start, fallbackModel, billingSize...) +func (g *OpenAIGateway) handleImagesResponse(resp *http.Response, w http.ResponseWriter, sseKA *ssePingKeepAlive, start time.Time, fallbackModel string, options ...imagesResponseOptions) (sdk.ForwardOutcome, error) { + return handleImagesResponseWithLogger(g.logger, resp, w, sseKA, start, fallbackModel, options...) } -func handleImagesResponseWithLogger(logger *slog.Logger, resp *http.Response, w http.ResponseWriter, sseKA *ssePingKeepAlive, start time.Time, fallbackModel string, billingSize ...string) (sdk.ForwardOutcome, error) { +func handleImagesResponseWithLogger(logger *slog.Logger, resp *http.Response, w http.ResponseWriter, sseKA *ssePingKeepAlive, start time.Time, fallbackModel string, options ...imagesResponseOptions) (sdk.ForwardOutcome, error) { body, err := io.ReadAll(resp.Body) if err != nil { reason := fmt.Sprintf("读取 Images 响应失败: %v", err) @@ -1667,8 +2209,31 @@ func handleImagesResponseWithLogger(logger *slog.Logger, resp *http.Response, w return transientOutcome(reason), fmt.Errorf("%s", reason) } + opts := imagesResponseOptions{} + if len(options) > 0 { + opts = options[0] + } + + modelName := strings.TrimSpace(gjson.GetBytes(body, "model").String()) + if modelName == "" { + modelName = fallbackModel + } + if logger == nil { + logger = slog.Default() + } + summary := summarizeImagesResponseForBilling(body, opts.BillingSize) + body = applyImagesResponseMetadata(body, opts, summary) + imageOutputTokens := calculateGPTImageOutputTokensForImages(modelName, summary.BillingSize, opts.RequestQuality, summary.NumImages) + if !isGPTImageTwoModel(modelName) { + opts.RequestImageInputTokens = 0 + } + body = applyImagesResponseUsage(body, opts, imageOutputTokens) + parsed := parseUsage(body) headers := resp.Header.Clone() + if headers.Get("Content-Length") != "" { + headers.Set("Content-Length", strconv.Itoa(len(body))) + } if sseKA != nil { sseKA.Stop() @@ -1681,45 +2246,17 @@ func handleImagesResponseWithLogger(logger *slog.Logger, resp *http.Response, w _, _ = w.Write(body) } - modelName := strings.TrimSpace(gjson.GetBytes(body, "model").String()) - if modelName == "" { - modelName = fallbackModel - } - - numImages := countUsableImages(body) - if logger == nil { - logger = slog.Default() - } - billSize := "" - if len(billingSize) > 0 { - billSize = billingSize[0] - } - // 与 OAuth 路径对齐:优先从响应体获取真实分辨率用于计费。 - // 优先级:响应 data[0].size → 解码 base64 图片实际宽高 → 请求 size 兜底。 - if dataArr := gjson.GetBytes(body, "data"); dataArr.Exists() && dataArr.IsArray() { - for _, item := range dataArr.Array() { - if sz := strings.TrimSpace(item.Get("size").String()); sz != "" { - billSize = sz - break - } - if b64 := item.Get("b64_json").String(); b64 != "" { - if sz, ok := imageActualSizeFromBase64(b64); ok { - billSize = sz - break - } - } - } - } logger.Debug("images_native_result_returned", "request_model", fallbackModel, sdk.LogFieldModel, modelName, - "num_images", numImages, - "billing_size", billSize, + "num_images", summary.NumImages, + "billing_size", summary.BillingSize, ) elapsed := time.Since(start) usage := newTokenUsage(modelName, "", parsed.inputTokens, parsed.outputTokens, parsed.cachedInputTokens, 0, elapsed.Milliseconds()) - fillUsageCostPerImageBySize(usage, numImages, billSize) + setUsageInputTokenDetails(usage, parsed.textInputTokens, parsed.imageInputTokens) + fillUsageCostPerImageBySize(usage, summary.NumImages, summary.BillingSize) outcome := sdk.ForwardOutcome{ Kind: sdk.OutcomeSuccess, @@ -1735,24 +2272,6 @@ func handleImagesResponseWithLogger(logger *slog.Logger, resp *http.Response, w return outcome, nil } -// countUsableImages 统计响应体中实际携带图片数据(b64_json 或 url)的条目数。 -// 不含可用图片时返回 0,避免空响应被错误计费。 -func countUsableImages(body []byte) int { - dataArr := gjson.GetBytes(body, "data") - if !dataArr.Exists() || !dataArr.IsArray() { - return 0 - } - n := 0 - for _, item := range dataArr.Array() { - if item.Get("b64_json").String() != "" { - n++ - } else if u := item.Get("url").String(); strings.HasPrefix(u, "http://") || strings.HasPrefix(u, "https://") { - n++ - } - } - return n -} - // ────────────────────────────────────────────────────── // 异步图片任务轮询(apimart.ai 等异步上游) // ────────────────────────────────────────────────────── diff --git a/backend/internal/gateway/images_test.go b/backend/internal/gateway/images_test.go index c5e8b70..4161a04 100644 --- a/backend/internal/gateway/images_test.go +++ b/backend/internal/gateway/images_test.go @@ -25,7 +25,7 @@ import ( "github.com/gorilla/websocket" "github.com/tidwall/gjson" - sdk "github.com/DouDOU-start/airgate-sdk/sdkgo" + sdk "github.com/DevilGenius/airgate-sdk/sdkgo" ) func testPNGDataURL(width, height int, pixel func(int, int) color.RGBA) string { @@ -226,7 +226,7 @@ func TestHandleImagesResponse_TokenAttribution(t *testing.T) { } w := httptest.NewRecorder() - outcome, err := handleImagesResponse(resp, w, nil, time.Now(), "gpt-image-1.5", "2048x2048") + outcome, err := handleImagesResponse(resp, w, nil, time.Now(), "gpt-image-1.5", imagesResponseOptions{BillingSize: "2048x2048"}) if err != nil { t.Fatalf("handleImagesResponse returned err: %v", err) } @@ -250,19 +250,28 @@ func TestHandleImagesResponse_TokenAttribution(t *testing.T) { t.Errorf("cached_input_tokens = %d, want 10", got) } - if got := usageCostByKey(u, usageCostInput); !almostEqual(got, 0, 1e-9) { - t.Errorf("input cost = %v, want 0 (per-image billing)", got) + if got := u.InputCost; !almostEqual(got, 0.0002, 1e-9) { + t.Errorf("input cost = %v, want token cost 0.0002", got) } - if !almostEqual(u.AccountCost, 0.20, 1e-9) { - t.Errorf("AccountCost = %v, want 0.20 (1 image × 2K tier $0.20)", u.AccountCost) + if !almostEqual(u.AccountCost, 0.125005, 1e-9) { + t.Errorf("AccountCost = %v, want token cost 0.125005", u.AccountCost) + } + if got, want := usageImageUnitPrice(u), "0.2"; got != want { + t.Errorf("image unit_price = %q, want %q", got, want) } if w.Code != http.StatusOK { t.Errorf("writer status = %d, want 200", w.Code) } gotBody, _ := io.ReadAll(w.Result().Body) - if len(gotBody) != len(body) { - t.Errorf("response body len = %d, want %d", len(gotBody), len(body)) + if got := gjson.GetBytes(gotBody, "quality").String(); got != "medium" { + t.Errorf("response quality = %q, want medium", got) + } + if gjson.GetBytes(gotBody, "data.0.quality").Exists() { + t.Errorf("data[0].quality should be omitted") + } + if got := gjson.GetBytes(gotBody, "size").String(); got != "2048x2048" { + t.Errorf("response size = %q, want 2048x2048", got) } } @@ -321,9 +330,20 @@ func TestHandleImagesResponse_StreamWrapsRESTJSONAsSSE(t *testing.T) { t.Fatalf("writer Content-Type = %q, want text/event-stream", got) } gotBody := w.Body.String() - wantBody := "data: " + body + "\n\ndata: [DONE]\n\n" - if gotBody != wantBody { - t.Fatalf("writer body = %q, want %q", gotBody, wantBody) + const prefix = "data: " + const suffix = "\n\ndata: [DONE]\n\n" + if !strings.HasPrefix(gotBody, prefix) || !strings.HasSuffix(gotBody, suffix) { + t.Fatalf("writer body = %q, want REST JSON event and DONE", gotBody) + } + eventBody := strings.TrimSuffix(strings.TrimPrefix(gotBody, prefix), suffix) + if got := gjson.Get(eventBody, "quality").String(); got != "medium" { + t.Fatalf("event quality = %q, want medium", got) + } + if gjson.Get(eventBody, "data.0.quality").Exists() { + t.Fatalf("event data[0].quality should be omitted") + } + if got := gjson.Get(eventBody, "output_format").String(); got != "png" { + t.Fatalf("event output_format = %q, want png", got) } } @@ -335,12 +355,12 @@ func TestHandleImagesResponse_APIKeyBillingUsesRequestSize(t *testing.T) { Body: ioNopCloserFromString(body), } - outcome, err := handleImagesResponse(resp, nil, nil, time.Now(), "gpt-image-1.5", "3840x2160") + outcome, err := handleImagesResponse(resp, nil, nil, time.Now(), "gpt-image-1.5", imagesResponseOptions{BillingSize: "3840x2160"}) if err != nil { t.Fatalf("handleImagesResponse returned err: %v", err) } - if got, want := outcome.Usage.AccountCost, 0.80; !almostEqual(got, want, 1e-9) { - t.Fatalf("AccountCost = %v, want %v (2 images × 4K tier $0.40)", got, want) + if got, want := outcome.Usage.AccountCost, 0.00305; !almostEqual(got, want, 1e-9) { + t.Fatalf("AccountCost = %v, want %v (token standard cost)", got, want) } if got, want := usageImageUnitPrice(outcome.Usage), "0.4"; got != want { t.Fatalf("image unit_price = %q, want %q", got, want) @@ -359,14 +379,182 @@ func TestHandleImagesResponse_NonStreamReturnsBodyWithoutWriter(t *testing.T) { if err != nil { t.Fatalf("handleImagesResponse returned err: %v", err) } - if len(outcome.Upstream.Body) != len(body) { - t.Fatalf("Upstream.Body len = %d, want %d", len(outcome.Upstream.Body), len(body)) + if got := gjson.GetBytes(outcome.Upstream.Body, "quality").String(); got != "medium" { + t.Fatalf("Upstream.Body quality = %q, want medium", got) + } + if gjson.GetBytes(outcome.Upstream.Body, "data.0.quality").Exists() { + t.Fatalf("Upstream.Body data[0].quality should be omitted") + } + if gjson.GetBytes(outcome.Upstream.Body, "background").Exists() { + t.Fatalf("Upstream.Body background should be omitted when upstream did not return it") } if got := outcome.Upstream.Headers.Get("Content-Type"); got != "application/json" { t.Fatalf("Content-Type = %q, want application/json", got) } } +func TestHandleImagesResponse_RequestQualityOverridesResponseEcho(t *testing.T) { + body := `{"quality":"low","data":[{"url":"https://example/a.png","quality":"low"},{"url":"https://example/b.png"}],"usage":{"input_tokens":10,"output_tokens":100}}` + resp := &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{ + "Content-Type": []string{"application/json"}, + "Content-Length": []string{fmt.Sprint(len(body))}, + }, + Body: ioNopCloserFromString(body), + } + + outcome, err := handleImagesResponse(resp, nil, nil, time.Now(), "gpt-image-2", imagesResponseOptions{RequestQuality: "high"}) + if err != nil { + t.Fatalf("handleImagesResponse returned err: %v", err) + } + + var got map[string]any + if err := json.Unmarshal(outcome.Upstream.Body, &got); err != nil { + t.Fatalf("unmarshal response body: %v", err) + } + if got["quality"] != "high" { + t.Fatalf("root quality = %v, want high", got["quality"]) + } + data := got["data"].([]any) + for i, item := range data { + obj := item.(map[string]any) + if _, ok := obj["quality"]; ok { + t.Fatalf("data[%d].quality should be omitted", i) + } + } + if got["output_format"] != "png" { + t.Fatalf("root output_format = %v, want png", got["output_format"]) + } + if got, want := outcome.Upstream.Headers.Get("Content-Length"), fmt.Sprint(len(outcome.Upstream.Body)); got != want { + t.Fatalf("Content-Length = %q, want rewritten length %q", got, want) + } +} + +func TestHandleImagesResponse_DefaultQualityEchoesMedium(t *testing.T) { + body := `{"quality":"low","data":[{"url":"https://example/a.png","quality":"low"}],"usage":{"input_tokens":10,"output_tokens":100}}` + resp := &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: ioNopCloserFromString(body), + } + + outcome, err := handleImagesResponse(resp, nil, nil, time.Now(), "gpt-image-2") + if err != nil { + t.Fatalf("handleImagesResponse returned err: %v", err) + } + if got := gjson.GetBytes(outcome.Upstream.Body, "quality").String(); got != "medium" { + t.Fatalf("root quality = %q, want medium", got) + } + if gjson.GetBytes(outcome.Upstream.Body, "data.0.quality").Exists() { + t.Fatalf("data[0].quality should be omitted") + } +} + +func TestHandleImagesResponse_AutoQualityEchoesMedium(t *testing.T) { + body := `{"quality":"low","data":[{"url":"https://example/a.png","quality":"low"}],"usage":{"input_tokens":10,"output_tokens":100}}` + resp := &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: ioNopCloserFromString(body), + } + + outcome, err := handleImagesResponse(resp, nil, nil, time.Now(), "gpt-image-2", imagesResponseOptions{RequestQuality: "auto"}) + if err != nil { + t.Fatalf("handleImagesResponse returned err: %v", err) + } + if got := gjson.GetBytes(outcome.Upstream.Body, "quality").String(); got != "medium" { + t.Fatalf("root quality = %q, want medium", got) + } + if gjson.GetBytes(outcome.Upstream.Body, "data.0.quality").Exists() { + t.Fatalf("data[0].quality should be omitted") + } +} + +func TestHandleImagesResponse_GPTImageAddsCalculatedOutputTokens(t *testing.T) { + body := `{"data":[{"url":"https://example/a.png"}]}` + resp := &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: ioNopCloserFromString(body), + } + + outcome, err := handleImagesResponse(resp, nil, nil, time.Now(), "gpt-image-2", imagesResponseOptions{ + BillingSize: "3840x2160", + RequestQuality: "low", + }) + if err != nil { + t.Fatalf("handleImagesResponse returned err: %v", err) + } + if got := int(gjson.GetBytes(outcome.Upstream.Body, "usage.output_tokens").Int()); got != 371 { + t.Fatalf("response usage.output_tokens = %d, want 371", got) + } + if got := int(gjson.GetBytes(outcome.Upstream.Body, "usage.total_tokens").Int()); got != 371 { + t.Fatalf("response usage.total_tokens = %d, want 371", got) + } + if got := usageMetricInt(outcome.Usage, usageMetricOutputTokens); got != 371 { + t.Fatalf("Outcome output_tokens = %d, want 371", got) + } + if got := gjson.GetBytes(outcome.Upstream.Body, "quality").String(); got != "low" { + t.Fatalf("response quality = %q, want low", got) + } + if got := gjson.GetBytes(outcome.Upstream.Body, "size").String(); got != "3840x2160" { + t.Fatalf("response size = %q, want 3840x2160", got) + } +} + +func TestHandleImagesResponse_GPTImageAddsCalculatedInputImageTokens(t *testing.T) { + body := `{ + "data": [], + "usage": { + "input_tokens": 5, + "input_tokens_details": { + "image_tokens": 0, + "text_tokens": 5 + }, + "output_tokens": 14272, + "total_tokens": 14277 + } + }` + resp := &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: ioNopCloserFromString(body), + } + + outcome, err := handleImagesResponse(resp, nil, nil, time.Now(), "gpt-image-2", imagesResponseOptions{ + RequestImageInputTokens: 7024, + }) + if err != nil { + t.Fatalf("handleImagesResponse returned err: %v", err) + } + + if got := int(gjson.GetBytes(outcome.Upstream.Body, "usage.input_tokens_details.text_tokens").Int()); got != 5 { + t.Fatalf("text_tokens = %d, want 5", got) + } + if got := int(gjson.GetBytes(outcome.Upstream.Body, "usage.input_tokens_details.image_tokens").Int()); got != 7024 { + t.Fatalf("image_tokens = %d, want 7024", got) + } + if got := int(gjson.GetBytes(outcome.Upstream.Body, "usage.input_tokens").Int()); got != 7029 { + t.Fatalf("input_tokens = %d, want 7029", got) + } + if got := int(gjson.GetBytes(outcome.Upstream.Body, "usage.total_tokens").Int()); got != 21301 { + t.Fatalf("total_tokens = %d, want 21301", got) + } + if got := usageMetricInt(outcome.Usage, usageMetricInputTokens); got != 7029 { + t.Fatalf("Outcome input_tokens = %d, want 7029", got) + } + if got := usageMetricInt(outcome.Usage, usageMetricTextInputTokens); got != 5 { + t.Fatalf("Outcome input_text_tokens = %d, want 5", got) + } + if got := usageMetricInt(outcome.Usage, usageMetricImageInputTokens); got != 7024 { + t.Fatalf("Outcome input_image_tokens = %d, want 7024", got) + } + if got := usageMetricInt(outcome.Usage, usageMetricOutputTokens); got != 14272 { + t.Fatalf("Outcome output_tokens = %d, want 14272", got) + } +} + // TestHandleImagesResponse_FallbackModelWhenBodyLacksModel 验证 Images 响应里 // 没有 model 字段时,会回退到请求侧传入的 fallbackModel,避免 fillUsageCost 查不到定价。 func TestHandleImagesResponse_FallbackModelWhenBodyLacksModel(t *testing.T) { @@ -385,8 +573,11 @@ func TestHandleImagesResponse_FallbackModelWhenBodyLacksModel(t *testing.T) { t.Fatalf("Usage.Model = %q, want gpt-image-1 (fallback)", outcome.Usage.Model) } // Writer 为 nil 时 Upstream.Body/Headers 应带回给 core - if len(outcome.Upstream.Body) != len(body) { - t.Errorf("Upstream.Body len = %d, want %d", len(outcome.Upstream.Body), len(body)) + if got := gjson.GetBytes(outcome.Upstream.Body, "quality").String(); got != "medium" { + t.Errorf("Upstream.Body quality = %q, want medium", got) + } + if gjson.GetBytes(outcome.Upstream.Body, "data.0.quality").Exists() { + t.Errorf("Upstream.Body data[0].quality should be omitted") } if outcome.Upstream.Headers.Get("Content-Type") != "application/json" { t.Errorf("Upstream.Headers Content-Type not preserved") @@ -396,13 +587,18 @@ func TestHandleImagesResponse_FallbackModelWhenBodyLacksModel(t *testing.T) { } } -// TestFillUsageCostPerImageBySize_1K 按尺寸分档计费 1K。 +// TestFillUsageCostPerImageBySize_1K 按尺寸分档写入图片 metadata。 func TestFillUsageCostPerImageBySize_1K(t *testing.T) { usage := &sdk.Usage{Model: "gpt-image-1"} fillUsageCostPerImageBySize(usage, 3, "1024x1024") - // 3 张 × $0.10 = 0.30 - if !almostEqual(usage.AccountCost, 0.30, 1e-9) { - t.Errorf("AccountCost = %v, want 0.30", usage.AccountCost) + if !almostEqual(usage.AccountCost, 0, 1e-9) { + t.Errorf("AccountCost = %v, want 0 when no token usage is present", usage.AccountCost) + } + if got := usage.Metadata["openai.image.count"]; got != "3" { + t.Errorf("image count = %q, want 3", got) + } + if got := usageImageUnitPrice(usage); got != "0.1" { + t.Errorf("image unit_price = %q, want 0.1", got) } } @@ -459,57 +655,43 @@ func TestImagePriceForSize(t *testing.T) { } } -// TestFillUsageCostPerImageBySize 验证按张 × 档位单价填到 OutputCost。 +// TestFillUsageCostPerImageBySize 验证按张 × 档位单价只写入 metadata。 func TestFillUsageCostPerImageBySize(t *testing.T) { cases := []struct { - name string - size string - numImages int - want float64 + name string + size string + numImages int + wantUnitPrice string }{ - {"1K single", "1024x1024", 1, 0.10}, - {"1K triple", "1536x1024", 3, 0.30}, - {"2K single", "2048x2048", 1, 0.20}, - {"4K double", "3840x2160", 2, 0.80}, - {"auto fallback to 1K", "auto", 4, 0.40}, - {"zero images skipped", "1024x1024", 0, 0}, + {"1K single", "1024x1024", 1, "0.1"}, + {"1K triple", "1536x1024", 3, "0.1"}, + {"2K single", "2048x2048", 1, "0.2"}, + {"4K double", "3840x2160", 2, "0.4"}, + {"auto fallback to 1K", "auto", 4, "0.1"}, + {"zero images skipped", "1024x1024", 0, ""}, } for _, tc := range cases { t.Run(tc.name, func(t *testing.T) { usage := &sdk.Usage{Model: "gpt-image-2"} fillUsageCostPerImageBySize(usage, tc.numImages, tc.size) - if !almostEqual(usage.AccountCost, tc.want, 1e-9) { - t.Errorf("AccountCost = %v, want %v", usage.AccountCost, tc.want) + if !almostEqual(usage.OutputCost, 0, 1e-9) { + t.Errorf("output cost = %v, want 0", usage.OutputCost) + } + if !almostEqual(usage.InputCost, 0, 1e-9) { + t.Errorf("input cost = %v, want 0", usage.InputCost) } - if !almostEqual(usageCostByKey(usage, usageCostInput), 0, 1e-9) { - t.Errorf("input cost = %v, want 0", usageCostByKey(usage, usageCostInput)) + if got := usageImageUnitPrice(usage); got != tc.wantUnitPrice { + t.Errorf("image unit_price = %q, want %q", got, tc.wantUnitPrice) } }) } } -func usageCostByKey(usage *sdk.Usage, key string) float64 { - if usage == nil { - return 0 - } - for _, detail := range usage.CostDetails { - if detail.Key == key { - return detail.AccountCost - } - } - return 0 -} - func usageImageUnitPrice(usage *sdk.Usage) string { if usage == nil { return "" } - for _, detail := range usage.CostDetails { - if detail.Key == usageCostImage { - return detail.Metadata["unit_price"] - } - } - return "" + return usage.Metadata["openai.image.unit_price"] } func almostEqual(a, b, eps float64) bool { @@ -541,6 +723,29 @@ func TestParseUsage_ToolImageGen(t *testing.T) { } } +func TestParseUsage_InputTokenDetails(t *testing.T) { + body := []byte(`{ + "usage": { + "input_tokens": 23342, + "output_tokens": 50, + "input_tokens_details": { + "text_tokens": 42, + "image_tokens": 23300 + } + } + }`) + got := parseUsage(body) + if got.inputTokens != 23342 { + t.Fatalf("inputTokens = %d, want 23342", got.inputTokens) + } + if got.textInputTokens != 42 { + t.Fatalf("textInputTokens = %d, want 42", got.textInputTokens) + } + if got.imageInputTokens != 23300 { + t.Fatalf("imageInputTokens = %d, want 23300", got.imageInputTokens) + } +} + func TestParseUsage_ImageGenerationCallSummary(t *testing.T) { b64 := strings.TrimPrefix(testPNGDataURL(1024, 1536, func(x, y int) color.RGBA { return color.RGBA{R: 1, G: 2, B: 3, A: 255} @@ -609,42 +814,52 @@ func TestParseSSEUsage_ToolImageGen(t *testing.T) { } } -// TestFillUsageCostWithImageTool 叠加计费:主 model (gpt-5.4) 的 chat token 按 -// 其单价、image tool 按尺寸分档计费。 +// TestFillUsageCostWithImageTool 验证主 model token 按其单价计费,image tool +// 的固定单张价只写入 metadata,由 Core 决定是否替代最终计费。 func TestFillUsageCostWithImageTool(t *testing.T) { usage := newTokenUsage("gpt-5.4", "", 1000, 500, 0, 0, 0) fillUsageCostWithImageTool(usage, 1, "1024x1024") // 主 gpt-5.4 standard: input=$2.5/1M → 0.0025, output=$15/1M → 0.0075 - // image tool: 1 张 × $0.10 (1K) = 0.10 - // total account cost = 0.0025 + 0.0075 + 0.10 = 0.1100 - if !almostEqual(usageCostByKey(usage, usageCostInput), 0.0025, 1e-9) { - t.Errorf("input cost = %v, want 0.0025", usageCostByKey(usage, usageCostInput)) + // image tool 单张价只进入 metadata;标准成本仍是 token 成本。 + // total account cost = 0.0025 + 0.0075 = 0.0100 + if !almostEqual(usage.InputCost, 0.0025, 1e-9) { + t.Errorf("input cost = %v, want 0.0025", usage.InputCost) + } + if !almostEqual(usage.AccountCost, 0.0100, 1e-9) { + t.Errorf("AccountCost = %v, want 0.0100", usage.AccountCost) } - if !almostEqual(usage.AccountCost, 0.1100, 1e-9) { - t.Errorf("AccountCost = %v, want 0.1100", usage.AccountCost) + if got := usage.InputPrice; !almostEqual(got, 2.5, 1e-9) { + t.Errorf("input unit price = %v, want 2.5", got) } - if got := usage.Metrics[0].Metadata["unit_price"]; got != "2.5" { - t.Errorf("input unit_price = %q, want 2.5", got) + if got := usageImageUnitPrice(usage); got != "0.1" { + t.Errorf("image unit price = %q, want 0.1", got) } } -func TestAddUsageCostForModel_CombinesResponsesContextAndImageCost(t *testing.T) { +func TestFillUsageCostPerImageBySize_ReplacesContextCostWithImageModelTokenCost(t *testing.T) { usage := newTokenUsage("gpt-image-2", "", 12, 0, 0, 0, 0) - addUsageCostForModel(usage, "gpt-5.4", "", 1000, 500, 0, 0, "responses_context", "上下文") + addUsageCostForModel(usage, "gpt-5.4", "", 1000, 500, 0, 0) + setUsageTokens(usage, 12, 1056, 0, 0) fillUsageCostPerImageBySize(usage, 1, "1024x1024") - if got := usageCostByKey(usage, "responses_context_"+usageCostInput); !almostEqual(got, 0.0025, 1e-9) { - t.Errorf("context input cost = %v, want 0.0025", got) + if got := usage.InputPrice; !almostEqual(got, 5, 1e-9) { + t.Errorf("image model input price = %v, want 5", got) } - if got := usageCostByKey(usage, "responses_context_"+usageCostOutput); !almostEqual(got, 0.0075, 1e-9) { - t.Errorf("context output cost = %v, want 0.0075", got) + if got := usage.OutputPrice; !almostEqual(got, 30, 1e-9) { + t.Errorf("image model output price = %v, want 30", got) } - if got := usageCostByKey(usage, usageCostImage); !almostEqual(got, 0.10, 1e-9) { - t.Errorf("image cost = %v, want 0.10", got) + if got := usage.InputCost; !almostEqual(got, 0.00006, 1e-9) { + t.Errorf("image model input cost = %v, want 0.00006", got) } - if !almostEqual(usage.AccountCost, 0.1100, 1e-9) { - t.Errorf("AccountCost = %v, want 0.1100", usage.AccountCost) + if got := usage.OutputCost; !almostEqual(got, 0.03168, 1e-9) { + t.Errorf("image model output cost = %v, want 0.03168", got) + } + if got := usageImageUnitPrice(usage); got != "0.1" { + t.Errorf("image unit price = %q, want 0.1", got) + } + if !almostEqual(usage.AccountCost, 0.03174, 1e-9) { + t.Errorf("AccountCost = %v, want 0.03174", usage.AccountCost) } if usage.Model != "gpt-image-2" { t.Errorf("Usage.Model = %q, want gpt-image-2", usage.Model) @@ -658,8 +873,8 @@ func TestFillUsageCostWithImageTool_NoToolUsage(t *testing.T) { if usageMetricInt(usage, usageMetricInputTokens) != 1000 || usageMetricInt(usage, usageMetricOutputTokens) != 500 { t.Errorf("token counts mutated when no image tool usage") } - if !almostEqual(usageCostByKey(usage, usageCostInput), 0.0025, 1e-9) { - t.Errorf("input cost = %v, want 0.0025", usageCostByKey(usage, usageCostInput)) + if !almostEqual(usage.InputCost, 0.0025, 1e-9) { + t.Errorf("input cost = %v, want 0.0025", usage.InputCost) } } @@ -823,7 +1038,7 @@ func TestImageTaskQualityEcho(t *testing.T) { // Responses body,tool 配置保持 Codex 对齐的极简 schema。 func TestBuildImagesToolCreateMsg(t *testing.T) { body := []byte(`{"model":"gpt-image-1.5","prompt":"a shiba","n":1,"size":"1024x1024","quality":"low","background":"transparent","output_format":"png"}`) - msg, n, promptTokens, err := buildImagesToolCreateMsg(body, "application/json", false, openAISessionResolution{}) + msg, n, inputTokens, err := buildImagesToolCreateMsg(body, "application/json", false, openAISessionResolution{}) if err != nil { t.Fatalf("buildImagesToolCreateMsg returned err: %v", err) } @@ -831,8 +1046,8 @@ func TestBuildImagesToolCreateMsg(t *testing.T) { t.Errorf("n = %d, want 1", n) } // "a shiba" = 7 runes → (7+2)/3 = 3 tokens - if promptTokens != 3 { - t.Errorf("promptTokens = %d, want 3", promptTokens) + if inputTokens != 3 { + t.Errorf("inputTokens = %d, want 3", inputTokens) } if gjson.GetBytes(msg, "type").Exists() { @@ -909,6 +1124,9 @@ func TestBuildImagesToolCreateMsg_ClampsOversizedSize(t *testing.T) { if got := gjson.GetBytes(msg, "tools.0.size").String(); got != "3840x2160" { t.Errorf("tools[0].size = %q, want clamped 3840x2160", got) } + if got := gjson.GetBytes(msg, "tools.0.quality").String(); got != "medium" { + t.Errorf("tools[0].quality = %q, want medium", got) + } } func TestBuildImagesToolCreateMsg_KeepsSessionFieldsWithoutEventWrapper(t *testing.T) { @@ -984,9 +1202,11 @@ func TestBuildImagesToolCreateMsg_Edit_JSON(t *testing.T) { if n != 1 { t.Errorf("n = %d, want 1", n) } - // text prompt "make it cyberpunk" = 17 runes → 6;reference image + region annotation = 2 * 272 → 550 - if inputTokens != 6+272*2 { - t.Errorf("inputTokens = %d, want %d", inputTokens, 6+272*2) + // text prompt "make it cyberpunk" = 17 runes → 6;2×2 reference image + region annotation + // 都按 gpt-image-2 high fidelity 计算,单张 4609。 + wantImageTokens := 4609 * 2 + if inputTokens != 6+wantImageTokens { + t.Errorf("inputTokens = %d, want %d", inputTokens, 6+wantImageTokens) } content := gjson.GetBytes(msg, "input.0.content") if !content.IsArray() || len(content.Array()) != 3 { @@ -1091,11 +1311,11 @@ func TestBuildImagesToolCreateMsg_Edit_MissingImage(t *testing.T) { } } -// TestBuildImagesToolCreateMsg_Edit_Img2ImgNoMask 钉死 OAuth Responses-tool 路径 +// TestBuildImagesToolCreateMsg_Edit_ImageToImageNoMask 钉死 OAuth Responses-tool 路径 // 上"纯图生图"(无 mask)的载荷形状:必须 1 张 input_image 而且不生成 region // annotation。这是 studio ComposerBar 图生图入口最常见的请求形态,是 OAuth // 路径上历史最容易"看上去通了但参考图没生效"的退化点。 -func TestBuildImagesToolCreateMsg_Edit_Img2ImgNoMask(t *testing.T) { +func TestBuildImagesToolCreateMsg_Edit_ImageToImageNoMask(t *testing.T) { imageRef := testPNGDataURL(2, 2, func(x, y int) color.RGBA { return color.RGBA{R: 120, G: 30, B: 30, A: 255} }) @@ -1533,6 +1753,106 @@ func TestParseImagesEditMultipart(t *testing.T) { } } +func TestImagesResponseOptionsFromMultipartDoesNotParseImagePart(t *testing.T) { + var buf bytes.Buffer + mw := multipart.NewWriter(&buf) + h := textproto.MIMEHeader{} + h.Set("Content-Disposition", `form-data; name="image"; filename="bad.bin"`) + h.Set("Content-Type", "application/octet-stream") + w, _ := mw.CreatePart(h) + _, _ = w.Write(bytes.Repeat([]byte{0x01, 0x02, 0x03, 0x04}, 1024)) + _ = mw.WriteField("quality", "high") + _ = mw.WriteField("size", "2048x2048") + _ = mw.WriteField("background", "opaque") + _ = mw.WriteField("output_format", "webp") + _ = mw.Close() + + opts := imagesResponseOptionsFromRequestBody(buf.Bytes(), mw.FormDataContentType(), true) + if opts.RequestQuality != "high" || opts.BillingSize != "2048x2048" || + opts.RequestOutputFormat != "webp" { + t.Fatalf("options = %+v, want scalar response options", opts) + } + if opts.RequestImageInputTokens != 0 { + t.Fatalf("RequestImageInputTokens = %d, want 0 for unparsable image part", opts.RequestImageInputTokens) + } +} + +func TestImagesResponseOptionsFromMultipartCalculatesImageInputTokens(t *testing.T) { + pngBytes := testPNGBytes(2, 2, func(x, y int) color.RGBA { + return color.RGBA{R: 10, G: 20, B: 30, A: 255} + }) + var buf bytes.Buffer + mw := multipart.NewWriter(&buf) + h := textproto.MIMEHeader{} + h.Set("Content-Disposition", `form-data; name="image"; filename="input.png"`) + h.Set("Content-Type", "image/png") + w, _ := mw.CreatePart(h) + _, _ = w.Write(pngBytes) + _ = mw.WriteField("prompt", "edit it") + _ = mw.WriteField("size", "1024x1024") + _ = mw.WriteField("quality", "auto") + _ = mw.Close() + + opts := imagesResponseOptionsFromRequestBody(buf.Bytes(), mw.FormDataContentType(), true) + if opts.RequestTextInputTokens != 3 { + t.Fatalf("RequestTextInputTokens = %d, want 3", opts.RequestTextInputTokens) + } + if opts.RequestImageInputTokens != 4609 { + t.Fatalf("RequestImageInputTokens = %d, want 4609", opts.RequestImageInputTokens) + } + if opts.RequestQuality != "medium" { + t.Fatalf("RequestQuality = %q, want medium", opts.RequestQuality) + } +} + +func TestImagesResponseOptionsFromJSONRemoteImageDoesNotDownloadForTokenEstimate(t *testing.T) { + hits := 0 + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + hits++ + w.Header().Set("Content-Type", "image/png") + _, _ = w.Write(testPNGBytes(2, 2, func(x, y int) color.RGBA { + return color.RGBA{R: 1, G: 2, B: 3, A: 255} + })) + })) + defer srv.Close() + + body := []byte(fmt.Sprintf(`{ + "prompt":"edit it", + "size":"1024x1024", + "image":%q + }`, srv.URL+"/input.png")) + + opts := imagesResponseOptionsFromRequestBody(body, "application/json", true) + if hits != 0 { + t.Fatalf("remote image was downloaded %d times while estimating tokens", hits) + } + if opts.RequestImageInputTokens != 7024 { + t.Fatalf("RequestImageInputTokens = %d, want 7024 fallback high-quality 1024x1024", opts.RequestImageInputTokens) + } +} + +func TestImagesResponseOptionsDefaultAndAutoQualityUseMedium(t *testing.T) { + opts := imagesResponseOptionsFromRequestBody([]byte(`{"prompt":"x"}`), "application/json", false) + if opts.RequestQuality != "medium" { + t.Fatalf("default JSON quality = %q, want medium", opts.RequestQuality) + } + + opts = imagesResponseOptionsFromRequestBody([]byte(`{"prompt":"x","quality":"auto"}`), "application/json", false) + if opts.RequestQuality != "medium" { + t.Fatalf("auto JSON quality = %q, want medium", opts.RequestQuality) + } + + var buf bytes.Buffer + mw := multipart.NewWriter(&buf) + _ = mw.WriteField("prompt", "x") + _ = mw.WriteField("quality", "auto") + _ = mw.Close() + opts = imagesResponseOptionsFromRequestBody(buf.Bytes(), mw.FormDataContentType(), true) + if opts.RequestQuality != "medium" { + t.Fatalf("auto multipart quality = %q, want medium", opts.RequestQuality) + } +} + func TestNormalizeImageRef(t *testing.T) { cases := map[string]string{ "data:image/png;base64,AAA=": "data:image/png;base64,AAA=", @@ -1578,7 +1898,7 @@ func TestEstimatePromptTokens(t *testing.T) { } // TestBuildImagesRESTResponse 把 WSResult 打包回 OpenAI Images REST 响应格式。 -// 计费口径对齐 OpenAI 官方:usage.input_tokens = prompt tokens、output_tokens = 图像 tokens、 +// 计费口径对齐 OpenAI 官方:usage.input_tokens = text + image input tokens、output_tokens = 图像 tokens、 // root 级 model 使用实际响应的图像模型。instructions / 工具包装的 chat tokens 不暴露。 func TestBuildImagesRESTResponse(t *testing.T) { ws := WSResult{ @@ -1588,13 +1908,17 @@ func TestBuildImagesRESTResponse(t *testing.T) { ToolImageOutputTokens: 4160, ToolImageModel: "gpt-image-2", ImageGenCalls: []ImageGenCall{ - {Result: "PNG_BASE64_A", RevisedPrompt: "revised a"}, + {Result: "PNG_BASE64_A", Background: "opaque", RevisedPrompt: "revised a"}, {Result: "PNG_BASE64_B"}, }, } promptTokens := 12 imageOut := 4160 - body := buildImagesRESTResponse(ws, promptTokens, imageOut, imageGenerationBillingModel(ws.ToolImageModel, "dall-e-3")) + body := buildImagesRESTResponse(ws, promptTokens, imageOut, imageGenerationBillingModel(ws.ToolImageModel, "dall-e-3"), imagesResponseOptions{ + BillingSize: "1024x1024", + RequestQuality: "high", + RequestOutputFormat: "webp", + }) var got map[string]any if err := json.Unmarshal(body, &got); err != nil { @@ -1603,20 +1927,41 @@ func TestBuildImagesRESTResponse(t *testing.T) { if got["model"] != "gpt-image-2" { t.Errorf("root model = %v, want gpt-image-2", got["model"]) } + if got["quality"] != "high" { + t.Errorf("root quality = %v, want high", got["quality"]) + } + if got["size"] != "1024x1024" { + t.Errorf("root size = %v, want 1024x1024", got["size"]) + } + if got["background"] != "opaque" { + t.Errorf("root background = %v, want opaque", got["background"]) + } + if got["output_format"] != "webp" { + t.Errorf("root output_format = %v, want webp", got["output_format"]) + } + if _, ok := got["revised_prompt"]; ok { + t.Errorf("root revised_prompt should be omitted") + } data, _ := got["data"].([]any) if len(data) != 2 { t.Fatalf("data len = %d, want 2", len(data)) } first, _ := data[0].(map[string]any) - if first["b64_json"] != "PNG_BASE64_A" || first["revised_prompt"] != "revised a" { + if first["b64_json"] != "PNG_BASE64_A" { t.Errorf("data[0] fields wrong: %+v", first) } + if _, ok := first["revised_prompt"]; ok { + t.Errorf("data[0].revised_prompt should be omitted") + } + if _, ok := first["quality"]; ok { + t.Errorf("data[0].quality should be omitted") + } second, _ := data[1].(map[string]any) if second["b64_json"] != "PNG_BASE64_B" { t.Errorf("data[1].b64_json = %v, want PNG_BASE64_B", second["b64_json"]) } if _, ok := second["revised_prompt"]; ok { - t.Errorf("empty revised_prompt should be omitted") + t.Errorf("data[1].revised_prompt should be omitted") } usage, ok := got["usage"].(map[string]any) if !ok { @@ -1633,6 +1978,56 @@ func TestBuildImagesRESTResponse(t *testing.T) { } } +func TestBuildImagesRESTResponse_IncludesImageInputTokens(t *testing.T) { + ws := WSResult{ImageGenCalls: []ImageGenCall{{Result: "PNG_BASE64"}}} + body := buildImagesRESTResponse(ws, 5, 14272, "gpt-image-2", imagesResponseOptions{ + RequestImageInputTokens: 7024, + }) + + if got := int(gjson.GetBytes(body, "usage.input_tokens_details.text_tokens").Int()); got != 5 { + t.Fatalf("text_tokens = %d, want 5", got) + } + if got := int(gjson.GetBytes(body, "usage.input_tokens_details.image_tokens").Int()); got != 7024 { + t.Fatalf("image_tokens = %d, want 7024", got) + } + if got := int(gjson.GetBytes(body, "usage.input_tokens").Int()); got != 7029 { + t.Fatalf("input_tokens = %d, want 7029", got) + } + if got := int(gjson.GetBytes(body, "usage.total_tokens").Int()); got != 21301 { + t.Fatalf("total_tokens = %d, want 21301", got) + } +} + +func TestBuildImagesRESTResponse_QualityEchoPrefersRequest(t *testing.T) { + ws := WSResult{ImageGenCalls: []ImageGenCall{{Result: "PNG_BASE64", Quality: "low"}}} + body := buildImagesRESTResponse(ws, 1, 2, "gpt-image-2", imagesResponseOptions{RequestQuality: "high"}) + + if got := gjson.GetBytes(body, "quality").String(); got != "high" { + t.Fatalf("root quality = %q, want high", got) + } + if gjson.GetBytes(body, "data.0.quality").Exists() { + t.Fatalf("data[0].quality should be omitted") + } + if gjson.GetBytes(body, "background").Exists() { + t.Fatalf("root background should be omitted when upstream did not return it") + } + + body = buildImagesRESTResponse(ws, 1, 2, "gpt-image-2") + if got := gjson.GetBytes(body, "quality").String(); got != "medium" { + t.Fatalf("root quality without request = %q, want medium", got) + } + + body = buildImagesRESTResponse(ws, 1, 2, "gpt-image-2", imagesResponseOptions{RequestQuality: ""}) + if got := gjson.GetBytes(body, "quality").String(); got != "medium" { + t.Fatalf("root quality with default request = %q, want medium", got) + } + + body = buildImagesRESTResponse(ws, 1, 2, "gpt-image-2", imagesResponseOptions{RequestQuality: "auto"}) + if got := gjson.GetBytes(body, "quality").String(); got != "medium" { + t.Fatalf("root quality with auto request = %q, want medium", got) + } +} + // TestBuildImagesRESTResponse_ChainedCostParity 验证 AirGate 套 AirGate 时两级 // 金额一致:下一级拿到 body 按 root model 单价重算,应等于本级结果。 func TestBuildImagesRESTResponse_ChainedCostParity(t *testing.T) { @@ -1918,10 +2313,10 @@ func TestForwardImagesViaResponsesTool_UsesHTTPSSEForLargeEditImage(t *testing.T } } -// TestBuildImagesToolCreateMsg_Edit_GPTImage2_SkipsInputFidelity gpt-image-2 始终 +// TestBuildImagesToolCreateMsg_Edit_GPTImageTwo_SkipsInputFidelity gpt-image-2 始终 // 用 high fidelity 处理输入图,input_fidelity 是 no-op,constraints 不应再追加。 // SKILL.md: "do not set input_fidelity with this model". -func TestBuildImagesToolCreateMsg_Edit_GPTImage2_SkipsInputFidelity(t *testing.T) { +func TestBuildImagesToolCreateMsg_Edit_GPTImageTwo_SkipsInputFidelity(t *testing.T) { imageRef := testPNGDataURL(2, 2, func(x, y int) color.RGBA { return color.RGBA{R: 80, G: 90, B: 100, A: 255} }) diff --git a/backend/internal/gateway/images_web_reverse.go b/backend/internal/gateway/images_web_reverse.go index c48492f..64d7149 100644 --- a/backend/internal/gateway/images_web_reverse.go +++ b/backend/internal/gateway/images_web_reverse.go @@ -16,9 +16,9 @@ import ( "strings" "time" - sdk "github.com/DouDOU-start/airgate-sdk/sdkgo" + sdk "github.com/DevilGenius/airgate-sdk/sdkgo" - "github.com/DouDOU-start/airgate-openai/backend/internal/gateway/imgen" + "github.com/DevilGenius/airgate-openai/backend/internal/gateway/imgen" ) // decodeImageRefs 把 parseImagesRequest 返回的 data URL / http URL 字符串列表 @@ -167,6 +167,11 @@ func (g *OpenAIGateway) forwardImagesViaWebReverse(ctx context.Context, req *sdk "n", imgReq.N, ) + inputEstimate := imagesInputTokenEstimate{TextTokens: estimatePromptTokens(imgReq.Prompt)} + if isEdit { + inputEstimate.ImageTokens = estimateGPTImageInputTokensForImages(imgReq.Images, imgReq.Size) + } + var imageInputs []imgen.ImageInput if isEdit && len(imgReq.Images) > 0 { imageInputs, err = decodeImageRefs(imgReq.Images) @@ -224,14 +229,15 @@ func (g *OpenAIGateway) forwardImagesViaWebReverse(ctx context.Context, req *sdk "num_images", numImages, ) - respBody := buildWebReverseImagesResponse(imgRes, 0, 0) + respBody := buildWebReverseImagesResponse(imgRes, inputEstimate.TextTokens, inputEstimate.ImageTokens, 0) if sseKA != nil { sseKA.Stop() writeImagesRESTSSE(req.Writer, respBody) } elapsed := time.Since(start) - usage := newTokenUsage(imagesWebReverseModel, "", 0, 0, 0, 0, elapsed.Milliseconds()) + usage := newTokenUsage(imagesWebReverseModel, "", inputEstimate.Total(), 0, 0, 0, elapsed.Milliseconds()) + setUsageInputTokenDetails(usage, inputEstimate.TextTokens, inputEstimate.ImageTokens) // Web 逆向上游不返 size 字段,直接解码生成的 PNG header 拿真实宽高(O(1))。 // 解码失败 fallback 到请求 size(auto/空时 imagePriceForSize 兜底 1K)。 billingSize := imgReq.Size @@ -240,7 +246,7 @@ func (g *OpenAIGateway) forwardImagesViaWebReverse(ctx context.Context, req *sdk billingSize = fmt.Sprintf("%dx%d", cfg.Width, cfg.Height) } } - // 图片尺寸作为通用 UsageAttribute 入库,后台费用明细可用它解释分档。 + // 图片尺寸写入 Usage 标准字段和 metadata,后台费用明细可用它解释分档。 fillUsageCostPerImageBySize(usage, numImages, billingSize) outcome := sdk.ForwardOutcome{ @@ -264,7 +270,7 @@ func (g *OpenAIGateway) forwardImagesViaWebReverse(ctx context.Context, req *sdk // - b64_json:PNG 二进制的 base64 // - revised_prompt:网页端没有 revised_prompt 字段暴露给下游,这里留空 // - model:"gpt-image-2" -func buildWebReverseImagesResponse(res *imgen.Result, promptTokens, outputTokens int) []byte { +func buildWebReverseImagesResponse(res *imgen.Result, textInputTokens, imageInputTokens, outputTokens int) []byte { data := make([]map[string]any, 0, len(res.Images)) for _, img := range res.Images { data = append(data, map[string]any{ @@ -278,14 +284,15 @@ func buildWebReverseImagesResponse(res *imgen.Result, promptTokens, outputTokens // root 级 model 供 handleImagesResponse 或 Core 做费用查价 "model": imagesWebReverseModel, } - if promptTokens+outputTokens > 0 { + inputTokens := textInputTokens + imageInputTokens + if inputTokens+outputTokens > 0 { payload["usage"] = map[string]any{ - "input_tokens": promptTokens, + "input_tokens": inputTokens, "output_tokens": outputTokens, - "total_tokens": promptTokens + outputTokens, + "total_tokens": inputTokens + outputTokens, "input_tokens_details": map[string]any{ - "text_tokens": promptTokens, - "image_tokens": 0, + "text_tokens": textInputTokens, + "image_tokens": imageInputTokens, }, } } diff --git a/backend/internal/gateway/images_web_reverse_test.go b/backend/internal/gateway/images_web_reverse_test.go index 7b67589..e6c9b57 100644 --- a/backend/internal/gateway/images_web_reverse_test.go +++ b/backend/internal/gateway/images_web_reverse_test.go @@ -6,7 +6,7 @@ import ( "testing" "time" - sdk "github.com/DouDOU-start/airgate-sdk/sdkgo" + sdk "github.com/DevilGenius/airgate-sdk/sdkgo" ) func TestWebReverseImagesErrorClientStatusReturnsNilErr(t *testing.T) { diff --git a/backend/internal/gateway/imgen/generate.go b/backend/internal/gateway/imgen/generate.go index 68da9bc..68b756d 100644 --- a/backend/internal/gateway/imgen/generate.go +++ b/backend/internal/gateway/imgen/generate.go @@ -6,7 +6,7 @@ import ( "strings" "time" - sdk "github.com/DouDOU-start/airgate-sdk/sdkgo" + sdk "github.com/DevilGenius/airgate-sdk/sdkgo" ) // Image 单张图像产物。 diff --git a/backend/internal/gateway/imgen/poll.go b/backend/internal/gateway/imgen/poll.go index 96ed8d7..ccf8244 100644 --- a/backend/internal/gateway/imgen/poll.go +++ b/backend/internal/gateway/imgen/poll.go @@ -10,7 +10,7 @@ import ( "strings" "time" - sdk "github.com/DouDOU-start/airgate-sdk/sdkgo" + sdk "github.com/DevilGenius/airgate-sdk/sdkgo" ) // ---------- stream_status ---------- diff --git a/backend/internal/gateway/imgen/sentinel.go b/backend/internal/gateway/imgen/sentinel.go index a06a230..a987cac 100644 --- a/backend/internal/gateway/imgen/sentinel.go +++ b/backend/internal/gateway/imgen/sentinel.go @@ -8,7 +8,7 @@ import ( "log/slog" "time" - sdk "github.com/DouDOU-start/airgate-sdk/sdkgo" + sdk "github.com/DevilGenius/airgate-sdk/sdkgo" ) // ChatRequirementsResult /backend-api/sentinel/chat-requirements 的产物。 diff --git a/backend/internal/gateway/imgen/upload.go b/backend/internal/gateway/imgen/upload.go index e57b897..bb040c6 100644 --- a/backend/internal/gateway/imgen/upload.go +++ b/backend/internal/gateway/imgen/upload.go @@ -14,7 +14,7 @@ import ( "net/http" "strings" - sdk "github.com/DouDOU-start/airgate-sdk/sdkgo" + sdk "github.com/DevilGenius/airgate-sdk/sdkgo" ) // ImageInput 调用方传入的待上传图片(内存中的原始二进制)。 diff --git a/backend/internal/gateway/metadata.go b/backend/internal/gateway/metadata.go index f1b7b15..785373b 100644 --- a/backend/internal/gateway/metadata.go +++ b/backend/internal/gateway/metadata.go @@ -1,6 +1,6 @@ package gateway -import sdk "github.com/DouDOU-start/airgate-sdk/sdkgo" +import sdk "github.com/DevilGenius/airgate-sdk/sdkgo" //go:generate go run ../../cmd/genmanifest @@ -18,7 +18,7 @@ const ( // // 默认值是开发态版本,正式 release 构建时由 GitHub Actions 通过 ldflags 注入: // -// go build -ldflags "-X 'github.com/DouDOU-start/airgate-openai/backend/internal/gateway.PluginVersion=0.1.42'" +// go build -ldflags "-X 'github.com/DevilGenius/airgate-openai/backend/internal/gateway.PluginVersion=0.1.42'" // // 这样 git tag 即唯一发版来源,无需手动维护 plugin.yaml / metadata.go 里的版本字段。 var PluginVersion = "dev" diff --git a/backend/internal/gateway/oauth.go b/backend/internal/gateway/oauth.go index 379dc0a..b445f08 100644 --- a/backend/internal/gateway/oauth.go +++ b/backend/internal/gateway/oauth.go @@ -14,7 +14,7 @@ import ( "sync" "time" - sdk "github.com/DouDOU-start/airgate-sdk/sdkgo" + sdk "github.com/DevilGenius/airgate-sdk/sdkgo" ) const chatGPTBrowserUserAgent = `Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/131.0.0.0 Safari/537.36 Edg/131.0.0.0` diff --git a/backend/internal/gateway/oauth_handler.go b/backend/internal/gateway/oauth_handler.go index 6845e1b..206d9f0 100644 --- a/backend/internal/gateway/oauth_handler.go +++ b/backend/internal/gateway/oauth_handler.go @@ -8,8 +8,8 @@ import ( "strconv" "strings" - "github.com/DouDOU-start/airgate-sdk/devkit/devserver" - sdk "github.com/DouDOU-start/airgate-sdk/sdkgo" + "github.com/DevilGenius/airgate-sdk/devkit/devserver" + sdk "github.com/DevilGenius/airgate-sdk/sdkgo" ) // OAuthDevHandler devserver 的 OAuth HTTP handler diff --git a/backend/internal/gateway/outcome.go b/backend/internal/gateway/outcome.go index 4436ca5..1970e32 100644 --- a/backend/internal/gateway/outcome.go +++ b/backend/internal/gateway/outcome.go @@ -3,38 +3,36 @@ package gateway import ( "fmt" "net/http" + "strconv" + "strings" "time" - sdk "github.com/DouDOU-start/airgate-sdk/sdkgo" + sdk "github.com/DevilGenius/airgate-sdk/sdkgo" - "github.com/DouDOU-start/airgate-openai/backend/internal/model" + "github.com/DevilGenius/airgate-openai/backend/internal/model" ) // 构造 ForwardOutcome 的小 helper,避免各路径散落一堆 struct literal。 // -// Success 必带 Usage,fillCost 会基于 Model / ServiceTier 补成本字段。 +// Success 必带 Usage,fillCost 会基于 Model / Metadata.service_tier 补成本字段。 // ClientError / AccountRateLimited / AccountDead / UpstreamTransient 都从 // Upstream + 上游错误消息推出 Kind。 const ( usageCurrencyUSD = "USD" - usageAttrModel = "model" usageAttrServiceTier = "service_tier" - usageAttrImageSize = "image_size" + usageAttrImageSize = "openai.image.size" + usageAttrResponseID = "openai.response_id" usageMetricInputTokens = "input_tokens" + usageMetricTextInputTokens = "openai.image.input_text_tokens" + usageMetricImageInputTokens = "openai.image.input_image_tokens" usageMetricCachedInputTokens = "cached_input_tokens" usageMetricOutputTokens = "output_tokens" usageMetricReasoningOutputTokens = "reasoning_output_tokens" usageMetricTotalTokens = "total_tokens" - usageMetricImages = "images" - - usageCostInput = "input_tokens" - usageCostCachedInput = "cached_input_tokens" - usageCostOutput = "output_tokens" - usageCostImage = "images" - usageCostImageTool = "image_tool" + usageMetricImages = "openai.image.count" ) // successOutcome 构造 Success 判决,Usage 由调用方填。Duration 由调用方填。 @@ -105,32 +103,16 @@ func newTokenUsage(modelID, serviceTier string, inputTokens, outputTokens, cache Currency: usageCurrencyUSD, FirstTokenMs: firstTokenMs, } - setUsageModelAttribute(usage, modelID) setUsageServiceTier(usage, serviceTier) setUsageTokens(usage, inputTokens, outputTokens, cachedInputTokens, reasoningOutputTokens) return usage } -func setUsageModelAttribute(usage *sdk.Usage, modelID string) { - if usage == nil || modelID == "" { - return - } - setUsageAttribute(usage, sdk.UsageAttribute{ - Key: usageAttrModel, - Label: "模型", - Kind: "model", - Value: modelID, - }) -} - func setUsageReasoningEffort(usage *sdk.Usage, effort string) { if usage == nil || effort == "" { return } - if usage.Metadata == nil { - usage.Metadata = map[string]string{} - } - usage.Metadata["reasoning_effort"] = effort + usage.ReasoningEffort = effort } func setUsageServiceTier(usage *sdk.Usage, tier string) { @@ -141,84 +123,47 @@ func setUsageServiceTier(usage *sdk.Usage, tier string) { if tier == "" { return } - setUsageAttribute(usage, sdk.UsageAttribute{ - Key: usageAttrServiceTier, - Label: "服务档位", - Kind: "tier", - Value: tier, - }) - if usage.Metadata == nil { - usage.Metadata = map[string]string{} - } - usage.Metadata[usageAttrServiceTier] = tier + setUsageMetadata(usage, usageAttrServiceTier, tier) } func usageServiceTier(usage *sdk.Usage) string { if usage == nil { return "" } - for _, attr := range usage.Attributes { - if attr.Key == usageAttrServiceTier { - return normalizeOpenAIServiceTier(attr.Value) - } - } - if usage.Metadata != nil { - return normalizeOpenAIServiceTier(usage.Metadata[usageAttrServiceTier]) - } - return "" + return normalizeOpenAIServiceTier(usageMetadataText(usage, usageAttrServiceTier)) } func setUsageImageSize(usage *sdk.Usage, size string) { if usage == nil || size == "" { return } - setUsageAttribute(usage, sdk.UsageAttribute{ - Key: usageAttrImageSize, - Label: "图片尺寸", - Kind: "resolution", - Value: size, - }) + setUsageMetadata(usage, usageAttrImageSize, size) +} + +func setUsageResponseID(usage *sdk.Usage, responseID string) { + responseID = strings.TrimSpace(responseID) + if usage == nil || !strings.HasPrefix(responseID, "resp_") { + return + } + setUsageMetadata(usage, usageAttrResponseID, responseID) } func setUsageTokens(usage *sdk.Usage, inputTokens, outputTokens, cachedInputTokens, reasoningOutputTokens int) { if usage == nil { return } - setUsageMetric(usage, sdk.UsageMetric{ - Key: usageMetricInputTokens, - Label: "输入 Token", - Kind: "token", - Unit: "token", - Value: float64(inputTokens), - }) - setUsageMetric(usage, sdk.UsageMetric{ - Key: usageMetricCachedInputTokens, - Label: "缓存输入 Token", - Kind: "token", - Unit: "token", - Value: float64(cachedInputTokens), - }) - setUsageMetric(usage, sdk.UsageMetric{ - Key: usageMetricOutputTokens, - Label: "输出 Token", - Kind: "token", - Unit: "token", - Value: float64(outputTokens), - }) - setUsageMetric(usage, sdk.UsageMetric{ - Key: usageMetricReasoningOutputTokens, - Label: "推理 Token", - Kind: "token", - Unit: "token", - Value: float64(reasoningOutputTokens), - }) - setUsageMetric(usage, sdk.UsageMetric{ - Key: usageMetricTotalTokens, - Label: "总 Token", - Kind: "token", - Unit: "token", - Value: float64(inputTokens + cachedInputTokens + outputTokens), - }) + usage.InputTokens = inputTokens + usage.OutputTokens = outputTokens + usage.CachedInputTokens = cachedInputTokens + usage.ReasoningOutputTokens = reasoningOutputTokens +} + +func setUsageInputTokenDetails(usage *sdk.Usage, textInputTokens, imageInputTokens int) { + if usage == nil || textInputTokens+imageInputTokens <= 0 { + return + } + setUsageMetadataInt(usage, usageMetricTextInputTokens, textInputTokens) + setUsageMetadataInt(usage, usageMetricImageInputTokens, imageInputTokens) } func usageMetricInt(usage *sdk.Usage, key string) int { @@ -229,71 +174,86 @@ func usageMetricValue(usage *sdk.Usage, key string) float64 { if usage == nil { return 0 } - for _, metric := range usage.Metrics { - if metric.Key == key { - return metric.Value - } + switch key { + case usageMetricInputTokens: + return float64(usage.InputTokens) + case usageMetricTextInputTokens: + return usageMetadataFloat(usage, usageMetricTextInputTokens) + case usageMetricImageInputTokens: + return usageMetadataFloat(usage, usageMetricImageInputTokens) + case usageMetricCachedInputTokens: + return float64(usage.CachedInputTokens) + case usageMetricOutputTokens: + return float64(usage.OutputTokens) + case usageMetricReasoningOutputTokens: + return float64(usage.ReasoningOutputTokens) + case usageMetricTotalTokens: + return float64(usage.InputTokens + usage.CachedInputTokens + usage.OutputTokens) + case usageMetricImages: + return usageMetadataFloat(usage, usageMetricImages) } return 0 } -func setUsageAttribute(usage *sdk.Usage, attr sdk.UsageAttribute) { - for i := range usage.Attributes { - if usage.Attributes[i].Key == attr.Key { - usage.Attributes[i] = attr - return - } +func setUsageMetadata(usage *sdk.Usage, key, value string) { + if usage == nil { + return + } + value = strings.TrimSpace(value) + if value == "" { + return + } + if usage.Metadata == nil { + usage.Metadata = map[string]string{} } - usage.Attributes = append(usage.Attributes, attr) + usage.Metadata[key] = value } -func setUsageMetric(usage *sdk.Usage, metric sdk.UsageMetric) { - for i := range usage.Metrics { - if usage.Metrics[i].Key == metric.Key { - usage.Metrics[i] = metric - return - } +func setUsageMetadataInt(usage *sdk.Usage, key string, value int) { + if value <= 0 { + return } - usage.Metrics = append(usage.Metrics, metric) + setUsageMetadata(usage, key, strconv.Itoa(value)) } -func setUsageCostDetail(usage *sdk.Usage, detail sdk.UsageCostDetail) { - if detail.AccountCost <= 0 { - removeUsageCostDetail(usage, detail.Key) +func setUsageMetadataFloat(usage *sdk.Usage, key string, value float64) { + if value <= 0 { return } - for i := range usage.CostDetails { - if usage.CostDetails[i].Key == detail.Key { - usage.CostDetails[i] = detail - recomputeUsageAccountCost(usage) - return - } + setUsageMetadata(usage, key, strconv.FormatFloat(value, 'f', -1, 64)) +} + +func addUsageMetadataInt(usage *sdk.Usage, key string, delta int) { + if usage == nil || delta <= 0 { + return } - usage.CostDetails = append(usage.CostDetails, detail) - recomputeUsageAccountCost(usage) + setUsageMetadataInt(usage, key, int(usageMetadataFloat(usage, key))+delta) } -func removeUsageCostDetail(usage *sdk.Usage, key string) { +func usageMetadataText(usage *sdk.Usage, key string) string { if usage == nil { - return + return "" } - for i := range usage.CostDetails { - if usage.CostDetails[i].Key == key { - usage.CostDetails = append(usage.CostDetails[:i], usage.CostDetails[i+1:]...) - recomputeUsageAccountCost(usage) - return - } + return strings.TrimSpace(usage.Metadata[key]) +} + +func usageMetadataFloat(usage *sdk.Usage, key string) float64 { + raw := usageMetadataText(usage, key) + if raw == "" { + return 0 + } + value, err := strconv.ParseFloat(raw, 64) + if err != nil { + return 0 } + return value } func recomputeUsageAccountCost(usage *sdk.Usage) { if usage == nil { return } - var total float64 - for _, detail := range usage.CostDetails { - total += detail.AccountCost - } + total := usage.InputCost + usage.OutputCost + usage.CachedInputCost + usage.CacheCreationCost usage.AccountCost = total if usage.Currency == "" { usage.Currency = usageCurrencyUSD @@ -368,20 +328,6 @@ func tokenCost(tokens int, pricePerMillion float64) float64 { return float64(tokens) * pricePerMillion / 1_000_000 } -func priceMetadata(price float64, tier string, longContext bool) map[string]string { - metadata := map[string]string{ - "unit_price": fmt.Sprintf("%.10g", price), - "unit": "USD/1M tokens", - } - if tier != "" { - metadata["service_tier"] = tier - } - if longContext { - metadata["long_context"] = "true" - } - return metadata -} - // fillUsageCost 用插件自己的模型规格填充 Usage 的平台标准成本。 // // SDK 只承载通用 Usage 结构;OpenAI 的标准价格、服务档位和长上下文阶梯都留在 @@ -395,7 +341,7 @@ func fillUsageCost(usage *sdk.Usage) { inputTokens := usageMetricInt(usage, usageMetricInputTokens) outputTokens := usageMetricInt(usage, usageMetricOutputTokens) cachedInputTokens := usageMetricInt(usage, usageMetricCachedInputTokens) - prices, longContext := applyLongContextPricing( + prices, _ := applyLongContextPricing( spec, pricesForServiceTier(spec, serviceTier), inputTokens, @@ -405,104 +351,55 @@ func fillUsageCost(usage *sdk.Usage) { inputCost := tokenCost(inputTokens, prices.input) cachedCost := tokenCost(cachedInputTokens, prices.cached) outputCost := tokenCost(outputTokens, prices.output) + usage.InputPrice = prices.input + usage.CachedInputPrice = prices.cached + usage.OutputPrice = prices.output + usage.InputCost = inputCost + usage.CachedInputCost = cachedCost + usage.OutputCost = outputCost + recomputeUsageAccountCost(usage) + +} - setUsageMetric(usage, sdk.UsageMetric{ - Key: usageMetricInputTokens, - Label: "输入 Token", - Kind: "token", - Unit: "token", - Value: float64(inputTokens), - AccountCost: inputCost, - Currency: usageCurrencyUSD, - Metadata: priceMetadata(prices.input, serviceTier, longContext), - }) - setUsageMetric(usage, sdk.UsageMetric{ - Key: usageMetricCachedInputTokens, - Label: "缓存输入 Token", - Kind: "token", - Unit: "token", - Value: float64(cachedInputTokens), - AccountCost: cachedCost, - Currency: usageCurrencyUSD, - Metadata: priceMetadata(prices.cached, serviceTier, longContext), - }) - setUsageMetric(usage, sdk.UsageMetric{ - Key: usageMetricOutputTokens, - Label: "输出 Token", - Kind: "token", - Unit: "token", - Value: float64(outputTokens), - AccountCost: outputCost, - Currency: usageCurrencyUSD, - Metadata: priceMetadata(prices.output, serviceTier, longContext), - }) - setUsageCostDetail(usage, sdk.UsageCostDetail{ - Key: usageCostInput, - Label: "输入 Token", - AccountCost: inputCost, - Currency: usageCurrencyUSD, - Metadata: priceMetadata(prices.input, serviceTier, longContext), - }) - setUsageCostDetail(usage, sdk.UsageCostDetail{ - Key: usageCostCachedInput, - Label: "缓存输入 Token", - AccountCost: cachedCost, - Currency: usageCurrencyUSD, - Metadata: priceMetadata(prices.cached, serviceTier, longContext), - }) - setUsageCostDetail(usage, sdk.UsageCostDetail{ - Key: usageCostOutput, - Label: "输出 Token", - AccountCost: outputCost, - Currency: usageCurrencyUSD, - Metadata: priceMetadata(prices.output, serviceTier, longContext), - }) -} - -// fillUsageCostPerImageBySize 按 1K/2K/4K size 分档填充 Usage(USD/张)。 -// 用于 OAuth → image_generation tool 路径,价格由 imagePriceForSize 硬编码(详见其注释)。 -// 跟 spec.ImagePrice 解耦:plugin.yaml 不需要登记 ImagePrice,分档定价完全由网关侧决定。 +// fillUsageCostPerImageBySize 记录 1K/2K/4K 图片 metadata,并按模型 token 标准价 +// 填充 Usage。分组固定单张价由 Core 用 metadata 替代最终计费。 func fillUsageCostPerImageBySize(usage *sdk.Usage, numImages int, size string) { if usage == nil || numImages <= 0 { return } price := imagePriceForSize(size) + fillUsageCost(usage) setUsageImageSize(usage, size) - addImageCost(usage, usageCostImage, "图片生成", numImages, price, size) + addImageCost(usage, numImages, price) } func addUsageCostForModel( usage *sdk.Usage, modelID, serviceTier string, inputTokens, outputTokens, cachedInputTokens, reasoningOutputTokens int, - keyPrefix, labelPrefix string, ) { if usage == nil || modelID == "" { return } source := newTokenUsage(modelID, serviceTier, inputTokens, outputTokens, cachedInputTokens, reasoningOutputTokens, 0) fillUsageCost(source) - for _, detail := range source.CostDetails { - copied := detail - if keyPrefix != "" { - copied.Key = keyPrefix + "_" + copied.Key - } - if labelPrefix != "" { - copied.Label = labelPrefix + copied.Label - } - if copied.Metadata != nil { - metadata := make(map[string]string, len(copied.Metadata)+1) - for k, v := range copied.Metadata { - metadata[k] = v - } - metadata["source_model"] = modelID - copied.Metadata = metadata - } - setUsageCostDetail(usage, copied) + if source.InputPrice > 0 && usage.InputPrice == 0 { + usage.InputPrice = source.InputPrice } + if source.CachedInputPrice > 0 && usage.CachedInputPrice == 0 { + usage.CachedInputPrice = source.CachedInputPrice + } + if source.OutputPrice > 0 && usage.OutputPrice == 0 { + usage.OutputPrice = source.OutputPrice + } + usage.InputCost += source.InputCost + usage.CachedInputCost += source.CachedInputCost + usage.OutputCost += source.OutputCost + usage.CacheCreationCost += source.CacheCreationCost + recomputeUsageAccountCost(usage) } -// fillUsageCostWithImageTool 先按主 model 定价算 token 成本,再按尺寸分档叠加图像费用。 +// fillUsageCostWithImageTool 先按主 model 定价算 token 成本,再写入图片分档 metadata。 func fillUsageCostWithImageTool(usage *sdk.Usage, numImages int, size string) { fillUsageCost(usage) if usage == nil || numImages <= 0 { @@ -510,36 +407,15 @@ func fillUsageCostWithImageTool(usage *sdk.Usage, numImages int, size string) { } price := imagePriceForSize(size) setUsageImageSize(usage, size) - addImageCost(usage, usageCostImageTool, "图片工具", numImages, price, size) + addImageCost(usage, numImages, price) } -func addImageCost(usage *sdk.Usage, key, label string, numImages int, pricePerImage float64, size string) { +func addImageCost(usage *sdk.Usage, numImages int, pricePerImage float64) { if usage == nil || numImages <= 0 || pricePerImage <= 0 { return } - cost := float64(numImages) * pricePerImage - metadata := map[string]string{ - "unit_price": fmt.Sprintf("%.10g", pricePerImage), - "unit": "USD/image", - } - if size != "" { - metadata["size"] = size - } - setUsageMetric(usage, sdk.UsageMetric{ - Key: usageMetricImages, - Label: "图片数量", - Kind: "image", - Unit: "image", - Value: float64(numImages), - AccountCost: cost, - Currency: usageCurrencyUSD, - Metadata: metadata, - }) - setUsageCostDetail(usage, sdk.UsageCostDetail{ - Key: key, - Label: label, - AccountCost: cost, - Currency: usageCurrencyUSD, - Metadata: metadata, - }) + addUsageMetadataInt(usage, usageMetricImages, numImages) + setUsageMetadataFloat(usage, "openai.image.unit_price", pricePerImage) + setUsageMetadata(usage, "openai.image.unit", "USD/image") + recomputeUsageAccountCost(usage) } diff --git a/backend/internal/gateway/outcome_test.go b/backend/internal/gateway/outcome_test.go index ab9df09..425518f 100644 --- a/backend/internal/gateway/outcome_test.go +++ b/backend/internal/gateway/outcome_test.go @@ -4,7 +4,7 @@ import ( "errors" "testing" - sdk "github.com/DouDOU-start/airgate-sdk/sdkgo" + sdk "github.com/DevilGenius/airgate-sdk/sdkgo" ) func TestForwardErrForOutcomeSuppressesClientError(t *testing.T) { diff --git a/backend/internal/gateway/persistence.go b/backend/internal/gateway/persistence.go index ec10b74..7edadaa 100644 --- a/backend/internal/gateway/persistence.go +++ b/backend/internal/gateway/persistence.go @@ -102,6 +102,7 @@ CREATE TABLE IF NOT EXISTS %s ( session_id text NOT NULL DEFAULT '', conversation_id text NOT NULL DEFAULT '', prompt_cache_key text NOT NULL DEFAULT '', + account_id bigint NOT NULL DEFAULT 0, last_response_id text NOT NULL DEFAULT '', last_turn_state text NOT NULL DEFAULT '', last_seen_at timestamptz NOT NULL DEFAULT NOW(), @@ -119,6 +120,9 @@ CREATE INDEX IF NOT EXISTS idx_%s_updated_at if _, err := s.db.ExecContext(ctx, sessionQuery); err != nil { return fmt.Errorf("ensure session state schema: %w", err) } + if _, err := s.db.ExecContext(ctx, fmt.Sprintf(`ALTER TABLE %s ADD COLUMN IF NOT EXISTS account_id bigint NOT NULL DEFAULT 0`, sessionStatePersistTable)); err != nil { + return fmt.Errorf("ensure session account_id column: %w", err) + } return nil } @@ -263,15 +267,16 @@ func (s *codexUsagePersistenceStore) upsertSessionState(ctx context.Context, sta query := fmt.Sprintf(` INSERT INTO %s ( plugin_id, session_key, session_id, conversation_id, prompt_cache_key, - last_response_id, last_turn_state, last_seen_at, last_updated_at, + account_id, last_response_id, last_turn_state, last_seen_at, last_updated_at, last_response_at, last_turn_state_at ) -VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10, $11) +VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10, $11, $12) ON CONFLICT (plugin_id, session_key) DO UPDATE SET session_id = EXCLUDED.session_id, conversation_id = EXCLUDED.conversation_id, prompt_cache_key = EXCLUDED.prompt_cache_key, + account_id = EXCLUDED.account_id, last_response_id = EXCLUDED.last_response_id, last_turn_state = EXCLUDED.last_turn_state, last_seen_at = EXCLUDED.last_seen_at, @@ -288,6 +293,7 @@ DO UPDATE SET strings.TrimSpace(state.SessionID), strings.TrimSpace(state.ConversationID), strings.TrimSpace(state.PromptCacheKey), + state.AccountID, strings.TrimSpace(state.LastResponseID), strings.TrimSpace(state.LastTurnState), state.LastSeenAt.UTC(), @@ -367,7 +373,7 @@ func (s *codexUsagePersistenceStore) WarmCache(ctx context.Context) error { } func (s *codexUsagePersistenceStore) warmSessionStateCache(ctx context.Context) (int, error) { - query := fmt.Sprintf(`SELECT session_key, session_id, conversation_id, prompt_cache_key, last_response_id, last_turn_state, last_seen_at, last_updated_at, last_response_at, last_turn_state_at FROM %s WHERE plugin_id = $1`, sessionStatePersistTable) + query := fmt.Sprintf(`SELECT session_key, session_id, conversation_id, prompt_cache_key, account_id, last_response_id, last_turn_state, last_seen_at, last_updated_at, last_response_at, last_turn_state_at FROM %s WHERE plugin_id = $1`, sessionStatePersistTable) rows, err := s.db.QueryContext(ctx, query, s.pluginID) if err != nil { return 0, fmt.Errorf("warm session state query: %w", err) @@ -383,6 +389,7 @@ func (s *codexUsagePersistenceStore) warmSessionStateCache(ctx context.Context) &state.SessionID, &state.ConversationID, &state.PromptCacheKey, + &state.AccountID, &state.LastResponseID, &state.LastTurnState, &state.LastSeenAt, diff --git a/backend/internal/gateway/request.go b/backend/internal/gateway/request.go index 4f7d0a9..d706f9d 100644 --- a/backend/internal/gateway/request.go +++ b/backend/internal/gateway/request.go @@ -11,9 +11,9 @@ import ( "github.com/tidwall/gjson" "github.com/tidwall/sjson" - "github.com/DouDOU-start/airgate-openai/backend/internal/model" - "github.com/DouDOU-start/airgate-openai/backend/resources" - sdk "github.com/DouDOU-start/airgate-sdk/sdkgo" + "github.com/DevilGenius/airgate-openai/backend/internal/model" + "github.com/DevilGenius/airgate-openai/backend/resources" + sdk "github.com/DevilGenius/airgate-sdk/sdkgo" ) // modelMetadataOverrides 仅用于 /v1/models 响应补齐。 @@ -196,7 +196,7 @@ func buildAPIKeyURL(account *sdk.Account, reqPath string) string { // 拿到的 body 格式一致。当前处理步骤: // 1. model 同步(body 中的 model 与 core 传入的 model 对齐) // 2. data:image 输入保持原样(对齐 Codex,不在网关内重采样用户图片) -// 3. 剔除客户端 previous_response_id(跨账号接续不可靠,会话接续由网关内部管理) +// 3. 保留 previous_response_id(Core 已按 response_id 做账号粘连) // 4. 上下文守卫(/v1/chat/completions 超长 messages 裁剪) // 5. input 规范化(/v1/responses 的 string input → list,messages → input 转换) // 6. Responses API 强制禁用上游存储(store=false) @@ -225,12 +225,6 @@ func preprocessRequestBody(body []byte, model, reqPath string) []byte { result = preserveOpenAIConversationImages(result) - // 剔除客户端传入的 previous_response_id。 - // AirGate 在多个上游账号之间做负载均衡,客户端的 previous_response_id - // 可能指向另一个账号的 response,上游会返回 "not found"。 - // 会话接续由网关内部的 session 机制(OAuth sessionState / Anthropic digestChain)管理。 - result, _ = dropPreviousResponseIDFromJSON(result) - result = applyContextGuard(result, reqPath) result = normalizeResponsesInput(result, reqPath) result = forceResponsesStoreFalse(result, reqPath) @@ -470,20 +464,11 @@ func applyContinuationState(reqData map[string]any, session openAISessionResolut return reqData } - // 不再从 session 回填 previous_response_id。 - // 跨账号接续时,上一轮 response 可能在另一个账号上,注入后上游会返回 "not found"; - // 且 function_call_output 自带 call_id,上游可以靠 call_id 匹配,不依赖 previous_response_id。 - // 客户端的 previous_response_id 已在 preprocessRequestBody 统一剔除。 - return reqData -} - -func dropPreviousResponseIDFromJSON(body []byte) ([]byte, bool) { - if len(body) == 0 || !gjson.GetBytes(body, "previous_response_id").Exists() { - return body, false + if _, ok := reqData["previous_response_id"].(string); ok { + return reqData } - next, err := sjson.DeleteBytes(body, "previous_response_id") - if err != nil { - return body, false + if requestNeedsPreviousResponseID(reqData) && strings.TrimSpace(session.PreviousRespID) != "" { + reqData["previous_response_id"] = strings.TrimSpace(session.PreviousRespID) } - return next, true + return reqData } diff --git a/backend/internal/gateway/request_test.go b/backend/internal/gateway/request_test.go index 8c11d35..b50f2d3 100644 --- a/backend/internal/gateway/request_test.go +++ b/backend/internal/gateway/request_test.go @@ -14,7 +14,7 @@ import ( "github.com/tidwall/gjson" - sdk "github.com/DouDOU-start/airgate-sdk/sdkgo" + sdk "github.com/DevilGenius/airgate-sdk/sdkgo" ) // TestIsAnthropicRequest 只认两个权威信号:X-Forwarded-Path + Anthropic-Version 头。 @@ -115,7 +115,7 @@ func TestIsAnthropicRequest(t *testing.T) { } } -func TestApplyContinuationStateDoesNotBackfillPreviousResponseID(t *testing.T) { +func TestApplyContinuationStateBackfillsPreviousResponseIDForToolOutput(t *testing.T) { reqBody := map[string]any{ "input": []any{ map[string]any{ @@ -128,18 +128,31 @@ func TestApplyContinuationStateDoesNotBackfillPreviousResponseID(t *testing.T) { session := openAISessionResolution{PreviousRespID: "resp_prev"} reqBody = applyContinuationState(reqBody, session) - if got, _ := reqBody["previous_response_id"].(string); got != "" { - t.Fatalf("expected previous_response_id to NOT be backfilled, got %q", got) + if got, _ := reqBody["previous_response_id"].(string); got != "resp_prev" { + t.Fatalf("previous_response_id = %q, want resp_prev", got) } } -func TestDropPreviousResponseIDFromJSON(t *testing.T) { - next, changed := dropPreviousResponseIDFromJSON([]byte(`{"model":"gpt-5.4","previous_response_id":"resp_old","input":[]}`)) - if !changed { - t.Fatalf("expected previous_response_id to be removed") +func TestApplyContinuationStateDoesNotBackfillWithToolCallContext(t *testing.T) { + reqBody := map[string]any{ + "input": []any{ + map[string]any{ + "type": "function_call", + "call_id": "call_1", + "name": "lookup", + }, + map[string]any{ + "type": "function_call_output", + "call_id": "call_1", + "output": "ok", + }, + }, } - if string(next) == `{"model":"gpt-5.4","previous_response_id":"resp_old","input":[]}` { - t.Fatalf("expected updated payload") + + session := openAISessionResolution{PreviousRespID: "resp_prev"} + reqBody = applyContinuationState(reqBody, session) + if got, _ := reqBody["previous_response_id"].(string); got != "" { + t.Fatalf("expected previous_response_id to stay empty when tool call context is present, got %q", got) } } @@ -326,6 +339,14 @@ func TestPreprocessRequestBody_ForcesResponsesStoreFalse(t *testing.T) { } } +func TestPreprocessRequestBodyPreservesPreviousResponseID(t *testing.T) { + body := []byte(`{"model":"gpt-5.4","previous_response_id":"resp_old","input":[]}`) + got := preprocessRequestBody(body, "gpt-5.4", "/v1/responses") + if previous := gjson.GetBytes(got, "previous_response_id").String(); previous != "resp_old" { + t.Fatalf("previous_response_id = %q, want resp_old; body=%s", previous, got) + } +} + func TestFirstNonEmptyTier_RequestFastFallsBackToUpstreamPriority(t *testing.T) { if got := firstNonEmptyTier("fast", "priority"); got != "priority" { t.Fatalf("firstNonEmptyTier(fast, priority) = %q, want %q", got, "priority") diff --git a/backend/internal/gateway/responses_failure.go b/backend/internal/gateway/responses_failure.go index 76607b5..f5626d9 100644 --- a/backend/internal/gateway/responses_failure.go +++ b/backend/internal/gateway/responses_failure.go @@ -8,7 +8,7 @@ import ( "github.com/tidwall/gjson" - sdk "github.com/DouDOU-start/airgate-sdk/sdkgo" + sdk "github.com/DevilGenius/airgate-sdk/sdkgo" ) type responsesFailureKind string diff --git a/backend/internal/gateway/responses_failure_test.go b/backend/internal/gateway/responses_failure_test.go index 250e8be..387d339 100644 --- a/backend/internal/gateway/responses_failure_test.go +++ b/backend/internal/gateway/responses_failure_test.go @@ -8,7 +8,7 @@ import ( "testing" "time" - sdk "github.com/DouDOU-start/airgate-sdk/sdkgo" + sdk "github.com/DevilGenius/airgate-sdk/sdkgo" ) func TestClassifyResponsesFailureContextWindow(t *testing.T) { @@ -110,6 +110,13 @@ func TestClassifyHTTPFailureKeepsDisabled403AsAccountDead(t *testing.T) { } } +func TestClassifyHTTPFailureTreatsPlain403AsAccountUnavailable(t *testing.T) { + got := classifyHTTPFailure(403, "访问被拒绝,账号可能已被禁用或无权限 (HTTP 403)") + if got != sdk.OutcomeAccountUnavailable { + t.Fatalf("expected AccountUnavailable, got %v", got) + } +} + func TestClassifyHTTPFailureTreatsDisabled400AsAccountDead(t *testing.T) { got := classifyHTTPFailure(400, "Organization disabled due to policy violation") if got != sdk.OutcomeAccountDead { diff --git a/backend/internal/gateway/session_state.go b/backend/internal/gateway/session_state.go index 75b201c..c98b822 100644 --- a/backend/internal/gateway/session_state.go +++ b/backend/internal/gateway/session_state.go @@ -17,6 +17,7 @@ type openAISessionState struct { SessionID string `json:"session_id,omitempty"` ConversationID string `json:"conversation_id,omitempty"` PromptCacheKey string `json:"prompt_cache_key,omitempty"` + AccountID int64 `json:"account_id,omitempty"` LastResponseID string `json:"last_response_id,omitempty"` LastTurnState string `json:"last_turn_state,omitempty"` LastSeenAt time.Time `json:"last_seen_at"` @@ -130,13 +131,14 @@ type openAISessionResolution struct { PromptCacheKey string PreviousRespID string LastTurnState string + AccountID int64 FromStoredState bool DigestChain string MatchedDigest string SessionSource string } -func resolveOpenAISession(headers http.Header, body []byte) openAISessionResolution { +func resolveOpenAISession(headers http.Header, body []byte, accountID int64) openAISessionResolution { promptCacheKey := resolvePromptCacheKeyFromBody(body) sessionID := "" conversationID := "" @@ -159,6 +161,7 @@ func resolveOpenAISession(headers http.Header, body []byte) openAISessionResolut SessionID: sessionID, ConversationID: conversationID, PromptCacheKey: promptCacheKey, + AccountID: accountID, } switch { case sessionID != "": @@ -184,7 +187,7 @@ func resolveOpenAISession(headers http.Header, body []byte) openAISessionResolut if resolution.PromptCacheKey == "" { resolution.PromptCacheKey = state.PromptCacheKey } - if previousResponseID == "" { + if previousResponseID == "" && (state.AccountID == 0 || state.AccountID == accountID) { previousResponseID = state.LastResponseID } resolution.LastTurnState = state.LastTurnState @@ -412,7 +415,7 @@ func touchSessionState(sessionKey string, update func(*openAISessionState)) { upsertSessionState(current) } -func updateSessionStateFromRequest(resolution openAISessionResolution) { +func updateSessionStateFromRequest(resolution openAISessionResolution, accountID int64) { if resolution.SessionKey == "" { return } @@ -429,16 +432,22 @@ func updateSessionStateFromRequest(resolution openAISessionResolution) { if pck := normalizeSessionValue(resolution.PromptCacheKey); pck != "" { state.PromptCacheKey = pck } + if accountID > 0 { + state.AccountID = accountID + } }) } -func updateSessionStateResponseID(sessionKey, responseID string) { +func updateSessionStateResponseID(sessionKey, responseID string, accountID int64) { responseID = strings.TrimSpace(responseID) if sessionKey == "" || responseID == "" { return } now := time.Now().UTC() touchSessionState(sessionKey, func(state *openAISessionState) { + if accountID > 0 { + state.AccountID = accountID + } state.LastResponseID = responseID state.LastResponseAt = now }) diff --git a/backend/internal/gateway/session_state_test.go b/backend/internal/gateway/session_state_test.go index 0fb1331..1b97781 100644 --- a/backend/internal/gateway/session_state_test.go +++ b/backend/internal/gateway/session_state_test.go @@ -7,7 +7,7 @@ import ( func TestResolveOpenAISessionUsesPromptCacheKeyAsFallback(t *testing.T) { headers := http.Header{} - resolution := resolveOpenAISession(headers, []byte(`{"prompt_cache_key":"pcache_123"}`)) + resolution := resolveOpenAISession(headers, []byte(`{"prompt_cache_key":"pcache_123"}`), 101) if resolution.SessionKey != "pcache:pcache_123" { t.Fatalf("expected session key from prompt_cache_key, got %q", resolution.SessionKey) } @@ -17,15 +17,17 @@ func TestResolveOpenAISessionUsesPromptCacheKeyAsFallback(t *testing.T) { } func TestResolveOpenAISessionReadsStoredState(t *testing.T) { + sessionStateStore.Delete("pcache:pcache_456") upsertSessionState(&openAISessionState{ SessionKey: "pcache:pcache_456", PromptCacheKey: "pcache_456", SessionID: "pcache_456", + AccountID: 202, LastResponseID: "resp_abc", LastTurnState: "turn_state_xyz", }) - resolution := resolveOpenAISession(http.Header{}, []byte(`{"prompt_cache_key":"pcache_456"}`)) + resolution := resolveOpenAISession(http.Header{}, []byte(`{"prompt_cache_key":"pcache_456"}`), 202) if resolution.PreviousRespID != "resp_abc" { t.Fatalf("expected previous response id from stored state, got %q", resolution.PreviousRespID) } @@ -34,6 +36,25 @@ func TestResolveOpenAISessionReadsStoredState(t *testing.T) { } } +func TestResolveOpenAISessionIgnoresStoredResponseFromDifferentAccount(t *testing.T) { + sessionStateStore.Delete("pcache:pcache_account_mismatch") + upsertSessionState(&openAISessionState{ + SessionKey: "pcache:pcache_account_mismatch", + PromptCacheKey: "pcache_account_mismatch", + SessionID: "pcache_account_mismatch", + AccountID: 301, + LastResponseID: "resp_wrong_account", + }) + + resolution := resolveOpenAISession(http.Header{}, []byte(`{"prompt_cache_key":"pcache_account_mismatch"}`), 302) + if resolution.PreviousRespID != "" { + t.Fatalf("expected previous response id to be ignored across accounts, got %q", resolution.PreviousRespID) + } + if !resolution.FromStoredState { + t.Fatalf("expected stored state to still be detected") + } +} + func TestDeriveAnthropicPromptCacheKey_IgnoresLaterUserEphemeralChanges(t *testing.T) { body1 := []byte(`{ "system":[{"type":"text","text":"stable system","cache_control":{"type":"ephemeral"}}], diff --git a/backend/internal/gateway/stream.go b/backend/internal/gateway/stream.go index d2b2311..70c9352 100644 --- a/backend/internal/gateway/stream.go +++ b/backend/internal/gateway/stream.go @@ -16,7 +16,7 @@ import ( "github.com/tidwall/gjson" - sdk "github.com/DouDOU-start/airgate-sdk/sdkgo" + sdk "github.com/DevilGenius/airgate-sdk/sdkgo" ) // upstreamSSEMaxLineBytes 是上游 SSE 单行最大字节数。 @@ -62,6 +62,7 @@ func handleStreamResponseWithLogger(logger *slog.Logger, resp *http.Response, w var toolImageIn, toolImageOut int // 接收 response.tool_usage.image_gen,用于图像工具计费。 var imageGenCount int var imageGenSize string + responseID := "" for scanner.Scan() { line := scanner.Text() @@ -77,6 +78,9 @@ func handleStreamResponseWithLogger(logger *slog.Logger, resp *http.Response, w completed = true diagnostics.completionEvent = "[DONE]" } else if data != "" { + if id := responseIDFromSSEData(data); id != "" { + responseID = id + } if streamErr = parseSSEFailureEvent([]byte(data)); streamErr != nil { logStreamFailure(logger, streamErr, resp, streamStarted, diagnostics) streamErrLogged = true @@ -174,6 +178,7 @@ func handleStreamResponseWithLogger(logger *slog.Logger, resp *http.Response, w numImages = estimateImageCountFromTokens(toolImageOut) } fillUsageCostWithImageTool(usage, numImages, imageGenSize) + setUsageResponseID(usage, responseID) return sdk.ForwardOutcome{ Kind: sdk.OutcomeSuccess, Upstream: sdk.UpstreamResponse{StatusCode: resp.StatusCode}, @@ -405,6 +410,7 @@ func handleNonStreamResponse(resp *http.Response, w http.ResponseWriter, start t numImages = estimateImageCountFromTokens(parsed.toolImageOutputTokens) } fillUsageCostWithImageTool(usage, numImages, parsed.imageGenCallSize) + setUsageResponseID(usage, responseIDFromBody(body)) outcome := sdk.ForwardOutcome{ Kind: sdk.OutcomeSuccess, @@ -680,6 +686,20 @@ func firstNonEmptyHeader(headers http.Header, keys ...string) string { return "" } +func responseIDFromSSEData(data string) string { + if id := strings.TrimSpace(gjson.Get(data, "response.id").String()); id != "" { + return id + } + return strings.TrimSpace(gjson.Get(data, "id").String()) +} + +func responseIDFromBody(body []byte) string { + if id := strings.TrimSpace(gjson.GetBytes(body, "id").String()); id != "" { + return id + } + return strings.TrimSpace(gjson.GetBytes(body, "response.id").String()) +} + // extractSSEData 从 SSE 行中提取 data 内容 func extractSSEData(line string) (string, bool) { if !strings.HasPrefix(line, "data:") { @@ -704,7 +724,6 @@ func parseSSEUsage(data []byte, out *sdk.Usage, toolImageIn, toolImageOut *int) return } out.Model = resp.Get("model").String() - setUsageModelAttribute(out, out.Model) if usageServiceTier(out) == "" { setUsageServiceTier(out, resp.Get("service_tier").String()) } @@ -738,7 +757,6 @@ func parseSSEUsage(data []byte, out *sdk.Usage, toolImageIn, toolImageOut *int) cachedInputTokens := int(usage.Get("prompt_tokens_details.cached_tokens").Int()) reasoningOutputTokens := int(usage.Get("completion_tokens_details.reasoning_tokens").Int()) out.Model = gjson.GetBytes(data, "model").String() - setUsageModelAttribute(out, out.Model) if cachedInputTokens > 0 && inputTokens >= cachedInputTokens { inputTokens -= cachedInputTokens } @@ -804,6 +822,8 @@ func parseSSEFailureEvent(data []byte) error { // openaiUsage 非流式响应的 usage 解析结果 type openaiUsage struct { inputTokens int + textInputTokens int + imageInputTokens int outputTokens int cachedInputTokens int reasoningOutputTokens int @@ -832,6 +852,17 @@ func parseUsage(body []byte) openaiUsage { if usage.outputTokens == 0 { usage.outputTokens = int(usageNode.Get("completion_tokens").Int()) } + rawInputTokens := usage.inputTokens + textNode := usageNode.Get("input_tokens_details.text_tokens") + imageNode := usageNode.Get("input_tokens_details.image_tokens") + if imageNode.Exists() { + usage.imageInputTokens = int(imageNode.Int()) + } + if textNode.Exists() { + usage.textInputTokens = int(textNode.Int()) + } else if usage.imageInputTokens > 0 && rawInputTokens >= usage.imageInputTokens { + usage.textInputTokens = rawInputTokens - usage.imageInputTokens + } // 仅提取 cache read(缓存命中)token,不含 cache creation // cache_creation 按正常输入价计费,已包含在 input_tokens 中无需额外处理 diff --git a/backend/internal/gateway/task_image.go b/backend/internal/gateway/task_image.go index e1b03d5..d241908 100644 --- a/backend/internal/gateway/task_image.go +++ b/backend/internal/gateway/task_image.go @@ -12,7 +12,7 @@ import ( "strconv" "strings" - sdk "github.com/DouDOU-start/airgate-sdk/sdkgo" + sdk "github.com/DevilGenius/airgate-sdk/sdkgo" ) const ( diff --git a/backend/internal/gateway/task_input_resolver_test.go b/backend/internal/gateway/task_input_resolver_test.go index b0aa2b0..7668ddd 100644 --- a/backend/internal/gateway/task_input_resolver_test.go +++ b/backend/internal/gateway/task_input_resolver_test.go @@ -7,7 +7,7 @@ import ( "strings" "testing" - sdk "github.com/DouDOU-start/airgate-sdk/sdkgo" + sdk "github.com/DevilGenius/airgate-sdk/sdkgo" ) func TestObjectKeyFromRuntimeAssetURL(t *testing.T) { diff --git a/backend/internal/gateway/task_registry.go b/backend/internal/gateway/task_registry.go index 67e74b1..097cb0a 100644 --- a/backend/internal/gateway/task_registry.go +++ b/backend/internal/gateway/task_registry.go @@ -4,7 +4,7 @@ import ( "context" "sort" - sdk "github.com/DouDOU-start/airgate-sdk/sdkgo" + sdk "github.com/DevilGenius/airgate-sdk/sdkgo" ) // TaskHandler 任务类型处理器。插件为每个任务类型实现此接口,注册到 TaskRegistry。 diff --git a/backend/internal/gateway/task_runner.go b/backend/internal/gateway/task_runner.go index 904ec8a..861800b 100644 --- a/backend/internal/gateway/task_runner.go +++ b/backend/internal/gateway/task_runner.go @@ -10,7 +10,7 @@ import ( "strconv" "strings" - sdk "github.com/DouDOU-start/airgate-sdk/sdkgo" + sdk "github.com/DevilGenius/airgate-sdk/sdkgo" ) const ( diff --git a/backend/internal/gateway/tool_continuation.go b/backend/internal/gateway/tool_continuation.go new file mode 100644 index 0000000..e67a2b7 --- /dev/null +++ b/backend/internal/gateway/tool_continuation.go @@ -0,0 +1,58 @@ +package gateway + +import "strings" + +type toolContinuationSignals struct { + hasToolOutput bool + hasToolCallContext bool +} + +func analyzeToolContinuationSignalsFromMap(reqData map[string]any) toolContinuationSignals { + var signals toolContinuationSignals + if reqData == nil { + return signals + } + input, ok := reqData["input"].([]any) + if !ok { + return signals + } + for _, item := range input { + itemMap, ok := item.(map[string]any) + if !ok { + continue + } + itemType, _ := itemMap["type"].(string) + switch { + case isOpenAIToolOutputItemType(itemType): + signals.hasToolOutput = true + case isOpenAIToolCallContextItemType(itemType): + if strings.TrimSpace(jsonString(itemMap["call_id"])) != "" { + signals.hasToolCallContext = true + } + } + } + return signals +} + +func isOpenAIToolCallContextItemType(itemType string) bool { + switch strings.TrimSpace(itemType) { + case "tool_call", "function_call", "local_shell_call", "tool_search_call", "custom_tool_call", "mcp_tool_call": + return true + default: + return false + } +} + +func isOpenAIToolOutputItemType(itemType string) bool { + switch strings.TrimSpace(itemType) { + case "function_call_output", "tool_search_output", "custom_tool_call_output", "mcp_tool_call_output": + return true + default: + return false + } +} + +func requestNeedsPreviousResponseID(reqData map[string]any) bool { + signals := analyzeToolContinuationSignalsFromMap(reqData) + return signals.hasToolOutput && !signals.hasToolCallContext +} diff --git a/backend/internal/gateway/ws_handler.go b/backend/internal/gateway/ws_handler.go index 369183b..8cb5447 100644 --- a/backend/internal/gateway/ws_handler.go +++ b/backend/internal/gateway/ws_handler.go @@ -8,7 +8,7 @@ import ( "github.com/gorilla/websocket" - sdk "github.com/DouDOU-start/airgate-sdk/sdkgo" + sdk "github.com/DevilGenius/airgate-sdk/sdkgo" ) // wsDialResult 封装 DialWebSocket 的认证失败信息 diff --git a/backend/internal/model/registry.go b/backend/internal/model/registry.go index 4a79100..1a474eb 100644 --- a/backend/internal/model/registry.go +++ b/backend/internal/model/registry.go @@ -4,7 +4,7 @@ import ( "sort" "strings" - sdk "github.com/DouDOU-start/airgate-sdk/sdkgo" + sdk "github.com/DevilGenius/airgate-sdk/sdkgo" ) // ────────────────────────────────────────────────────── @@ -25,7 +25,8 @@ type Spec struct { ContextWindow int MaxOutputTokens int - // 按张计费($/张)。> 0 时图像生成走固定单价,不按 token 估算。 + // 图像模型能力标记价($/张)。> 0 表示具备图像生成能力;实际 token 标准价仍 + // 使用 InputPrice / OutputPrice,分组固定单张价由 Core 通过 metadata 覆盖。 ImagePrice float64 // 标准档单价($/1M tokens) @@ -81,13 +82,12 @@ func withPriorityMultiplier(s Spec, multiplier float64) Spec { return s } -// imgSpec 构造按张计费的图像模型 Spec。 -func imgSpec(name string, pricePerImage float64) Spec { - return Spec{ - Name: name, - ContextWindow: 32000, - ImagePrice: pricePerImage, - } +// imgSpec 构造图像模型 Spec。ImagePrice 只作为图像能力标记; +// 标准计费仍按 token 单价计算,Core 可用分组图片单张价替代最终计费。 +func imgSpec(name string, input, cached, output, pricePerImage float64) Spec { + spec := std(name, 32000, 0, input, cached, output) + spec.ImagePrice = pricePerImage + return spec } // withLongCtx 在已构造的 Spec 基础上附加 gpt-5.4 家族的长上下文阶梯。 @@ -116,10 +116,10 @@ var registry = map[string]Spec{ // ── GPT 基础系列 ── "gpt-5.2": std("GPT 5.2", 272000, 128000, 1.75, 0.175, 14.0), - // ── 图像生成(按张计费 $0.20/张)── - "gpt-image-1": imgSpec("GPT Image 1", 0.20), - "gpt-image-1.5": imgSpec("GPT Image 1.5", 0.20), - "gpt-image-2": imgSpec("GPT Image 2", 0.20), + // ── 图像生成:标准成本按 token 计,分组图片固定单价由 Core 替代最终计费 ── + "gpt-image-1": imgSpec("GPT Image 1", 5.0, 0.5, 30.0, 0.20), + "gpt-image-1.5": imgSpec("GPT Image 1.5", 5.0, 0.5, 30.0, 0.20), + "gpt-image-2": imgSpec("GPT Image 2", 5.0, 0.5, 30.0, 0.20), } // DefaultSpec 未注册模型的最终兜底值。按 gpt-5.4 标准档计价——宁可略高也不能 0。 diff --git a/backend/internal/model/registry_test.go b/backend/internal/model/registry_test.go index 9dd2858..54bd277 100644 --- a/backend/internal/model/registry_test.go +++ b/backend/internal/model/registry_test.go @@ -64,4 +64,14 @@ func TestLookup_KnownModelUnchanged(t *testing.T) { t.Errorf("gpt-5.4 定价变化: In=%v Out=%v", spec.InputPrice, spec.OutputPrice) } }) + + t.Run("gpt-image-2", func(t *testing.T) { + spec := Lookup("gpt-image-2") + if spec.InputPrice != 5.0 || spec.OutputPrice != 30.0 || spec.CachedPrice != 0.5 { + t.Errorf("gpt-image-2 token 定价变化: In=%v Out=%v Cached=%v", spec.InputPrice, spec.OutputPrice, spec.CachedPrice) + } + if spec.ImagePrice <= 0 { + t.Errorf("gpt-image-2 ImagePrice = %v, want image capability marker", spec.ImagePrice) + } + }) } diff --git a/backend/main.go b/backend/main.go index c446eed..672b086 100644 --- a/backend/main.go +++ b/backend/main.go @@ -1,8 +1,8 @@ package main import ( - "github.com/DouDOU-start/airgate-openai/backend/internal/gateway" - sdkgrpc "github.com/DouDOU-start/airgate-sdk/runtimego/grpc" + "github.com/DevilGenius/airgate-openai/backend/internal/gateway" + sdkgrpc "github.com/DevilGenius/airgate-sdk/runtimego/grpc" ) func main() { diff --git a/web/package.json b/web/package.json index 7ff55e9..e6d3524 100644 --- a/web/package.json +++ b/web/package.json @@ -11,7 +11,7 @@ "type-check": "tsc --noEmit" }, "dependencies": { - "@doudou-start/airgate-theme": "https://github.com/DouDOU-start/airgate-sdk/releases/download/v0.2.1/airgate-theme-v0.2.1.tgz", + "@devilgenius/airgate-theme": "https://github.com/DevilGenius/airgate-sdk/releases/download/v0.2.1/airgate-theme-v0.2.1.tgz", "react": "^19.0.0", "react-dom": "^19.0.0" }, diff --git a/web/pnpm-lock.yaml b/web/pnpm-lock.yaml index 747f1f8..14fb5a7 100644 --- a/web/pnpm-lock.yaml +++ b/web/pnpm-lock.yaml @@ -8,9 +8,9 @@ importers: .: dependencies: - '@doudou-start/airgate-theme': - specifier: https://github.com/DouDOU-start/airgate-sdk/releases/download/v0.2.1/airgate-theme-v0.2.1.tgz - version: https://github.com/DouDOU-start/airgate-sdk/releases/download/v0.2.1/airgate-theme-v0.2.1.tgz(react@19.2.5) + '@devilgenius/airgate-theme': + specifier: https://github.com/DevilGenius/airgate-sdk/releases/download/v0.2.1/airgate-theme-v0.2.1.tgz + version: https://github.com/DevilGenius/airgate-sdk/releases/download/v0.2.1/airgate-theme-v0.2.1.tgz(react@19.2.5) react: specifier: ^19.0.0 version: 19.2.5 @@ -128,8 +128,8 @@ packages: resolution: {integrity: sha512-LwdZHpScM4Qz8Xw2iKSzS+cfglZzJGvofQICy7W7v4caru4EaAmyUuO6BGrbyQ2mYV11W0U8j5mBhd14dd3B0A==} engines: {node: '>=6.9.0'} - '@doudou-start/airgate-theme@https://github.com/DouDOU-start/airgate-sdk/releases/download/v0.2.1/airgate-theme-v0.2.1.tgz': - resolution: {tarball: https://github.com/DouDOU-start/airgate-sdk/releases/download/v0.2.1/airgate-theme-v0.2.1.tgz} + '@devilgenius/airgate-theme@https://github.com/DevilGenius/airgate-sdk/releases/download/v0.2.1/airgate-theme-v0.2.1.tgz': + resolution: {tarball: https://github.com/DevilGenius/airgate-sdk/releases/download/v0.2.1/airgate-theme-v0.2.1.tgz} version: 0.2.1 peerDependencies: react: ^19.0.0 @@ -1123,7 +1123,7 @@ snapshots: '@babel/helper-string-parser': 7.27.1 '@babel/helper-validator-identifier': 7.28.5 - '@doudou-start/airgate-theme@https://github.com/DouDOU-start/airgate-sdk/releases/download/v0.2.1/airgate-theme-v0.2.1.tgz(react@19.2.5)': + '@devilgenius/airgate-theme@https://github.com/DevilGenius/airgate-sdk/releases/download/v0.2.1/airgate-theme-v0.2.1.tgz(react@19.2.5)': dependencies: react: 19.2.5 diff --git a/web/src/components/AccountForm.tsx b/web/src/components/AccountForm.tsx index bd2fe54..fda68c3 100644 --- a/web/src/components/AccountForm.tsx +++ b/web/src/components/AccountForm.tsx @@ -1,12 +1,12 @@ import { useCallback, useEffect, useMemo, useState } from 'react'; -import { cssVar } from '@doudou-start/airgate-theme'; +import { cssVar } from '@devilgenius/airgate-theme'; import type { AccountFormProps, PluginBatchAccountInput, PluginOAuthBatchExchangeResult, PluginOAuthBridge, PluginOAuthExchangeResult, -} from '@doudou-start/airgate-theme/plugin'; +} from '@devilgenius/airgate-theme/plugin'; type BatchExchangeResult = PluginOAuthBatchExchangeResult; type BatchAccountInput = PluginBatchAccountInput; diff --git a/web/src/components/AccountIdentity.tsx b/web/src/components/AccountIdentity.tsx index ed9060c..45c7f6e 100644 --- a/web/src/components/AccountIdentity.tsx +++ b/web/src/components/AccountIdentity.tsx @@ -1,5 +1,5 @@ import type { CSSProperties } from 'react'; -import type { AccountSurfaceProps } from '@doudou-start/airgate-theme/plugin'; +import type { AccountSurfaceProps } from '@devilgenius/airgate-theme/plugin'; type AccountLike = { type?: string; diff --git a/web/src/components/UsageCostDetail.tsx b/web/src/components/UsageCostDetail.tsx index 2e1a123..420fd40 100644 --- a/web/src/components/UsageCostDetail.tsx +++ b/web/src/components/UsageCostDetail.tsx @@ -1,5 +1,5 @@ import type { CSSProperties, ReactNode } from 'react'; -import type { UsageRecordSurfaceProps } from '@doudou-start/airgate-theme/plugin'; +import type { UsageRecordSurfaceProps } from '@devilgenius/airgate-theme/plugin'; interface UsageCostDetailItem { key?: string; @@ -8,7 +8,6 @@ interface UsageCostDetailItem { user_cost?: number; billing_multiplier?: number; currency?: string; - metadata?: Record