feat(openai): support Fast mode service_tier across responses/chat/WS paths

- Accept fast|priority (canonical priority), flex|auto|default|scale on
  /v1/responses and /v1/chat/completions; reject unknown/empty/non-string
  with HTTP 400; omitted and null stay compatible.
- Propagate service_tier through JSON/SSE, Responses<->Chat conversions,
  fallback paths and HTTP->upstream WebSocket bridge.
- Billing prefers the upstream terminal tier; the outbound (policy-
  transformed) tier is used only when upstream omits the field.
  Explicit upstream default bills Standard even when Fast was requested.
- Pricing: Fast premium 2x Standard for gpt-5.6-sol/terra/luna and
  gpt-5.4; 2.5x for gpt-5.5; channel FastMultiplier stays authoritative.
- Live verification (official Codex 0.149.0 + gateway, HTTP & WS):
  upstream ChatGPT backend may return terminal default even when the
  account catalog advertises priority; billing follows the actual tier.
This commit is contained in:
alfadb
2026-08-24 11:52:48 +08:00
parent 7075ae0d82
commit f06bf181d2
25 changed files with 1518 additions and 82 deletions
@@ -89,6 +89,10 @@ func (h *OpenAIGatewayHandler) ChatCompletions(c *gin.Context) {
h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", invalidStreamFieldTypeMessage)
return
}
if _, err := service.ValidateOpenAIServiceTierField(body); err != nil {
h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", err.Error())
return
}
if service.IsGPTImageGenerationModel(reqModel) {
h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "This model is not supported on the Chat Completions endpoint")
return
@@ -392,6 +392,10 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) {
h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", invalidStreamFieldTypeMessage)
return
}
if _, err := service.ValidateOpenAIServiceTierField(body); err != nil {
h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", err.Error())
return
}
reqLog = reqLog.With(zap.String("model", reqModel), zap.Bool("stream", reqStream))
previousResponseID := strings.TrimSpace(gjson.GetBytes(body, "previous_response_id").String())
if previousResponseID != "" {
@@ -0,0 +1,127 @@
package handler
import (
"net/http"
"net/http/httptest"
"strings"
"testing"
"github.com/Wei-Shaw/sub2api/internal/config"
middleware2 "github.com/Wei-Shaw/sub2api/internal/server/middleware"
"github.com/Wei-Shaw/sub2api/internal/service"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/require"
)
// 非法 service_tier 必须在两个 OpenAI 端点(/v1/responses、/v1/chat/completions)
// 上以 OpenAI 兼容错误结构返回 HTTP 400;合法值(fast/priority)不被拒绝。
func newServiceTierHandlerTest(t *testing.T) *OpenAIGatewayHandler {
t.Helper()
return &OpenAIGatewayHandler{
gatewayService: &service.OpenAIGatewayService{},
billingCacheService: service.NewBillingCacheService(nil, nil, nil, nil, nil, nil, &config.Config{RunMode: config.RunModeSimple}, nil),
apiKeyService: &service.APIKeyService{},
concurrencyHelper: &ConcurrencyHelper{concurrencyService: service.NewConcurrencyService(
&helperConcurrencyCacheStub{userSeq: []bool{true}},
)},
cfg: &config.Config{},
imageLimiter: &imageConcurrencyLimiter{},
}
}
func runOpenAIHandlerServiceTierTest(t *testing.T, path, body string, handler func(h *OpenAIGatewayHandler, c *gin.Context)) *httptest.ResponseRecorder {
t.Helper()
gin.SetMode(gin.TestMode)
rec := httptest.NewRecorder()
c, _ := gin.CreateTestContext(rec)
c.Request = httptest.NewRequest(http.MethodPost, path, strings.NewReader(body))
c.Request.Header.Set("Content-Type", "application/json")
groupID := int64(6401)
userID := int64(6402)
c.Set(string(middleware2.ContextKeyAPIKey), &service.APIKey{
ID: 6403,
GroupID: &groupID,
Group: &service.Group{
ID: groupID,
Platform: service.PlatformOpenAI,
},
User: &service.User{ID: userID, Status: service.StatusActive},
})
c.Set(string(middleware2.ContextKeyUser), middleware2.AuthSubject{UserID: userID, Concurrency: 1})
handler(newServiceTierHandlerTest(t), c)
return rec
}
func TestOpenAIGatewayHandlerResponses_InvalidServiceTierRejected400(t *testing.T) {
for _, body := range []string{
`{"model":"gpt-5.5","input":"hi","service_tier":"turbo"}`,
`{"model":"gpt-5.5","input":"hi","service_tier":"SPEED"}`,
`{"model":"gpt-5.5","input":"hi","service_tier":""}`,
`{"model":"gpt-5.5","input":"hi","service_tier":123}`,
`{"model":"gpt-5.5","input":"hi","service_tier":{}}`,
} {
rec := runOpenAIHandlerServiceTierTest(t, "/v1/responses", body, func(h *OpenAIGatewayHandler, c *gin.Context) {
h.Responses(c)
})
require.Equal(t, http.StatusBadRequest, rec.Code, "body=%s", body)
require.Contains(t, rec.Body.String(), "invalid_request_error", "body=%s", body)
require.Contains(t, rec.Body.String(), "invalid service_tier", "body=%s", body)
}
}
func TestOpenAIGatewayHandlerResponses_ValidServiceTierNotRejected(t *testing.T) {
for _, tier := range []string{"fast", "priority", "flex"} {
body := `{"model":"gpt-5.5","input":"hi","service_tier":"` + tier + `"}`
rec := runOpenAIHandlerServiceTierTest(t, "/v1/responses", body, func(h *OpenAIGatewayHandler, c *gin.Context) {
h.Responses(c)
})
require.NotEqual(t, http.StatusBadRequest, rec.Code, "tier=%s must not be rejected as invalid", tier)
require.NotContains(t, rec.Body.String(), "invalid service_tier", "tier=%s", tier)
}
}
func TestOpenAIGatewayHandlerResponses_ServiceTierOmittedKeepsCurrentBehavior(t *testing.T) {
body := `{"model":"gpt-5.5","input":"hi"}`
rec := runOpenAIHandlerServiceTierTest(t, "/v1/responses", body, func(h *OpenAIGatewayHandler, c *gin.Context) {
h.Responses(c)
})
require.NotContains(t, rec.Body.String(), "invalid service_tier")
}
func TestOpenAIGatewayHandlerChatCompletions_InvalidServiceTierRejected400(t *testing.T) {
for _, body := range []string{
`{"model":"gpt-5.5","messages":[{"role":"user","content":"hi"}],"service_tier":"turbo"}`,
`{"model":"gpt-5.5","messages":[{"role":"user","content":"hi"}],"service_tier":"ultra"}`,
`{"model":"gpt-5.5","messages":[{"role":"user","content":"hi"}],"service_tier":""}`,
`{"model":"gpt-5.5","messages":[{"role":"user","content":"hi"}],"service_tier":["priority"]}`,
} {
rec := runOpenAIHandlerServiceTierTest(t, "/v1/chat/completions", body, func(h *OpenAIGatewayHandler, c *gin.Context) {
h.ChatCompletions(c)
})
require.Equal(t, http.StatusBadRequest, rec.Code, "body=%s", body)
require.Contains(t, rec.Body.String(), "invalid_request_error", "body=%s", body)
require.Contains(t, rec.Body.String(), "invalid service_tier", "body=%s", body)
}
}
func TestOpenAIGatewayHandlerChatCompletions_ValidServiceTierNotRejected(t *testing.T) {
for _, tier := range []string{"fast", "priority", "auto", "default", "scale", "flex"} {
body := `{"model":"gpt-5.5","messages":[{"role":"user","content":"hi"}],"service_tier":"` + tier + `"}`
rec := runOpenAIHandlerServiceTierTest(t, "/v1/chat/completions", body, func(h *OpenAIGatewayHandler, c *gin.Context) {
h.ChatCompletions(c)
})
require.NotEqual(t, http.StatusBadRequest, rec.Code, "tier=%q must not be rejected as invalid", tier)
require.NotContains(t, rec.Body.String(), "invalid service_tier", "tier=%q", tier)
}
}
func TestOpenAIGatewayHandlerChatCompletions_ServiceTierOmittedKeepsCurrentBehavior(t *testing.T) {
body := `{"model":"gpt-5.5","messages":[{"role":"user","content":"hi"}]}`
rec := runOpenAIHandlerServiceTierTest(t, "/v1/chat/completions", body, func(h *OpenAIGatewayHandler, c *gin.Context) {
h.ChatCompletions(c)
})
require.NotContains(t, rec.Body.String(), "invalid service_tier")
}
@@ -1212,10 +1212,11 @@ func ChatCompletionsResponseToResponses(resp *ChatCompletionsResponse, model str
}
out := &ResponsesResponse{
ID: id,
Object: "response",
Model: model,
Status: "completed",
ID: id,
Object: "response",
Model: model,
Status: "completed",
ServiceTier: chatServiceTier(resp),
}
if resp == nil {
out.Output = []ResponsesOutput{emptyResponsesMessageOutput()}
@@ -1242,6 +1243,13 @@ func ChatCompletionsResponseToResponses(resp *ChatCompletionsResponse, model str
return out
}
func chatServiceTier(resp *ChatCompletionsResponse) string {
if resp == nil {
return ""
}
return resp.ServiceTier
}
func chatMessageToResponsesOutput(message ChatMessage, customTools, functionTools map[string]bool, toolSearch bool, namespaceTools map[string]NamespacedToolName) []ResponsesOutput {
var outputs []ResponsesOutput
reasoning := message.reasoningText()
@@ -1414,6 +1422,7 @@ type ChatCompletionsToResponsesStreamState struct {
ResponseID string
Model string
Created int64
ServiceTier string // upstream Chat chunk service_tier, echoed on response events
SequenceNumber int
CreatedSent bool
CompletedSent bool
@@ -1545,6 +1554,9 @@ func ChatCompletionsChunkToResponsesEvents(
if state.Model == "" && chunk.Model != "" {
state.Model = chunk.Model
}
if chunk.ServiceTier != "" {
state.ServiceTier = chunk.ServiceTier
}
if chunk.Usage != nil {
state.Usage = ChatUsageToResponsesUsage(chunk.Usage)
}
@@ -1701,6 +1713,7 @@ func FinalizeChatCompletionsResponsesStream(state *ChatCompletionsToResponsesStr
Object: "response",
Model: state.Model,
Status: status,
ServiceTier: state.ServiceTier,
Output: state.chatOutput(),
Usage: state.Usage,
IncompleteDetails: incompleteDetails,
@@ -1716,11 +1729,12 @@ func ensureChatToResponsesCreated(state *ChatCompletionsToResponsesStreamState)
state.CreatedSent = true
return []ResponsesStreamEvent{chatToResponsesEvent(state, "response.created", &ResponsesStreamEvent{
Response: &ResponsesResponse{
ID: state.ResponseID,
Object: "response",
Model: state.Model,
Status: "in_progress",
Output: []ResponsesOutput{},
ID: state.ResponseID,
Object: "response",
Model: state.Model,
Status: "in_progress",
ServiceTier: state.ServiceTier,
Output: []ResponsesOutput{},
},
})}
}
@@ -23,10 +23,11 @@ func ResponsesToChatCompletions(resp *ResponsesResponse, model string) *ChatComp
}
out := &ChatCompletionsResponse{
ID: id,
Object: "chat.completion",
Created: time.Now().Unix(),
Model: model,
ID: id,
Object: "chat.completion",
Created: time.Now().Unix(),
Model: model,
ServiceTier: resp.ServiceTier,
}
var contentText string
@@ -118,6 +119,7 @@ type ResponsesEventToChatState struct {
ID string
Model string
Created int64
ServiceTier string // upstream tier observed on response events; echoed on chunks
SentRole bool
SawToolCall bool
SawText bool
@@ -187,12 +189,13 @@ func FinalizeResponsesChatStream(state *ResponsesEventToChatState) []ChatComplet
if state.IncludeUsage && state.Usage != nil {
chunks = append(chunks, ChatCompletionsChunk{
ID: state.ID,
Object: "chat.completion.chunk",
Created: state.Created,
Model: state.Model,
Choices: []ChatChunkChoice{},
Usage: state.Usage,
ID: state.ID,
Object: "chat.completion.chunk",
Created: state.Created,
Model: state.Model,
ServiceTier: state.ServiceTier,
Choices: []ChatChunkChoice{},
Usage: state.Usage,
})
}
@@ -218,6 +221,9 @@ func resToChatHandleCreated(evt *ResponsesStreamEvent, state *ResponsesEventToCh
if state.Model == "" && evt.Response.Model != "" {
state.Model = evt.Response.Model
}
if evt.Response.ServiceTier != "" {
state.ServiceTier = evt.Response.ServiceTier
}
}
// Emit the role chunk.
if state.SentRole {
@@ -301,6 +307,9 @@ func resToChatHandleCompleted(evt *ResponsesStreamEvent, state *ResponsesEventTo
if evt.Response.Usage != nil {
state.Usage = chatUsageFromResponsesUsage(evt.Response.Usage)
}
if evt.Response.ServiceTier != "" {
state.ServiceTier = evt.Response.ServiceTier
}
switch evt.Response.Status {
case "incomplete":
@@ -326,12 +335,13 @@ func resToChatHandleCompleted(evt *ResponsesStreamEvent, state *ResponsesEventTo
if state.IncludeUsage && state.Usage != nil {
chunks = append(chunks, ChatCompletionsChunk{
ID: state.ID,
Object: "chat.completion.chunk",
Created: state.Created,
Model: state.Model,
Choices: []ChatChunkChoice{},
Usage: state.Usage,
ID: state.ID,
Object: "chat.completion.chunk",
Created: state.Created,
Model: state.Model,
ServiceTier: state.ServiceTier,
Choices: []ChatChunkChoice{},
Usage: state.Usage,
})
}
@@ -401,10 +411,11 @@ func completionDetailsFromResponses(src *ResponsesOutputTokensDetails) *ChatToke
func makeChatDeltaChunk(state *ResponsesEventToChatState, delta ChatDelta) ChatCompletionsChunk {
return ChatCompletionsChunk{
ID: state.ID,
Object: "chat.completion.chunk",
Created: state.Created,
Model: state.Model,
ID: state.ID,
Object: "chat.completion.chunk",
Created: state.Created,
Model: state.Model,
ServiceTier: state.ServiceTier,
Choices: []ChatChunkChoice{{
Index: 0,
Delta: delta,
@@ -416,10 +427,11 @@ func makeChatDeltaChunk(state *ResponsesEventToChatState, delta ChatDelta) ChatC
func makeChatFinishChunk(state *ResponsesEventToChatState, finishReason string) ChatCompletionsChunk {
empty := ""
return ChatCompletionsChunk{
ID: state.ID,
Object: "chat.completion.chunk",
Created: state.Created,
Model: state.Model,
ID: state.ID,
Object: "chat.completion.chunk",
Created: state.Created,
Model: state.Model,
ServiceTier: state.ServiceTier,
Choices: []ChatChunkChoice{{
Index: 0,
Delta: ChatDelta{Content: &empty},
@@ -0,0 +1,175 @@
package apicompat
import (
"encoding/json"
"testing"
"github.com/stretchr/testify/require"
)
// 上游响应中的 service_tier 必须如实回传,不被重写/丢弃:
// 非流式(ResponsesResponse → ChatCompletionsResponse)与流式 chunk 均覆盖。
func TestResponsesToChatCompletions_PreservesUpstreamServiceTier(t *testing.T) {
resp := &ResponsesResponse{
ID: "resp_1",
Object: "response",
Model: "gpt-5.5",
Status: "completed",
ServiceTier: "priority",
Output: []ResponsesOutput{{
Type: "message",
Role: "assistant",
Content: []ResponsesContentPart{{
Type: "output_text",
Text: "hi",
}},
}},
Usage: &ResponsesUsage{InputTokens: 1, OutputTokens: 1},
}
chat := ResponsesToChatCompletions(resp, "gpt-5.5")
require.Equal(t, "priority", chat.ServiceTier)
// 序列化后字段仍在(omitempty 不丢非空值)。
raw, err := json.Marshal(chat)
require.NoError(t, err)
require.Contains(t, string(raw), `"service_tier":"priority"`)
}
func TestResponsesToChatCompletions_OmitsMissingServiceTier(t *testing.T) {
resp := &ResponsesResponse{ID: "resp_1", Model: "gpt-5.5", Status: "completed"}
chat := ResponsesToChatCompletions(resp, "gpt-5.5")
require.Empty(t, chat.ServiceTier)
raw, err := json.Marshal(chat)
require.NoError(t, err)
require.NotContains(t, string(raw), "service_tier")
}
func TestResponsesEventToChatChunks_PreservesUpstreamServiceTier(t *testing.T) {
state := NewResponsesEventToChatState()
state.IncludeUsage = true
created := &ResponsesStreamEvent{Type: "response.created"}
require.NoError(t, json.Unmarshal([]byte(`{"type":"response.created","response":{"id":"resp_s1","model":"gpt-5.5","service_tier":"priority","status":"in_progress"}}`), created))
chunks := ResponsesEventToChatChunks(created, state)
require.NotEmpty(t, chunks)
for _, chunk := range chunks {
require.Equal(t, "priority", chunk.ServiceTier)
}
// 后续 delta chunk 继续携带(OpenAI 流式 chunk 的 service_tier 语义)。
delta := &ResponsesStreamEvent{Type: "response.output_text.delta", Delta: "hi"}
chunks = ResponsesEventToChatChunks(delta, state)
require.NotEmpty(t, chunks)
require.Equal(t, "priority", chunks[0].ServiceTier)
// 终止事件同样携带。
completed := &ResponsesStreamEvent{Type: "response.completed", Response: &ResponsesResponse{
ID: "resp", Model: "gpt-5.5", Status: "completed",
Usage: &ResponsesUsage{InputTokens: 1, OutputTokens: 1},
}}
chunks = ResponsesEventToChatChunks(completed, state)
require.NotEmpty(t, chunks)
for _, chunk := range chunks {
require.Equal(t, "priority", chunk.ServiceTier)
}
}
func TestResponsesEventToChatChunks_NoServiceTierStaysClean(t *testing.T) {
state := NewResponsesEventToChatState()
created := &ResponsesStreamEvent{Type: "response.created", Response: &ResponsesResponse{ID: "resp", Model: "gpt-5.5"}}
chunks := ResponsesEventToChatChunks(created, state)
require.NotEmpty(t, chunks)
require.Empty(t, chunks[0].ServiceTier)
raw, err := json.Marshal(chunks[0])
require.NoError(t, err)
require.NotContains(t, string(raw), "service_tier")
}
// 上游 JSON 反序列化时 service_tier 进入 ResponsesResponse(缓冲桥读取链路)。
func TestResponsesResponse_UnmarshalPreservesServiceTier(t *testing.T) {
var resp ResponsesResponse
require.NoError(t, json.Unmarshal([]byte(`{"id":"resp_1","object":"response","model":"gpt-5.5","status":"completed","service_tier":"flex","output":[]}`), &resp))
require.Equal(t, "flex", resp.ServiceTier)
}
// ---------------------------------------------------------------------------
// 反向转换(Chat-only fallback):CC 响应/流 chunk 的 service_tier 保留到
// Responses 形态,客户端与计费都能拿到上游回显。
// ---------------------------------------------------------------------------
func TestChatCompletionsResponseToResponses_PreservesServiceTier(t *testing.T) {
cc := &ChatCompletionsResponse{
ID: "chatcmpl-1",
Model: "gpt-5.5",
ServiceTier: "default",
Choices: []ChatChoice{{
Index: 0,
Message: ChatMessage{Role: "assistant", Content: json.RawMessage(`"hi"`)},
FinishReason: "stop",
}},
Usage: &ChatUsage{PromptTokens: 1, CompletionTokens: 1, TotalTokens: 2},
}
resp := ChatCompletionsResponseToResponses(cc, "gpt-5.5", nil, false, nil)
require.Equal(t, "default", resp.ServiceTier)
raw, err := json.Marshal(resp)
require.NoError(t, err)
require.Contains(t, string(raw), `"service_tier":"default"`)
}
func TestChatCompletionsResponseToResponses_NilRespOmitsServiceTier(t *testing.T) {
resp := ChatCompletionsResponseToResponses(nil, "gpt-5.5", nil, false, nil)
require.Empty(t, resp.ServiceTier)
}
func TestChatCompletionsChunkToResponsesEvents_PreservesServiceTier(t *testing.T) {
state := NewChatCompletionsToResponsesStreamState("gpt-5.5")
chunk := &ChatCompletionsChunk{
ID: "chatcmpl-2",
Model: "gpt-5.5",
ServiceTier: "flex",
Choices: []ChatChunkChoice{{
Index: 0,
Delta: ChatDelta{Content: strPtr("hi")},
}},
}
events := ChatCompletionsChunkToResponsesEvents(chunk, state)
require.NotEmpty(t, events)
// response.created 携带 service_tier。
created := findEvent(events, "response.created")
require.NotNil(t, created)
require.NotNil(t, created.Response)
require.Equal(t, "flex", created.Response.ServiceTier)
// 终止事件同样携带。
final := FinalizeChatCompletionsResponsesStream(state)
completed := findEvent(final, "response.completed")
require.NotNil(t, completed)
require.NotNil(t, completed.Response)
require.Equal(t, "flex", completed.Response.ServiceTier)
}
func TestChatCompletionsChunkToResponsesEvents_NoTierStaysClean(t *testing.T) {
state := NewChatCompletionsToResponsesStreamState("gpt-5.5")
chunk := &ChatCompletionsChunk{ID: "chatcmpl-3", Model: "gpt-5.5"}
events := ChatCompletionsChunkToResponsesEvents(chunk, state)
created := findEvent(events, "response.created")
require.NotNil(t, created)
require.Empty(t, created.Response.ServiceTier)
raw, err := json.Marshal(created)
require.NoError(t, err)
require.NotContains(t, string(raw), "service_tier")
}
func findEvent(events []ResponsesStreamEvent, eventType string) *ResponsesStreamEvent {
for i := range events {
if events[i].Type == eventType {
return &events[i]
}
}
return nil
}
+7 -6
View File
@@ -344,12 +344,13 @@ func (t *ResponsesTool) UnmarshalJSON(data []byte) error {
// ResponsesResponse is the non-streaming response from POST /v1/responses.
type ResponsesResponse struct {
ID string `json:"id"`
Object string `json:"object"` // "response"
Model string `json:"model"`
Status string `json:"status"` // "completed" | "incomplete" | "failed"
Output []ResponsesOutput `json:"output"`
Usage *ResponsesUsage `json:"usage,omitempty"`
ID string `json:"id"`
Object string `json:"object"` // "response"
Model string `json:"model"`
Status string `json:"status"` // "completed" | "incomplete" | "failed"
Output []ResponsesOutput `json:"output"`
Usage *ResponsesUsage `json:"usage,omitempty"`
ServiceTier string `json:"service_tier,omitempty"` // upstream tier, echoed back verbatim
// incomplete_details is present when status="incomplete"
IncompleteDetails *ResponsesIncompleteDetails `json:"incomplete_details,omitempty"`
+38 -1
View File
@@ -1475,7 +1475,8 @@ func (s *BillingService) applyModelSpecificPricingPolicy(model string, pricing *
(pricing.LongContextInputThreshold <= 0 || pricing.LongContextInputMultiplier <= 0 || pricing.LongContextOutputMultiplier <= 0)
needsCacheCreationPolicy := isGPT56 && !pricing.CacheCreationPriceExplicit && (pricing.CacheCreationPricePerToken <= 0 ||
(pricing.InputPricePerTokenPriority > 0 && pricing.CacheCreationPricePerTokenPriority <= 0))
if !needsLongContextPolicy && !needsCacheCreationPolicy {
fastRatio := openAIModelFastPricingRatio(normalized)
if !needsLongContextPolicy && !needsCacheCreationPolicy && fastRatio <= 0 {
return pricing
}
cloned := *pricing
@@ -1498,9 +1499,45 @@ func (s *BillingService) applyModelSpecificPricingPolicy(model string, pricing *
cloned.LongContextOutputMultiplier = openAIGPT54LongContextOutputMultiplier
}
}
if fastRatio > 0 {
enforceOpenAIFastPricingRatio(&cloned, fastRatio)
}
return &cloned
}
// openAIModelFastPricingRatio 返回业务口径下 OpenAI GPT-5.x 模型 Fast/priority
// 的标准价倍率:gpt-5.6 系列与 gpt-5.4 为 2x,gpt-5.5 为 2.5x。未定义 Fast
// 档的模型(如 gpt-5.5-pro、gpt-5.4-mini/nano)返回 0。
func openAIModelFastPricingRatio(normalized string) float64 {
switch normalized {
case "gpt-5.4", "gpt-5.6-sol", "gpt-5.6-terra", "gpt-5.6-luna":
return 2.0
case "gpt-5.5":
return 2.5
default:
return 0
}
}
// enforceOpenAIFastPricingRatio 把 priority 档价格改写为「标准价 × ratio」。
// 本地/远程 LiteLLM 目录可能只带官方旧口径(如 gpt-5.5 priority 仍标 2x),
// 直接采用会导致 Fast 模式少计费;这里按业务倍率兜底修正,且对已正确的
// fallback 条目(2x/2.5x)是幂等的。computeTokenBreakdown 在 priority 价格
// 存在时走显式档位价、不再叠加通用 tier 倍率,因此不会重复乘价。
func enforceOpenAIFastPricingRatio(pricing *ModelPricing, ratio float64) {
if pricing == nil || ratio <= 0 {
return
}
pricing.InputPricePerTokenPriority = pricing.InputPricePerToken * ratio
pricing.OutputPricePerTokenPriority = pricing.OutputPricePerToken * ratio
if pricing.CacheReadPricePerToken > 0 {
pricing.CacheReadPricePerTokenPriority = pricing.CacheReadPricePerToken * ratio
}
if pricing.CacheCreationPricePerToken > 0 {
pricing.CacheCreationPricePerTokenPriority = pricing.CacheCreationPricePerToken * ratio
}
}
func (s *BillingService) shouldApplySessionLongContextPricing(tokens UsageTokens, pricing *ModelPricing) bool {
if pricing == nil || pricing.LongContextInputThreshold <= 0 {
return false
@@ -0,0 +1,834 @@
package service
import (
"bytes"
"context"
"encoding/json"
"errors"
"io"
"net/http"
"net/http/httptest"
"strings"
"testing"
"github.com/Wei-Shaw/sub2api/internal/config"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/require"
"github.com/tidwall/gjson"
)
// ---------------------------------------------------------------------------
// 请求侧:service_tier 校验(fast/priority 等价、非法值拒绝、省略保持现状)
// ---------------------------------------------------------------------------
func TestValidateOpenAIServiceTierField(t *testing.T) {
t.Parallel()
t.Run("fast normalizes to priority", func(t *testing.T) {
norm, err := ValidateOpenAIServiceTierField([]byte(`{"model":"gpt-5.5","service_tier":"fast"}`))
require.NoError(t, err)
require.Equal(t, "priority", norm)
})
t.Run("priority passes through", func(t *testing.T) {
norm, err := ValidateOpenAIServiceTierField([]byte(`{"model":"gpt-5.5","service_tier":"priority"}`))
require.NoError(t, err)
require.Equal(t, "priority", norm)
})
t.Run("case and whitespace insensitive", func(t *testing.T) {
norm, err := ValidateOpenAIServiceTierField([]byte(`{"model":"gpt-5.5","service_tier":" FAST "}`))
require.NoError(t, err)
require.Equal(t, "priority", norm)
})
t.Run("official tiers pass through", func(t *testing.T) {
for _, tier := range []string{"flex", "auto", "default", "scale"} {
norm, err := ValidateOpenAIServiceTierField([]byte(`{"model":"gpt-5.5","service_tier":"` + tier + `"}`))
require.NoError(t, err, "tier %q must be accepted", tier)
require.Equal(t, tier, norm)
}
})
t.Run("invalid tier rejected", func(t *testing.T) {
_, err := ValidateOpenAIServiceTierField([]byte(`{"model":"gpt-5.5","service_tier":"turbo"}`))
require.Error(t, err)
var invalid *ErrInvalidOpenAIServiceTier
require.True(t, errors.As(err, &invalid))
require.Equal(t, "turbo", invalid.Value)
require.Contains(t, err.Error(), "invalid service_tier")
require.Contains(t, err.Error(), "fast", "allowed-value hint must mention fast")
})
t.Run("omitted field stays valid", func(t *testing.T) {
norm, err := ValidateOpenAIServiceTierField([]byte(`{"model":"gpt-5.5","input":"hi"}`))
require.NoError(t, err)
require.Empty(t, norm)
})
t.Run("null value keeps omission semantics", func(t *testing.T) {
norm, err := ValidateOpenAIServiceTierField([]byte(`{"model":"gpt-5.5","service_tier":null}`))
require.NoError(t, err)
require.Empty(t, norm)
})
t.Run("explicit empty string rejected as invalid enum value", func(t *testing.T) {
_, err := ValidateOpenAIServiceTierField([]byte(`{"model":"gpt-5.5","service_tier":""}`))
require.Error(t, err)
var invalid *ErrInvalidOpenAIServiceTier
require.True(t, errors.As(err, &invalid))
})
t.Run("non-string service_tier rejected", func(t *testing.T) {
// service_tier 必须为字符串;数字/布尔/对象/数组等类型同样按非法值拒绝。
for _, raw := range []string{
`{"model":"gpt-5.5","service_tier":123}`,
`{"model":"gpt-5.5","service_tier":true}`,
`{"model":"gpt-5.5","service_tier":{}}`,
`{"model":"gpt-5.5","service_tier":["priority"]}`,
} {
_, err := ValidateOpenAIServiceTierField([]byte(raw))
require.Error(t, err, "raw=%s must be rejected", raw)
var invalid *ErrInvalidOpenAIServiceTier
require.True(t, errors.As(err, &invalid), "raw=%s", raw)
require.Equal(t, "<non-string>", invalid.Value, "raw=%s", raw)
require.Contains(t, err.Error(), "invalid service_tier")
}
})
t.Run("oversized unknown string is truncated", func(t *testing.T) {
blob := strings.Repeat("z", 4096)
_, err := ValidateOpenAIServiceTierField([]byte(`{"model":"gpt-5.5","service_tier":"` + blob + `"}`))
require.Error(t, err)
require.Contains(t, err.Error(), "invalid service_tier")
require.NotContains(t, err.Error(), blob)
require.Less(t, len(err.Error()), 200)
var invalid *ErrInvalidOpenAIServiceTier
require.True(t, errors.As(err, &invalid))
require.Equal(t, strings.Repeat("z", 64)+"...", invalid.Value)
})
t.Run("non-string large object/array is not echoed", func(t *testing.T) {
blob := strings.Repeat("x", 4096)
payloads := []string{
`{"model":"gpt-5.5","service_tier":{"blob":"` + blob + `"}}`,
`{"model":"gpt-5.5","service_tier":["` + blob + `"]}`,
}
for _, raw := range payloads {
_, err := ValidateOpenAIServiceTierField([]byte(raw))
require.Error(t, err)
require.Contains(t, err.Error(), "invalid service_tier")
require.NotContains(t, err.Error(), blob)
require.Less(t, len(err.Error()), 200)
var invalid *ErrInvalidOpenAIServiceTier
require.True(t, errors.As(err, &invalid))
require.Equal(t, "<non-string>", invalid.Value)
}
})
}
// ---------------------------------------------------------------------------
// 计费:gpt-5.6 系列 / gpt-5.4 按标准价 2x,gpt-5.5 按标准价 2.5x
// ---------------------------------------------------------------------------
func TestApplyModelSpecificPricingPolicy_EnforcesOpenAIFastRatios(t *testing.T) {
t.Parallel()
svc := &BillingService{}
t.Run("gpt-5.5 catalog 2x priority is corrected to 2.5x", func(t *testing.T) {
// 模拟本地 LiteLLM 目录仍携带官方旧口径(gpt-5.5 priority = 2x)。
catalog := &ModelPricing{
InputPricePerToken: 5e-6,
InputPricePerTokenPriority: 10e-6,
OutputPricePerToken: 30e-6,
OutputPricePerTokenPriority: 60e-6,
CacheReadPricePerToken: 0.5e-6,
CacheReadPricePerTokenPriority: 1e-6,
}
got := svc.applyModelSpecificPricingPolicy("gpt-5.5", catalog)
require.InDelta(t, 12.5e-6, got.InputPricePerTokenPriority, 1e-12)
require.InDelta(t, 75e-6, got.OutputPricePerTokenPriority, 1e-12)
require.InDelta(t, 1.25e-6, got.CacheReadPricePerTokenPriority, 1e-12)
// 标准价不被改动。
require.InDelta(t, 5e-6, got.InputPricePerToken, 1e-12)
// 原始指针不被污染。
require.InDelta(t, 10e-6, catalog.InputPricePerTokenPriority, 1e-12)
})
t.Run("gpt-5.4 keeps 2x", func(t *testing.T) {
got := svc.applyModelSpecificPricingPolicy("gpt-5.4", &ModelPricing{
InputPricePerToken: 2.5e-6,
InputPricePerTokenPriority: 5e-6,
OutputPricePerToken: 15e-6,
OutputPricePerTokenPriority: 30e-6,
})
require.InDelta(t, 5e-6, got.InputPricePerTokenPriority, 1e-12)
require.InDelta(t, 30e-6, got.OutputPricePerTokenPriority, 1e-12)
})
t.Run("gpt-5.6 family keeps 2x", func(t *testing.T) {
for _, model := range []string{"gpt-5.6", "gpt-5.6-sol", "gpt-5.6-terra", "gpt-5.6-luna", "gpt-5.6-max", "gpt-5.6-sol-preview"} {
got := svc.applyModelSpecificPricingPolicy(model, &ModelPricing{
InputPricePerToken: 5e-6,
InputPricePerTokenPriority: 10e-6,
OutputPricePerToken: 30e-6,
OutputPricePerTokenPriority: 60e-6,
CacheReadPricePerToken: 0.5e-6,
CacheReadPricePerTokenPriority: 1e-6,
})
require.InDelta(t, 10e-6, got.InputPricePerTokenPriority, 1e-12, "model %s", model)
require.InDelta(t, 60e-6, got.OutputPricePerTokenPriority, 1e-12, "model %s", model)
}
})
t.Run("missing priority prices are backfilled from standard", func(t *testing.T) {
got := svc.applyModelSpecificPricingPolicy("gpt-5.5", &ModelPricing{
InputPricePerToken: 5e-6,
OutputPricePerToken: 30e-6,
CacheReadPricePerToken: 0.5e-6,
CacheCreationPricePerToken: 5e-6,
})
require.InDelta(t, 12.5e-6, got.InputPricePerTokenPriority, 1e-12)
require.InDelta(t, 75e-6, got.OutputPricePerTokenPriority, 1e-12)
require.InDelta(t, 1.25e-6, got.CacheReadPricePerTokenPriority, 1e-12)
require.InDelta(t, 12.5e-6, got.CacheCreationPricePerTokenPriority, 1e-12)
})
t.Run("gpt-5.5-pro has no mandated fast tier", func(t *testing.T) {
got := svc.applyModelSpecificPricingPolicy("gpt-5.5-pro", &ModelPricing{
InputPricePerToken: 30e-6,
InputPricePerTokenPriority: 60e-6,
OutputPricePerToken: 180e-6,
})
require.InDelta(t, 60e-6, got.InputPricePerTokenPriority, 1e-12)
})
t.Run("unrelated models untouched", func(t *testing.T) {
got := svc.applyModelSpecificPricingPolicy("claude-opus-5", &ModelPricing{InputPricePerToken: 1, OutputPricePerToken: 2})
require.InDelta(t, 1, got.InputPricePerToken, 1e-12)
require.Zero(t, got.InputPricePerTokenPriority)
})
}
func TestOpenAIFastBillingMultiplier_2xAnd25x(t *testing.T) {
t.Parallel()
// 目录数据携带官方旧口径(gpt-5.5 priority=2x);修正后 fast 必须按 2.5x 计费。
catalog := map[string]*LiteLLMModelPricing{
"gpt-5.4": {
InputCostPerToken: 2.5e-6,
InputCostPerTokenPriority: 5e-6,
OutputCostPerToken: 15e-6,
OutputCostPerTokenPriority: 30e-6,
CacheReadInputTokenCost: 0.25e-6,
CacheReadInputTokenCostPriority: 0.5e-6,
},
"gpt-5.5": {
InputCostPerToken: 5e-6,
InputCostPerTokenPriority: 10e-6,
OutputCostPerToken: 30e-6,
OutputCostPerTokenPriority: 60e-6,
CacheReadInputTokenCost: 0.5e-6,
CacheReadInputTokenCostPriority: 1e-6,
},
"gpt-5.6-sol": {
InputCostPerToken: 5e-6,
InputCostPerTokenPriority: 10e-6,
OutputCostPerToken: 30e-6,
OutputCostPerTokenPriority: 60e-6,
CacheReadInputTokenCost: 0.5e-6,
CacheReadInputTokenCostPriority: 1e-6,
},
}
billing := NewBillingService(&config.Config{}, &PricingService{pricingData: catalog})
tokens := UsageTokens{InputTokens: 1_000_000, OutputTokens: 1_000_000}
standard := func(model string) *CostBreakdown {
cost, err := billing.CalculateCost(model, tokens, 1)
require.NoError(t, err)
return cost
}
fast := func(model, tier string) *CostBreakdown {
cost, err := billing.CalculateCostWithServiceTier(model, tokens, 1, tier)
require.NoError(t, err)
return cost
}
tests := []struct {
model string
ratio float64
}{
{model: "gpt-5.4", ratio: 2.0},
{model: "gpt-5.5", ratio: 2.5},
{model: "gpt-5.6-sol", ratio: 2.0},
{model: "gpt-5.6-terra", ratio: 2.0},
{model: "gpt-5.6-luna", ratio: 2.0},
}
for _, tt := range tests {
t.Run(tt.model+"/fast", func(t *testing.T) {
base := standard(tt.model)
fastCost := fast(tt.model, "fast")
require.InDelta(t, base.TotalCost*tt.ratio, fastCost.TotalCost, 1e-9,
"fast total must be %.1fx standard", tt.ratio)
})
t.Run(tt.model+"/priority_alias", func(t *testing.T) {
fastCost := fast(tt.model, "fast")
priorityCost := fast(tt.model, "priority")
require.InDelta(t, fastCost.TotalCost, priorityCost.TotalCost, 1e-12,
"client alias fast must bill identically to priority")
require.InDelta(t, standard(tt.model).TotalCost*tt.ratio, priorityCost.TotalCost, 1e-9)
})
t.Run(tt.model+"/no_tier_unchanged", func(t *testing.T) {
base := standard(tt.model)
noTier, err := billing.CalculateCostWithServiceTier(tt.model, tokens, 1, "")
require.NoError(t, err)
require.InDelta(t, base.TotalCost, noTier.TotalCost, 1e-12)
})
t.Run(tt.model+"/default_equals_standard", func(t *testing.T) {
base := standard(tt.model)
defaultCost, err := billing.CalculateCostWithServiceTier(tt.model, tokens, 1, "default")
require.NoError(t, err)
require.InDelta(t, base.TotalCost, defaultCost.TotalCost, 1e-12)
require.InDelta(t, base.InputCost, defaultCost.InputCost, 1e-12)
require.InDelta(t, base.OutputCost, defaultCost.OutputCost, 1e-12)
require.InDelta(t, base.CacheReadCost, defaultCost.CacheReadCost, 1e-12)
})
}
}
func TestOpenAIFastBilling_FastMultiplierOverridesEnforcedRatio(t *testing.T) {
t.Parallel()
svc := &BillingService{}
catalog := &ModelPricing{
InputPricePerToken: 5e-6,
InputPricePerTokenPriority: 10e-6,
OutputPricePerToken: 30e-6,
OutputPricePerTokenPriority: 60e-6,
CacheReadPricePerToken: 0.5e-6,
CacheReadPricePerTokenPriority: 1e-6,
}
pricing := svc.applyModelSpecificPricingPolicy("gpt-5.5", catalog)
require.InDelta(t, 12.5e-6, pricing.InputPricePerTokenPriority, 1e-12, "enforce must still write 2.5x priority prices")
require.InDelta(t, 75e-6, pricing.OutputPricePerTokenPriority, 1e-12)
multiplier := 1.7
pricing.FastMultiplier = &multiplier
tokens := UsageTokens{InputTokens: 1_000_000, OutputTokens: 1_000_000, CacheReadTokens: 1_000_000}
standard := svc.computeTokenBreakdown(pricing, tokens, 1, "", false)
fast := svc.computeTokenBreakdown(pricing, tokens, 1, "fast", false)
priority := svc.computeTokenBreakdown(pricing, tokens, 1, "priority", false)
require.InDelta(t, standard.TotalCost*1.7, fast.TotalCost, 1e-9)
require.InDelta(t, fast.TotalCost, priority.TotalCost, 1e-12)
withoutOverride := *pricing
withoutOverride.FastMultiplier = nil
enforced := svc.computeTokenBreakdown(&withoutOverride, tokens, 1, "fast", false)
require.InDelta(t, standard.TotalCost*2.5, enforced.TotalCost, 1e-9,
"without FastMultiplier the same enforced prices still bill 2.5x")
}
// ---------------------------------------------------------------------------
// 上游 payload:fast 归一化为 priority 并确实到达上游
// ---------------------------------------------------------------------------
func TestForwardAsChatCompletions_ServiceTierFastNormalizedToPriorityUpstream(t *testing.T) {
gin.SetMode(gin.TestMode)
rec := httptest.NewRecorder()
c, _ := gin.CreateTestContext(rec)
body := []byte(`{"model":"gpt-5.5","messages":[{"role":"user","content":"hello"}],"service_tier":"fast","stream":false}`)
c.Request = httptest.NewRequest(http.MethodPost, "/v1/chat/completions", bytes.NewReader(body))
c.Request.Header.Set("Content-Type", "application/json")
upstream := &httpUpstreamRecorder{resp: &http.Response{
StatusCode: http.StatusBadRequest,
Header: http.Header{"Content-Type": []string{"application/json"}, "x-request-id": []string{"rid-chat-st"}},
Body: io.NopCloser(strings.NewReader(`{"error":{"type":"invalid_request_error","message":"stop"}}`)),
}}
svc := &OpenAIGatewayService{
cfg: &config.Config{},
httpUpstream: upstream,
settingService: NewSettingService(&openAIFastPolicyRepoStub{values: map[string]string{}}, &config.Config{}),
}
account := &Account{
ID: 21,
Name: "openai-compatible",
Platform: PlatformOpenAI,
Type: AccountTypeAPIKey,
Concurrency: 1,
Credentials: map[string]any{"api_key": "sk-compatible"},
Extra: map[string]any{"openai_responses_supported": true},
}
_, err := svc.ForwardAsChatCompletions(context.Background(), c, account, body, "", "gpt-5.5")
require.Error(t, err) // upstream 400 → 错误返回,但请求体已被 recorder 捕获
require.NotNil(t, upstream.lastBody)
require.Equal(t, "priority", gjson.GetBytes(upstream.lastBody, "service_tier").String(),
"client alias fast must reach upstream as priority")
}
func TestForwardAsChatCompletions_ServiceTierPriorityPreservedUpstream(t *testing.T) {
gin.SetMode(gin.TestMode)
rec := httptest.NewRecorder()
c, _ := gin.CreateTestContext(rec)
body := []byte(`{"model":"gpt-5.5","messages":[{"role":"user","content":"hello"}],"service_tier":"priority","stream":false}`)
c.Request = httptest.NewRequest(http.MethodPost, "/v1/chat/completions", bytes.NewReader(body))
c.Request.Header.Set("Content-Type", "application/json")
upstream := &httpUpstreamRecorder{resp: &http.Response{
StatusCode: http.StatusBadRequest,
Header: http.Header{"Content-Type": []string{"application/json"}, "x-request-id": []string{"rid-chat-st2"}},
Body: io.NopCloser(strings.NewReader(`{"error":{"type":"invalid_request_error","message":"stop"}}`)),
}}
svc := &OpenAIGatewayService{
cfg: &config.Config{},
httpUpstream: upstream,
settingService: NewSettingService(&openAIFastPolicyRepoStub{values: map[string]string{}}, &config.Config{}),
}
account := &Account{
ID: 2,
Name: "openai-compatible",
Platform: PlatformOpenAI,
Type: AccountTypeAPIKey,
Concurrency: 1,
Credentials: map[string]any{"api_key": "sk-compatible"},
Extra: map[string]any{"openai_responses_supported": true},
}
_, err := svc.ForwardAsChatCompletions(context.Background(), c, account, body, "", "gpt-5.5")
require.Error(t, err)
require.NotNil(t, upstream.lastBody)
require.Equal(t, "priority", gjson.GetBytes(upstream.lastBody, "service_tier").String())
}
func TestForward_ResponsesServiceTierFastNormalizedToPriorityUpstream(t *testing.T) {
gin.SetMode(gin.TestMode)
rec := httptest.NewRecorder()
c, _ := gin.CreateTestContext(rec)
body := []byte(`{"model":"gpt-5.5","service_tier":"fast","input":"hello","stream":false}`)
c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", bytes.NewReader(body))
c.Request.Header.Set("Content-Type", "application/json")
upstream := &httpUpstreamRecorder{resp: &http.Response{
StatusCode: http.StatusOK,
Header: http.Header{"Content-Type": []string{"application/json"}, "x-request-id": []string{"rid-resp-st"}},
Body: io.NopCloser(strings.NewReader(
`{"id":"resp_1","object":"response","status":"completed","model":"gpt-5.5","output":[],"usage":{"input_tokens":1,"output_tokens":1}}`,
)),
}}
svc := &OpenAIGatewayService{
cfg: &config.Config{Security: config.SecurityConfig{URLAllowlist: config.URLAllowlistConfig{Enabled: false}}},
httpUpstream: upstream,
settingService: NewSettingService(&openAIFastPolicyRepoStub{values: map[string]string{}}, &config.Config{}),
}
account := &Account{
ID: 7,
Name: "openai-apikey",
Platform: PlatformOpenAI,
Type: AccountTypeAPIKey,
Concurrency: 1,
Credentials: map[string]any{"api_key": "sk-test"},
Extra: map[string]any{"openai_responses_supported": true},
Status: StatusActive,
Schedulable: true,
}
result, err := svc.Forward(context.Background(), c, account, body)
require.NoError(t, err)
require.NotNil(t, result)
require.NotNil(t, upstream.lastBody)
require.Equal(t, "priority", gjson.GetBytes(upstream.lastBody, "service_tier").String(),
"client alias fast must reach the upstream as priority")
// 计费上下文:result 携带归一化后的 tier。
require.NotNil(t, result.ServiceTier)
require.Equal(t, "priority", *result.ServiceTier)
}
func TestForward_ResponsesServiceTierOmittedStaysOmitted(t *testing.T) {
gin.SetMode(gin.TestMode)
rec := httptest.NewRecorder()
c, _ := gin.CreateTestContext(rec)
body := []byte(`{"model":"gpt-5.5","input":"hello","stream":false}`)
c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", bytes.NewReader(body))
c.Request.Header.Set("Content-Type", "application/json")
upstream := &httpUpstreamRecorder{resp: &http.Response{
StatusCode: http.StatusOK,
Header: http.Header{"Content-Type": []string{"application/json"}, "x-request-id": []string{"rid-resp-st2"}},
Body: io.NopCloser(strings.NewReader(
`{"id":"resp_2","object":"response","status":"completed","model":"gpt-5.5","output":[],"usage":{"input_tokens":1,"output_tokens":1}}`,
)),
}}
svc := &OpenAIGatewayService{
cfg: &config.Config{Security: config.SecurityConfig{URLAllowlist: config.URLAllowlistConfig{Enabled: false}}},
httpUpstream: upstream,
settingService: NewSettingService(&openAIFastPolicyRepoStub{values: map[string]string{}}, &config.Config{}),
}
account := &Account{
ID: 7,
Name: "openai-apikey",
Platform: PlatformOpenAI,
Type: AccountTypeAPIKey,
Concurrency: 1,
Credentials: map[string]any{"api_key": "sk-test"},
Extra: map[string]any{"openai_responses_supported": true},
Status: StatusActive,
Schedulable: true,
}
result, err := svc.Forward(context.Background(), c, account, body)
require.NoError(t, err)
require.NotNil(t, result)
require.NotNil(t, upstream.lastBody)
require.False(t, gjson.GetBytes(upstream.lastBody, "service_tier").Exists(),
"omitted service_tier must stay omitted")
require.Nil(t, result.ServiceTier)
}
// ---------------------------------------------------------------------------
// 流式计费上下文:service_tier 需要从请求体传到 usage 计费
// ---------------------------------------------------------------------------
func TestForwardStreaming_ServiceTierPropagatedToResult(t *testing.T) {
gin.SetMode(gin.TestMode)
rec := httptest.NewRecorder()
c, _ := gin.CreateTestContext(rec)
body := []byte(`{"model":"gpt-5.5","service_tier":"fast","input":"hello","stream":true}`)
c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", bytes.NewReader(body))
c.Request.Header.Set("Content-Type", "application/json")
streamPayload := "data: {\"type\":\"response.created\",\"response\":{\"id\":\"resp_s1\",\"object\":\"response\",\"model\":\"gpt-5.5\",\"status\":\"in_progress\"}}\n\n" +
"data: {\"type\":\"response.output_text.delta\",\"item_id\":\"it_1\",\"output_index\":0,\"delta\":\"hi\"}\n\n" +
"data: {\"type\":\"response.completed\",\"response\":{\"id\":\"resp_s1\",\"object\":\"response\",\"model\":\"gpt-5.5\",\"status\":\"completed\",\"output\":[],\"usage\":{\"input_tokens\":1,\"output_tokens\":1}}}\n\n" +
"data: [DONE]\n\n"
upstream := &httpUpstreamRecorder{resp: &http.Response{
StatusCode: http.StatusOK,
Header: http.Header{"Content-Type": []string{"text/event-stream"}, "x-request-id": []string{"rid-resp-stream-st"}},
Body: io.NopCloser(strings.NewReader(streamPayload)),
}}
svc := &OpenAIGatewayService{
cfg: &config.Config{Security: config.SecurityConfig{URLAllowlist: config.URLAllowlistConfig{Enabled: false}}},
httpUpstream: upstream,
settingService: NewSettingService(&openAIFastPolicyRepoStub{values: map[string]string{}}, &config.Config{}),
}
account := &Account{
ID: 7,
Name: "openai-apikey",
Platform: PlatformOpenAI,
Type: AccountTypeAPIKey,
Concurrency: 1,
Credentials: map[string]any{"api_key": "sk-test"},
Extra: map[string]any{"openai_responses_supported": true},
Status: StatusActive,
Schedulable: true,
}
result, err := svc.Forward(context.Background(), c, account, body)
require.NoError(t, err)
require.NotNil(t, result)
require.NotNil(t, result.ServiceTier)
require.Equal(t, "priority", *result.ServiceTier, "streaming billing context must carry the normalized tier")
// /v1/responses 流是上游 SSE 原样透传:上游没回 service_tier 就不该出现;
// 网关只在计费结果里携带请求侧 tier,不往下游流里注入。
require.Contains(t, rec.Body.String(), `"delta":"hi"`, "streamed content must reach the client")
require.NotContains(t, rec.Body.String(), `"service_tier"`, "upstream did not return service_tier, client stream must stay untouched")
}
// ---------------------------------------------------------------------------
// 上游回显优先:请求 fast 但上游真实返回 default → 计费按标准价
// ---------------------------------------------------------------------------
func TestForward_ResponsesUpstreamEchoesDefault_OverridesRequestFast(t *testing.T) {
gin.SetMode(gin.TestMode)
rec := httptest.NewRecorder()
c, _ := gin.CreateTestContext(rec)
body := []byte(`{"model":"gpt-5.5","service_tier":"fast","input":"hello","stream":false}`)
c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", bytes.NewReader(body))
c.Request.Header.Set("Content-Type", "application/json")
// 上游回显 service_tier=default(例如请求实际被降级)。
upstream := &httpUpstreamRecorder{resp: &http.Response{
StatusCode: http.StatusOK,
Header: http.Header{"Content-Type": []string{"application/json"}, "x-request-id": []string{"rid-resp-echo"}},
Body: io.NopCloser(strings.NewReader(
`{"id":"resp_1","object":"response","status":"completed","model":"gpt-5.5","service_tier":"default","output":[],"usage":{"input_tokens":1,"output_tokens":1}}`,
)),
}}
svc := &OpenAIGatewayService{
cfg: &config.Config{Security: config.SecurityConfig{URLAllowlist: config.URLAllowlistConfig{Enabled: false}}},
httpUpstream: upstream,
settingService: NewSettingService(&openAIFastPolicyRepoStub{values: map[string]string{}}, &config.Config{}),
}
account := &Account{
ID: 7,
Name: "openai-apikey",
Platform: PlatformOpenAI,
Type: AccountTypeAPIKey,
Concurrency: 1,
Credentials: map[string]any{"api_key": "sk-test"},
Extra: map[string]any{"openai_responses_supported": true},
Status: StatusActive,
Schedulable: true,
}
result, err := svc.Forward(context.Background(), c, account, body)
require.NoError(t, err)
require.NotNil(t, result)
require.NotNil(t, result.ServiceTier)
require.Equal(t, "default", *result.ServiceTier,
"upstream-echoed default must override the client-requested fast tier for billing")
// 非流式响应原样透传:客户端同样看到 default。
require.Contains(t, rec.Body.String(), `"service_tier":"default"`)
require.NotContains(t, rec.Body.String(), `"service_tier":"priority"`)
}
func TestForwardStreaming_UpstreamEchoesDefault_OverridesRequestFast(t *testing.T) {
gin.SetMode(gin.TestMode)
rec := httptest.NewRecorder()
c, _ := gin.CreateTestContext(rec)
body := []byte(`{"model":"gpt-5.5","service_tier":"fast","input":"hello","stream":true}`)
c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", bytes.NewReader(body))
c.Request.Header.Set("Content-Type", "application/json")
streamPayload := "data: {\"type\":\"response.created\",\"response\":{\"id\":\"resp_s1\",\"object\":\"response\",\"model\":\"gpt-5.5\",\"status\":\"in_progress\"}}\n\n" +
"data: {\"type\":\"response.completed\",\"response\":{\"id\":\"resp_s1\",\"object\":\"response\",\"model\":\"gpt-5.5\",\"status\":\"completed\",\"service_tier\":\"default\",\"output\":[],\"usage\":{\"input_tokens\":1,\"output_tokens\":1}}}\n\n" +
"data: [DONE]\n\n"
upstream := &httpUpstreamRecorder{resp: &http.Response{
StatusCode: http.StatusOK,
Header: http.Header{"Content-Type": []string{"text/event-stream"}, "x-request-id": []string{"rid-resp-echo-stream"}},
Body: io.NopCloser(strings.NewReader(streamPayload)),
}}
svc := &OpenAIGatewayService{
cfg: &config.Config{Security: config.SecurityConfig{URLAllowlist: config.URLAllowlistConfig{Enabled: false}}},
httpUpstream: upstream,
settingService: NewSettingService(&openAIFastPolicyRepoStub{values: map[string]string{}}, &config.Config{}),
}
account := &Account{
ID: 7,
Name: "openai-apikey",
Platform: PlatformOpenAI,
Type: AccountTypeAPIKey,
Concurrency: 1,
Credentials: map[string]any{"api_key": "sk-test"},
Extra: map[string]any{"openai_responses_supported": true},
Status: StatusActive,
Schedulable: true,
}
result, err := svc.Forward(context.Background(), c, account, body)
require.NoError(t, err)
require.NotNil(t, result)
require.NotNil(t, result.ServiceTier)
require.Equal(t, "default", *result.ServiceTier,
"terminal SSE event's upstream-echoed default must win for billing")
// 流式原样透传:客户端在终止事件里看到 default。
require.Contains(t, rec.Body.String(), `"service_tier":"default"`)
}
func TestForwardAsChatCompletions_UpstreamEchoesDefault_BillsStandard(t *testing.T) {
gin.SetMode(gin.TestMode)
rec := httptest.NewRecorder()
c, _ := gin.CreateTestContext(rec)
body := []byte(`{"model":"gpt-5.5","messages":[{"role":"user","content":"hello"}],"service_tier":"fast","stream":false}`)
c.Request = httptest.NewRequest(http.MethodPost, "/v1/chat/completions", bytes.NewReader(body))
c.Request.Header.Set("Content-Type", "application/json")
streamPayload := "data: {\"type\":\"response.created\",\"response\":{\"id\":\"resp_c1\",\"object\":\"response\",\"model\":\"gpt-5.5\",\"status\":\"in_progress\"}}\n\n" +
"data: {\"type\":\"response.output_text.delta\",\"item_id\":\"it_1\",\"output_index\":0,\"delta\":\"hi\"}\n\n" +
"data: {\"type\":\"response.completed\",\"response\":{\"id\":\"resp_c1\",\"object\":\"response\",\"model\":\"gpt-5.5\",\"status\":\"completed\",\"service_tier\":\"default\",\"output\":[],\"usage\":{\"input_tokens\":1,\"output_tokens\":1}}}\n\n" +
"data: [DONE]\n\n"
upstream := &httpUpstreamRecorder{resp: &http.Response{
StatusCode: http.StatusOK,
Header: http.Header{"Content-Type": []string{"text/event-stream"}, "x-request-id": []string{"rid-chat-echo"}},
Body: io.NopCloser(strings.NewReader(streamPayload)),
}}
svc := &OpenAIGatewayService{
cfg: &config.Config{Security: config.SecurityConfig{URLAllowlist: config.URLAllowlistConfig{Enabled: false}}},
httpUpstream: upstream,
settingService: NewSettingService(&openAIFastPolicyRepoStub{values: map[string]string{}}, &config.Config{}),
}
account := &Account{
ID: 21,
Name: "openai-compatible",
Platform: PlatformOpenAI,
Type: AccountTypeAPIKey,
Concurrency: 1,
Credentials: map[string]any{"api_key": "sk-compatible"},
Extra: map[string]any{"openai_responses_supported": true},
Status: StatusActive,
Schedulable: true,
}
result, err := svc.ForwardAsChatCompletions(context.Background(), c, account, body, "", "gpt-5.5")
require.NoError(t, err)
require.NotNil(t, result)
require.NotNil(t, result.ServiceTier)
require.Equal(t, "default", *result.ServiceTier,
"CC bridge must bill on the upstream-echoed default, not the requested fast")
// 缓冲转回 Chat Completions:客户端响应里如实回显 default。
require.Contains(t, rec.Body.String(), `"service_tier":"default"`)
require.NotContains(t, rec.Body.String(), `"service_tier":"priority"`)
}
// ---------------------------------------------------------------------------
// policy filter:删除 service_tier 后不得再按原请求 Fast 计费
// ---------------------------------------------------------------------------
func TestForward_ServiceTierFilteredByPolicyBillsStandard(t *testing.T) {
gin.SetMode(gin.TestMode)
rec := httptest.NewRecorder()
c, _ := gin.CreateTestContext(rec)
body := []byte(`{"model":"gpt-5.5","service_tier":"priority","input":"hello","stream":false}`)
c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", bytes.NewReader(body))
c.Request.Header.Set("Content-Type", "application/json")
// 管理员配置 priority → filter:字段在出站前被删除。
settings := &OpenAIFastPolicySettings{Rules: []OpenAIFastPolicyRule{{
ServiceTier: OpenAIFastTierPriority,
Action: BetaPolicyActionFilter,
Scope: BetaPolicyScopeAll,
}}}
raw, err := json.Marshal(settings)
require.NoError(t, err)
repo := &openAIFastPolicyRepoStub{values: map[string]string{SettingKeyOpenAIFastPolicySettings: string(raw)}}
upstream := &httpUpstreamRecorder{resp: &http.Response{
StatusCode: http.StatusOK,
Header: http.Header{"Content-Type": []string{"application/json"}, "x-request-id": []string{"rid-resp-filter"}},
Body: io.NopCloser(strings.NewReader(
`{"id":"resp_1","object":"response","status":"completed","model":"gpt-5.5","output":[],"usage":{"input_tokens":1,"output_tokens":1}}`,
)),
}}
svc := &OpenAIGatewayService{
cfg: &config.Config{Security: config.SecurityConfig{URLAllowlist: config.URLAllowlistConfig{Enabled: false}}},
httpUpstream: upstream,
settingService: NewSettingService(repo, &config.Config{}),
}
account := &Account{
ID: 7,
Name: "openai-apikey",
Platform: PlatformOpenAI,
Type: AccountTypeAPIKey,
Concurrency: 1,
Credentials: map[string]any{"api_key": "sk-test"},
Extra: map[string]any{"openai_responses_supported": true},
Status: StatusActive,
Schedulable: true,
}
result, err := svc.Forward(context.Background(), c, account, body)
require.NoError(t, err)
require.NotNil(t, result)
// 出站 body 已剥离 service_tier、上游也未回显 → 无 tier → 按标准价计费。
require.False(t, gjson.GetBytes(upstream.lastBody, "service_tier").Exists(),
"policy filter must strip service_tier from the outbound body")
require.Nil(t, result.ServiceTier, "filtered request must not bill as fast")
}
// ---------------------------------------------------------------------------
// 上游回显观察与解析器单测
// ---------------------------------------------------------------------------
func TestUpstreamResponseModelObserver_ObservesServiceTier(t *testing.T) {
t.Parallel()
observer := &upstreamResponseModelObserver{}
observer.ObserveOpenAI([]byte(`{"type":"response.created","response":{"model":"gpt-5.5","service_tier":"flex"}}`), "response.created")
require.Equal(t, "flex", observer.ServiceTier())
// terminal 声明优先。
observer.ObserveOpenAI([]byte(`{"type":"response.completed","response":{"model":"gpt-5.5","service_tier":"default"}}`), "response.completed")
require.Equal(t, "default", observer.ServiceTier())
// Chat Completions 顶层 service_tier 同样可观察。
ccObserver := &upstreamResponseModelObserver{}
ccObserver.ObserveOpenAI([]byte(`{"id":"chatcmpl-1","model":"gpt-5.5","service_tier":"priority","choices":[]}`), "chat.completion")
require.Equal(t, "priority", ccObserver.ServiceTier())
}
func TestResolvedOpenAIUpstreamServiceTier(t *testing.T) {
t.Parallel()
priority := func() *string { v := "priority"; return &v }()
t.Run("upstream echo wins over outbound tier", func(t *testing.T) {
gin.SetMode(gin.TestMode)
c, _ := gin.CreateTestContext(nil)
observer := beginUpstreamResponseModelObservation(c)
observer.ObserveOpenAI([]byte(`{"type":"response.completed","response":{"service_tier":"default"}}`), "response.completed")
got := resolvedOpenAIUpstreamServiceTier(c, priority)
require.NotNil(t, got)
require.Equal(t, "default", *got)
})
t.Run("no upstream echo falls back to outbound tier", func(t *testing.T) {
gin.SetMode(gin.TestMode)
c, _ := gin.CreateTestContext(nil)
beginUpstreamResponseModelObservation(c)
got := resolvedOpenAIUpstreamServiceTier(c, priority)
require.NotNil(t, got)
require.Equal(t, "priority", *got)
})
t.Run("upstream alias fast normalizes to priority", func(t *testing.T) {
gin.SetMode(gin.TestMode)
c, _ := gin.CreateTestContext(nil)
observer := beginUpstreamResponseModelObservation(c)
observer.ObserveOpenAI([]byte(`{"type":"response.completed","response":{"service_tier":"fast"}}`), "response.completed")
got := resolvedOpenAIUpstreamServiceTier(c, nil)
require.NotNil(t, got)
require.Equal(t, "priority", *got)
})
t.Run("no observer keeps outbound tier", func(t *testing.T) {
got := resolvedOpenAIUpstreamServiceTier(nil, priority)
require.NotNil(t, got)
require.Equal(t, "priority", *got)
})
t.Run("no observer and no outbound tier stays nil", func(t *testing.T) {
require.Nil(t, resolvedOpenAIUpstreamServiceTier(nil, nil))
})
t.Run("local observer wins without gin context", func(t *testing.T) {
observer := &upstreamResponseModelObserver{}
observer.ObserveOpenAI([]byte(`{"type":"response.completed","response":{"service_tier":"default"}}`), "response.completed")
got := resolvedOpenAIUpstreamServiceTierFromObserver(observer, priority)
require.NotNil(t, got)
require.Equal(t, "default", *got)
})
t.Run("nil local observer falls back to outbound tier", func(t *testing.T) {
got := resolvedOpenAIUpstreamServiceTierFromObserver(nil, priority)
require.NotNil(t, got)
require.Equal(t, "priority", *got)
})
}
@@ -249,6 +249,7 @@ type ccStreamScanState struct {
// emit 回调做各自的协议转换与写出。读错误按既有约定过滤 context 取消类噪声后
// 记入 Warn 日志。
func (s *OpenAIGatewayService) scanCCStream(
c *gin.Context,
resp *http.Response,
logPrefix string,
requestID string,
@@ -272,6 +273,10 @@ func (s *OpenAIGatewayService) scanCCStream(
st.SawDone = true
break
}
// 观察上游 CC chunk 回显的 model / service_tier(计费以回显为准)。
if observer := upstreamResponseModelObserverFromContext(c); observer != nil {
observer.ObserveOpenAI([]byte(payload), "chat.completion.chunk")
}
if u := extractCCStreamUsage(payload); u != nil {
st.Usage = *u
@@ -331,6 +336,10 @@ func (s *OpenAIGatewayService) readCCUpstreamJSONResponse(
writeError(c, http.StatusBadGateway, "api_error", "Failed to parse upstream response")
return nil, OpenAIUsage{}, fmt.Errorf("parse chat completions response: %w", err)
}
// 观察上游 CC JSON 回显的 model / service_tier(计费以回显为准)。
if observer := upstreamResponseModelObserverFromContext(c); observer != nil {
observer.ObserveOpenAI(respBody, "chat.completion")
}
usage := OpenAIUsage{}
if parsed, ok := extractOpenAIUsageFromJSONBytes(respBody); ok {
@@ -190,7 +190,8 @@ func (s *OpenAIGatewayService) ForwardAsChatCompletions(
responsesBody = stripped
}
}
responsesBody, normalizedServiceTier, err := normalizeResponsesBodyServiceTier(responsesBody)
var normalizedServiceTier string
responsesBody, normalizedServiceTier, err = normalizeResponsesBodyServiceTier(responsesBody)
if err != nil {
return nil, fmt.Errorf("normalize service_tier in responses-shape body: %w", err)
}
@@ -371,11 +372,13 @@ func (s *OpenAIGatewayService) ForwardAsChatCompletions(
return nil, handleErr
}
// Propagate ServiceTier and ReasoningEffort to result for billing
// Propagate ServiceTier and ReasoningEffort to result for billing.
// 计费 tier 优先采用上游回显值;上游未回显时回退到最终出站 body(经过
// fast policy filter/force 之后)里的 tier,policy filter 删掉字段后不再
// 按原请求 Fast 计费。
if handleErr == nil && result != nil {
if responsesReq.ServiceTier != "" {
st := responsesReq.ServiceTier
result.ServiceTier = &st
if tier := resolvedOpenAIUpstreamServiceTier(c, extractOpenAIServiceTierFromBody(responsesBody)); tier != nil {
result.ServiceTier = tier
}
if responsesReq.Reasoning != nil && responsesReq.Reasoning.Effort != "" {
re := responsesReq.Reasoning.Effort
@@ -475,6 +478,7 @@ func (s *OpenAIGatewayService) handleChatBufferedStreamingResponse(
observer = beginUpstreamResponseModelObservation(c)
}
observer.Observe(finalResponse.Model, true)
observer.ObserveServiceTier(finalResponse.ServiceTier, true)
if strings.TrimSpace(finalResponse.Status) == "failed" {
payload, _ := json.Marshal(gin.H{"type": "response.failed", "response": finalResponse})
// cyber_policy 致命不可重试:不 failover,以 Chat Completions 错误格式回写(F4),
@@ -69,9 +69,6 @@ func (s *OpenAIGatewayService) forwardAsRawChatCompletions(
}
clientStream := gjson.GetBytes(body, "stream").Bool()
// 1b. Extract service tier from the raw body before any transformation.
serviceTier := extractOpenAIServiceTierFromBody(body)
// 2. Resolve model mapping (same as ForwardAsChatCompletions)
billingModel := resolveOpenAIForwardModel(account, originalModel, defaultMappedModel)
upstreamModel := normalizeOpenAIModelForUpstream(account, billingModel)
@@ -106,6 +103,9 @@ func (s *OpenAIGatewayService) forwardAsRawChatCompletions(
return nil, policyErr
}
upstreamBody = updatedBody
// 计费兜底 tier = 最终出站 body(policy filter/force 后)里的 tier;
// 最终值由 resolvedOpenAIUpstreamServiceTier 决定(上游回显优先)。
serviceTier := extractOpenAIServiceTierFromBody(upstreamBody)
if account.Platform == PlatformGrok {
strippedBody, stripErr := stripRedundantGrokChatViewImageTool(upstreamBody)
if stripErr != nil {
@@ -390,7 +390,7 @@ func (s *OpenAIGatewayService) streamRawChatCompletions(
UpstreamResponseModelConflict: observedUpstreamResponseModelConflict(c),
UpstreamResponseServiceTier: observedUpstreamResponseServiceTier(c),
ReasoningEffort: reasoningEffort,
ServiceTier: serviceTier,
ServiceTier: resolvedOpenAIUpstreamServiceTier(c, serviceTier),
Stream: true,
Duration: time.Since(startTime),
FirstTokenMs: firstTokenMs,
@@ -492,7 +492,7 @@ func (s *OpenAIGatewayService) bufferRawChatCompletions(
UpstreamResponseModelConflict: observedUpstreamResponseModelConflict(c),
UpstreamResponseServiceTier: observedUpstreamResponseServiceTier(c),
ReasoningEffort: reasoningEffort,
ServiceTier: serviceTier,
ServiceTier: resolvedOpenAIUpstreamServiceTier(c, serviceTier),
Stream: false,
Duration: time.Since(startTime),
}, nil
@@ -1180,7 +1180,7 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco
UpstreamResponseModel: observedUpstreamResponseModel(c),
UpstreamResponseModelConflict: observedUpstreamResponseModelConflict(c),
UpstreamResponseServiceTier: observedUpstreamResponseServiceTier(c),
ServiceTier: serviceTier,
ServiceTier: resolvedOpenAIUpstreamServiceTier(c, serviceTier),
ReasoningEffort: reasoningEffort,
Stream: reqStream,
OpenAIWSMode: false,
@@ -492,9 +492,10 @@ func (s *OpenAIGatewayService) ForwardAsAnthropic(
if promptCacheKey != "" && anthropicDigestChain != "" {
s.bindOpenAICompatAnthropicDigestPromptCacheKey(account, apiKeyID, anthropicDigestChain, promptCacheKey, anthropicMatchedDigestChain)
}
if responsesReq.ServiceTier != "" {
st := responsesReq.ServiceTier
result.ServiceTier = &st
// 计费 tier 优先采用上游回显值;上游未回显时回退到最终出站 body(经过
// fast policy filter/force 之后)里的 tier。
if tier := resolvedOpenAIUpstreamServiceTier(c, extractOpenAIServiceTierFromBody(responsesBody)); tier != nil {
result.ServiceTier = tier
}
if responsesReq.Reasoning != nil && responsesReq.Reasoning.Effort != "" {
re := responsesReq.Reasoning.Effort
@@ -572,6 +573,7 @@ func (s *OpenAIGatewayService) handleAnthropicBufferedStreamingResponse(
observer = beginUpstreamResponseModelObservation(c)
}
observer.Observe(finalResponse.Model, true)
observer.ObserveServiceTier(finalResponse.ServiceTier, true)
if strings.TrimSpace(finalResponse.Status) == "failed" {
payload, _ := json.Marshal(gin.H{"type": "response.failed", "response": finalResponse})
@@ -158,7 +158,7 @@ func (s *OpenAIGatewayService) bufferChatCompletionsAsAnthropic(
BillingModel: billingModel,
UpstreamModel: upstreamModel,
ReasoningEffort: reasoningEffort,
ServiceTier: serviceTier,
ServiceTier: resolvedOpenAIUpstreamServiceTier(c, serviceTier),
Stream: false,
Duration: time.Since(startTime),
}, nil
@@ -204,7 +204,7 @@ func (s *OpenAIGatewayService) streamChatCompletionsAsAnthropic(
}
}
scan := s.scanCCStream(resp, "openai messages chat fallback", requestID, startTime, emitChunk)
scan := s.scanCCStream(c, resp, "openai messages chat fallback", requestID, startTime, emitChunk)
usage := scan.Usage
if scan.Err != nil {
@@ -218,7 +218,7 @@ func (s *OpenAIGatewayService) streamChatCompletionsAsAnthropic(
BillingModel: billingModel,
UpstreamModel: upstreamModel,
ReasoningEffort: reasoningEffort,
ServiceTier: serviceTier,
ServiceTier: resolvedOpenAIUpstreamServiceTier(c, serviceTier),
Stream: true,
Duration: time.Since(startTime),
FirstTokenMs: scan.FirstTokenMs,
@@ -253,7 +253,7 @@ func (s *OpenAIGatewayService) streamChatCompletionsAsAnthropic(
BillingModel: billingModel,
UpstreamModel: upstreamModel,
ReasoningEffort: reasoningEffort,
ServiceTier: serviceTier,
ServiceTier: resolvedOpenAIUpstreamServiceTier(c, serviceTier),
Stream: true,
Duration: time.Since(startTime),
FirstTokenMs: scan.FirstTokenMs,
@@ -515,7 +515,7 @@ func (s *OpenAIGatewayService) forwardOpenAIPassthrough(
UpstreamResponseModel: observedUpstreamResponseModel(c),
UpstreamResponseModelConflict: observedUpstreamResponseModelConflict(c),
UpstreamResponseServiceTier: observedUpstreamResponseServiceTier(c),
ServiceTier: serviceTier,
ServiceTier: resolvedOpenAIUpstreamServiceTier(c, serviceTier),
ReasoningEffort: reasoningEffort,
Stream: reqStream,
OpenAIWSMode: false,
@@ -1240,6 +1240,57 @@ func normalizeOpenAIServiceTier(raw string) *string {
}
}
// ErrInvalidOpenAIServiceTier indicates a request carried a service_tier value
// that is not a known OpenAI tier. HTTP handlers translate it into a 400
// invalid_request_error so malformed values are rejected up front instead of
// being silently stripped (which would mask the user's intent to use fast
// mode).
type ErrInvalidOpenAIServiceTier struct {
Value string
}
func (e *ErrInvalidOpenAIServiceTier) Error() string {
return fmt.Sprintf("invalid service_tier %q: must be one of auto, default, fast, flex, priority, scale", e.Value)
}
const invalidOpenAIServiceTierValueMaxLen = 64
func boundInvalidOpenAIServiceTierValue(raw string) string {
if len(raw) <= invalidOpenAIServiceTierValueMaxLen {
return raw
}
return raw[:invalidOpenAIServiceTierValueMaxLen] + "..."
}
// ValidateOpenAIServiceTierField validates the service_tier field of a raw
// OpenAI-compatible request body (/v1/responses and /v1/chat/completions).
//
// - absent / null → valid, returns "" (field omitted keeps current behavior)
// - "fast" → normalized to "priority" (the two are equivalent; the canonical
// value is what reaches the OpenAI upstream)
// - "priority" / "flex" / "auto" / "default" / "scale" → valid, returned as-is
// - an explicitly present non-string value, an empty string, or any other
// unknown value → *ErrInvalidOpenAIServiceTier (handler maps to HTTP 400),
// matching OpenAI's enum validation semantics
func ValidateOpenAIServiceTierField(body []byte) (string, error) {
tierResult := gjson.GetBytes(body, "service_tier")
if !tierResult.Exists() || tierResult.Type == gjson.Null {
return "", nil
}
if tierResult.Type != gjson.String {
return "", &ErrInvalidOpenAIServiceTier{Value: "<non-string>"}
}
raw := strings.TrimSpace(tierResult.String())
if raw == "" {
return "", &ErrInvalidOpenAIServiceTier{Value: raw}
}
norm := normalizedOpenAIServiceTierValue(raw)
if norm == "" {
return "", &ErrInvalidOpenAIServiceTier{Value: boundInvalidOpenAIServiceTierValue(raw)}
}
return norm, nil
}
// OpenAIFastBlockedError indicates a request was rejected by the OpenAI fast
// policy (action=block). Mirrors BetaBlockedError on the Claude side.
type OpenAIFastBlockedError struct {
@@ -39,7 +39,6 @@ func (s *OpenAIGatewayService) forwardResponsesViaRawChatCompletions(
}
clientStream := responsesReq.Stream
serviceTier := extractOpenAIServiceTierFromBody(body)
// custom 工具(如 codex 的 exec)降级为 function 工具转发,回程需按名字还原为
// custom_tool_call 项,先记下名字集合;tool_search 工具同理,回程还原为
// tool_search_call 项;namespace 子工具(如 MCP 工具)摊平转发,回程按映射还原
@@ -88,9 +87,10 @@ func (s *OpenAIGatewayService) forwardResponsesViaRawChatCompletions(
}
return nil, err
}
if serviceTier == nil {
serviceTier = extractOpenAIServiceTierFromBody(chatBody)
}
// 计费兜底 tier = 最终出站 body(policy filter/force 后)里的 tier;最终值由
// resolvedOpenAIUpstreamServiceTier 决定(上游回显优先)。filter 删掉字段后
// 这里取到 nil,不再按原请求 Fast 计费。
serviceTier := extractOpenAIServiceTierFromBody(chatBody)
logger.L().Debug("openai responses: forwarding via raw chat completions",
zap.Int64("account_id", account.ID),
@@ -160,7 +160,7 @@ func (s *OpenAIGatewayService) bufferChatCompletionsAsResponses(
BillingModel: billingModel,
UpstreamModel: upstreamModel,
ReasoningEffort: reasoningEffort,
ServiceTier: serviceTier,
ServiceTier: resolvedOpenAIUpstreamServiceTier(c, serviceTier),
Stream: false,
Duration: time.Since(startTime),
}, nil
@@ -216,7 +216,7 @@ func (s *OpenAIGatewayService) streamChatCompletionsAsResponses(
c.Writer.Flush()
}
scan := s.scanCCStream(resp, "openai responses chat fallback", requestID, startTime, func(chunk *apicompat.ChatCompletionsChunk) {
scan := s.scanCCStream(c, resp, "openai responses chat fallback", requestID, startTime, func(chunk *apicompat.ChatCompletionsChunk) {
events := apicompat.ChatCompletionsChunkToResponsesEvents(chunk, state)
s.cacheReasoningItemsFromEvents(events)
writeEvents(events)
@@ -230,7 +230,7 @@ func (s *OpenAIGatewayService) streamChatCompletionsAsResponses(
BillingModel: billingModel,
UpstreamModel: upstreamModel,
ReasoningEffort: reasoningEffort,
ServiceTier: serviceTier,
ServiceTier: resolvedOpenAIUpstreamServiceTier(c, serviceTier),
Stream: true,
Duration: time.Since(startTime),
FirstTokenMs: scan.FirstTokenMs,
@@ -244,7 +244,7 @@ func (s *OpenAIGatewayService) streamChatCompletionsAsResponses(
BillingModel: billingModel,
UpstreamModel: upstreamModel,
ReasoningEffort: reasoningEffort,
ServiceTier: serviceTier,
ServiceTier: resolvedOpenAIUpstreamServiceTier(c, serviceTier),
Stream: true,
Duration: time.Since(startTime),
FirstTokenMs: scan.FirstTokenMs,
@@ -274,7 +274,7 @@ func (s *OpenAIGatewayService) streamChatCompletionsAsResponses(
BillingModel: billingModel,
UpstreamModel: upstreamModel,
ReasoningEffort: reasoningEffort,
ServiceTier: serviceTier,
ServiceTier: resolvedOpenAIUpstreamServiceTier(c, serviceTier),
Stream: true,
Duration: time.Since(startTime),
FirstTokenMs: scan.FirstTokenMs,
@@ -253,10 +253,8 @@ type OpenAIForwardResult struct {
// UpstreamEndpoint is the actual upstream API path used for this request.
// It avoids guessing when one downstream protocol can use multiple upstream endpoints.
UpstreamEndpoint string
// ServiceTier records the OpenAI Responses API service tier requested by the
// client, e.g. "priority" / "flex". Nil means the request did not specify a
// recognized tier. Usage recording lowers it to UpstreamResponseServiceTier
// when the upstream reports a cheaper tier (see ResolveBillingServiceTier).
// ServiceTier 优先取上游实际响应回显的 tier;缺失时回退到最终出站 body 的
// tier。nil 表示两者都无识别 tier。
ServiceTier *string
// ReasoningEffort is extracted from request body (reasoning.effort) or derived from model suffix.
// Stored for usage records display; nil means not provided / not applicable.
@@ -776,7 +776,7 @@ func (s *OpenAIGatewayService) forwardOpenAIWSV2(
UpstreamResponseServiceTier: responseModelObserver.ServiceTier(),
ImageCount: imageCounter.Count(),
ImageOutputSizes: imageCounter.Sizes(),
ServiceTier: extractOpenAIServiceTier(reqBody),
ServiceTier: resolvedOpenAIUpstreamServiceTierFromObserver(responseModelObserver, extractOpenAIServiceTier(reqBody)),
ReasoningEffort: extractOpenAIReasoningEffort(reqBody, mappedModel, originalModel),
Stream: reqStream,
OpenAIWSMode: true,
@@ -0,0 +1,96 @@
package service
import (
"context"
"fmt"
"net/http"
"net/http/httptest"
"testing"
"github.com/Wei-Shaw/sub2api/internal/config"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/require"
)
// HTTP POST /v1/responses → forwardOpenAIWSV2 共用 stream/non-stream 的
// OpenAIForwardResult:上游 response.completed.service_tier 必须覆盖请求
// fast/priority,不能只读 reqBody。
func TestForwardOpenAIWSV2_UpstreamDefaultServiceTierWinsOverRequest(t *testing.T) {
gin.SetMode(gin.TestMode)
cases := []struct {
name string
requestTier string
stream bool
}{
{name: "priority_nonstream", requestTier: "priority", stream: false},
{name: "fast_stream", requestTier: "fast", stream: true},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
rec := httptest.NewRecorder()
c, _ := gin.CreateTestContext(rec)
c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", nil)
c.Request.Header.Set("User-Agent", "unit-test-agent/1.0")
cfg := &config.Config{}
cfg.Security.URLAllowlist.Enabled = false
cfg.Security.URLAllowlist.AllowInsecureHTTP = true
cfg.Gateway.OpenAIWS.Enabled = true
cfg.Gateway.OpenAIWS.APIKeyEnabled = true
cfg.Gateway.OpenAIWS.ResponsesWebsocketsV2 = true
cfg.Gateway.OpenAIWS.MaxConnsPerAccount = 1
cfg.Gateway.OpenAIWS.MinIdlePerAccount = 0
cfg.Gateway.OpenAIWS.MaxIdlePerAccount = 1
cfg.Gateway.OpenAIWS.QueueLimitPerConn = 8
cfg.Gateway.OpenAIWS.DialTimeoutSeconds = 3
cfg.Gateway.OpenAIWS.ReadTimeoutSeconds = 5
cfg.Gateway.OpenAIWS.WriteTimeoutSeconds = 3
captureConn := &openAIWSCaptureConn{
events: [][]byte{
[]byte(`{"type":"response.completed","response":{"id":"resp_tier_v2","status":"completed","service_tier":"default","usage":{"input_tokens":1,"output_tokens":1}}}`),
},
}
captureDialer := &openAIWSCaptureDialer{conn: captureConn}
pool := newOpenAIWSConnPool(cfg)
pool.setClientDialerForTest(captureDialer)
svc := &OpenAIGatewayService{
cfg: cfg,
httpUpstream: &httpUpstreamRecorder{},
cache: &stubGatewayCache{},
openaiWSResolver: NewOpenAIWSProtocolResolver(cfg),
toolCorrector: NewCodexToolCorrector(),
openaiWSPool: pool,
}
account := &Account{
ID: 5882,
Name: "openai-ws-v2-tier",
Platform: PlatformOpenAI,
Type: AccountTypeAPIKey,
Status: StatusActive,
Schedulable: true,
Concurrency: 1,
Credentials: map[string]any{"api_key": "sk-test"},
Extra: map[string]any{"responses_websockets_v2_enabled": true},
}
body := []byte(fmt.Sprintf(
`{"model":"gpt-5.5","stream":%t,"service_tier":%q,"input":[{"type":"input_text","text":"hi"}]}`,
tc.stream, tc.requestTier,
))
result, err := svc.Forward(context.Background(), c, account, body)
require.NoError(t, err)
require.NotNil(t, result)
require.True(t, result.OpenAIWSMode, "must take HTTP POST → forwardOpenAIWSV2, not HTTP fallback")
require.Equal(t, tc.stream, result.Stream)
require.Equal(t, "resp_tier_v2", result.RequestID)
require.NotNil(t, result.ServiceTier)
require.Equal(t, "default", *result.ServiceTier)
require.Equal(t, "priority", captureConn.lastWrite["service_tier"],
"outbound WS payload still carries the requested Fast tier")
})
}
}
@@ -514,7 +514,7 @@ func (s *OpenAIGatewayService) proxyOpenAIWSHTTPBridgeTurn(
UpstreamResponseModel: responseModelObserver.Model(),
UpstreamResponseModelConflict: responseModelObserver.Conflict(),
UpstreamResponseServiceTier: responseModelObserver.ServiceTier(),
ServiceTier: extractOpenAIServiceTierFromBody(body),
ServiceTier: resolvedOpenAIUpstreamServiceTierFromObserver(responseModelObserver, extractOpenAIServiceTierFromBody(body)),
ReasoningEffort: ApplyThinkingEnabledFallback(extractOpenAIReasoningEffortFromBody(body, mappedModel, originalModel), body, mappedModel),
Stream: reqStream,
OpenAIWSMode: true,
@@ -44,6 +44,45 @@ func TestPrepareOpenAIWSHTTPBridgeBodyStripsWSFields(t *testing.T) {
require.Error(t, err)
}
func TestProxyOpenAIWSHTTPBridgeTurn_UpstreamDefaultServiceTierWinsOverRequest(t *testing.T) {
gin.SetMode(gin.TestMode)
// proxyOpenAIWSHTTPBridgeTurn 是 client WS→HTTP bridge,本身不 canonicalize
// fast→priority;生产入口的归一化在 openai_ws_forwarder_ingress.go 的 fast
// policy。本测试只覆盖局部 observer:canonical 请求 priority 被上游
// response.completed service_tier=default 覆盖。
sse := strings.Join([]string{
`data: {"type":"response.completed","response":{"id":"resp_tier","status":"completed","service_tier":"default","usage":{"input_tokens":1,"output_tokens":1}}}`,
``,
}, "\n")
upstream := &httpUpstreamRecorder{resp: &http.Response{
StatusCode: http.StatusOK,
Header: http.Header{"Content-Type": []string{"text/event-stream"}},
Body: io.NopCloser(strings.NewReader(sse)),
}}
svc := &OpenAIGatewayService{
cfg: &config.Config{Gateway: config.GatewayConfig{MaxLineSize: defaultMaxLineSize}},
httpUpstream: upstream,
}
account := &Account{ID: 5881, Platform: PlatformOpenAI, Type: AccountTypeAPIKey, Concurrency: 1}
payload := []byte(`{"type":"response.create","model":"gpt-5.5","stream":true,"service_tier":"priority","input":"hi"}`)
recorder := httptest.NewRecorder()
c, _ := gin.CreateTestContext(recorder)
c.Request = httptest.NewRequest(http.MethodGet, "/v1/responses", nil)
result, err := svc.proxyOpenAIWSHTTPBridgeTurn(
context.Background(), c, account, "test-token", payload, len(payload),
"gpt-5.5", "", "", "", "", 1,
func([]byte) error { return nil },
)
require.NoError(t, err)
require.NotNil(t, result)
require.Equal(t, "priority", gjson.GetBytes(upstream.lastBody, "service_tier").String())
require.NotNil(t, result.ServiceTier)
require.Equal(t, "default", *result.ServiceTier)
}
func TestProxyOpenAIWSHTTPBridgeTurnAPIKeyAdaptsClientTools(t *testing.T) {
gin.SetMode(gin.TestMode)
@@ -22,8 +22,10 @@ const (
// (see responseModelBillingDeclaration).
//
// The same observer also records the service tier the upstream reports having
// used (OpenAI service_tier, Anthropic usage.speed). Billing consumes it through
// ResolveBillingServiceTier, which only ever lowers the tier a request asked for.
// used (OpenAI service_tier, Anthropic usage.speed). The billable tier is
// resolved by resolvedOpenAIUpstreamServiceTierFromObserver (upstream echo
// first, outbound body tier as fallback); the upstream ResolveBillingServiceTier
// only-lowers path additionally audits downgrades at usage-record time.
type upstreamResponseModelObserver struct {
first string
terminal string
@@ -216,6 +218,33 @@ func observedUpstreamResponseServiceTier(c *gin.Context) string {
return upstreamResponseModelObserverFromContext(c).ServiceTier()
}
// resolvedOpenAIUpstreamServiceTierFromObserver 返回计费/用量日志实际使用的
// service tier:
//
// 1. observer 记录到的上游真实回显优先——只有上游实际给了 priority/fast 才按
// Fast 计费;上游回显 default/flex/auto 等则如实采用并据此计费;
// 2. 上游未回显时,回退到「最终出站 body」里的 tier(经过 fast policy
// filter/force 之后),保证 policy filter 删掉字段后不再按原请求 Fast 计费。
//
// HTTP→WS 等使用局部 observer 的路径必须把该 observer 传进来,不能只读
// Gin context——局部 observer 不会自动写入 context。
func resolvedOpenAIUpstreamServiceTierFromObserver(observer *upstreamResponseModelObserver, outboundBodyTier *string) *string {
if observer != nil {
if tier := strings.TrimSpace(observer.ServiceTier()); tier != "" {
return normalizeOpenAIServiceTier(tier)
}
}
return outboundBodyTier
}
// resolvedOpenAIUpstreamServiceTier 读取 Gin context 上的 observer 后委托
// resolvedOpenAIUpstreamServiceTierFromObserver。标准 HTTP 转发路径通过
// beginUpstreamResponseModelObservation 把 observer 挂到 context;局部
// observer 路径应直接调用 FromObserver。
func resolvedOpenAIUpstreamServiceTier(c *gin.Context, outboundBodyTier *string) *string {
return resolvedOpenAIUpstreamServiceTierFromObserver(upstreamResponseModelObserverFromContext(c), outboundBodyTier)
}
func observeOpenAISSEBody(observer *upstreamResponseModelObserver, body string) {
if observer == nil || strings.TrimSpace(body) == "" {
return
@@ -5162,12 +5162,12 @@
"cache_read_input_token_cost": 5e-07,
"cache_read_input_token_cost_above_272k_tokens": 1e-06,
"cache_read_input_token_cost_flex": 2.5e-07,
"cache_read_input_token_cost_priority": 1e-06,
"cache_read_input_token_cost_priority": 1.25e-6,
"input_cost_per_token": 5e-06,
"input_cost_per_token_above_272k_tokens": 1e-05,
"input_cost_per_token_batches": 2.5e-06,
"input_cost_per_token_flex": 2.5e-06,
"input_cost_per_token_priority": 1e-05,
"input_cost_per_token_priority": 12.5e-6,
"litellm_provider": "openai",
"max_input_tokens": 1050000,
"max_output_tokens": 128000,
@@ -5177,7 +5177,7 @@
"output_cost_per_token_above_272k_tokens": 4.5e-05,
"output_cost_per_token_batches": 1.5e-05,
"output_cost_per_token_flex": 1.5e-05,
"output_cost_per_token_priority": 6e-05,
"output_cost_per_token_priority": 75e-6,
"supported_endpoints": [
"/v1/chat/completions",
"/v1/batch",
@@ -5210,12 +5210,12 @@
"cache_read_input_token_cost": 5e-07,
"cache_read_input_token_cost_above_272k_tokens": 1e-06,
"cache_read_input_token_cost_flex": 2.5e-07,
"cache_read_input_token_cost_priority": 1e-06,
"cache_read_input_token_cost_priority": 1.25e-6,
"input_cost_per_token": 5e-06,
"input_cost_per_token_above_272k_tokens": 1e-05,
"input_cost_per_token_batches": 2.5e-06,
"input_cost_per_token_flex": 2.5e-06,
"input_cost_per_token_priority": 1e-05,
"input_cost_per_token_priority": 12.5e-6,
"litellm_provider": "openai",
"max_input_tokens": 1050000,
"max_output_tokens": 128000,
@@ -5225,7 +5225,7 @@
"output_cost_per_token_above_272k_tokens": 4.5e-05,
"output_cost_per_token_batches": 1.5e-05,
"output_cost_per_token_flex": 1.5e-05,
"output_cost_per_token_priority": 6e-05,
"output_cost_per_token_priority": 75e-6,
"supported_endpoints": [
"/v1/chat/completions",
"/v1/batch",