fix(billing): 默认关闭 OpenAI 长上下文计费

This commit is contained in:
benjamin
2026-07-13 23:32:16 +08:00
parent f63d168ae0
commit e9fb5983cd
19 changed files with 182 additions and 73 deletions
@@ -1303,6 +1303,10 @@ func (h *AccountHandler) ApplyOAuthCredentials(c *gin.Context) {
response.ErrorFrom(c, infraerrors.BadRequest("NOT_OAUTH", "cannot apply oauth credentials to non-OAuth account"))
return
}
if err := service.ValidateOpenAILongContextBillingExtra(existing.Platform, req.Extra); err != nil {
response.ErrorFrom(c, err)
return
}
updatedAccount, err := h.adminService.UpdateAccount(ctx, accountID, &service.UpdateAccountInput{
Type: req.Type,
@@ -8,6 +8,7 @@ import (
"testing"
infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors"
"github.com/Wei-Shaw/sub2api/internal/service"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/require"
)
@@ -112,6 +113,35 @@ func TestAccountCreateBoundaryDoesNotApplyOpenAIValidationToOtherPlatforms(t *te
require.Equal(t, http.StatusOK, recorder.Code)
}
func TestApplyOAuthCredentialsRejectsMalformedOpenAILongContextBillingBeforeMutation(t *testing.T) {
gin.SetMode(gin.TestMode)
stub := newStubAdminService()
stub.getAccountResult = &service.Account{
ID: 1,
Platform: service.PlatformOpenAI,
Type: service.AccountTypeOAuth,
}
handler := NewAccountHandler(stub, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil)
router := gin.New()
router.POST("/accounts/:id/apply-oauth-credentials", handler.ApplyOAuthCredentials)
recorder := httptest.NewRecorder()
request := httptest.NewRequest(http.MethodPost, "/accounts/1/apply-oauth-credentials", bytes.NewBufferString(
`{"type":"oauth","credentials":{"access_token":"new-token"},"extra":{"openai_long_context_billing_enabled":"true"}}`,
))
request.Header.Set("Content-Type", "application/json")
router.ServeHTTP(recorder, request)
require.Equal(t, http.StatusBadRequest, recorder.Code)
var responseBody struct {
Reason string `json:"reason"`
}
require.NoError(t, json.Unmarshal(recorder.Body.Bytes(), &responseBody))
require.Equal(t, "OPENAI_LONG_CONTEXT_BILLING_INVALID", responseBody.Reason)
require.Zero(t, stub.updateAccountCalls)
require.Zero(t, stub.updateAccountExtraCalls)
}
func TestOpenAIOAuthCodexPATBoundaryRejectsMalformedOpenAILongContextBillingValueBeforeTokenValidation(t *testing.T) {
gin.SetMode(gin.TestMode)
handler := NewOpenAIOAuthHandler(nil, newStubAdminService(), nil)
@@ -33,6 +33,9 @@ type stubAdminService struct {
createSparkShadowErr error
updateAccountErr error
bulkUpdateAccountErr error
getAccountResult *service.Account
updateAccountCalls int
updateAccountExtraCalls int
checkMixedErr error
lastMixedCheck struct {
accountID int64
@@ -388,6 +391,9 @@ func (s *stubAdminService) ListOpenAISchedulableAccountsForSchedulerScore(_ cont
}
func (s *stubAdminService) GetAccount(ctx context.Context, id int64) (*service.Account, error) {
if s.getAccountResult != nil {
return s.getAccountResult, nil
}
account := service.Account{ID: id, Name: "account", Status: service.StatusActive}
return &account, nil
}
@@ -413,6 +419,7 @@ func (s *stubAdminService) CreateAccount(ctx context.Context, input *service.Cre
}
func (s *stubAdminService) UpdateAccount(ctx context.Context, id int64, input *service.UpdateAccountInput) (*service.Account, error) {
s.updateAccountCalls++
if s.updateAccountErr != nil {
return nil, s.updateAccountErr
}
@@ -421,6 +428,7 @@ func (s *stubAdminService) UpdateAccount(ctx context.Context, id int64, input *s
}
func (s *stubAdminService) UpdateAccountExtra(ctx context.Context, id int64, updates map[string]any) error {
s.updateAccountExtraCalls++
return nil
}
@@ -60,7 +60,7 @@ SELECT (extra->>'openai_long_context_billing_enabled')::boolean
FROM accounts
WHERE id = $1
`, ordinaryID).Scan(&ordinaryEnabled))
require.True(t, ordinaryEnabled)
require.False(t, ordinaryEnabled)
var shadowEnabled bool
require.NoError(t, tx.QueryRowContext(ctx, `
@@ -126,7 +126,7 @@ INSERT INTO accounts (name, platform, type, extra)
VALUES ('migration-175-rolling-writer', 'openai', 'oauth', '{}'::jsonb)
RETURNING (extra->>'openai_long_context_billing_enabled')::boolean
`).Scan(&ordinaryEnabled))
require.True(t, ordinaryEnabled)
require.False(t, ordinaryEnabled)
_, err = tx.ExecContext(ctx, "TRUNCATE scheduler_outbox")
require.NoError(t, err)
+2 -6
View File
@@ -1195,14 +1195,10 @@ func (a *Account) IsOpenAI() bool {
}
func (a *Account) IsOpenAILongContextBillingEnabled() bool {
if a == nil || !a.IsOpenAI() {
if a == nil || !a.IsOpenAI() || a.Extra == nil {
return false
}
raw, exists := a.Extra[openAILongContextBillingEnabledKey]
if !exists {
return true
}
enabled, ok := raw.(bool)
enabled, ok := a.Extra[openAILongContextBillingEnabledKey].(bool)
return ok && enabled
}
@@ -19,8 +19,8 @@ func TestAccountIsOpenAILongContextBillingEnabled(t *testing.T) {
}{
{name: "nil account is disabled", account: nil, want: false},
{name: "non OpenAI account is disabled", account: &Account{Platform: PlatformGrok}, want: false},
{name: "missing extra defaults enabled", account: &Account{Platform: PlatformOpenAI}, want: true},
{name: "missing key defaults enabled", account: &Account{Platform: PlatformOpenAI, Extra: map[string]any{}}, want: true},
{name: "missing extra defaults disabled", account: &Account{Platform: PlatformOpenAI}, want: false},
{name: "missing key defaults disabled", account: &Account{Platform: PlatformOpenAI, Extra: map[string]any{}}, want: false},
{name: "explicit true is enabled", account: &Account{Platform: PlatformOpenAI, Extra: map[string]any{"openai_long_context_billing_enabled": true}}, want: true},
{name: "explicit false is disabled", account: &Account{Platform: PlatformOpenAI, Extra: map[string]any{"openai_long_context_billing_enabled": false}}, want: false},
{name: "malformed value is disabled", account: &Account{Platform: PlatformOpenAI, Extra: map[string]any{"openai_long_context_billing_enabled": "false"}}, want: false},
@@ -34,11 +34,11 @@ func TestAccountIsOpenAILongContextBillingEnabled(t *testing.T) {
}
func TestNormalizeOpenAILongContextBillingExtra(t *testing.T) {
t.Run("OpenAI missing key persists enabled default", func(t *testing.T) {
t.Run("OpenAI missing key persists disabled default", func(t *testing.T) {
extra, err := normalizeOpenAILongContextBillingExtra(PlatformOpenAI, nil)
require.NoError(t, err)
require.Equal(t, true, extra["openai_long_context_billing_enabled"])
require.Equal(t, false, extra["openai_long_context_billing_enabled"])
})
t.Run("OpenAI explicit false is preserved", func(t *testing.T) {
@@ -116,7 +116,7 @@ func (r *longContextBillingRepoStub) BulkUpdate(_ context.Context, _ []int64, _
return 1, nil
}
func TestAdminServiceCreateAccountDefaultsOpenAILongContextBillingEnabled(t *testing.T) {
func TestAdminServiceCreateAccountDefaultsOpenAILongContextBillingDisabled(t *testing.T) {
repo := &longContextBillingRepoStub{}
svc := &adminServiceImpl{accountRepo: repo}
@@ -130,7 +130,7 @@ func TestAdminServiceCreateAccountDefaultsOpenAILongContextBillingEnabled(t *tes
require.NoError(t, err)
require.Same(t, account, repo.createdAccount)
require.Equal(t, true, account.Extra[openAILongContextBillingEnabledKey])
require.Equal(t, false, account.Extra[openAILongContextBillingEnabledKey])
}
func TestAdminServiceCreateAccountRejectsMalformedOpenAILongContextBillingValue(t *testing.T) {
@@ -162,7 +162,7 @@ func TestAdminServiceUpdateAccountPreservesOpenAILongContextBillingOptOutWhenOmi
require.Equal(t, false, account.Extra[openAILongContextBillingEnabledKey])
}
func TestAdminServiceUpdateAccountPreservesCodexImportOptOutWhenIncomingDefaultsTrue(t *testing.T) {
func TestAdminServiceUpdateAccountAllowsExplicitCodexImportOptIn(t *testing.T) {
repo := &longContextBillingRepoStub{account: &Account{
ID: 1,
Platform: PlatformOpenAI,
@@ -184,7 +184,7 @@ func TestAdminServiceUpdateAccountPreservesCodexImportOptOutWhenIncomingDefaults
})
require.NoError(t, err)
require.Equal(t, false, account.Extra[openAILongContextBillingEnabledKey])
require.Equal(t, true, account.Extra[openAILongContextBillingEnabledKey])
}
func TestAdminServiceUpdateAccountAllowsExplicitOptInOutsideCodexImport(t *testing.T) {
+1 -16
View File
@@ -101,7 +101,7 @@ func normalizeOpenAILongContextBillingExtra(platform string, extra map[string]an
}
_, exists := normalized[openAILongContextBillingEnabledKey]
if !exists {
normalized[openAILongContextBillingEnabledKey] = true
normalized[openAILongContextBillingEnabledKey] = false
}
return normalized, nil
}
@@ -118,21 +118,6 @@ func normalizeOpenAILongContextBillingUpdateExtra(account *Account, input *Updat
if hasCurrent {
normalized[openAILongContextBillingEnabledKey] = current
}
return normalized, nil
}
incoming, ok := normalized[openAILongContextBillingEnabledKey].(bool)
if !ok {
return nil, infraerrors.BadRequest(
"OPENAI_LONG_CONTEXT_BILLING_INVALID",
"openai_long_context_billing_enabled must be a boolean",
)
}
importSource, _ := input.Extra["import_source"].(string)
accessToken, _ := input.Credentials["access_token"].(string)
isCodexSessionImport := importSource == "codex_session" && strings.TrimSpace(accessToken) != ""
if hasCurrent && !current && incoming && isCodexSessionImport {
normalized[openAILongContextBillingEnabledKey] = false
}
return normalized, nil
}
@@ -163,7 +163,7 @@ func TestCreateShadowInheritsParentEffectiveOpenAILongContextBillingValue(t *tes
parentExtra map[string]any
want bool
}{
{name: "missing parent value defaults enabled", want: true},
{name: "missing parent value defaults disabled", want: false},
{name: "explicit parent opt-out is inherited", parentExtra: map[string]any{openAILongContextBillingEnabledKey: false}, want: false},
}
+1 -1
View File
@@ -1283,7 +1283,7 @@ func (s *BillingService) CalculateCostWithLongContext(model string, tokens Usage
CacheReadCost: inRangeCost.CacheReadCost + outRangeCost.CacheReadCost,
TotalCost: inRangeCost.TotalCost + outRangeCost.TotalCost,
ActualCost: inRangeCost.ActualCost + outRangeCost.ActualCost,
LongContextBillingApplied: true,
LongContextBillingApplied: outRangeCost.ActualCost > 0,
}, nil
}
@@ -848,6 +848,17 @@ func TestCalculateCostWithLongContext_AboveThreshold_CacheBelowThreshold(t *test
require.True(t, cost.ActualCost > normalCost.ActualCost, "长上下文费用应高于正常费用")
}
func TestCalculateCostWithLongContext_MarkerRequiresActualCostIncrease(t *testing.T) {
svc := newTestBillingService()
tokens := UsageTokens{InputTokens: 300000}
cost, err := svc.CalculateCostWithLongContext("claude-sonnet-4", tokens, 0, 200000, 2.0)
require.NoError(t, err)
require.Zero(t, cost.ActualCost)
require.False(t, cost.LongContextBillingApplied)
}
func TestCalculateCostWithLongContext_DisabledThreshold(t *testing.T) {
svc := newTestBillingService()
@@ -72,18 +72,24 @@ func TestCRSSyncOpenAILongContextBilling(t *testing.T) {
wantAction string
wantEnabled bool
}{
{name: "OAuth create defaults missing value enabled", collection: "openaiOAuthAccounts", credentials: map[string]any{"access_token": "oauth-token"}, wantAction: "created", wantEnabled: true},
{name: "OAuth create defaults missing value disabled", collection: "openaiOAuthAccounts", credentials: map[string]any{"access_token": "oauth-token"}, wantAction: "created"},
{name: "OAuth create preserves source true", collection: "openaiOAuthAccounts", credentials: map[string]any{"access_token": "oauth-token"}, sourceExtra: map[string]any{openAILongContextBillingEnabledKey: true}, wantAction: "created", wantEnabled: true},
{name: "OAuth create preserves source false", collection: "openaiOAuthAccounts", credentials: map[string]any{"access_token": "oauth-token"}, sourceExtra: map[string]any{openAILongContextBillingEnabledKey: false}, wantAction: "created"},
{name: "OAuth update defaults missing value enabled", collection: "openaiOAuthAccounts", credentials: map[string]any{"access_token": "oauth-token"}, existingExtra: map[string]any{"existing": true}, wantAction: "updated", wantEnabled: true},
{name: "OAuth update defaults missing value disabled", collection: "openaiOAuthAccounts", credentials: map[string]any{"access_token": "oauth-token"}, existingExtra: map[string]any{"existing": true}, wantAction: "updated"},
{name: "OAuth update preserves existing true when source omits value", collection: "openaiOAuthAccounts", credentials: map[string]any{"access_token": "oauth-token"}, existingExtra: map[string]any{openAILongContextBillingEnabledKey: true}, wantAction: "updated", wantEnabled: true},
{name: "OAuth update preserves existing false when source omits value", collection: "openaiOAuthAccounts", credentials: map[string]any{"access_token": "oauth-token"}, existingExtra: map[string]any{openAILongContextBillingEnabledKey: false}, wantAction: "updated"},
{name: "OAuth update preserves source true over existing false", collection: "openaiOAuthAccounts", credentials: map[string]any{"access_token": "oauth-token"}, sourceExtra: map[string]any{openAILongContextBillingEnabledKey: true}, existingExtra: map[string]any{openAILongContextBillingEnabledKey: false}, wantAction: "updated", wantEnabled: true},
{name: "OAuth update preserves source false over existing true", collection: "openaiOAuthAccounts", credentials: map[string]any{"access_token": "oauth-token"}, sourceExtra: map[string]any{openAILongContextBillingEnabledKey: false}, existingExtra: map[string]any{openAILongContextBillingEnabledKey: true}, wantAction: "updated"},
{name: "OAuth rejects malformed source value", collection: "openaiOAuthAccounts", credentials: map[string]any{"access_token": "oauth-token"}, sourceExtra: map[string]any{openAILongContextBillingEnabledKey: "false"}, wantAction: "failed"},
{name: "OAuth rejects malformed existing value", collection: "openaiOAuthAccounts", credentials: map[string]any{"access_token": "oauth-token"}, existingExtra: map[string]any{openAILongContextBillingEnabledKey: "false"}, wantAction: "failed"},
{name: "OAuth update rejects malformed source value", collection: "openaiOAuthAccounts", credentials: map[string]any{"access_token": "oauth-token"}, sourceExtra: map[string]any{openAILongContextBillingEnabledKey: "false"}, existingExtra: map[string]any{openAILongContextBillingEnabledKey: true}, wantAction: "failed"},
{name: "API key create defaults missing value enabled", collection: "openaiResponsesAccounts", credentials: map[string]any{"api_key": "sk-test"}, wantAction: "created", wantEnabled: true},
{name: "API key create defaults missing value disabled", collection: "openaiResponsesAccounts", credentials: map[string]any{"api_key": "sk-test"}, wantAction: "created"},
{name: "API key create preserves source true", collection: "openaiResponsesAccounts", credentials: map[string]any{"api_key": "sk-test"}, sourceExtra: map[string]any{openAILongContextBillingEnabledKey: true}, wantAction: "created", wantEnabled: true},
{name: "API key create preserves source false", collection: "openaiResponsesAccounts", credentials: map[string]any{"api_key": "sk-test"}, sourceExtra: map[string]any{openAILongContextBillingEnabledKey: false}, wantAction: "created"},
{name: "API key update defaults missing value enabled", collection: "openaiResponsesAccounts", credentials: map[string]any{"api_key": "sk-test"}, existingExtra: map[string]any{"existing": true}, wantAction: "updated", wantEnabled: true},
{name: "API key update defaults missing value disabled", collection: "openaiResponsesAccounts", credentials: map[string]any{"api_key": "sk-test"}, existingExtra: map[string]any{"existing": true}, wantAction: "updated"},
{name: "API key update preserves existing true when source omits value", collection: "openaiResponsesAccounts", credentials: map[string]any{"api_key": "sk-test"}, existingExtra: map[string]any{openAILongContextBillingEnabledKey: true}, wantAction: "updated", wantEnabled: true},
{name: "API key update preserves existing false when source omits value", collection: "openaiResponsesAccounts", credentials: map[string]any{"api_key": "sk-test"}, existingExtra: map[string]any{openAILongContextBillingEnabledKey: false}, wantAction: "updated"},
{name: "API key update preserves source true over existing false", collection: "openaiResponsesAccounts", credentials: map[string]any{"api_key": "sk-test"}, sourceExtra: map[string]any{openAILongContextBillingEnabledKey: true}, existingExtra: map[string]any{openAILongContextBillingEnabledKey: false}, wantAction: "updated", wantEnabled: true},
{name: "API key update preserves source false over existing true", collection: "openaiResponsesAccounts", credentials: map[string]any{"api_key": "sk-test"}, sourceExtra: map[string]any{openAILongContextBillingEnabledKey: false}, existingExtra: map[string]any{openAILongContextBillingEnabledKey: true}, wantAction: "updated"},
{name: "API key rejects malformed source value", collection: "openaiResponsesAccounts", credentials: map[string]any{"api_key": "sk-test"}, sourceExtra: map[string]any{openAILongContextBillingEnabledKey: "false"}, wantAction: "failed"},
{name: "API key rejects malformed existing value", collection: "openaiResponsesAccounts", credentials: map[string]any{"api_key": "sk-test"}, existingExtra: map[string]any{openAILongContextBillingEnabledKey: "false"}, wantAction: "failed"},
@@ -1056,7 +1056,7 @@ func TestOpenAIGatewayServiceRecordUsage_GPT56SeparatesCacheWriteForBillingAndSt
require.InDelta(t, usageRepo.lastLog.TotalCost*1.1, usageRepo.lastLog.ActualCost, 1e-12)
}
func TestOpenAIGatewayServiceRecordUsage_Gpt54LongContextBillingEnabledByDefault(t *testing.T) {
func TestOpenAIGatewayServiceRecordUsage_Gpt54LongContextBillingDisabledByDefault(t *testing.T) {
usageRepo := &openAIRecordUsageLogRepoStub{inserted: true}
userRepo := &openAIRecordUsageUserRepoStub{}
subRepo := &openAIRecordUsageSubRepoStub{}
@@ -1080,17 +1080,17 @@ func TestOpenAIGatewayServiceRecordUsage_Gpt54LongContextBillingEnabledByDefault
require.NoError(t, err)
require.NotNil(t, usageRepo.lastLog)
expectedInput := 300000 * 2.5e-6 * 2.0
expectedOutput := 2000 * 15e-6 * 1.5
expectedInput := 300000 * 2.5e-6
expectedOutput := 2000 * 15e-6
require.InDelta(t, expectedInput, usageRepo.lastLog.InputCost, 1e-10)
require.InDelta(t, expectedOutput, usageRepo.lastLog.OutputCost, 1e-10)
require.InDelta(t, expectedInput+expectedOutput, usageRepo.lastLog.TotalCost, 1e-10)
require.InDelta(t, (expectedInput+expectedOutput)*1.1, usageRepo.lastLog.ActualCost, 1e-10)
require.True(t, usageRepo.lastLog.LongContextBillingApplied)
require.False(t, usageRepo.lastLog.LongContextBillingApplied)
require.Equal(t, 1, userRepo.deductCalls)
}
func TestOpenAIGatewayServiceRecordUsage_Gpt54LongContextBillingDisabledPerAccount(t *testing.T) {
func TestOpenAIGatewayServiceRecordUsage_Gpt54LongContextBillingEnabledPerAccount(t *testing.T) {
usageRepo := &openAIRecordUsageLogRepoStub{inserted: true}
userRepo := &openAIRecordUsageUserRepoStub{}
subRepo := &openAIRecordUsageSubRepoStub{}
@@ -1111,20 +1111,20 @@ func TestOpenAIGatewayServiceRecordUsage_Gpt54LongContextBillingDisabledPerAccou
Account: &Account{
ID: 3015,
Platform: PlatformOpenAI,
Extra: map[string]any{"openai_long_context_billing_enabled": false},
Extra: map[string]any{"openai_long_context_billing_enabled": true},
},
})
require.NoError(t, err)
require.NotNil(t, usageRepo.lastLog)
expectedInput := 300000 * 2.5e-6
expectedOutput := 2000 * 15e-6
expectedInput := 300000 * 2.5e-6 * 2.0
expectedOutput := 2000 * 15e-6 * 1.5
require.InDelta(t, expectedInput, usageRepo.lastLog.InputCost, 1e-10)
require.InDelta(t, expectedOutput, usageRepo.lastLog.OutputCost, 1e-10)
require.InDelta(t, expectedInput+expectedOutput, usageRepo.lastLog.TotalCost, 1e-10)
require.InDelta(t, (expectedInput+expectedOutput)*1.1, usageRepo.lastLog.ActualCost, 1e-10)
require.False(t, usageRepo.lastLog.LongContextBillingApplied)
require.True(t, usageRepo.lastLog.LongContextBillingApplied)
}
func TestOpenAIGatewayServiceRecordUsage_SparkShadowUsesCurrentParentBillingSetting(t *testing.T) {
@@ -14,7 +14,7 @@ BEGIN
IF NEW.parent_account_id IS NOT NULL AND NEW.quota_dimension = 'spark' THEN
SELECT CASE
WHEN parent.platform IS DISTINCT FROM 'openai' THEN 'false'::jsonb
WHEN NOT (COALESCE(parent.extra, '{}'::jsonb) ? 'openai_long_context_billing_enabled') THEN 'true'::jsonb
WHEN NOT (COALESCE(parent.extra, '{}'::jsonb) ? 'openai_long_context_billing_enabled') THEN 'false'::jsonb
WHEN jsonb_typeof(parent.extra->'openai_long_context_billing_enabled') = 'boolean'
THEN parent.extra->'openai_long_context_billing_enabled'
ELSE 'false'::jsonb
@@ -43,7 +43,7 @@ BEGIN
NEW.extra := jsonb_set(
NEW.extra,
'{openai_long_context_billing_enabled}',
'true'::jsonb,
'false'::jsonb,
true
);
END IF;
@@ -121,7 +121,7 @@ UPDATE accounts
SET extra = jsonb_set(
COALESCE(extra, '{}'::jsonb),
'{openai_long_context_billing_enabled}',
'true'::jsonb,
'false'::jsonb,
true
)
WHERE platform = 'openai'
@@ -133,7 +133,7 @@ WITH shadow_values AS (
shadow.id,
CASE
WHEN parent.platform IS DISTINCT FROM 'openai' THEN 'false'::jsonb
WHEN NOT (COALESCE(parent.extra, '{}'::jsonb) ? 'openai_long_context_billing_enabled') THEN 'true'::jsonb
WHEN NOT (COALESCE(parent.extra, '{}'::jsonb) ? 'openai_long_context_billing_enabled') THEN 'false'::jsonb
WHEN jsonb_typeof(parent.extra->'openai_long_context_billing_enabled') = 'boolean'
THEN parent.extra->'openai_long_context_billing_enabled'
ELSE 'false'::jsonb