mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-10-07 11:43:23 +08:00
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:
@@ -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
|
||||
}
|
||||
@@ -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"`
|
||||
|
||||
@@ -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",
|
||||
|
||||
Reference in New Issue
Block a user