diff --git a/backend/cmd/server/wire_gen.go b/backend/cmd/server/wire_gen.go
index e4ef733b32..ba5b78a72f 100644
--- a/backend/cmd/server/wire_gen.go
+++ b/backend/cmd/server/wire_gen.go
@@ -190,9 +190,10 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) {
claudeUsageFetcher := repository.NewClaudeUsageFetcher(httpUpstream)
antigravityQuotaFetcher := service.NewAntigravityQuotaFetcher(proxyRepository)
grokQuotaFetcher := service.NewGrokQuotaFetcher()
+ grokQuotaService := service.ProvideGrokQuotaService(accountRepository, proxyRepository, grokTokenProvider, httpUpstream, usageLogRepository)
openAIQuotaService := service.ProvideOpenAIQuotaService(accountRepository, proxyRepository, openAITokenProvider, privacyClientFactory)
usageCache := service.NewUsageCache()
- accountUsageService := service.NewAccountUsageService(accountRepository, usageLogRepository, claudeUsageFetcher, geminiQuotaService, antigravityQuotaFetcher, grokQuotaFetcher, openAIQuotaService, usageCache, identityCache, tlsFingerprintProfileService)
+ accountUsageService := service.NewAccountUsageService(accountRepository, usageLogRepository, claudeUsageFetcher, geminiQuotaService, antigravityQuotaFetcher, grokQuotaFetcher, grokQuotaService, openAIQuotaService, usageCache, identityCache, tlsFingerprintProfileService)
accountTestService := service.NewAccountTestService(accountRepository, geminiTokenProvider, claudeTokenProvider, grokTokenProvider, antigravityGatewayService, httpUpstream, configConfig, tlsFingerprintProfileService)
crsSyncService := service.NewCRSSyncService(accountRepository, proxyRepository, oAuthService, openAIOAuthService, geminiOAuthService, configConfig)
accountHandler := admin.NewAccountHandler(adminService, oAuthService, openAIOAuthService, geminiOAuthService, antigravityOAuthService, rateLimitService, accountUsageService, accountTestService, concurrencyService, crsSyncService, sessionLimitCache, rpmCache, compositeTokenCacheInvalidator)
@@ -207,7 +208,6 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) {
openAIOAuthHandler := admin.NewOpenAIOAuthHandler(openAIOAuthService, adminService, openAIQuotaService)
geminiOAuthHandler := admin.NewGeminiOAuthHandler(geminiOAuthService)
antigravityOAuthHandler := admin.NewAntigravityOAuthHandler(antigravityOAuthService)
- grokQuotaService := service.ProvideGrokQuotaService(accountRepository, proxyRepository, grokTokenProvider, httpUpstream)
grokOAuthHandler := admin.NewGrokOAuthHandler(grokOAuthService, adminService, grokQuotaService)
proxyHandler := admin.NewProxyHandler(adminService)
adminRedeemHandler := admin.NewRedeemHandler(adminService, redeemService)
diff --git a/backend/ent/migrate/schema.go b/backend/ent/migrate/schema.go
index 3441afec04..cac559d535 100644
--- a/backend/ent/migrate/schema.go
+++ b/backend/ent/migrate/schema.go
@@ -1560,6 +1560,7 @@ var (
{Name: "total_cost", Type: field.TypeFloat64, Default: 0, SchemaType: map[string]string{"postgres": "decimal(20,10)"}},
{Name: "actual_cost", Type: field.TypeFloat64, Default: 0, SchemaType: map[string]string{"postgres": "decimal(20,10)"}},
{Name: "rate_multiplier", Type: field.TypeFloat64, Default: 1, SchemaType: map[string]string{"postgres": "decimal(10,4)"}},
+ {Name: "long_context_billing_applied", Type: field.TypeBool, Default: false},
{Name: "account_rate_multiplier", Type: field.TypeFloat64, Nullable: true, SchemaType: map[string]string{"postgres": "decimal(10,4)"}},
{Name: "billing_type", Type: field.TypeInt8, Default: 0},
{Name: "stream", Type: field.TypeBool, Default: false},
@@ -1592,31 +1593,31 @@ var (
ForeignKeys: []*schema.ForeignKey{
{
Symbol: "usage_logs_api_keys_usage_logs",
- Columns: []*schema.Column{UsageLogsColumns[40]},
+ Columns: []*schema.Column{UsageLogsColumns[41]},
RefColumns: []*schema.Column{APIKeysColumns[0]},
OnDelete: schema.NoAction,
},
{
Symbol: "usage_logs_accounts_usage_logs",
- Columns: []*schema.Column{UsageLogsColumns[41]},
+ Columns: []*schema.Column{UsageLogsColumns[42]},
RefColumns: []*schema.Column{AccountsColumns[0]},
OnDelete: schema.NoAction,
},
{
Symbol: "usage_logs_groups_usage_logs",
- Columns: []*schema.Column{UsageLogsColumns[42]},
+ Columns: []*schema.Column{UsageLogsColumns[43]},
RefColumns: []*schema.Column{GroupsColumns[0]},
OnDelete: schema.SetNull,
},
{
Symbol: "usage_logs_users_usage_logs",
- Columns: []*schema.Column{UsageLogsColumns[43]},
+ Columns: []*schema.Column{UsageLogsColumns[44]},
RefColumns: []*schema.Column{UsersColumns[0]},
OnDelete: schema.NoAction,
},
{
Symbol: "usage_logs_user_subscriptions_usage_logs",
- Columns: []*schema.Column{UsageLogsColumns[44]},
+ Columns: []*schema.Column{UsageLogsColumns[45]},
RefColumns: []*schema.Column{UserSubscriptionsColumns[0]},
OnDelete: schema.SetNull,
},
@@ -1625,32 +1626,32 @@ var (
{
Name: "usagelog_user_id",
Unique: false,
- Columns: []*schema.Column{UsageLogsColumns[43]},
+ Columns: []*schema.Column{UsageLogsColumns[44]},
},
{
Name: "usagelog_api_key_id",
Unique: false,
- Columns: []*schema.Column{UsageLogsColumns[40]},
+ Columns: []*schema.Column{UsageLogsColumns[41]},
},
{
Name: "usagelog_account_id",
Unique: false,
- Columns: []*schema.Column{UsageLogsColumns[41]},
+ Columns: []*schema.Column{UsageLogsColumns[42]},
},
{
Name: "usagelog_group_id",
Unique: false,
- Columns: []*schema.Column{UsageLogsColumns[42]},
+ Columns: []*schema.Column{UsageLogsColumns[43]},
},
{
Name: "usagelog_subscription_id",
Unique: false,
- Columns: []*schema.Column{UsageLogsColumns[44]},
+ Columns: []*schema.Column{UsageLogsColumns[45]},
},
{
Name: "usagelog_created_at",
Unique: false,
- Columns: []*schema.Column{UsageLogsColumns[39]},
+ Columns: []*schema.Column{UsageLogsColumns[40]},
},
{
Name: "usagelog_model",
@@ -1670,17 +1671,17 @@ var (
{
Name: "usagelog_user_id_created_at",
Unique: false,
- Columns: []*schema.Column{UsageLogsColumns[43], UsageLogsColumns[39]},
+ Columns: []*schema.Column{UsageLogsColumns[44], UsageLogsColumns[40]},
},
{
Name: "usagelog_api_key_id_created_at",
Unique: false,
- Columns: []*schema.Column{UsageLogsColumns[40], UsageLogsColumns[39]},
+ Columns: []*schema.Column{UsageLogsColumns[41], UsageLogsColumns[40]},
},
{
Name: "usagelog_group_id_created_at",
Unique: false,
- Columns: []*schema.Column{UsageLogsColumns[42], UsageLogsColumns[39]},
+ Columns: []*schema.Column{UsageLogsColumns[43], UsageLogsColumns[40]},
},
},
}
diff --git a/backend/ent/mutation.go b/backend/ent/mutation.go
index ab7c424a47..fb35531878 100644
--- a/backend/ent/mutation.go
+++ b/backend/ent/mutation.go
@@ -41763,83 +41763,84 @@ func (m *UsageCleanupTaskMutation) ResetEdge(name string) error {
// UsageLogMutation represents an operation that mutates the UsageLog nodes in the graph.
type UsageLogMutation struct {
config
- op Op
- typ string
- id *int64
- request_id *string
- model *string
- requested_model *string
- upstream_model *string
- channel_id *int64
- addchannel_id *int64
- model_mapping_chain *string
- billing_tier *string
- billing_mode *string
- input_tokens *int
- addinput_tokens *int
- output_tokens *int
- addoutput_tokens *int
- cache_creation_tokens *int
- addcache_creation_tokens *int
- cache_read_tokens *int
- addcache_read_tokens *int
- cache_creation_5m_tokens *int
- addcache_creation_5m_tokens *int
- cache_creation_1h_tokens *int
- addcache_creation_1h_tokens *int
- input_cost *float64
- addinput_cost *float64
- output_cost *float64
- addoutput_cost *float64
- cache_creation_cost *float64
- addcache_creation_cost *float64
- cache_read_cost *float64
- addcache_read_cost *float64
- total_cost *float64
- addtotal_cost *float64
- actual_cost *float64
- addactual_cost *float64
- rate_multiplier *float64
- addrate_multiplier *float64
- account_rate_multiplier *float64
- addaccount_rate_multiplier *float64
- billing_type *int8
- addbilling_type *int8
- stream *bool
- duration_ms *int
- addduration_ms *int
- first_token_ms *int
- addfirst_token_ms *int
- user_agent *string
- ip_address *string
- image_count *int
- addimage_count *int
- image_size *string
- image_input_size *string
- image_output_size *string
- image_size_source *string
- image_size_breakdown *map[string]int
- video_count *int
- addvideo_count *int
- video_resolution *string
- video_duration_seconds *int
- addvideo_duration_seconds *int
- cache_ttl_overridden *bool
- created_at *time.Time
- clearedFields map[string]struct{}
- user *int64
- cleareduser bool
- api_key *int64
- clearedapi_key bool
- account *int64
- clearedaccount bool
- group *int64
- clearedgroup bool
- subscription *int64
- clearedsubscription bool
- done bool
- oldValue func(context.Context) (*UsageLog, error)
- predicates []predicate.UsageLog
+ op Op
+ typ string
+ id *int64
+ request_id *string
+ model *string
+ requested_model *string
+ upstream_model *string
+ channel_id *int64
+ addchannel_id *int64
+ model_mapping_chain *string
+ billing_tier *string
+ billing_mode *string
+ input_tokens *int
+ addinput_tokens *int
+ output_tokens *int
+ addoutput_tokens *int
+ cache_creation_tokens *int
+ addcache_creation_tokens *int
+ cache_read_tokens *int
+ addcache_read_tokens *int
+ cache_creation_5m_tokens *int
+ addcache_creation_5m_tokens *int
+ cache_creation_1h_tokens *int
+ addcache_creation_1h_tokens *int
+ input_cost *float64
+ addinput_cost *float64
+ output_cost *float64
+ addoutput_cost *float64
+ cache_creation_cost *float64
+ addcache_creation_cost *float64
+ cache_read_cost *float64
+ addcache_read_cost *float64
+ total_cost *float64
+ addtotal_cost *float64
+ actual_cost *float64
+ addactual_cost *float64
+ rate_multiplier *float64
+ addrate_multiplier *float64
+ long_context_billing_applied *bool
+ account_rate_multiplier *float64
+ addaccount_rate_multiplier *float64
+ billing_type *int8
+ addbilling_type *int8
+ stream *bool
+ duration_ms *int
+ addduration_ms *int
+ first_token_ms *int
+ addfirst_token_ms *int
+ user_agent *string
+ ip_address *string
+ image_count *int
+ addimage_count *int
+ image_size *string
+ image_input_size *string
+ image_output_size *string
+ image_size_source *string
+ image_size_breakdown *map[string]int
+ video_count *int
+ addvideo_count *int
+ video_resolution *string
+ video_duration_seconds *int
+ addvideo_duration_seconds *int
+ cache_ttl_overridden *bool
+ created_at *time.Time
+ clearedFields map[string]struct{}
+ user *int64
+ cleareduser bool
+ api_key *int64
+ clearedapi_key bool
+ account *int64
+ clearedaccount bool
+ group *int64
+ clearedgroup bool
+ subscription *int64
+ clearedsubscription bool
+ done bool
+ oldValue func(context.Context) (*UsageLog, error)
+ predicates []predicate.UsageLog
}
var _ ent.Mutation = (*UsageLogMutation)(nil)
@@ -43261,6 +43262,42 @@ func (m *UsageLogMutation) ResetRateMultiplier() {
m.addrate_multiplier = nil
}
+// SetLongContextBillingApplied sets the "long_context_billing_applied" field.
+func (m *UsageLogMutation) SetLongContextBillingApplied(b bool) {
+ m.long_context_billing_applied = &b
+}
+
+// LongContextBillingApplied returns the value of the "long_context_billing_applied" field in the mutation.
+func (m *UsageLogMutation) LongContextBillingApplied() (r bool, exists bool) {
+ v := m.long_context_billing_applied
+ if v == nil {
+ return
+ }
+ return *v, true
+}
+
+// OldLongContextBillingApplied returns the old "long_context_billing_applied" field's value of the UsageLog entity.
+// If the UsageLog object wasn't provided to the builder, the object is fetched from the database.
+// An error is returned if the mutation operation is not UpdateOne, or the database query fails.
+func (m *UsageLogMutation) OldLongContextBillingApplied(ctx context.Context) (v bool, err error) {
+ if !m.op.Is(OpUpdateOne) {
+ return v, errors.New("OldLongContextBillingApplied is only allowed on UpdateOne operations")
+ }
+ if m.id == nil || m.oldValue == nil {
+ return v, errors.New("OldLongContextBillingApplied requires an ID field in the mutation")
+ }
+ oldValue, err := m.oldValue(ctx)
+ if err != nil {
+ return v, fmt.Errorf("querying old value for OldLongContextBillingApplied: %w", err)
+ }
+ return oldValue.LongContextBillingApplied, nil
+}
+
+// ResetLongContextBillingApplied resets all changes to the "long_context_billing_applied" field.
+func (m *UsageLogMutation) ResetLongContextBillingApplied() {
+ m.long_context_billing_applied = nil
+}
+
// SetAccountRateMultiplier sets the "account_rate_multiplier" field.
func (m *UsageLogMutation) SetAccountRateMultiplier(f float64) {
m.account_rate_multiplier = &f
@@ -44378,7 +44415,7 @@ func (m *UsageLogMutation) Type() string {
// order to get all numeric fields that were incremented/decremented, call
// AddedFields().
func (m *UsageLogMutation) Fields() []string {
- fields := make([]string, 0, 44)
+ fields := make([]string, 0, 45)
if m.user != nil {
fields = append(fields, usagelog.FieldUserID)
}
@@ -44457,6 +44494,9 @@ func (m *UsageLogMutation) Fields() []string {
if m.rate_multiplier != nil {
fields = append(fields, usagelog.FieldRateMultiplier)
}
+ if m.long_context_billing_applied != nil {
+ fields = append(fields, usagelog.FieldLongContextBillingApplied)
+ }
if m.account_rate_multiplier != nil {
fields = append(fields, usagelog.FieldAccountRateMultiplier)
}
@@ -44571,6 +44611,8 @@ func (m *UsageLogMutation) Field(name string) (ent.Value, bool) {
return m.ActualCost()
case usagelog.FieldRateMultiplier:
return m.RateMultiplier()
+ case usagelog.FieldLongContextBillingApplied:
+ return m.LongContextBillingApplied()
case usagelog.FieldAccountRateMultiplier:
return m.AccountRateMultiplier()
case usagelog.FieldBillingType:
@@ -44668,6 +44710,8 @@ func (m *UsageLogMutation) OldField(ctx context.Context, name string) (ent.Value
return m.OldActualCost(ctx)
case usagelog.FieldRateMultiplier:
return m.OldRateMultiplier(ctx)
+ case usagelog.FieldLongContextBillingApplied:
+ return m.OldLongContextBillingApplied(ctx)
case usagelog.FieldAccountRateMultiplier:
return m.OldAccountRateMultiplier(ctx)
case usagelog.FieldBillingType:
@@ -44895,6 +44939,13 @@ func (m *UsageLogMutation) SetField(name string, value ent.Value) error {
}
m.SetRateMultiplier(v)
return nil
+ case usagelog.FieldLongContextBillingApplied:
+ v, ok := value.(bool)
+ if !ok {
+ return fmt.Errorf("unexpected type %T for field %s", value, name)
+ }
+ m.SetLongContextBillingApplied(v)
+ return nil
case usagelog.FieldAccountRateMultiplier:
v, ok := value.(float64)
if !ok {
@@ -45526,6 +45577,9 @@ func (m *UsageLogMutation) ResetField(name string) error {
case usagelog.FieldRateMultiplier:
m.ResetRateMultiplier()
return nil
+ case usagelog.FieldLongContextBillingApplied:
+ m.ResetLongContextBillingApplied()
+ return nil
case usagelog.FieldAccountRateMultiplier:
m.ResetAccountRateMultiplier()
return nil
diff --git a/backend/ent/runtime/runtime.go b/backend/ent/runtime/runtime.go
index 4cb3f800f8..867f1cbdbd 100644
--- a/backend/ent/runtime/runtime.go
+++ b/backend/ent/runtime/runtime.go
@@ -1940,56 +1940,60 @@ func init() {
usagelogDescRateMultiplier := usagelogFields[25].Descriptor()
// usagelog.DefaultRateMultiplier holds the default value on creation for the rate_multiplier field.
usagelog.DefaultRateMultiplier = usagelogDescRateMultiplier.Default.(float64)
+ // usagelogDescLongContextBillingApplied is the schema descriptor for long_context_billing_applied field.
+ usagelogDescLongContextBillingApplied := usagelogFields[26].Descriptor()
+ // usagelog.DefaultLongContextBillingApplied holds the default value on creation for the long_context_billing_applied field.
+ usagelog.DefaultLongContextBillingApplied = usagelogDescLongContextBillingApplied.Default.(bool)
// usagelogDescBillingType is the schema descriptor for billing_type field.
- usagelogDescBillingType := usagelogFields[27].Descriptor()
+ usagelogDescBillingType := usagelogFields[28].Descriptor()
// usagelog.DefaultBillingType holds the default value on creation for the billing_type field.
usagelog.DefaultBillingType = usagelogDescBillingType.Default.(int8)
// usagelogDescStream is the schema descriptor for stream field.
- usagelogDescStream := usagelogFields[28].Descriptor()
+ usagelogDescStream := usagelogFields[29].Descriptor()
// usagelog.DefaultStream holds the default value on creation for the stream field.
usagelog.DefaultStream = usagelogDescStream.Default.(bool)
// usagelogDescUserAgent is the schema descriptor for user_agent field.
- usagelogDescUserAgent := usagelogFields[31].Descriptor()
+ usagelogDescUserAgent := usagelogFields[32].Descriptor()
// usagelog.UserAgentValidator is a validator for the "user_agent" field. It is called by the builders before save.
usagelog.UserAgentValidator = usagelogDescUserAgent.Validators[0].(func(string) error)
// usagelogDescIPAddress is the schema descriptor for ip_address field.
- usagelogDescIPAddress := usagelogFields[32].Descriptor()
+ usagelogDescIPAddress := usagelogFields[33].Descriptor()
// usagelog.IPAddressValidator is a validator for the "ip_address" field. It is called by the builders before save.
usagelog.IPAddressValidator = usagelogDescIPAddress.Validators[0].(func(string) error)
// usagelogDescImageCount is the schema descriptor for image_count field.
- usagelogDescImageCount := usagelogFields[33].Descriptor()
+ usagelogDescImageCount := usagelogFields[34].Descriptor()
// usagelog.DefaultImageCount holds the default value on creation for the image_count field.
usagelog.DefaultImageCount = usagelogDescImageCount.Default.(int)
// usagelogDescImageSize is the schema descriptor for image_size field.
- usagelogDescImageSize := usagelogFields[34].Descriptor()
+ usagelogDescImageSize := usagelogFields[35].Descriptor()
// usagelog.ImageSizeValidator is a validator for the "image_size" field. It is called by the builders before save.
usagelog.ImageSizeValidator = usagelogDescImageSize.Validators[0].(func(string) error)
// usagelogDescImageInputSize is the schema descriptor for image_input_size field.
- usagelogDescImageInputSize := usagelogFields[35].Descriptor()
+ usagelogDescImageInputSize := usagelogFields[36].Descriptor()
// usagelog.ImageInputSizeValidator is a validator for the "image_input_size" field. It is called by the builders before save.
usagelog.ImageInputSizeValidator = usagelogDescImageInputSize.Validators[0].(func(string) error)
// usagelogDescImageOutputSize is the schema descriptor for image_output_size field.
- usagelogDescImageOutputSize := usagelogFields[36].Descriptor()
+ usagelogDescImageOutputSize := usagelogFields[37].Descriptor()
// usagelog.ImageOutputSizeValidator is a validator for the "image_output_size" field. It is called by the builders before save.
usagelog.ImageOutputSizeValidator = usagelogDescImageOutputSize.Validators[0].(func(string) error)
// usagelogDescImageSizeSource is the schema descriptor for image_size_source field.
- usagelogDescImageSizeSource := usagelogFields[37].Descriptor()
+ usagelogDescImageSizeSource := usagelogFields[38].Descriptor()
// usagelog.ImageSizeSourceValidator is a validator for the "image_size_source" field. It is called by the builders before save.
usagelog.ImageSizeSourceValidator = usagelogDescImageSizeSource.Validators[0].(func(string) error)
// usagelogDescVideoCount is the schema descriptor for video_count field.
- usagelogDescVideoCount := usagelogFields[39].Descriptor()
+ usagelogDescVideoCount := usagelogFields[40].Descriptor()
// usagelog.DefaultVideoCount holds the default value on creation for the video_count field.
usagelog.DefaultVideoCount = usagelogDescVideoCount.Default.(int)
// usagelogDescVideoResolution is the schema descriptor for video_resolution field.
- usagelogDescVideoResolution := usagelogFields[40].Descriptor()
+ usagelogDescVideoResolution := usagelogFields[41].Descriptor()
// usagelog.VideoResolutionValidator is a validator for the "video_resolution" field. It is called by the builders before save.
usagelog.VideoResolutionValidator = usagelogDescVideoResolution.Validators[0].(func(string) error)
// usagelogDescCacheTTLOverridden is the schema descriptor for cache_ttl_overridden field.
- usagelogDescCacheTTLOverridden := usagelogFields[42].Descriptor()
+ usagelogDescCacheTTLOverridden := usagelogFields[43].Descriptor()
// usagelog.DefaultCacheTTLOverridden holds the default value on creation for the cache_ttl_overridden field.
usagelog.DefaultCacheTTLOverridden = usagelogDescCacheTTLOverridden.Default.(bool)
// usagelogDescCreatedAt is the schema descriptor for created_at field.
- usagelogDescCreatedAt := usagelogFields[43].Descriptor()
+ usagelogDescCreatedAt := usagelogFields[44].Descriptor()
// usagelog.DefaultCreatedAt holds the default value on creation for the created_at field.
usagelog.DefaultCreatedAt = usagelogDescCreatedAt.Default.(func() time.Time)
userMixin := schema.User{}.Mixin()
diff --git a/backend/ent/schema/usage_log.go b/backend/ent/schema/usage_log.go
index e84cc1c140..6d8c2d4191 100644
--- a/backend/ent/schema/usage_log.go
+++ b/backend/ent/schema/usage_log.go
@@ -100,6 +100,9 @@ func (UsageLog) Fields() []ent.Field {
field.Float("rate_multiplier").
Default(1).
SchemaType(map[string]string{dialect.Postgres: "decimal(10,4)"}),
+ field.Bool("long_context_billing_applied").
+ Default(false).
+ Comment("Whether long-context pricing changed token prices for this request"),
// account_rate_multiplier: 账号计费倍率快照(NULL 表示按 1.0 处理)
field.Float("account_rate_multiplier").
diff --git a/backend/ent/usagelog.go b/backend/ent/usagelog.go
index 4d374a8495..b13e29b2f7 100644
--- a/backend/ent/usagelog.go
+++ b/backend/ent/usagelog.go
@@ -75,6 +75,8 @@ type UsageLog struct {
ActualCost float64 `json:"actual_cost,omitempty"`
// RateMultiplier holds the value of the "rate_multiplier" field.
RateMultiplier float64 `json:"rate_multiplier,omitempty"`
+ // Whether long-context pricing changed token prices for this request
+ LongContextBillingApplied bool `json:"long_context_billing_applied,omitempty"`
// AccountRateMultiplier holds the value of the "account_rate_multiplier" field.
AccountRateMultiplier *float64 `json:"account_rate_multiplier,omitempty"`
// BillingType holds the value of the "billing_type" field.
@@ -196,7 +198,7 @@ func (*UsageLog) scanValues(columns []string) ([]any, error) {
switch columns[i] {
case usagelog.FieldImageSizeBreakdown:
values[i] = new([]byte)
- case usagelog.FieldStream, usagelog.FieldCacheTTLOverridden:
+ case usagelog.FieldLongContextBillingApplied, usagelog.FieldStream, usagelog.FieldCacheTTLOverridden:
values[i] = new(sql.NullBool)
case usagelog.FieldInputCost, usagelog.FieldOutputCost, usagelog.FieldCacheCreationCost, usagelog.FieldCacheReadCost, usagelog.FieldTotalCost, usagelog.FieldActualCost, usagelog.FieldRateMultiplier, usagelog.FieldAccountRateMultiplier:
values[i] = new(sql.NullFloat64)
@@ -391,6 +393,12 @@ func (_m *UsageLog) assignValues(columns []string, values []any) error {
} else if value.Valid {
_m.RateMultiplier = value.Float64
}
+ case usagelog.FieldLongContextBillingApplied:
+ if value, ok := values[i].(*sql.NullBool); !ok {
+ return fmt.Errorf("unexpected type %T for field long_context_billing_applied", values[i])
+ } else if value.Valid {
+ _m.LongContextBillingApplied = value.Bool
+ }
case usagelog.FieldAccountRateMultiplier:
if value, ok := values[i].(*sql.NullFloat64); !ok {
return fmt.Errorf("unexpected type %T for field account_rate_multiplier", values[i])
@@ -667,6 +675,9 @@ func (_m *UsageLog) String() string {
builder.WriteString("rate_multiplier=")
builder.WriteString(fmt.Sprintf("%v", _m.RateMultiplier))
builder.WriteString(", ")
+ builder.WriteString("long_context_billing_applied=")
+ builder.WriteString(fmt.Sprintf("%v", _m.LongContextBillingApplied))
+ builder.WriteString(", ")
if v := _m.AccountRateMultiplier; v != nil {
builder.WriteString("account_rate_multiplier=")
builder.WriteString(fmt.Sprintf("%v", *v))
diff --git a/backend/ent/usagelog/usagelog.go b/backend/ent/usagelog/usagelog.go
index a74a92c40f..a87d937195 100644
--- a/backend/ent/usagelog/usagelog.go
+++ b/backend/ent/usagelog/usagelog.go
@@ -66,6 +66,8 @@ const (
FieldActualCost = "actual_cost"
// FieldRateMultiplier holds the string denoting the rate_multiplier field in the database.
FieldRateMultiplier = "rate_multiplier"
+ // FieldLongContextBillingApplied holds the string denoting the long_context_billing_applied field in the database.
+ FieldLongContextBillingApplied = "long_context_billing_applied"
// FieldAccountRateMultiplier holds the string denoting the account_rate_multiplier field in the database.
FieldAccountRateMultiplier = "account_rate_multiplier"
// FieldBillingType holds the string denoting the billing_type field in the database.
@@ -180,6 +182,7 @@ var Columns = []string{
FieldTotalCost,
FieldActualCost,
FieldRateMultiplier,
+ FieldLongContextBillingApplied,
FieldAccountRateMultiplier,
FieldBillingType,
FieldStream,
@@ -251,6 +254,8 @@ var (
DefaultActualCost float64
// DefaultRateMultiplier holds the default value on creation for the "rate_multiplier" field.
DefaultRateMultiplier float64
+ // DefaultLongContextBillingApplied holds the default value on creation for the "long_context_billing_applied" field.
+ DefaultLongContextBillingApplied bool
// DefaultBillingType holds the default value on creation for the "billing_type" field.
DefaultBillingType int8
// DefaultStream holds the default value on creation for the "stream" field.
@@ -417,6 +422,11 @@ func ByRateMultiplier(opts ...sql.OrderTermOption) OrderOption {
return sql.OrderByField(FieldRateMultiplier, opts...).ToFunc()
}
+// ByLongContextBillingApplied orders the results by the long_context_billing_applied field.
+func ByLongContextBillingApplied(opts ...sql.OrderTermOption) OrderOption {
+ return sql.OrderByField(FieldLongContextBillingApplied, opts...).ToFunc()
+}
+
// ByAccountRateMultiplier orders the results by the account_rate_multiplier field.
func ByAccountRateMultiplier(opts ...sql.OrderTermOption) OrderOption {
return sql.OrderByField(FieldAccountRateMultiplier, opts...).ToFunc()
diff --git a/backend/ent/usagelog/where.go b/backend/ent/usagelog/where.go
index 4b08cc3425..a9462e0d0e 100644
--- a/backend/ent/usagelog/where.go
+++ b/backend/ent/usagelog/where.go
@@ -185,6 +185,11 @@ func RateMultiplier(v float64) predicate.UsageLog {
return predicate.UsageLog(sql.FieldEQ(FieldRateMultiplier, v))
}
+// LongContextBillingApplied applies equality check predicate on the "long_context_billing_applied" field. It's identical to LongContextBillingAppliedEQ.
+func LongContextBillingApplied(v bool) predicate.UsageLog {
+ return predicate.UsageLog(sql.FieldEQ(FieldLongContextBillingApplied, v))
+}
+
// AccountRateMultiplier applies equality check predicate on the "account_rate_multiplier" field. It's identical to AccountRateMultiplierEQ.
func AccountRateMultiplier(v float64) predicate.UsageLog {
return predicate.UsageLog(sql.FieldEQ(FieldAccountRateMultiplier, v))
@@ -1465,6 +1470,16 @@ func RateMultiplierLTE(v float64) predicate.UsageLog {
return predicate.UsageLog(sql.FieldLTE(FieldRateMultiplier, v))
}
+// LongContextBillingAppliedEQ applies the EQ predicate on the "long_context_billing_applied" field.
+func LongContextBillingAppliedEQ(v bool) predicate.UsageLog {
+ return predicate.UsageLog(sql.FieldEQ(FieldLongContextBillingApplied, v))
+}
+
+// LongContextBillingAppliedNEQ applies the NEQ predicate on the "long_context_billing_applied" field.
+func LongContextBillingAppliedNEQ(v bool) predicate.UsageLog {
+ return predicate.UsageLog(sql.FieldNEQ(FieldLongContextBillingApplied, v))
+}
+
// AccountRateMultiplierEQ applies the EQ predicate on the "account_rate_multiplier" field.
func AccountRateMultiplierEQ(v float64) predicate.UsageLog {
return predicate.UsageLog(sql.FieldEQ(FieldAccountRateMultiplier, v))
diff --git a/backend/ent/usagelog_create.go b/backend/ent/usagelog_create.go
index 3326f72fc0..31cf45328e 100644
--- a/backend/ent/usagelog_create.go
+++ b/backend/ent/usagelog_create.go
@@ -351,6 +351,20 @@ func (_c *UsageLogCreate) SetNillableRateMultiplier(v *float64) *UsageLogCreate
return _c
}
+// SetLongContextBillingApplied sets the "long_context_billing_applied" field.
+func (_c *UsageLogCreate) SetLongContextBillingApplied(v bool) *UsageLogCreate {
+ _c.mutation.SetLongContextBillingApplied(v)
+ return _c
+}
+
+// SetNillableLongContextBillingApplied sets the "long_context_billing_applied" field if the given value is not nil.
+func (_c *UsageLogCreate) SetNillableLongContextBillingApplied(v *bool) *UsageLogCreate {
+ if v != nil {
+ _c.SetLongContextBillingApplied(*v)
+ }
+ return _c
+}
+
// SetAccountRateMultiplier sets the "account_rate_multiplier" field.
func (_c *UsageLogCreate) SetAccountRateMultiplier(v float64) *UsageLogCreate {
_c.mutation.SetAccountRateMultiplier(v)
@@ -707,6 +721,10 @@ func (_c *UsageLogCreate) defaults() {
v := usagelog.DefaultRateMultiplier
_c.mutation.SetRateMultiplier(v)
}
+ if _, ok := _c.mutation.LongContextBillingApplied(); !ok {
+ v := usagelog.DefaultLongContextBillingApplied
+ _c.mutation.SetLongContextBillingApplied(v)
+ }
if _, ok := _c.mutation.BillingType(); !ok {
v := usagelog.DefaultBillingType
_c.mutation.SetBillingType(v)
@@ -824,6 +842,9 @@ func (_c *UsageLogCreate) check() error {
if _, ok := _c.mutation.RateMultiplier(); !ok {
return &ValidationError{Name: "rate_multiplier", err: errors.New(`ent: missing required field "UsageLog.rate_multiplier"`)}
}
+ if _, ok := _c.mutation.LongContextBillingApplied(); !ok {
+ return &ValidationError{Name: "long_context_billing_applied", err: errors.New(`ent: missing required field "UsageLog.long_context_billing_applied"`)}
+ }
if _, ok := _c.mutation.BillingType(); !ok {
return &ValidationError{Name: "billing_type", err: errors.New(`ent: missing required field "UsageLog.billing_type"`)}
}
@@ -997,6 +1018,10 @@ func (_c *UsageLogCreate) createSpec() (*UsageLog, *sqlgraph.CreateSpec) {
_spec.SetField(usagelog.FieldRateMultiplier, field.TypeFloat64, value)
_node.RateMultiplier = value
}
+ if value, ok := _c.mutation.LongContextBillingApplied(); ok {
+ _spec.SetField(usagelog.FieldLongContextBillingApplied, field.TypeBool, value)
+ _node.LongContextBillingApplied = value
+ }
if value, ok := _c.mutation.AccountRateMultiplier(); ok {
_spec.SetField(usagelog.FieldAccountRateMultiplier, field.TypeFloat64, value)
_node.AccountRateMultiplier = &value
@@ -1650,6 +1675,18 @@ func (u *UsageLogUpsert) AddRateMultiplier(v float64) *UsageLogUpsert {
return u
}
+// SetLongContextBillingApplied sets the "long_context_billing_applied" field.
+func (u *UsageLogUpsert) SetLongContextBillingApplied(v bool) *UsageLogUpsert {
+ u.Set(usagelog.FieldLongContextBillingApplied, v)
+ return u
+}
+
+// UpdateLongContextBillingApplied sets the "long_context_billing_applied" field to the value that was provided on create.
+func (u *UsageLogUpsert) UpdateLongContextBillingApplied() *UsageLogUpsert {
+ u.SetExcluded(usagelog.FieldLongContextBillingApplied)
+ return u
+}
+
// SetAccountRateMultiplier sets the "account_rate_multiplier" field.
func (u *UsageLogUpsert) SetAccountRateMultiplier(v float64) *UsageLogUpsert {
u.Set(usagelog.FieldAccountRateMultiplier, v)
@@ -2531,6 +2568,20 @@ func (u *UsageLogUpsertOne) UpdateRateMultiplier() *UsageLogUpsertOne {
})
}
+// SetLongContextBillingApplied sets the "long_context_billing_applied" field.
+func (u *UsageLogUpsertOne) SetLongContextBillingApplied(v bool) *UsageLogUpsertOne {
+ return u.Update(func(s *UsageLogUpsert) {
+ s.SetLongContextBillingApplied(v)
+ })
+}
+
+// UpdateLongContextBillingApplied sets the "long_context_billing_applied" field to the value that was provided on create.
+func (u *UsageLogUpsertOne) UpdateLongContextBillingApplied() *UsageLogUpsertOne {
+ return u.Update(func(s *UsageLogUpsert) {
+ s.UpdateLongContextBillingApplied()
+ })
+}
+
// SetAccountRateMultiplier sets the "account_rate_multiplier" field.
func (u *UsageLogUpsertOne) SetAccountRateMultiplier(v float64) *UsageLogUpsertOne {
return u.Update(func(s *UsageLogUpsert) {
@@ -3631,6 +3682,20 @@ func (u *UsageLogUpsertBulk) UpdateRateMultiplier() *UsageLogUpsertBulk {
})
}
+// SetLongContextBillingApplied sets the "long_context_billing_applied" field.
+func (u *UsageLogUpsertBulk) SetLongContextBillingApplied(v bool) *UsageLogUpsertBulk {
+ return u.Update(func(s *UsageLogUpsert) {
+ s.SetLongContextBillingApplied(v)
+ })
+}
+
+// UpdateLongContextBillingApplied sets the "long_context_billing_applied" field to the value that was provided on create.
+func (u *UsageLogUpsertBulk) UpdateLongContextBillingApplied() *UsageLogUpsertBulk {
+ return u.Update(func(s *UsageLogUpsert) {
+ s.UpdateLongContextBillingApplied()
+ })
+}
+
// SetAccountRateMultiplier sets the "account_rate_multiplier" field.
func (u *UsageLogUpsertBulk) SetAccountRateMultiplier(v float64) *UsageLogUpsertBulk {
return u.Update(func(s *UsageLogUpsert) {
diff --git a/backend/ent/usagelog_update.go b/backend/ent/usagelog_update.go
index 00a65ccff1..2a60d6f44d 100644
--- a/backend/ent/usagelog_update.go
+++ b/backend/ent/usagelog_update.go
@@ -542,6 +542,20 @@ func (_u *UsageLogUpdate) AddRateMultiplier(v float64) *UsageLogUpdate {
return _u
}
+// SetLongContextBillingApplied sets the "long_context_billing_applied" field.
+func (_u *UsageLogUpdate) SetLongContextBillingApplied(v bool) *UsageLogUpdate {
+ _u.mutation.SetLongContextBillingApplied(v)
+ return _u
+}
+
+// SetNillableLongContextBillingApplied sets the "long_context_billing_applied" field if the given value is not nil.
+func (_u *UsageLogUpdate) SetNillableLongContextBillingApplied(v *bool) *UsageLogUpdate {
+ if v != nil {
+ _u.SetLongContextBillingApplied(*v)
+ }
+ return _u
+}
+
// SetAccountRateMultiplier sets the "account_rate_multiplier" field.
func (_u *UsageLogUpdate) SetAccountRateMultiplier(v float64) *UsageLogUpdate {
_u.mutation.ResetAccountRateMultiplier()
@@ -1199,6 +1213,9 @@ func (_u *UsageLogUpdate) sqlSave(ctx context.Context) (_node int, err error) {
if value, ok := _u.mutation.AddedRateMultiplier(); ok {
_spec.AddField(usagelog.FieldRateMultiplier, field.TypeFloat64, value)
}
+ if value, ok := _u.mutation.LongContextBillingApplied(); ok {
+ _spec.SetField(usagelog.FieldLongContextBillingApplied, field.TypeBool, value)
+ }
if value, ok := _u.mutation.AccountRateMultiplier(); ok {
_spec.SetField(usagelog.FieldAccountRateMultiplier, field.TypeFloat64, value)
}
@@ -1982,6 +1999,20 @@ func (_u *UsageLogUpdateOne) AddRateMultiplier(v float64) *UsageLogUpdateOne {
return _u
}
+// SetLongContextBillingApplied sets the "long_context_billing_applied" field.
+func (_u *UsageLogUpdateOne) SetLongContextBillingApplied(v bool) *UsageLogUpdateOne {
+ _u.mutation.SetLongContextBillingApplied(v)
+ return _u
+}
+
+// SetNillableLongContextBillingApplied sets the "long_context_billing_applied" field if the given value is not nil.
+func (_u *UsageLogUpdateOne) SetNillableLongContextBillingApplied(v *bool) *UsageLogUpdateOne {
+ if v != nil {
+ _u.SetLongContextBillingApplied(*v)
+ }
+ return _u
+}
+
// SetAccountRateMultiplier sets the "account_rate_multiplier" field.
func (_u *UsageLogUpdateOne) SetAccountRateMultiplier(v float64) *UsageLogUpdateOne {
_u.mutation.ResetAccountRateMultiplier()
@@ -2669,6 +2700,9 @@ func (_u *UsageLogUpdateOne) sqlSave(ctx context.Context) (_node *UsageLog, err
if value, ok := _u.mutation.AddedRateMultiplier(); ok {
_spec.AddField(usagelog.FieldRateMultiplier, field.TypeFloat64, value)
}
+ if value, ok := _u.mutation.LongContextBillingApplied(); ok {
+ _spec.SetField(usagelog.FieldLongContextBillingApplied, field.TypeBool, value)
+ }
if value, ok := _u.mutation.AccountRateMultiplier(); ok {
_spec.SetField(usagelog.FieldAccountRateMultiplier, field.TypeFloat64, value)
}
diff --git a/backend/internal/config/config.go b/backend/internal/config/config.go
index 8e081bac34..ea9169e0b7 100644
--- a/backend/internal/config/config.go
+++ b/backend/internal/config/config.go
@@ -601,6 +601,7 @@ type ServerConfig struct {
Host string `mapstructure:"host"`
Port int `mapstructure:"port"`
Mode string `mapstructure:"mode"` // debug/release
+ EnableServerTiming bool `mapstructure:"enable_server_timing"` // Admin UI Server-Timing response header
FrontendURL string `mapstructure:"frontend_url"` // 前端基础 URL,用于生成邮件中的外部链接
ReadHeaderTimeout int `mapstructure:"read_header_timeout"` // 读取请求头超时(秒)
IdleTimeout int `mapstructure:"idle_timeout"` // 空闲连接超时(秒)
@@ -1459,6 +1460,9 @@ func load(allowMissingJWTSecret bool) (*Config, error) {
// 环境变量支持
viper.AutomaticEnv()
viper.SetEnvKeyReplacer(strings.NewReplacer(".", "_"))
+ if err := viper.BindEnv("server.enable_server_timing", "ENABLE_SERVER_TIMING"); err != nil {
+ return nil, fmt.Errorf("bind ENABLE_SERVER_TIMING: %w", err)
+ }
// 默认值
setDefaults()
@@ -1614,6 +1618,7 @@ func setDefaults() {
viper.SetDefault("server.host", "0.0.0.0")
viper.SetDefault("server.port", 8080)
viper.SetDefault("server.mode", "release")
+ viper.SetDefault("server.enable_server_timing", false)
viper.SetDefault("server.frontend_url", "")
viper.SetDefault("server.read_header_timeout", 30) // 30秒读取请求头
viper.SetDefault("server.idle_timeout", 120) // 120秒空闲超时
diff --git a/backend/internal/config/config_test.go b/backend/internal/config/config_test.go
index 489bc346b8..4eea3a2840 100644
--- a/backend/internal/config/config_test.go
+++ b/backend/internal/config/config_test.go
@@ -17,6 +17,23 @@ func resetViperWithJWTSecret(t *testing.T) {
t.Setenv("JWT_SECRET", strings.Repeat("x", 32))
}
+func TestLoadServerTimingConfig(t *testing.T) {
+ t.Run("disabled by default", func(t *testing.T) {
+ resetViperWithJWTSecret(t)
+ cfg, err := Load()
+ require.NoError(t, err)
+ require.False(t, cfg.Server.EnableServerTiming)
+ })
+
+ t.Run("enabled by exact environment variable", func(t *testing.T) {
+ resetViperWithJWTSecret(t)
+ t.Setenv("ENABLE_SERVER_TIMING", "true")
+ cfg, err := Load()
+ require.NoError(t, err)
+ require.True(t, cfg.Server.EnableServerTiming)
+ })
+}
+
func TestLoadForBootstrapAllowsMissingJWTSecret(t *testing.T) {
viper.Reset()
t.Setenv("JWT_SECRET", "")
diff --git a/backend/internal/handler/admin/account_codex_import.go b/backend/internal/handler/admin/account_codex_import.go
index 01a5fbfa1c..271bd19470 100644
--- a/backend/internal/handler/admin/account_codex_import.go
+++ b/backend/internal/handler/admin/account_codex_import.go
@@ -115,6 +115,10 @@ func (h *AccountHandler) ImportCodexSession(c *gin.Context) {
response.BadRequest(c, "Invalid request: "+err.Error())
return
}
+ if err := service.ValidateOpenAILongContextBillingExtra(service.PlatformOpenAI, req.Extra); err != nil {
+ response.ErrorFrom(c, err)
+ return
+ }
if req.Concurrency != nil && *req.Concurrency < 0 {
response.BadRequest(c, "concurrency must be >= 0")
return
diff --git a/backend/internal/handler/admin/account_codex_import_test.go b/backend/internal/handler/admin/account_codex_import_test.go
index a52463aa86..96a033d8c3 100644
--- a/backend/internal/handler/admin/account_codex_import_test.go
+++ b/backend/internal/handler/admin/account_codex_import_test.go
@@ -630,6 +630,7 @@ func TestImportCodexSessionsAccessTokenOnlySameUserUpdatesExisting(t *testing.T)
"chatgpt_user_id": "user-1",
"access_token": existingToken,
},
+ Extra: map[string]any{"openai_long_context_billing_enabled": false},
}})
handler := NewAccountHandler(svc, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil)
req := CodexSessionImportRequest{SkipDefaultGroupBind: boolPtr(true)}
@@ -650,6 +651,9 @@ func TestImportCodexSessionsAccessTokenOnlySameUserUpdatesExisting(t *testing.T)
if len(svc.updatedAccounts) != 1 || svc.updatedAccounts[0].id != 10 {
t.Fatalf("updated accounts = %+v, want account 10", svc.updatedAccounts)
}
+ if got := svc.updatedAccounts[0].input.Extra["openai_long_context_billing_enabled"]; got != false {
+ t.Fatalf("openai_long_context_billing_enabled = %v, want false", got)
+ }
}
func TestImportCodexSessionsUpgradesAccessTokenOnlyAccountWithRefreshToken(t *testing.T) {
diff --git a/backend/internal/handler/admin/account_handler.go b/backend/internal/handler/admin/account_handler.go
index a4b0773999..b886728159 100644
--- a/backend/internal/handler/admin/account_handler.go
+++ b/backend/internal/handler/admin/account_handler.go
@@ -784,6 +784,10 @@ func (h *AccountHandler) Create(c *gin.Context) {
response.BadRequest(c, "Invalid request: "+err.Error())
return
}
+ if err := service.ValidateOpenAILongContextBillingExtra(req.Platform, req.Extra); err != nil {
+ response.ErrorFrom(c, err)
+ return
+ }
if req.RateMultiplier != nil && *req.RateMultiplier < 0 {
response.BadRequest(c, "rate_multiplier must be >= 0")
return
@@ -1299,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,
@@ -1592,6 +1600,12 @@ func (h *AccountHandler) BatchCreate(c *gin.Context) {
response.BadRequest(c, "Invalid request: "+err.Error())
return
}
+ for _, item := range req.Accounts {
+ if err := service.ValidateOpenAILongContextBillingExtra(item.Platform, item.Extra); err != nil {
+ response.ErrorFrom(c, err)
+ return
+ }
+ }
executeAdminIdempotentJSON(c, "admin.accounts.batch_create", req, service.DefaultWriteIdempotencyTTL(), func(ctx context.Context) (any, error) {
success := 0
diff --git a/backend/internal/handler/admin/account_handler_long_context_billing_test.go b/backend/internal/handler/admin/account_handler_long_context_billing_test.go
new file mode 100644
index 0000000000..d50513a3e8
--- /dev/null
+++ b/backend/internal/handler/admin/account_handler_long_context_billing_test.go
@@ -0,0 +1,165 @@
+package admin
+
+import (
+ "bytes"
+ "encoding/json"
+ "net/http"
+ "net/http/httptest"
+ "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"
+)
+
+func TestAccountAdminBoundariesRejectMalformedOpenAILongContextBillingValue(t *testing.T) {
+ const malformedExtra = `"extra":{"openai_long_context_billing_enabled":"true"}`
+
+ tests := []struct {
+ name string
+ method string
+ path string
+ body string
+ mount func(*gin.Engine, *AccountHandler)
+ setup func(*stubAdminService)
+ }{
+ {
+ name: "create",
+ method: http.MethodPost,
+ path: "/accounts",
+ body: `{"name":"account","platform":"openai","type":"apikey","credentials":{"api_key":"test"},` + malformedExtra + `}`,
+ mount: func(router *gin.Engine, handler *AccountHandler) { router.POST("/accounts", handler.Create) },
+ },
+ {
+ name: "update",
+ method: http.MethodPut,
+ path: "/accounts/1",
+ body: `{` + malformedExtra + `}`,
+ mount: func(router *gin.Engine, handler *AccountHandler) { router.PUT("/accounts/:id", handler.Update) },
+ setup: func(stub *stubAdminService) {
+ stub.updateAccountErr = infraerrors.BadRequest("OPENAI_LONG_CONTEXT_BILLING_INVALID", "invalid")
+ },
+ },
+ {
+ name: "bulk update",
+ method: http.MethodPost,
+ path: "/accounts/bulk-update",
+ body: `{"account_ids":[1],` + malformedExtra + `}`,
+ mount: func(router *gin.Engine, handler *AccountHandler) {
+ router.POST("/accounts/bulk-update", handler.BulkUpdate)
+ },
+ setup: func(stub *stubAdminService) {
+ stub.bulkUpdateAccountErr = infraerrors.BadRequest("OPENAI_LONG_CONTEXT_BILLING_INVALID", "invalid")
+ },
+ },
+ {
+ name: "batch create",
+ method: http.MethodPost,
+ path: "/accounts/batch",
+ body: `{"accounts":[{"name":"account","platform":"openai","type":"apikey","credentials":{"api_key":"test"},` + malformedExtra + `}]}`,
+ mount: func(router *gin.Engine, handler *AccountHandler) { router.POST("/accounts/batch", handler.BatchCreate) },
+ },
+ {
+ name: "Codex session import",
+ method: http.MethodPost,
+ path: "/accounts/import-codex-session",
+ body: `{"content":"token",` + malformedExtra + `}`,
+ mount: func(router *gin.Engine, handler *AccountHandler) {
+ router.POST("/accounts/import-codex-session", handler.ImportCodexSession)
+ },
+ },
+ }
+
+ for _, tt := range tests {
+ t.Run(tt.name, func(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ stub := newStubAdminService()
+ if tt.setup != nil {
+ tt.setup(stub)
+ }
+ handler := NewAccountHandler(stub, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil)
+ router := gin.New()
+ tt.mount(router, handler)
+ recorder := httptest.NewRecorder()
+ request := httptest.NewRequest(tt.method, tt.path, bytes.NewBufferString(tt.body))
+ 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)
+ })
+ }
+}
+
+func TestAccountCreateBoundaryDoesNotApplyOpenAIValidationToOtherPlatforms(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ handler := NewAccountHandler(newStubAdminService(), nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil)
+ router := gin.New()
+ router.POST("/accounts", handler.Create)
+ recorder := httptest.NewRecorder()
+ request := httptest.NewRequest(http.MethodPost, "/accounts", bytes.NewBufferString(
+ `{"name":"account","platform":"anthropic","type":"apikey","credentials":{"api_key":"test"},"extra":{"openai_long_context_billing_enabled":"provider-owned"}}`,
+ ))
+ request.Header.Set("Content-Type", "application/json")
+
+ router.ServeHTTP(recorder, request)
+
+ 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)
+ router := gin.New()
+ router.Use(gin.Recovery())
+ router.POST("/openai/create-from-codex-pat", handler.CreateAccountFromCodexPAT)
+ recorder := httptest.NewRecorder()
+ request := httptest.NewRequest(http.MethodPost, "/openai/create-from-codex-pat", bytes.NewBufferString(
+ `{"access_token":"token","extra":{"openai_long_context_billing_enabled":1}}`,
+ ))
+ 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)
+}
diff --git a/backend/internal/handler/admin/admin_service_stub_test.go b/backend/internal/handler/admin/admin_service_stub_test.go
index 7a7cbb473e..5e9c4d517e 100644
--- a/backend/internal/handler/admin/admin_service_stub_test.go
+++ b/backend/internal/handler/admin/admin_service_stub_test.go
@@ -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
}
diff --git a/backend/internal/handler/admin/grok_oauth_handler.go b/backend/internal/handler/admin/grok_oauth_handler.go
index a5841b9d5d..dfe5632e30 100644
--- a/backend/internal/handler/admin/grok_oauth_handler.go
+++ b/backend/internal/handler/admin/grok_oauth_handler.go
@@ -454,7 +454,7 @@ func (h *GrokOAuthHandler) QueryQuota(c *gin.Context) {
response.BadRequest(c, "grok quota service is not enabled")
return
}
- result, err := h.quotaService.ProbeUsage(c.Request.Context(), accountID)
+ result, err := h.quotaService.QueryQuota(c.Request.Context(), accountID)
if err != nil {
response.ErrorFrom(c, err)
return
diff --git a/backend/internal/handler/admin/grok_oauth_handler_test.go b/backend/internal/handler/admin/grok_oauth_handler_test.go
index 6101a25d35..64ea044aa3 100644
--- a/backend/internal/handler/admin/grok_oauth_handler_test.go
+++ b/backend/internal/handler/admin/grok_oauth_handler_test.go
@@ -8,6 +8,7 @@ import (
"net/http"
"net/http/httptest"
"strings"
+ "sync"
"testing"
"time"
@@ -41,17 +42,35 @@ func (r *grokQuotaHandlerAccountRepo) UpdateExtra(_ context.Context, id int64, u
}
type grokQuotaHandlerUpstream struct {
- resp *http.Response
- lastReq *http.Request
- lastBody []byte
+ mu sync.Mutex
+ requests []*http.Request
+ bodies [][]byte
}
func (u *grokQuotaHandlerUpstream) Do(req *http.Request, _ string, _ int64, _ int) (*http.Response, error) {
- u.lastReq = req
+ var body []byte
if req.Body != nil {
- u.lastBody, _ = io.ReadAll(req.Body)
+ body, _ = io.ReadAll(req.Body)
}
- return u.resp, nil
+ u.mu.Lock()
+ u.requests = append(u.requests, req)
+ u.bodies = append(u.bodies, body)
+ u.mu.Unlock()
+ if req.URL.Path == "/v1/responses" {
+ return &http.Response{
+ StatusCode: http.StatusOK,
+ Header: http.Header{
+ "X-Ratelimit-Limit-Requests": []string{"10"},
+ "X-Ratelimit-Remaining-Requests": []string{"8"},
+ },
+ Body: io.NopCloser(strings.NewReader(`{"id":"resp_probe"}`)),
+ }, nil
+ }
+ payload := `{"config":{"billingPeriodStart":"2026-07-01T00:00:00Z","billingPeriodEnd":"2026-08-01T00:00:00Z"}}`
+ if req.URL.RawQuery == "format=credits" {
+ payload = `{"config":{"currentPeriod":{"type":"WEEKLY","start":"2026-07-09T03:25:00Z","end":"2026-07-16T03:25:00Z"}}}`
+ }
+ return &http.Response{StatusCode: http.StatusOK, Header: make(http.Header), Body: io.NopCloser(strings.NewReader(payload))}, nil
}
func (u *grokQuotaHandlerUpstream) DoWithTLS(
@@ -77,14 +96,7 @@ func TestGrokOAuthHandlerQueryQuotaProbesUpstream(t *testing.T) {
"expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339),
},
}}
- upstream := &grokQuotaHandlerUpstream{resp: &http.Response{
- StatusCode: http.StatusOK,
- Header: http.Header{
- "X-Ratelimit-Limit-Requests": []string{"10"},
- "X-Ratelimit-Remaining-Requests": []string{"8"},
- },
- Body: io.NopCloser(strings.NewReader(`{"id":"resp_probe"}`)),
- }}
+ upstream := &grokQuotaHandlerUpstream{}
quotaService := service.NewGrokQuotaService(repo, nil, service.NewGrokTokenProvider(repo, nil), upstream)
handler := NewGrokOAuthHandler(nil, nil, quotaService)
@@ -95,12 +107,23 @@ func TestGrokOAuthHandlerQueryQuotaProbesUpstream(t *testing.T) {
router.ServeHTTP(rec, req)
require.Equal(t, http.StatusOK, rec.Code)
- require.Contains(t, rec.Body.String(), `"source":"active_probe"`)
+ require.Contains(t, rec.Body.String(), `"source":"hybrid_probe"`)
+ require.Contains(t, rec.Body.String(), `"billing":`)
+ require.Contains(t, rec.Body.String(), `"snapshot":`)
require.Contains(t, rec.Body.String(), `"headers_observed":true`)
require.NotContains(t, rec.Body.String(), "access-token")
- require.Equal(t, xai.DefaultCLIBaseURL+"/responses", upstream.lastReq.URL.String())
- require.Equal(t, "Bearer access-token", upstream.lastReq.Header.Get("Authorization"))
- require.Contains(t, string(upstream.lastBody), `"store":false`)
+ upstream.mu.Lock()
+ requests := append([]*http.Request(nil), upstream.requests...)
+ bodies := append([][]byte(nil), upstream.bodies...)
+ upstream.mu.Unlock()
+ require.Len(t, requests, 3)
+ for i, upstreamReq := range requests {
+ require.Equal(t, "Bearer access-token", upstreamReq.Header.Get("Authorization"))
+ if upstreamReq.URL.String() == xai.DefaultCLIBaseURL+"/responses" {
+ require.Contains(t, string(bodies[i]), `"model":"grok-4.5"`)
+ require.Contains(t, string(bodies[i]), `"store":false`)
+ }
+ }
require.NotNil(t, repo.updates[42])
}
diff --git a/backend/internal/handler/admin/openai_oauth_handler.go b/backend/internal/handler/admin/openai_oauth_handler.go
index d7a756bd00..78d57299b6 100644
--- a/backend/internal/handler/admin/openai_oauth_handler.go
+++ b/backend/internal/handler/admin/openai_oauth_handler.go
@@ -304,6 +304,10 @@ func (h *OpenAIOAuthHandler) CreateAccountFromCodexPAT(c *gin.Context) {
response.BadRequest(c, "Invalid request: "+err.Error())
return
}
+ if err := service.ValidateOpenAILongContextBillingExtra(service.PlatformOpenAI, req.Extra); err != nil {
+ response.ErrorFrom(c, err)
+ return
+ }
if req.Concurrency != nil && *req.Concurrency < 0 {
response.BadRequest(c, "concurrency must be >= 0")
return
diff --git a/backend/internal/handler/admin/ops_system_log_handler.go b/backend/internal/handler/admin/ops_system_log_handler.go
index 9f3c8b893a..1b6af45976 100644
--- a/backend/internal/handler/admin/ops_system_log_handler.go
+++ b/backend/internal/handler/admin/ops_system_log_handler.go
@@ -15,6 +15,7 @@ import (
type opsSystemLogCleanupRequest struct {
StartTime string `json:"start_time"`
EndTime string `json:"end_time"`
+ Host string `json:"host"`
Level string `json:"level"`
Component string `json:"component"`
@@ -56,6 +57,7 @@ func (h *OpsHandler) ListSystemLogs(c *gin.Context) {
PageSize: pageSize,
StartTime: &start,
EndTime: &end,
+ Host: strings.TrimSpace(c.Query("host")),
Level: strings.TrimSpace(c.Query("level")),
Component: strings.TrimSpace(c.Query("component")),
RequestID: strings.TrimSpace(c.Query("request_id")),
@@ -153,6 +155,7 @@ func (h *OpsHandler) CleanupSystemLogs(c *gin.Context) {
filter := &service.OpsSystemLogCleanupFilter{
StartTime: start,
EndTime: end,
+ Host: strings.TrimSpace(req.Host),
Level: strings.TrimSpace(req.Level),
Component: strings.TrimSpace(req.Component),
RequestID: strings.TrimSpace(req.RequestID),
diff --git a/backend/internal/handler/admin/ops_system_log_handler_test.go b/backend/internal/handler/admin/ops_system_log_handler_test.go
index 9557fce442..3390fbe3cb 100644
--- a/backend/internal/handler/admin/ops_system_log_handler_test.go
+++ b/backend/internal/handler/admin/ops_system_log_handler_test.go
@@ -2,6 +2,7 @@ package admin
import (
"bytes"
+ "context"
"encoding/json"
"net/http"
"net/http/httptest"
@@ -19,6 +20,26 @@ type responseEnvelope struct {
Data json.RawMessage `json:"data"`
}
+type opsSystemLogCaptureRepo struct {
+ service.OpsRepository
+ listFilter *service.OpsSystemLogFilter
+ cleanupFilter *service.OpsSystemLogCleanupFilter
+}
+
+func (r *opsSystemLogCaptureRepo) ListSystemLogs(_ context.Context, filter *service.OpsSystemLogFilter) (*service.OpsSystemLogList, error) {
+ r.listFilter = filter
+ return &service.OpsSystemLogList{Logs: []*service.OpsSystemLog{}, Page: filter.Page, PageSize: filter.PageSize}, nil
+}
+
+func (r *opsSystemLogCaptureRepo) DeleteSystemLogs(_ context.Context, filter *service.OpsSystemLogCleanupFilter) (int64, error) {
+ r.cleanupFilter = filter
+ return 1, nil
+}
+
+func (r *opsSystemLogCaptureRepo) InsertSystemLogCleanupAudit(_ context.Context, _ *service.OpsSystemLogCleanupAudit) error {
+ return nil
+}
+
func newOpsSystemLogTestRouter(handler *OpsHandler, withUser bool) *gin.Engine {
gin.SetMode(gin.TestMode)
r := gin.New()
@@ -121,6 +142,23 @@ func TestOpsSystemLogHandler_ListSuccess(t *testing.T) {
}
}
+func TestOpsSystemLogHandler_ListAcceptsHost(t *testing.T) {
+ repo := &opsSystemLogCaptureRepo{}
+ svc := service.NewOpsService(repo, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil)
+ h := NewOpsHandler(svc)
+ r := newOpsSystemLogTestRouter(h, false)
+
+ w := httptest.NewRecorder()
+ req := httptest.NewRequest(http.MethodGet, "/logs?host=api-node-1", nil)
+ r.ServeHTTP(w, req)
+ if w.Code != http.StatusOK {
+ t.Fatalf("status=%d, want 200", w.Code)
+ }
+ if repo.listFilter == nil || repo.listFilter.Host != "api-node-1" {
+ t.Fatalf("host filter = %+v, want api-node-1", repo.listFilter)
+ }
+}
+
func TestOpsSystemLogHandler_CleanupUnauthorized(t *testing.T) {
svc := service.NewOpsService(nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil)
h := NewOpsHandler(svc)
@@ -205,6 +243,24 @@ func TestOpsSystemLogHandler_CleanupAcceptsAPIKeyID(t *testing.T) {
}
}
+func TestOpsSystemLogHandler_CleanupAcceptsHost(t *testing.T) {
+ repo := &opsSystemLogCaptureRepo{}
+ svc := service.NewOpsService(repo, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil)
+ h := NewOpsHandler(svc)
+ r := newOpsSystemLogTestRouter(h, true)
+
+ w := httptest.NewRecorder()
+ req := httptest.NewRequest(http.MethodPost, "/logs/cleanup", bytes.NewBufferString(`{"host":"api-node-1"}`))
+ req.Header.Set("Content-Type", "application/json")
+ r.ServeHTTP(w, req)
+ if w.Code != http.StatusOK {
+ t.Fatalf("status=%d, want 200", w.Code)
+ }
+ if repo.cleanupFilter == nil || repo.cleanupFilter.Host != "api-node-1" {
+ t.Fatalf("host filter = %+v, want api-node-1", repo.cleanupFilter)
+ }
+}
+
func TestOpsSystemLogHandler_CleanupInvalidAPIKeyID(t *testing.T) {
svc := service.NewOpsService(nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil)
h := NewOpsHandler(svc)
diff --git a/backend/internal/handler/admin/ops_ws_handler.go b/backend/internal/handler/admin/ops_ws_handler.go
index 75fd7ea002..e4c42cc9c0 100644
--- a/backend/internal/handler/admin/ops_ws_handler.go
+++ b/backend/internal/handler/admin/ops_ws_handler.go
@@ -16,6 +16,7 @@ import (
"time"
"github.com/Wei-Shaw/sub2api/internal/pkg/logger"
+ servermiddleware "github.com/Wei-Shaw/sub2api/internal/server/middleware"
"github.com/Wei-Shaw/sub2api/internal/service"
"github.com/gin-gonic/gin"
"github.com/gorilla/websocket"
@@ -323,7 +324,7 @@ func (h *OpsHandler) QPSWSHandler(c *gin.Context) {
// If realtime monitoring is disabled, prefer a successful WS upgrade followed by a clean close
// with a deterministic close code. This prevents clients from spinning on 404/1006 reconnect loops.
if !h.opsService.IsRealtimeMonitoringEnabled(c.Request.Context()) {
- conn, err := upgrader.Upgrade(c.Writer, c.Request, nil)
+ conn, err := upgrader.Upgrade(c.Writer, c.Request, servermiddleware.ServerTimingResponseHeader(c))
if err != nil {
c.JSON(http.StatusNotFound, gin.H{"error": "ops realtime monitoring is disabled"})
return
@@ -358,7 +359,7 @@ func (h *OpsHandler) QPSWSHandler(c *gin.Context) {
defer releaseOpsWSIPSlot(clientIP)
}
- conn, err := upgrader.Upgrade(c.Writer, c.Request, nil)
+ conn, err := upgrader.Upgrade(c.Writer, c.Request, servermiddleware.ServerTimingResponseHeader(c))
if err != nil {
logger.LegacyPrintf("handler.admin.ops_ws", "[OpsWS] upgrade failed: %v", err)
return
diff --git a/backend/internal/handler/dto/mappers.go b/backend/internal/handler/dto/mappers.go
index e770bcf036..3c45c3b95e 100644
--- a/backend/internal/handler/dto/mappers.go
+++ b/backend/internal/handler/dto/mappers.go
@@ -599,54 +599,55 @@ func usageLogFromServiceUser(l *service.UsageLog) UsageLog {
requestedModel = l.Model
}
return UsageLog{
- ID: l.ID,
- UserID: l.UserID,
- APIKeyID: l.APIKeyID,
- AccountID: l.AccountID,
- RequestID: l.RequestID,
- Model: requestedModel,
- ServiceTier: l.ServiceTier,
- ReasoningEffort: l.ReasoningEffort,
- InboundEndpoint: l.InboundEndpoint,
- GroupID: l.GroupID,
- SubscriptionID: l.SubscriptionID,
- InputTokens: l.InputTokens,
- OutputTokens: l.OutputTokens,
- CacheCreationTokens: l.CacheCreationTokens,
- CacheReadTokens: l.CacheReadTokens,
- CacheCreation5mTokens: l.CacheCreation5mTokens,
- CacheCreation1hTokens: l.CacheCreation1hTokens,
- InputCost: l.InputCost,
- OutputCost: l.OutputCost,
- CacheCreationCost: l.CacheCreationCost,
- CacheReadCost: l.CacheReadCost,
- TotalCost: l.TotalCost,
- ActualCost: l.ActualCost,
- RateMultiplier: l.RateMultiplier,
- BillingType: l.BillingType,
- RequestType: requestType.String(),
- Stream: stream,
- OpenAIWSMode: openAIWSMode,
- DurationMs: l.DurationMs,
- FirstTokenMs: l.FirstTokenMs,
- ImageCount: l.ImageCount,
- ImageSize: l.ImageSize,
- ImageInputSize: l.ImageInputSize,
- ImageOutputSize: l.ImageOutputSize,
- ImageOutputTokens: l.ImageOutputTokens,
- ImageOutputCost: l.ImageOutputCost,
- ImageSizeSource: l.ImageSizeSource,
- ImageSizeBreakdown: l.ImageSizeBreakdown,
- MediaType: l.MediaType,
- UserAgent: l.UserAgent,
- IPAddress: l.IPAddress,
- CacheTTLOverridden: l.CacheTTLOverridden,
- BillingMode: l.BillingMode,
- CreatedAt: l.CreatedAt,
- User: UserFromServiceShallow(l.User),
- APIKey: APIKeyFromService(l.APIKey),
- Group: GroupFromServiceShallow(l.Group),
- Subscription: UserSubscriptionFromService(l.Subscription),
+ ID: l.ID,
+ UserID: l.UserID,
+ APIKeyID: l.APIKeyID,
+ AccountID: l.AccountID,
+ RequestID: l.RequestID,
+ Model: requestedModel,
+ ServiceTier: l.ServiceTier,
+ ReasoningEffort: l.ReasoningEffort,
+ InboundEndpoint: l.InboundEndpoint,
+ GroupID: l.GroupID,
+ SubscriptionID: l.SubscriptionID,
+ InputTokens: l.InputTokens,
+ OutputTokens: l.OutputTokens,
+ CacheCreationTokens: l.CacheCreationTokens,
+ CacheReadTokens: l.CacheReadTokens,
+ CacheCreation5mTokens: l.CacheCreation5mTokens,
+ CacheCreation1hTokens: l.CacheCreation1hTokens,
+ InputCost: l.InputCost,
+ OutputCost: l.OutputCost,
+ CacheCreationCost: l.CacheCreationCost,
+ CacheReadCost: l.CacheReadCost,
+ TotalCost: l.TotalCost,
+ ActualCost: l.ActualCost,
+ RateMultiplier: l.RateMultiplier,
+ LongContextBillingApplied: l.LongContextBillingApplied,
+ BillingType: l.BillingType,
+ RequestType: requestType.String(),
+ Stream: stream,
+ OpenAIWSMode: openAIWSMode,
+ DurationMs: l.DurationMs,
+ FirstTokenMs: l.FirstTokenMs,
+ ImageCount: l.ImageCount,
+ ImageSize: l.ImageSize,
+ ImageInputSize: l.ImageInputSize,
+ ImageOutputSize: l.ImageOutputSize,
+ ImageOutputTokens: l.ImageOutputTokens,
+ ImageOutputCost: l.ImageOutputCost,
+ ImageSizeSource: l.ImageSizeSource,
+ ImageSizeBreakdown: l.ImageSizeBreakdown,
+ MediaType: l.MediaType,
+ UserAgent: l.UserAgent,
+ IPAddress: l.IPAddress,
+ CacheTTLOverridden: l.CacheTTLOverridden,
+ BillingMode: l.BillingMode,
+ CreatedAt: l.CreatedAt,
+ User: UserFromServiceShallow(l.User),
+ APIKey: APIKeyFromService(l.APIKey),
+ Group: GroupFromServiceShallow(l.Group),
+ Subscription: UserSubscriptionFromService(l.Subscription),
}
}
diff --git a/backend/internal/handler/dto/types.go b/backend/internal/handler/dto/types.go
index 0418e5bc3e..619926c1e4 100644
--- a/backend/internal/handler/dto/types.go
+++ b/backend/internal/handler/dto/types.go
@@ -482,13 +482,14 @@ type UsageLog struct {
CacheCreation5mTokens int `json:"cache_creation_5m_tokens"`
CacheCreation1hTokens int `json:"cache_creation_1h_tokens"`
- InputCost float64 `json:"input_cost"`
- OutputCost float64 `json:"output_cost"`
- CacheCreationCost float64 `json:"cache_creation_cost"`
- CacheReadCost float64 `json:"cache_read_cost"`
- TotalCost float64 `json:"total_cost"`
- ActualCost float64 `json:"actual_cost"`
- RateMultiplier float64 `json:"rate_multiplier"`
+ InputCost float64 `json:"input_cost"`
+ OutputCost float64 `json:"output_cost"`
+ CacheCreationCost float64 `json:"cache_creation_cost"`
+ CacheReadCost float64 `json:"cache_read_cost"`
+ TotalCost float64 `json:"total_cost"`
+ ActualCost float64 `json:"actual_cost"`
+ RateMultiplier float64 `json:"rate_multiplier"`
+ LongContextBillingApplied bool `json:"long_context_billing_applied"`
BillingType int8 `json:"billing_type"`
RequestType string `json:"request_type"`
diff --git a/backend/internal/handler/openai_codex_models_handler.go b/backend/internal/handler/openai_codex_models_handler.go
index e64c555d14..1c1357cbfa 100644
--- a/backend/internal/handler/openai_codex_models_handler.go
+++ b/backend/internal/handler/openai_codex_models_handler.go
@@ -15,11 +15,13 @@ import (
// Codex CLI and the Codex desktop app refresh their model picker from
// GET {base_url}/models?client_version=... (custom provider mode) or
// GET /backend-api/codex/models (chatgpt_base_url mode). Both routes land
-// here. The manifest is proxied verbatim from the ChatGPT backend with a
-// schedulable OAuth account's credentials, so clients pointed at the gateway
-// see the account's real, always-current model entitlements instead of a
-// frozen local cache.
+// here. The manifest is proxied verbatim from the selected account's ChatGPT
+// backend or custom API key upstream. API key manifests use a short-lived,
+// asynchronously revalidated cache to tolerate canceled client requests.
func (h *OpenAIGatewayHandler) CodexModels(c *gin.Context) {
+ if c.Request.Context().Err() != nil {
+ return
+ }
apiKey, ok := middleware2.GetAPIKeyFromContext(c)
if !ok || apiKey.Group == nil {
h.errorResponse(c, http.StatusUnauthorized, "invalid_request_error", "API key group is required")
@@ -30,24 +32,54 @@ func (h *OpenAIGatewayHandler) CodexModels(c *gin.Context) {
return
}
- account, err := h.gatewayService.SelectAccountForModel(c.Request.Context(), apiKey.GroupID, "", "")
- if err != nil {
- h.errorResponse(c, http.StatusServiceUnavailable, "upstream_error", "No available OpenAI accounts")
- return
+ maxAccountSwitches := h.maxAccountSwitches
+ if maxAccountSwitches <= 0 {
+ maxAccountSwitches = 3
}
+ failedAccountIDs := make(map[int64]struct{})
+ switchCount := 0
+ var lastUpstreamErr error
- manifest, err := h.gatewayService.FetchCodexModelsManifest(c.Request.Context(), account, c.Query("client_version"), c.GetHeader("If-None-Match"))
- if err != nil {
- h.errorResponse(c, infraerrors.Code(err), "upstream_error", infraerrors.Message(err))
- return
- }
+ for {
+ account, err := h.gatewayService.SelectAccountForModelWithExclusions(c.Request.Context(), apiKey.GroupID, "", "", failedAccountIDs)
+ if err != nil {
+ if c.Request.Context().Err() != nil {
+ return
+ }
+ if lastUpstreamErr != nil {
+ h.errorResponse(c, infraerrors.Code(lastUpstreamErr), "upstream_error", infraerrors.Message(lastUpstreamErr))
+ return
+ }
+ h.errorResponse(c, http.StatusServiceUnavailable, "upstream_error", "No available OpenAI accounts")
+ return
+ }
- if manifest.ETag != "" {
- c.Header("ETag", manifest.ETag)
- }
- if manifest.NotModified {
- c.Status(http.StatusNotModified)
+ manifest, err := h.gatewayService.FetchCodexModelsManifest(c.Request.Context(), account, c.Query("client_version"), c.GetHeader("If-None-Match"))
+ if err != nil {
+ if c.Request.Context().Err() != nil {
+ return
+ }
+ if service.IsRetryableCodexModelsManifestError(err) && switchCount < maxAccountSwitches {
+ failedAccountIDs[account.ID] = struct{}{}
+ switchCount++
+ lastUpstreamErr = err
+ continue
+ }
+ h.errorResponse(c, infraerrors.Code(err), "upstream_error", infraerrors.Message(err))
+ return
+ }
+ if c.Request.Context().Err() != nil {
+ return
+ }
+
+ if manifest.ETag != "" {
+ c.Header("ETag", manifest.ETag)
+ }
+ if manifest.NotModified {
+ c.Status(http.StatusNotModified)
+ return
+ }
+ c.Data(http.StatusOK, "application/json", manifest.Body)
return
}
- c.Data(http.StatusOK, "application/json", manifest.Body)
}
diff --git a/backend/internal/handler/openai_codex_models_handler_test.go b/backend/internal/handler/openai_codex_models_handler_test.go
new file mode 100644
index 0000000000..ba74a5869f
--- /dev/null
+++ b/backend/internal/handler/openai_codex_models_handler_test.go
@@ -0,0 +1,288 @@
+package handler
+
+import (
+ "context"
+ "errors"
+ "fmt"
+ "io"
+ "net"
+ "net/http"
+ "net/http/httptest"
+ "strings"
+ "sync"
+ "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"
+)
+
+type codexModelsFailoverAccountRepo struct {
+ service.AccountRepository
+ accounts []service.Account
+}
+
+func (r codexModelsFailoverAccountRepo) GetByID(_ context.Context, id int64) (*service.Account, error) {
+ for i := range r.accounts {
+ if r.accounts[i].ID == id {
+ account := r.accounts[i]
+ return &account, nil
+ }
+ }
+ return nil, service.ErrNoAvailableAccounts
+}
+
+func (r codexModelsFailoverAccountRepo) ListSchedulableByPlatform(_ context.Context, platform string) ([]service.Account, error) {
+ accounts := make([]service.Account, 0, len(r.accounts))
+ for _, account := range r.accounts {
+ if account.Platform == platform {
+ accounts = append(accounts, account)
+ }
+ }
+ return accounts, nil
+}
+
+type codexModelsFailoverHTTPUpstream struct {
+ service.HTTPUpstream
+ mu sync.Mutex
+ accountIDs []int64
+ firstErr error
+ firstStatus int
+ statuses map[int64]int
+}
+
+func (u *codexModelsFailoverHTTPUpstream) Do(_ *http.Request, _ string, accountID int64, _ int) (*http.Response, error) {
+ u.mu.Lock()
+ u.accountIDs = append(u.accountIDs, accountID)
+ u.mu.Unlock()
+
+ status, hasStatus := u.statuses[accountID]
+ if accountID == 1 || hasStatus {
+ if u.firstErr != nil {
+ return nil, u.firstErr
+ }
+ if !hasStatus {
+ status = u.firstStatus
+ }
+ if status == 0 {
+ status = http.StatusServiceUnavailable
+ }
+ return &http.Response{
+ StatusCode: status,
+ Status: http.StatusText(status),
+ Header: make(http.Header),
+ Body: io.NopCloser(strings.NewReader(
+ `{"error":{"message":"No available OpenAI accounts","type":"upstream_error"}}`,
+ )),
+ }, nil
+ }
+ return &http.Response{
+ StatusCode: http.StatusOK,
+ Status: "200 OK",
+ Header: make(http.Header),
+ Body: io.NopCloser(strings.NewReader(`{"models":[{"slug":"gpt-5.6-sol"}]}`)),
+ }, nil
+}
+
+func (u *codexModelsFailoverHTTPUpstream) calls() []int64 {
+ u.mu.Lock()
+ defer u.mu.Unlock()
+ return append([]int64(nil), u.accountIDs...)
+}
+
+func TestCodexModelsCanceledRequestDoesNotWriteResponse(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ recorder := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(recorder)
+ ctx, cancel := context.WithCancel(context.Background())
+ cancel()
+ c.Request = httptest.NewRequest(http.MethodGet, "/v1/models", nil).WithContext(ctx)
+
+ h := &OpenAIGatewayHandler{}
+ h.CodexModels(c)
+
+ if c.Writer.Written() {
+ t.Fatalf("canceled request wrote an HTTP response: status=%d body=%q", recorder.Code, recorder.Body.String())
+ }
+}
+
+func TestCodexModelsFailsOverFromRetryableUpstreamStatus(t *testing.T) {
+ retryableStatuses := []int{
+ http.StatusTooManyRequests,
+ http.StatusInternalServerError,
+ http.StatusBadGateway,
+ http.StatusServiceUnavailable,
+ http.StatusGatewayTimeout,
+ }
+ for _, status := range retryableStatuses {
+ t.Run(http.StatusText(status), func(t *testing.T) {
+ handler, upstream, groupID := newCodexModelsFailoverTestHandler(status)
+ recorder := performCodexModelsRequest(t, handler, groupID)
+
+ if got, want := upstream.calls(), []int64{1, 2}; !equalInt64Slices(got, want) {
+ t.Fatalf("upstream account calls: got %v, want %v", got, want)
+ }
+ if recorder.Code != http.StatusOK {
+ t.Fatalf("status: got %d, want %d; body=%s", recorder.Code, http.StatusOK, recorder.Body.String())
+ }
+ if got, want := recorder.Body.String(), `{"models":[{"slug":"gpt-5.6-sol"}]}`; got != want {
+ t.Fatalf("body: got %q, want %q", got, want)
+ }
+ })
+ }
+}
+
+func TestCodexModelsFailsOverFromUpstreamTransportError(t *testing.T) {
+ handler, upstream, groupID := newCodexModelsFailoverTestHandler(http.StatusServiceUnavailable)
+ upstream.firstErr = &net.OpError{
+ Op: "read",
+ Net: "tcp",
+ Err: errors.New("connection reset"),
+ }
+ recorder := performCodexModelsRequest(t, handler, groupID)
+
+ if got, want := upstream.calls(), []int64{1, 2}; !equalInt64Slices(got, want) {
+ t.Fatalf("upstream account calls: got %v, want %v", got, want)
+ }
+ if recorder.Code != http.StatusOK {
+ t.Fatalf("status: got %d, want %d; body=%s", recorder.Code, http.StatusOK, recorder.Body.String())
+ }
+}
+
+func TestCodexModelsDoesNotFailOverFromPermanentUpstreamStatus(t *testing.T) {
+ statuses := []int{
+ http.StatusBadRequest,
+ http.StatusUnauthorized,
+ http.StatusForbidden,
+ http.StatusNotFound,
+ 600,
+ }
+ for _, status := range statuses {
+ t.Run(fmt.Sprintf("status_%d", status), func(t *testing.T) {
+ handler, upstream, groupID := newCodexModelsFailoverTestHandler(status)
+ recorder := performCodexModelsRequest(t, handler, groupID)
+
+ if got, want := upstream.calls(), []int64{1}; !equalInt64Slices(got, want) {
+ t.Fatalf("upstream account calls: got %v, want %v", got, want)
+ }
+ if recorder.Code != http.StatusBadGateway {
+ t.Fatalf("status: got %d, want %d; body=%s", recorder.Code, http.StatusBadGateway, recorder.Body.String())
+ }
+ })
+ }
+}
+
+func TestCodexModelsDoesNotFailOverFromUpstreamConfigurationError(t *testing.T) {
+ handler, upstream, groupID := newCodexModelsFailoverTestHandler(http.StatusServiceUnavailable)
+ upstream.firstErr = errors.New("invalid proxy URL")
+ recorder := performCodexModelsRequest(t, handler, groupID)
+
+ if got, want := upstream.calls(), []int64{1}; !equalInt64Slices(got, want) {
+ t.Fatalf("upstream account calls: got %v, want %v", got, want)
+ }
+ if recorder.Code != http.StatusBadGateway {
+ t.Fatalf("status: got %d, want %d; body=%s", recorder.Code, http.StatusBadGateway, recorder.Body.String())
+ }
+}
+
+func TestCodexModelsReturnsLastUpstreamErrorWhenAccountsAreExhausted(t *testing.T) {
+ handler, upstream, groupID := newCodexModelsFailoverTestHandler(http.StatusServiceUnavailable)
+ upstream.statuses = map[int64]int{
+ 1: http.StatusServiceUnavailable,
+ 2: http.StatusGatewayTimeout,
+ }
+ recorder := performCodexModelsRequest(t, handler, groupID)
+
+ if got, want := upstream.calls(), []int64{1, 2}; !equalInt64Slices(got, want) {
+ t.Fatalf("upstream account calls: got %v, want %v", got, want)
+ }
+ if recorder.Code != http.StatusBadGateway {
+ t.Fatalf("status: got %d, want %d; body=%s", recorder.Code, http.StatusBadGateway, recorder.Body.String())
+ }
+ if body := recorder.Body.String(); !strings.Contains(body, "upstream error 504") {
+ t.Fatalf("body does not preserve the last upstream error: %s", body)
+ }
+}
+
+func TestCodexModelsHonorsAccountSwitchLimit(t *testing.T) {
+ handler, upstream, groupID := newCodexModelsFailoverTestHandlerWithAccountCount(http.StatusServiceUnavailable, 4, 2)
+ upstream.statuses = map[int64]int{
+ 1: http.StatusServiceUnavailable,
+ 2: http.StatusBadGateway,
+ 3: http.StatusGatewayTimeout,
+ 4: http.StatusInternalServerError,
+ }
+ recorder := performCodexModelsRequest(t, handler, groupID)
+
+ if got, want := upstream.calls(), []int64{1, 2, 3}; !equalInt64Slices(got, want) {
+ t.Fatalf("upstream account calls: got %v, want %v", got, want)
+ }
+ if recorder.Code != http.StatusBadGateway {
+ t.Fatalf("status: got %d, want %d; body=%s", recorder.Code, http.StatusBadGateway, recorder.Body.String())
+ }
+ if body := recorder.Body.String(); !strings.Contains(body, "upstream error 504") {
+ t.Fatalf("body does not preserve the limit-ending upstream error: %s", body)
+ }
+}
+
+func newCodexModelsFailoverTestHandler(firstStatus int) (*OpenAIGatewayHandler, *codexModelsFailoverHTTPUpstream, int64) {
+ return newCodexModelsFailoverTestHandlerWithAccountCount(firstStatus, 2, 3)
+}
+
+func newCodexModelsFailoverTestHandlerWithAccountCount(firstStatus, accountCount, maxSwitches int) (*OpenAIGatewayHandler, *codexModelsFailoverHTTPUpstream, int64) {
+ gin.SetMode(gin.TestMode)
+ groupID := int64(42)
+ accounts := make([]service.Account, 0, accountCount)
+ for i := 1; i <= accountCount; i++ {
+ accounts = append(accounts, service.Account{
+ ID: int64(i),
+ Name: fmt.Sprintf("upstream-%d", i),
+ Platform: service.PlatformOpenAI,
+ Type: service.AccountTypeAPIKey,
+ Status: service.StatusActive,
+ Schedulable: true,
+ Priority: i - 1,
+ Concurrency: 1,
+ Credentials: map[string]any{
+ "api_key": fmt.Sprintf("sk-%d", i),
+ "base_url": fmt.Sprintf("https://upstream-%d.example/v1", i),
+ },
+ })
+ }
+ upstream := &codexModelsFailoverHTTPUpstream{firstStatus: firstStatus}
+ cfg := &config.Config{RunMode: config.RunModeSimple}
+ gatewayService := service.NewOpenAIGatewayService(
+ codexModelsFailoverAccountRepo{accounts: accounts},
+ nil, nil, nil, nil, nil, nil, cfg, nil, nil, nil, nil, nil,
+ upstream,
+ nil, nil, nil, nil, nil, nil, nil, nil,
+ )
+ return &OpenAIGatewayHandler{gatewayService: gatewayService, maxAccountSwitches: maxSwitches}, upstream, groupID
+}
+
+func performCodexModelsRequest(t *testing.T, handler *OpenAIGatewayHandler, groupID int64) *httptest.ResponseRecorder {
+ t.Helper()
+ recorder := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(recorder)
+ c.Request = httptest.NewRequest(http.MethodGet, "/v1/models?client_version=0.144.0", nil)
+ c.Set(string(middleware2.ContextKeyAPIKey), &service.APIKey{
+ GroupID: &groupID,
+ Group: &service.Group{ID: groupID, Platform: service.PlatformOpenAI},
+ })
+
+ handler.CodexModels(c)
+ return recorder
+}
+
+func equalInt64Slices(got, want []int64) bool {
+ if len(got) != len(want) {
+ return false
+ }
+ for i := range got {
+ if got[i] != want[i] {
+ return false
+ }
+ }
+ return true
+}
diff --git a/backend/internal/handler/openai_gateway_handler_test.go b/backend/internal/handler/openai_gateway_handler_test.go
index 781c16b392..e4b594c0b8 100644
--- a/backend/internal/handler/openai_gateway_handler_test.go
+++ b/backend/internal/handler/openai_gateway_handler_test.go
@@ -415,11 +415,13 @@ func TestResolveOpenAIMessagesDispatchMappedModel(t *testing.T) {
SonnetMappedModel: "gpt-5.2",
ExactModelMappings: map[string]string{
"claude-sonnet-4-5-20250929": "gpt-5.4-mini-high",
+ "claude-fable-5": "gpt-5.6-sol",
},
},
},
}
require.Equal(t, "gpt-5.4-mini", resolveOpenAIMessagesDispatchMappedModel(apiKey, "claude-sonnet-4-5-20250929"))
+ require.Equal(t, "gpt-5.6-sol", resolveOpenAIMessagesDispatchMappedModel(apiKey, "claude-fable-5"))
})
t.Run("uses_family_default_when_no_override", func(t *testing.T) {
diff --git a/backend/internal/pkg/antigravity/client.go b/backend/internal/pkg/antigravity/client.go
index e318d1cdaf..39b6d2c90c 100644
--- a/backend/internal/pkg/antigravity/client.go
+++ b/backend/internal/pkg/antigravity/client.go
@@ -17,6 +17,7 @@ import (
"github.com/Wei-Shaw/sub2api/internal/pkg/proxyurl"
"github.com/Wei-Shaw/sub2api/internal/pkg/proxyutil"
+ "github.com/Wei-Shaw/sub2api/internal/pkg/servertiming"
)
// ForbiddenError 表示上游返回 403 Forbidden
@@ -279,7 +280,6 @@ func NewClient(proxyURL string) (*Client, error) {
}
client.Transport = transport
}
-
return &Client{
httpClient: client,
}, nil
@@ -341,7 +341,7 @@ func (c *Client) ExchangeCode(ctx context.Context, code, codeVerifier string) (*
}
req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
- resp, err := c.httpClient.Do(req)
+ resp, err := servertiming.Do(c.httpClient, req)
if err != nil {
return nil, fmt.Errorf("token 交换请求失败: %w", err)
}
@@ -383,7 +383,7 @@ func (c *Client) RefreshToken(ctx context.Context, refreshToken string) (*TokenR
}
req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
- resp, err := c.httpClient.Do(req)
+ resp, err := servertiming.Do(c.httpClient, req)
if err != nil {
return nil, fmt.Errorf("token 刷新请求失败: %w", err)
}
@@ -414,7 +414,7 @@ func (c *Client) GetUserInfo(ctx context.Context, accessToken string) (*UserInfo
}
req.Header.Set("Authorization", "Bearer "+accessToken)
- resp, err := c.httpClient.Do(req)
+ resp, err := servertiming.Do(c.httpClient, req)
if err != nil {
return nil, fmt.Errorf("用户信息请求失败: %w", err)
}
@@ -465,7 +465,7 @@ func (c *Client) LoadCodeAssist(ctx context.Context, accessToken string) (*LoadC
req.Header.Set("Content-Type", "application/json")
req.Header.Set("User-Agent", GetUserAgentForContext(ctx))
- resp, err := c.httpClient.Do(req)
+ resp, err := servertiming.Do(c.httpClient, req)
if err != nil {
lastErr = fmt.Errorf("loadCodeAssist 请求失败: %w", err)
if shouldFallbackToNextURL(err, 0) && urlIdx < len(availableURLs)-1 {
@@ -544,7 +544,7 @@ func (c *Client) OnboardUser(ctx context.Context, accessToken, tierID string) (s
req.Header.Set("Content-Type", "application/json")
req.Header.Set("User-Agent", GetUserAgentForContext(ctx))
- resp, err := c.httpClient.Do(req)
+ resp, err := servertiming.Do(c.httpClient, req)
if err != nil {
lastErr = fmt.Errorf("onboardUser 请求失败: %w", err)
if shouldFallbackToNextURL(err, 0) && urlIdx < len(availableURLs)-1 {
@@ -683,7 +683,7 @@ func (c *Client) FetchAvailableModels(ctx context.Context, accessToken, projectI
req.Header.Set("Content-Type", "application/json")
req.Header.Set("User-Agent", GetUserAgentForContext(ctx))
- resp, err := fetchClient.Do(req)
+ resp, err := servertiming.Do(fetchClient, req)
if err != nil {
lastErr = fmt.Errorf("fetchAvailableModels 请求失败: %w", err)
if shouldFallbackToNextURL(err, 0) && urlIdx < len(availableURLs)-1 {
@@ -842,7 +842,7 @@ func (c *Client) SetUserSettings(ctx context.Context, accessToken string) (*SetU
req.Header.Set("X-Goog-Api-Client", "gl-node/22.21.1")
req.Host = "daily-cloudcode-pa.googleapis.com"
- resp, err := c.httpClient.Do(req)
+ resp, err := servertiming.Do(c.httpClient, req)
if err != nil {
return nil, fmt.Errorf("setUserSettings 请求失败: %w", err)
}
@@ -885,7 +885,7 @@ func (c *Client) FetchUserInfo(ctx context.Context, accessToken, projectID strin
req.Header.Set("X-Goog-Api-Client", "gl-node/22.21.1")
req.Host = "daily-cloudcode-pa.googleapis.com"
- resp, err := c.httpClient.Do(req)
+ resp, err := servertiming.Do(c.httpClient, req)
if err != nil {
return nil, fmt.Errorf("fetchUserInfo 请求失败: %w", err)
}
diff --git a/backend/internal/pkg/httpclient/pool.go b/backend/internal/pkg/httpclient/pool.go
index 12804cc67d..22d3c65feb 100644
--- a/backend/internal/pkg/httpclient/pool.go
+++ b/backend/internal/pkg/httpclient/pool.go
@@ -25,6 +25,7 @@ import (
"github.com/Wei-Shaw/sub2api/internal/pkg/proxyurl"
"github.com/Wei-Shaw/sub2api/internal/pkg/proxyutil"
+ "github.com/Wei-Shaw/sub2api/internal/pkg/servertiming"
"github.com/Wei-Shaw/sub2api/internal/util/urlvalidator"
)
@@ -92,6 +93,7 @@ func buildClient(opts Options) (*http.Client, error) {
if opts.ValidateResolvedIP && !opts.AllowPrivateHosts {
rt = newValidatedTransport(transport)
}
+ rt = servertiming.WrapRoundTripper(rt)
return &http.Client{
Transport: rt,
Timeout: opts.Timeout,
diff --git a/backend/internal/pkg/servertiming/collector.go b/backend/internal/pkg/servertiming/collector.go
new file mode 100644
index 0000000000..553edede31
--- /dev/null
+++ b/backend/internal/pkg/servertiming/collector.go
@@ -0,0 +1,348 @@
+package servertiming
+
+import (
+ "context"
+ "fmt"
+ "sort"
+ "strconv"
+ "strings"
+ "sync"
+ "time"
+)
+
+const (
+ HeaderName = "Server-Timing"
+ AdminUIHeader = "X-Admin-UI-Request"
+ MetricDatabase = "db"
+ MetricRedis = "redis"
+ dependencyPrefix = "dep_"
+
+ maxMetricNameLength = 48
+ maxIntervals = 2048
+ maxHeaderLength = 4096
+)
+
+type contextKey struct{}
+
+type interval struct {
+ start time.Time
+ end time.Time
+}
+
+type metric struct {
+ count int64
+ intervals []interval
+}
+
+// Collector stores request-scoped timing samples. It is safe for concurrent use.
+type Collector struct {
+ startedAt time.Time
+
+ mu sync.Mutex
+ metrics map[string]*metric
+ cacheStatus string
+}
+
+// New creates a collector whose total duration starts at startedAt.
+func New(startedAt time.Time) *Collector {
+ if startedAt.IsZero() {
+ startedAt = time.Now()
+ }
+ return &Collector{
+ startedAt: startedAt,
+ metrics: make(map[string]*metric),
+ }
+}
+
+// WithCollector attaches a collector to a context.
+func WithCollector(ctx context.Context, collector *Collector) context.Context {
+ if ctx == nil {
+ ctx = context.Background()
+ }
+ if collector == nil {
+ return ctx
+ }
+ return context.WithValue(ctx, contextKey{}, collector)
+}
+
+// FromContext returns the request timing collector, when one is active.
+func FromContext(ctx context.Context) (*Collector, bool) {
+ if ctx == nil {
+ return nil, false
+ }
+ collector, ok := ctx.Value(contextKey{}).(*Collector)
+ return collector, ok && collector != nil
+}
+
+// Active reports whether timing collection is enabled for this request.
+func Active(ctx context.Context) bool {
+ _, ok := FromContext(ctx)
+ return ok
+}
+
+// Record adds a completed interval and operation count to a metric.
+func Record(ctx context.Context, name string, startedAt, endedAt time.Time, count int) {
+ collector, ok := FromContext(ctx)
+ if !ok {
+ return
+ }
+ collector.Record(name, startedAt, endedAt, count)
+}
+
+// RecordInterval adds timing without incrementing the operation count. It is
+// useful when one logical operation has multiple blocking driver calls.
+func RecordInterval(ctx context.Context, name string, startedAt, endedAt time.Time) {
+ collector, ok := FromContext(ctx)
+ if !ok {
+ return
+ }
+ collector.record(name, startedAt, endedAt, 0)
+}
+
+// Record adds a completed interval directly to the collector.
+func (c *Collector) Record(name string, startedAt, endedAt time.Time, count int) {
+ if count <= 0 {
+ count = 1
+ }
+ c.record(name, startedAt, endedAt, count)
+}
+
+func (c *Collector) record(name string, startedAt, endedAt time.Time, count int) {
+ name = normalizeMetricName(name)
+ if c == nil || name == "" || startedAt.IsZero() || endedAt.Before(startedAt) {
+ return
+ }
+ if count < 0 {
+ count = 0
+ }
+
+ c.mu.Lock()
+ m := c.metrics[name]
+ if m == nil {
+ m = &metric{}
+ c.metrics[name] = m
+ }
+ m.count += int64(count)
+ if len(m.intervals) < maxIntervals {
+ m.intervals = append(m.intervals, interval{start: startedAt, end: endedAt})
+ }
+ c.mu.Unlock()
+}
+
+// Observe starts a metric span and returns an idempotent completion function.
+func Observe(ctx context.Context, name string) func() {
+ collector, ok := FromContext(ctx)
+ name = normalizeMetricName(name)
+ if !ok || name == "" {
+ return func() {}
+ }
+ startedAt := time.Now()
+ var once sync.Once
+ return func() {
+ once.Do(func() {
+ collector.Record(name, startedAt, time.Now(), 1)
+ })
+ }
+}
+
+// ObserveDependency starts a named external dependency span.
+func ObserveDependency(ctx context.Context, module string) func() {
+ return Observe(ctx, dependencyMetricName(module))
+}
+
+// RecordDependency records a completed external dependency interval.
+func RecordDependency(ctx context.Context, module string, startedAt, endedAt time.Time) {
+ Record(ctx, dependencyMetricName(module), startedAt, endedAt, 1)
+}
+
+// SetCacheStatus records the response-cache outcome for the request.
+func SetCacheStatus(ctx context.Context, status string) {
+ collector, ok := FromContext(ctx)
+ if !ok {
+ return
+ }
+ status = normalizeCacheStatus(status)
+ if status == "" {
+ return
+ }
+ collector.mu.Lock()
+ collector.cacheStatus = status
+ collector.mu.Unlock()
+}
+
+// HeaderValue renders a bounded, deterministic Server-Timing header.
+func HeaderValue(ctx context.Context, endedAt time.Time, cacheStatus string) string {
+ collector, ok := FromContext(ctx)
+ if !ok {
+ return ""
+ }
+ return collector.HeaderValue(endedAt, cacheStatus)
+}
+
+// HeaderValue renders a bounded, deterministic Server-Timing header.
+func (c *Collector) HeaderValue(endedAt time.Time, cacheStatus string) string {
+ if c == nil {
+ return ""
+ }
+ if endedAt.IsZero() {
+ endedAt = time.Now()
+ }
+ if endedAt.Before(c.startedAt) {
+ endedAt = c.startedAt
+ }
+
+ c.mu.Lock()
+ metrics := make(map[string]metric, len(c.metrics))
+ allIntervals := make([]interval, 0)
+ dependencyIntervals := make([]interval, 0)
+ var dependencyCount int64
+ for name, source := range c.metrics {
+ copied := metric{count: source.count, intervals: append([]interval(nil), source.intervals...)}
+ metrics[name] = copied
+ allIntervals = append(allIntervals, copied.intervals...)
+ if strings.HasPrefix(name, dependencyPrefix) {
+ dependencyIntervals = append(dependencyIntervals, copied.intervals...)
+ dependencyCount += copied.count
+ }
+ }
+ storedCacheStatus := c.cacheStatus
+ c.mu.Unlock()
+
+ total := endedAt.Sub(c.startedAt)
+ blocked := unionDuration(allIntervals, c.startedAt, endedAt)
+ app := total - blocked
+ if app < 0 {
+ app = 0
+ }
+
+ cacheStatus = normalizeCacheStatus(cacheStatus)
+ if cacheStatus == "" {
+ cacheStatus = normalizeCacheStatus(storedCacheStatus)
+ }
+ if cacheStatus == "" {
+ cacheStatus = "bypass"
+ }
+
+ database := metrics[MetricDatabase]
+ redisMetric := metrics[MetricRedis]
+ parts := []string{
+ "total;dur=" + formatDuration(total),
+ "app;dur=" + formatDuration(app),
+ fmt.Sprintf("db;dur=%s;desc=\"queries=%d\"", formatDuration(unionDuration(database.intervals, c.startedAt, endedAt)), database.count),
+ fmt.Sprintf("redis;dur=%s;desc=\"commands=%d\"", formatDuration(unionDuration(redisMetric.intervals, c.startedAt, endedAt)), redisMetric.count),
+ "cache;desc=\"" + cacheStatus + "\"",
+ fmt.Sprintf("deps;dur=%s;desc=\"calls=%d\"", formatDuration(unionDuration(dependencyIntervals, c.startedAt, endedAt)), dependencyCount),
+ }
+
+ dependencyNames := make([]string, 0)
+ for name := range metrics {
+ if strings.HasPrefix(name, dependencyPrefix) {
+ dependencyNames = append(dependencyNames, name)
+ }
+ }
+ sort.Strings(dependencyNames)
+ for _, name := range dependencyNames {
+ m := metrics[name]
+ part := fmt.Sprintf("%s;dur=%s;desc=\"calls=%d\"", name, formatDuration(unionDuration(m.intervals, c.startedAt, endedAt)), m.count)
+ candidate := strings.Join(append(parts, part), ", ")
+ if len(candidate) > maxHeaderLength {
+ break
+ }
+ parts = append(parts, part)
+ }
+
+ return strings.Join(parts, ", ")
+}
+
+func dependencyMetricName(module string) string {
+ module = normalizeMetricName(module)
+ module = strings.TrimPrefix(module, dependencyPrefix)
+ if module == "" {
+ module = "http"
+ }
+ return dependencyPrefix + module
+}
+
+func normalizeMetricName(name string) string {
+ name = strings.ToLower(strings.TrimSpace(name))
+ if name == "" {
+ return ""
+ }
+ var b strings.Builder
+ b.Grow(min(len(name), maxMetricNameLength))
+ for _, r := range name {
+ if b.Len() >= maxMetricNameLength {
+ break
+ }
+ switch {
+ case r >= 'a' && r <= 'z', r >= '0' && r <= '9':
+ _, _ = b.WriteRune(r)
+ case r == '_' || r == '-':
+ _ = b.WriteByte('_')
+ }
+ }
+ return strings.Trim(b.String(), "_")
+}
+
+func normalizeCacheStatus(status string) string {
+ switch strings.ToLower(strings.TrimSpace(status)) {
+ case "hit":
+ return "hit"
+ case "miss":
+ return "miss"
+ case "bypass":
+ return "bypass"
+ default:
+ return ""
+ }
+}
+
+func unionDuration(intervals []interval, lowerBound, upperBound time.Time) time.Duration {
+ if len(intervals) == 0 || !upperBound.After(lowerBound) {
+ return 0
+ }
+ normalized := make([]interval, 0, len(intervals))
+ for _, item := range intervals {
+ start := item.start
+ end := item.end
+ if start.Before(lowerBound) {
+ start = lowerBound
+ }
+ if end.After(upperBound) {
+ end = upperBound
+ }
+ if end.After(start) {
+ normalized = append(normalized, interval{start: start, end: end})
+ }
+ }
+ if len(normalized) == 0 {
+ return 0
+ }
+ sort.Slice(normalized, func(i, j int) bool {
+ return normalized[i].start.Before(normalized[j].start)
+ })
+
+ currentStart := normalized[0].start
+ currentEnd := normalized[0].end
+ var total time.Duration
+ for _, item := range normalized[1:] {
+ if !item.start.After(currentEnd) {
+ if item.end.After(currentEnd) {
+ currentEnd = item.end
+ }
+ continue
+ }
+ total += currentEnd.Sub(currentStart)
+ currentStart = item.start
+ currentEnd = item.end
+ }
+ total += currentEnd.Sub(currentStart)
+ return total
+}
+
+func formatDuration(value time.Duration) string {
+ if value < 0 {
+ value = 0
+ }
+ return strconv.FormatFloat(float64(value)/float64(time.Millisecond), 'f', 1, 64)
+}
diff --git a/backend/internal/pkg/servertiming/collector_test.go b/backend/internal/pkg/servertiming/collector_test.go
new file mode 100644
index 0000000000..1bb809f9bc
--- /dev/null
+++ b/backend/internal/pkg/servertiming/collector_test.go
@@ -0,0 +1,129 @@
+package servertiming
+
+import (
+ "context"
+ "fmt"
+ "strings"
+ "sync"
+ "testing"
+ "time"
+)
+
+func TestCollectorHeaderValueAggregatesIntervals(t *testing.T) {
+ startedAt := time.Unix(100, 0)
+ collector := New(startedAt)
+ collector.Record(MetricDatabase, startedAt.Add(10*time.Millisecond), startedAt.Add(40*time.Millisecond), 2)
+ collector.Record(MetricRedis, startedAt.Add(30*time.Millisecond), startedAt.Add(50*time.Millisecond), 3)
+ collector.Record(dependencyMetricName("openai"), startedAt.Add(70*time.Millisecond), startedAt.Add(100*time.Millisecond), 1)
+ collector.Record(dependencyMetricName("github"), startedAt.Add(60*time.Millisecond), startedAt.Add(90*time.Millisecond), 1)
+
+ got := collector.HeaderValue(startedAt.Add(120*time.Millisecond), "miss")
+ want := `total;dur=120.0, app;dur=40.0, db;dur=30.0;desc="queries=2", redis;dur=20.0;desc="commands=3", cache;desc="miss", deps;dur=40.0;desc="calls=2", dep_github;dur=30.0;desc="calls=1", dep_openai;dur=30.0;desc="calls=1"`
+ if got != want {
+ t.Fatalf("HeaderValue() = %q, want %q", got, want)
+ }
+}
+
+func TestRecordIntervalDoesNotIncrementCount(t *testing.T) {
+ startedAt := time.Unix(200, 0)
+ collector := New(startedAt)
+ ctx := WithCollector(context.Background(), collector)
+
+ Record(ctx, MetricDatabase, startedAt.Add(10*time.Millisecond), startedAt.Add(20*time.Millisecond), 1)
+ RecordInterval(ctx, MetricDatabase, startedAt.Add(30*time.Millisecond), startedAt.Add(40*time.Millisecond))
+
+ header := HeaderValue(ctx, startedAt.Add(100*time.Millisecond), "hit")
+ if !strings.Contains(header, `db;dur=20.0;desc="queries=1"`) {
+ t.Fatalf("header %q does not contain one query with both blocking intervals", header)
+ }
+ if !strings.Contains(header, "app;dur=80.0") {
+ t.Fatalf("header %q does not subtract the interval union from app time", header)
+ }
+}
+
+func TestCollectorCacheStatusFallback(t *testing.T) {
+ startedAt := time.Unix(300, 0)
+ collector := New(startedAt)
+ ctx := WithCollector(context.Background(), collector)
+
+ SetCacheStatus(ctx, " HIT ")
+ if got := HeaderValue(ctx, startedAt.Add(time.Millisecond), "invalid"); !strings.Contains(got, `cache;desc="hit"`) {
+ t.Fatalf("HeaderValue() = %q, want stored cache hit", got)
+ }
+
+ other := New(startedAt)
+ if got := other.HeaderValue(startedAt.Add(time.Millisecond), "invalid"); !strings.Contains(got, `cache;desc="bypass"`) {
+ t.Fatalf("HeaderValue() = %q, want cache bypass", got)
+ }
+}
+
+func TestCollectorSanitizesDependencyMetric(t *testing.T) {
+ startedAt := time.Unix(400, 0)
+ collector := New(startedAt)
+ ctx := WithCollector(context.Background(), collector)
+ RecordDependency(ctx, "GitHub API\r\nInjected;dur=999", startedAt, startedAt.Add(time.Millisecond))
+
+ header := HeaderValue(ctx, startedAt.Add(2*time.Millisecond), "bypass")
+ if strings.ContainsAny(header, "\r\n") || strings.Contains(header, ";dur=999") {
+ t.Fatalf("unsafe metric content reached header: %q", header)
+ }
+ if !strings.Contains(header, "dep_githubapiinjecteddur999;dur=1.0") {
+ t.Fatalf("sanitized dependency metric missing from header: %q", header)
+ }
+}
+
+func TestCollectorBoundsHeaderLength(t *testing.T) {
+ startedAt := time.Unix(500, 0)
+ collector := New(startedAt)
+ for i := 0; i < 300; i++ {
+ collector.Record(
+ dependencyMetricName(fmt.Sprintf("module_%03d_with_a_deliberately_long_name", i)),
+ startedAt,
+ startedAt.Add(time.Millisecond),
+ 1,
+ )
+ }
+
+ header := collector.HeaderValue(startedAt.Add(2*time.Millisecond), "bypass")
+ if len(header) > maxHeaderLength {
+ t.Fatalf("header length = %d, want <= %d", len(header), maxHeaderLength)
+ }
+ if !strings.Contains(header, "total;dur=2.0") || !strings.Contains(header, "deps;dur=1.0") {
+ t.Fatalf("bounded header lost fixed metrics: %q", header)
+ }
+}
+
+func TestCollectorConcurrentRecording(t *testing.T) {
+ startedAt := time.Now()
+ collector := New(startedAt)
+ ctx := WithCollector(context.Background(), collector)
+
+ const workers = 25
+ const recordsPerWorker = 100
+ var wg sync.WaitGroup
+ wg.Add(workers)
+ for i := 0; i < workers; i++ {
+ go func() {
+ defer wg.Done()
+ for j := 0; j < recordsPerWorker; j++ {
+ Record(ctx, MetricDatabase, startedAt, startedAt.Add(time.Microsecond), 1)
+ }
+ }()
+ }
+ wg.Wait()
+
+ header := HeaderValue(ctx, startedAt.Add(time.Millisecond), "bypass")
+ want := fmt.Sprintf(`queries=%d`, workers*recordsPerWorker)
+ if !strings.Contains(header, want) {
+ t.Fatalf("header %q does not contain %q", header, want)
+ }
+}
+
+func TestContextHelpersHandleMissingCollector(t *testing.T) {
+ if Active(context.Background()) {
+ t.Fatal("context without collector reported active")
+ }
+ if got := HeaderValue(context.Background(), time.Now(), "hit"); got != "" {
+ t.Fatalf("HeaderValue() = %q without collector, want empty", got)
+ }
+}
diff --git a/backend/internal/pkg/servertiming/http.go b/backend/internal/pkg/servertiming/http.go
new file mode 100644
index 0000000000..e326e24302
--- /dev/null
+++ b/backend/internal/pkg/servertiming/http.go
@@ -0,0 +1,104 @@
+package servertiming
+
+import (
+ "context"
+ "net/http"
+ "strings"
+ "time"
+)
+
+type dependencyModuleKey struct{}
+
+type timingRoundTripper struct {
+ base http.RoundTripper
+}
+
+// WithDependencyModule overrides the safe module name used for an outbound call.
+func WithDependencyModule(ctx context.Context, module string) context.Context {
+ if ctx == nil {
+ ctx = context.Background()
+ }
+ module = strings.TrimPrefix(normalizeMetricName(module), dependencyPrefix)
+ if module == "" {
+ return ctx
+ }
+ return context.WithValue(ctx, dependencyModuleKey{}, module)
+}
+
+// WrapRoundTripper records outbound response-header latency for active requests.
+func WrapRoundTripper(base http.RoundTripper) http.RoundTripper {
+ if base == nil {
+ base = http.DefaultTransport
+ }
+ if _, ok := base.(*timingRoundTripper); ok {
+ return base
+ }
+ return &timingRoundTripper{base: base}
+}
+
+// InstrumentClient returns a shallow client copy with an instrumented transport.
+func InstrumentClient(client *http.Client) *http.Client {
+ if client == nil {
+ client = &http.Client{}
+ }
+ copyClient := *client
+ copyClient.Transport = WrapRoundTripper(copyClient.Transport)
+ return ©Client
+}
+
+// Do records response-header latency without changing the client's transport
+// type. Use it for clients whose callers inspect or configure *http.Transport.
+func Do(client *http.Client, req *http.Request) (*http.Response, error) {
+ if client == nil {
+ client = http.DefaultClient
+ }
+ if req == nil || !Active(req.Context()) {
+ return client.Do(req)
+ }
+ startedAt := time.Now()
+ response, err := client.Do(req)
+ RecordDependency(req.Context(), dependencyModule(req), startedAt, time.Now())
+ return response, err
+}
+
+func (t *timingRoundTripper) RoundTrip(req *http.Request) (*http.Response, error) {
+ if req == nil || !Active(req.Context()) {
+ return t.base.RoundTrip(req)
+ }
+ startedAt := time.Now()
+ response, err := t.base.RoundTrip(req)
+ RecordDependency(req.Context(), dependencyModule(req), startedAt, time.Now())
+ return response, err
+}
+
+func dependencyModule(req *http.Request) string {
+ if req != nil {
+ if module, ok := req.Context().Value(dependencyModuleKey{}).(string); ok && module != "" {
+ return module
+ }
+ }
+ if req == nil || req.URL == nil {
+ return "http"
+ }
+ host := strings.ToLower(req.URL.Hostname())
+ switch {
+ case strings.Contains(host, "github"):
+ return "github"
+ case strings.Contains(host, "openai"):
+ return "openai"
+ case strings.Contains(host, "anthropic"):
+ return "anthropic"
+ case strings.Contains(host, "generativelanguage") || strings.Contains(host, "gemini"):
+ return "gemini"
+ case strings.Contains(host, "cloudcode") || strings.Contains(host, "antigravity"):
+ return "antigravity"
+ case strings.Contains(host, "googleapis") || strings.Contains(host, "google"):
+ return "google"
+ case strings.Contains(host, "amazonaws") || strings.Contains(host, "cloudflarestorage") || strings.Contains(host, "s3"):
+ return "s3"
+ case strings.Contains(host, "stripe") || strings.Contains(host, "airwallex") || strings.Contains(host, "alipay") || strings.Contains(host, "wechatpay") || strings.Contains(host, "paypal"):
+ return "payment"
+ default:
+ return "http"
+ }
+}
diff --git a/backend/internal/pkg/servertiming/http_test.go b/backend/internal/pkg/servertiming/http_test.go
new file mode 100644
index 0000000000..d37f378414
--- /dev/null
+++ b/backend/internal/pkg/servertiming/http_test.go
@@ -0,0 +1,168 @@
+package servertiming
+
+import (
+ "context"
+ "io"
+ "net/http"
+ "strings"
+ "testing"
+ "time"
+)
+
+type roundTripFunc func(*http.Request) (*http.Response, error)
+
+func (f roundTripFunc) RoundTrip(req *http.Request) (*http.Response, error) {
+ return f(req)
+}
+
+type trackingBody struct {
+ read bool
+}
+
+func (b *trackingBody) Read(_ []byte) (int, error) {
+ b.read = true
+ return 0, io.EOF
+}
+
+func (b *trackingBody) Close() error { return nil }
+
+func TestWrapRoundTripperRecordsResponseHeaderLatency(t *testing.T) {
+ startedAt := time.Now()
+ collector := New(startedAt)
+ body := &trackingBody{}
+ baseCalled := false
+ base := roundTripFunc(func(req *http.Request) (*http.Response, error) {
+ baseCalled = true
+ return &http.Response{
+ StatusCode: http.StatusOK,
+ Body: body,
+ Header: make(http.Header),
+ Request: req,
+ }, nil
+ })
+ req, err := http.NewRequestWithContext(WithCollector(context.Background(), collector), http.MethodGet, "https://api.github.com/repos/example/project", nil)
+ if err != nil {
+ t.Fatal(err)
+ }
+
+ resp, err := WrapRoundTripper(base).RoundTrip(req)
+ if err != nil {
+ t.Fatal(err)
+ }
+ defer func() { _ = resp.Body.Close() }()
+ if !baseCalled {
+ t.Fatal("base RoundTripper was not called")
+ }
+ if body.read {
+ t.Fatal("RoundTripper instrumentation read the response body; timing must stop at response headers")
+ }
+ header := collector.HeaderValue(time.Now(), "bypass")
+ if !strings.Contains(header, `dep_github;dur=`) || !strings.Contains(header, `deps;dur=`) {
+ t.Fatalf("dependency metrics missing from header: %q", header)
+ }
+}
+
+func TestWrapRoundTripperUsesContextModuleOverride(t *testing.T) {
+ collector := New(time.Now())
+ ctx := WithDependencyModule(WithCollector(context.Background(), collector), "data-managementd")
+ req, err := http.NewRequestWithContext(ctx, http.MethodGet, "https://private.example.test/path", nil)
+ if err != nil {
+ t.Fatal(err)
+ }
+ base := roundTripFunc(func(req *http.Request) (*http.Response, error) {
+ return &http.Response{StatusCode: http.StatusOK, Body: http.NoBody, Header: make(http.Header), Request: req}, nil
+ })
+
+ if _, err := WrapRoundTripper(base).RoundTrip(req); err != nil {
+ t.Fatal(err)
+ }
+ header := collector.HeaderValue(time.Now(), "bypass")
+ if !strings.Contains(header, "dep_data_managementd") {
+ t.Fatalf("module override missing from header: %q", header)
+ }
+ if strings.Contains(header, "private.example") {
+ t.Fatalf("raw host leaked into header: %q", header)
+ }
+}
+
+func TestWrapRoundTripperSkipsInactiveContext(t *testing.T) {
+ called := false
+ base := roundTripFunc(func(req *http.Request) (*http.Response, error) {
+ called = true
+ return &http.Response{StatusCode: http.StatusNoContent, Body: http.NoBody, Header: make(http.Header), Request: req}, nil
+ })
+ req, err := http.NewRequest(http.MethodGet, "https://api.openai.com/v1/models", nil)
+ if err != nil {
+ t.Fatal(err)
+ }
+ if _, err := WrapRoundTripper(base).RoundTrip(req); err != nil {
+ t.Fatal(err)
+ }
+ if !called {
+ t.Fatal("inactive request did not reach base RoundTripper")
+ }
+}
+
+func TestDoRecordsWithoutChangingTransportType(t *testing.T) {
+ base := roundTripFunc(func(req *http.Request) (*http.Response, error) {
+ return &http.Response{StatusCode: http.StatusOK, Body: http.NoBody, Header: make(http.Header), Request: req}, nil
+ })
+ client := &http.Client{Transport: base}
+ collector := New(time.Now())
+ req, err := http.NewRequestWithContext(WithCollector(context.Background(), collector), http.MethodGet, "https://api.openai.com/v1/models", nil)
+ if err != nil {
+ t.Fatal(err)
+ }
+ if _, err := Do(client, req); err != nil {
+ t.Fatal(err)
+ }
+ if _, ok := client.Transport.(roundTripFunc); !ok {
+ t.Fatalf("Do changed client transport type to %T", client.Transport)
+ }
+ if header := collector.HeaderValue(time.Now(), "bypass"); !strings.Contains(header, "dep_openai;dur=") {
+ t.Fatalf("dependency metric missing from header: %q", header)
+ }
+}
+
+func TestDependencyModuleClassification(t *testing.T) {
+ tests := map[string]string{
+ "https://api.github.com/repos/a/b": "github",
+ "https://api.openai.com/v1/models": "openai",
+ "https://api.anthropic.com/v1/messages": "anthropic",
+ "https://generativelanguage.googleapis.com/v1/models": "gemini",
+ "https://cloudcode-pa.googleapis.com/v1internal": "antigravity",
+ "https://storage.googleapis.com/bucket/object": "google",
+ "https://bucket.s3.amazonaws.com/object": "s3",
+ "https://api.stripe.com/v1/refunds": "payment",
+ "https://dependency.example.test/path": "http",
+ }
+ for rawURL, want := range tests {
+ req, err := http.NewRequest(http.MethodGet, rawURL, nil)
+ if err != nil {
+ t.Fatalf("NewRequest(%q): %v", rawURL, err)
+ }
+ if got := dependencyModule(req); got != want {
+ t.Errorf("dependencyModule(%q) = %q, want %q", rawURL, got, want)
+ }
+ }
+}
+
+func TestClientInstrumentationDoesNotMutateOriginal(t *testing.T) {
+ base := roundTripFunc(func(req *http.Request) (*http.Response, error) {
+ return &http.Response{StatusCode: http.StatusOK, Body: http.NoBody, Header: make(http.Header), Request: req}, nil
+ })
+ original := &http.Client{Transport: base, Timeout: time.Second}
+ instrumented := InstrumentClient(original)
+ if instrumented == original {
+ t.Fatal("InstrumentClient returned the original client")
+ }
+ if _, ok := original.Transport.(roundTripFunc); !ok {
+ t.Fatalf("InstrumentClient mutated the original transport to %T", original.Transport)
+ }
+ if instrumented.Timeout != original.Timeout {
+ t.Fatal("InstrumentClient did not preserve client settings")
+ }
+ if WrapRoundTripper(instrumented.Transport) != instrumented.Transport {
+ t.Fatal("WrapRoundTripper wrapped an already instrumented transport twice")
+ }
+}
diff --git a/backend/internal/pkg/xai/billing.go b/backend/internal/pkg/xai/billing.go
new file mode 100644
index 0000000000..15b9c7e50e
--- /dev/null
+++ b/backend/internal/pkg/xai/billing.go
@@ -0,0 +1,372 @@
+package xai
+
+import (
+ "encoding/json"
+ "fmt"
+ "math"
+ "net/http"
+ "strconv"
+ "strings"
+ "time"
+)
+
+const (
+ // CLI client identity required by cli-chat-proxy billing endpoints.
+ CLITokenAuthHeader = "x-xai-token-auth"
+ CLITokenAuthValue = "xai-grok-cli"
+ CLIClientVersionHeader = "x-grok-client-version"
+ // Keep in sync with https://x.ai/cli/stable.
+ CLIClientVersion = "0.2.93"
+ CLIUserAgent = "grok-pager/" + CLIClientVersion + " grok-shell/" + CLIClientVersion + " (macos; aarch64)"
+
+ BillingWeeklyPath = "/billing?format=credits"
+ BillingMonthlyPath = "/billing"
+
+ SuperGrokLimitCents = 15_000 // $150.00
+ SuperGrokHeavyLimitCents = 150_000 // $1,500.00
+)
+
+// BillingPeriod describes the current weekly/monthly window.
+type BillingPeriod struct {
+ Type string `json:"type,omitempty"`
+ Start string `json:"start,omitempty"`
+ End string `json:"end,omitempty"`
+}
+
+// BillingProductUsage is per-product usage inside the weekly credits window.
+type BillingProductUsage struct {
+ Product string `json:"product,omitempty"`
+ UsagePercent *float64 `json:"usagePercent,omitempty"`
+}
+
+// BillingConfig is the nested config object from /v1/billing responses.
+type BillingConfig struct {
+ CurrentPeriod *BillingPeriod `json:"currentPeriod,omitempty"`
+ CreditUsagePercent *float64 `json:"creditUsagePercent,omitempty"`
+ ProductUsage []BillingProductUsage `json:"productUsage,omitempty"`
+ MonthlyLimit json.RawMessage `json:"monthlyLimit,omitempty"`
+ Used json.RawMessage `json:"used,omitempty"`
+ BillingPeriodStart string `json:"billingPeriodStart,omitempty"`
+ BillingPeriodEnd string `json:"billingPeriodEnd,omitempty"`
+}
+
+// BillingPayload is the top-level body from /v1/billing.
+type BillingPayload struct {
+ Config *BillingConfig `json:"config,omitempty"`
+}
+
+// BillingProductSummary is a normalized product usage row for UI.
+type BillingProductSummary struct {
+ Product string `json:"product"`
+ UsagePercent *float64 `json:"usage_percent,omitempty"`
+}
+
+// BillingSummary is the merged weekly + monthly billing view.
+type BillingSummary struct {
+ PeriodType string `json:"period_type,omitempty"` // weekly | monthly | unknown
+ UsagePercent *float64 `json:"usage_percent,omitempty"`
+ PeriodStart string `json:"period_start,omitempty"`
+ PeriodEnd string `json:"period_end,omitempty"`
+ ProductUsage []BillingProductSummary `json:"product_usage,omitempty"`
+ MonthlyLimitCents *float64 `json:"monthly_limit_cents,omitempty"`
+ UsedCents *float64 `json:"used_cents,omitempty"`
+ IncludedUsedCents *float64 `json:"included_used_cents,omitempty"`
+ BillingPeriodStart string `json:"billing_period_start,omitempty"`
+ BillingPeriodEnd string `json:"billing_period_end,omitempty"`
+ UsedPercent *float64 `json:"used_percent,omitempty"`
+ Plan string `json:"plan,omitempty"` // SuperGrok | SuperGrok Heavy | ""
+ StatusCode int `json:"status_code,omitempty"`
+ Source string `json:"source,omitempty"`
+ FetchedAt string `json:"fetched_at,omitempty"`
+ UpdatedAt string `json:"updated_at,omitempty"`
+ WeeklyUpdatedAt string `json:"weekly_updated_at,omitempty"`
+ MonthlyUpdatedAt string `json:"monthly_updated_at,omitempty"`
+ Partial bool `json:"partial,omitempty"`
+ FailedWindows []string `json:"failed_windows,omitempty"`
+}
+
+// BuildBillingURL builds weekly or monthly billing URL against the CLI chat proxy.
+func BuildBillingURL(formatCredits bool) string {
+ base := strings.TrimRight(DefaultCLIBaseURL, "/")
+ if formatCredits {
+ return base + BillingWeeklyPath
+ }
+ return base + BillingMonthlyPath
+}
+
+// ApplyCLIBillingHeaders sets Authorization + CLI identity headers for billing GETs.
+func ApplyCLIBillingHeaders(req *http.Request, accessToken string) {
+ if req == nil {
+ return
+ }
+ token := strings.TrimSpace(accessToken)
+ if token != "" {
+ req.Header.Set("Authorization", "Bearer "+token)
+ }
+ req.Header.Set("Accept", "application/json")
+ req.Header.Set("Content-Type", "application/json")
+ req.Header.Set(CLITokenAuthHeader, CLITokenAuthValue)
+ req.Header.Set(CLIClientVersionHeader, CLIClientVersion)
+ req.Header.Set("User-Agent", CLIUserAgent)
+}
+
+// ParseBillingPayload unmarshals a billing API response body.
+func ParseBillingPayload(body []byte) (*BillingPayload, error) {
+ if len(body) == 0 {
+ return nil, fmt.Errorf("empty billing body")
+ }
+ var payload BillingPayload
+ if err := json.Unmarshal(body, &payload); err != nil {
+ return nil, err
+ }
+ return &payload, nil
+}
+
+// BuildBillingSummary normalizes a billing config into a UI-friendly summary.
+func BuildBillingSummary(config *BillingConfig) *BillingSummary {
+ if config == nil {
+ return nil
+ }
+ summary := &BillingSummary{}
+ period := config.CurrentPeriod
+ periodType := resolvePeriodType(period)
+ creditUsage := cloneFloat(config.CreditUsagePercent)
+
+ periodStart := ""
+ periodEnd := ""
+ if period != nil {
+ periodStart = strings.TrimSpace(period.Start)
+ periodEnd = strings.TrimSpace(period.End)
+ }
+ if periodStart == "" {
+ periodStart = strings.TrimSpace(config.BillingPeriodStart)
+ }
+ if periodEnd == "" {
+ periodEnd = strings.TrimSpace(config.BillingPeriodEnd)
+ }
+
+ products := make([]BillingProductSummary, 0, len(config.ProductUsage))
+ for _, item := range config.ProductUsage {
+ product := strings.TrimSpace(item.Product)
+ if product == "" {
+ continue
+ }
+ products = append(products, BillingProductSummary{
+ Product: product,
+ UsagePercent: cloneFloat(item.UsagePercent),
+ })
+ }
+
+ monthlyLimit := parseCentValue(config.MonthlyLimit)
+ used := parseCentValue(config.Used)
+ billingStart := strings.TrimSpace(config.BillingPeriodStart)
+ billingEnd := strings.TrimSpace(config.BillingPeriodEnd)
+
+ var includedUsed *float64
+ if used != nil {
+ if monthlyLimit != nil && *monthlyLimit > 0 {
+ v := math.Min(*used, *monthlyLimit)
+ includedUsed = &v
+ } else {
+ includedUsed = cloneFloat(used)
+ }
+ }
+
+ var usedPercent *float64
+ if monthlyLimit != nil && *monthlyLimit > 0 && includedUsed != nil {
+ v := (*includedUsed / *monthlyLimit) * 100
+ usedPercent = &v
+ }
+
+ hasWeekly := creditUsage != nil || periodType == "weekly" || len(products) > 0
+ hasMonthly := monthlyLimit != nil || used != nil || (!hasWeekly && billingEnd != "")
+ if !hasWeekly && !hasMonthly {
+ return nil
+ }
+
+ if hasWeekly {
+ if periodType == "unknown" {
+ periodType = "weekly"
+ }
+ summary.PeriodType = periodType
+ summary.UsagePercent = creditUsage
+ summary.PeriodStart = periodStart
+ summary.PeriodEnd = periodEnd
+ } else {
+ // Monthly-only: do not put monthly % into UsagePercent (weekly bar field).
+ // Frontend weekly bar only renders when PeriodType == weekly.
+ summary.PeriodType = "monthly"
+ summary.PeriodStart = billingStart
+ summary.PeriodEnd = billingEnd
+ }
+ summary.ProductUsage = products
+ summary.MonthlyLimitCents = monthlyLimit
+ summary.UsedCents = used
+ summary.IncludedUsedCents = includedUsed
+ if hasMonthly {
+ summary.BillingPeriodStart = billingStart
+ summary.BillingPeriodEnd = billingEnd
+ }
+ summary.UsedPercent = usedPercent
+ summary.Plan = resolvePlan(monthlyLimit)
+ return summary
+}
+
+// MergeBillingProbeResult updates successful billing domains while retaining
+// the previous value for any domain that could not be refreshed.
+func MergeBillingProbeResult(previous, weekly, monthly *BillingSummary, weeklyOK, monthlyOK bool) *BillingSummary {
+ var out BillingSummary
+ if previous != nil {
+ out = *previous
+ previousUpdatedAt := previous.UpdatedAt
+ if previousUpdatedAt == "" {
+ previousUpdatedAt = previous.FetchedAt
+ }
+ if out.WeeklyUpdatedAt == "" && (out.UsagePercent != nil || len(out.ProductUsage) > 0) {
+ out.WeeklyUpdatedAt = previousUpdatedAt
+ }
+ if out.MonthlyUpdatedAt == "" && (out.MonthlyLimitCents != nil || out.UsedPercent != nil) {
+ out.MonthlyUpdatedAt = previousUpdatedAt
+ }
+ }
+ now := time.Now().UTC().Format(time.RFC3339)
+
+ if weeklyOK && weekly != nil {
+ out.PeriodType = weekly.PeriodType
+ out.UsagePercent = weekly.UsagePercent
+ out.PeriodStart = weekly.PeriodStart
+ out.PeriodEnd = weekly.PeriodEnd
+ out.ProductUsage = weekly.ProductUsage
+ out.WeeklyUpdatedAt = now
+ }
+ if monthlyOK && monthly != nil {
+ if out.PeriodType == "" {
+ out.PeriodType = "monthly"
+ }
+ out.MonthlyLimitCents = monthly.MonthlyLimitCents
+ out.UsedCents = monthly.UsedCents
+ out.IncludedUsedCents = monthly.IncludedUsedCents
+ out.BillingPeriodStart = monthly.BillingPeriodStart
+ out.BillingPeriodEnd = monthly.BillingPeriodEnd
+ out.UsedPercent = monthly.UsedPercent
+ out.Plan = monthly.Plan
+ out.MonthlyUpdatedAt = now
+ }
+
+ out.Partial = !weeklyOK || !monthlyOK
+ out.FailedWindows = nil
+ if !weeklyOK {
+ out.FailedWindows = append(out.FailedWindows, "weekly")
+ }
+ if !monthlyOK {
+ out.FailedWindows = append(out.FailedWindows, "monthly")
+ }
+ if !weeklyOK && !monthlyOK && previous == nil {
+ return nil
+ }
+ return &out
+}
+
+// StampBillingSummary sets fetch metadata.
+func StampBillingSummary(summary *BillingSummary, statusCode int, source string) *BillingSummary {
+ if summary == nil {
+ return nil
+ }
+ now := time.Now().UTC().Format(time.RFC3339)
+ summary.StatusCode = statusCode
+ summary.Source = source
+ summary.FetchedAt = now
+ summary.UpdatedAt = now
+ return summary
+}
+
+func resolvePeriodType(period *BillingPeriod) string {
+ if period == nil {
+ return "unknown"
+ }
+ raw := strings.ToLower(strings.TrimSpace(period.Type))
+ if strings.Contains(raw, "weekly") {
+ return "weekly"
+ }
+ if strings.Contains(raw, "monthly") {
+ return "monthly"
+ }
+ return "unknown"
+}
+
+func resolvePlan(monthlyLimitCents *float64) string {
+ if monthlyLimitCents == nil {
+ return ""
+ }
+ // Allow small float noise.
+ limit := math.Round(*monthlyLimitCents)
+ switch limit {
+ case SuperGrokLimitCents:
+ return "SuperGrok"
+ case SuperGrokHeavyLimitCents:
+ return "SuperGrok Heavy"
+ default:
+ return ""
+ }
+}
+
+func parseCentValue(raw json.RawMessage) *float64 {
+ if len(raw) == 0 || string(raw) == "null" {
+ return nil
+ }
+ // Object form: {"val": 123}
+ var obj struct {
+ Val any `json:"val"`
+ }
+ if err := json.Unmarshal(raw, &obj); err == nil && obj.Val != nil {
+ return anyToFloat(obj.Val)
+ }
+ // Bare number / string
+ var n any
+ if err := json.Unmarshal(raw, &n); err != nil {
+ return nil
+ }
+ return anyToFloat(n)
+}
+
+func anyToFloat(v any) *float64 {
+ switch n := v.(type) {
+ case float64:
+ return &n
+ case float32:
+ f := float64(n)
+ return &f
+ case int:
+ f := float64(n)
+ return &f
+ case int64:
+ f := float64(n)
+ return &f
+ case json.Number:
+ f, err := n.Float64()
+ if err != nil {
+ return nil
+ }
+ return &f
+ case string:
+ s := strings.TrimSpace(n)
+ if s == "" {
+ return nil
+ }
+ f, err := strconv.ParseFloat(s, 64)
+ if err != nil {
+ return nil
+ }
+ return &f
+ default:
+ return nil
+ }
+}
+
+func cloneFloat(v *float64) *float64 {
+ if v == nil {
+ return nil
+ }
+ f := *v
+ return &f
+}
diff --git a/backend/internal/pkg/xai/billing_test.go b/backend/internal/pkg/xai/billing_test.go
new file mode 100644
index 0000000000..1d863f6a39
--- /dev/null
+++ b/backend/internal/pkg/xai/billing_test.go
@@ -0,0 +1,127 @@
+package xai
+
+import (
+ "encoding/json"
+ "net/http"
+ "testing"
+
+ "github.com/stretchr/testify/require"
+)
+
+func TestBuildBillingURL(t *testing.T) {
+ t.Parallel()
+ require.Equal(t, "https://cli-chat-proxy.grok.com/v1/billing?format=credits", BuildBillingURL(true))
+ require.Equal(t, "https://cli-chat-proxy.grok.com/v1/billing", BuildBillingURL(false))
+}
+
+func TestApplyCLIBillingHeaders(t *testing.T) {
+ t.Parallel()
+ req, err := http.NewRequest(http.MethodGet, BuildBillingURL(true), nil)
+ require.NoError(t, err)
+
+ ApplyCLIBillingHeaders(req, " token ")
+
+ require.Equal(t, "Bearer token", req.Header.Get("Authorization"))
+ require.Equal(t, CLITokenAuthValue, req.Header.Get(CLITokenAuthHeader))
+ require.Equal(t, CLIClientVersion, req.Header.Get(CLIClientVersionHeader))
+ require.Equal(t, "grok-pager/"+CLIClientVersion+" grok-shell/"+CLIClientVersion+" (macos; aarch64)", req.UserAgent())
+}
+
+func TestBuildBillingSummaryWeeklyAndMonthly(t *testing.T) {
+ t.Parallel()
+
+ weeklyBody := []byte(`{
+ "config": {
+ "currentPeriod": {"type":"WEEKLY","start":"2026-07-09T03:25:00Z","end":"2026-07-16T03:25:00Z"},
+ "creditUsagePercent": 2.0,
+ "productUsage": [{"product":"Api","usagePercent":2.0}]
+ }
+ }`)
+ monthlyBody := []byte(`{
+ "config": {
+ "monthlyLimit": {"val": 15000},
+ "used": {"val": 78},
+ "billingPeriodStart": "2026-07-01T00:00:00Z",
+ "billingPeriodEnd": "2026-08-01T00:00:00Z"
+ }
+ }`)
+
+ weeklyPayload, err := ParseBillingPayload(weeklyBody)
+ require.NoError(t, err)
+ monthlyPayload, err := ParseBillingPayload(monthlyBody)
+ require.NoError(t, err)
+
+ weekly := BuildBillingSummary(weeklyPayload.Config)
+ monthly := BuildBillingSummary(monthlyPayload.Config)
+ require.NotNil(t, weekly)
+ require.NotNil(t, monthly)
+ require.Equal(t, "weekly", weekly.PeriodType)
+ require.InDelta(t, 2.0, *weekly.UsagePercent, 1e-9)
+ require.Equal(t, "Api", weekly.ProductUsage[0].Product)
+ require.Equal(t, "SuperGrok", monthly.Plan)
+ require.InDelta(t, 15000, *monthly.MonthlyLimitCents, 1e-9)
+ require.InDelta(t, 78, *monthly.UsedCents, 1e-9)
+ require.InDelta(t, 0.52, *monthly.UsedPercent, 1e-2)
+
+ merged := MergeBillingProbeResult(nil, weekly, monthly, true, true)
+ require.Equal(t, "weekly", merged.PeriodType)
+ require.InDelta(t, 2.0, *merged.UsagePercent, 1e-9)
+ require.Equal(t, "SuperGrok", merged.Plan)
+ require.InDelta(t, 15000, *merged.MonthlyLimitCents, 1e-9)
+ require.Equal(t, "2026-08-01T00:00:00Z", merged.BillingPeriodEnd)
+}
+
+func TestParseCentValueBareNumber(t *testing.T) {
+ t.Parallel()
+ raw, _ := json.Marshal(15000)
+ v := parseCentValue(raw)
+ require.NotNil(t, v)
+ require.InDelta(t, 15000, *v, 1e-9)
+}
+
+func TestBuildBillingSummaryMonthlyOnlyKeepsWeeklyUsageEmpty(t *testing.T) {
+ t.Parallel()
+ payload, err := ParseBillingPayload([]byte(`{"config":{"monthlyLimit":{"val":15000},"used":{"val":7500},"billingPeriodStart":"2026-07-01T00:00:00Z","billingPeriodEnd":"2026-08-01T00:00:00Z"}}`))
+ require.NoError(t, err)
+
+ summary := BuildBillingSummary(payload.Config)
+ require.NotNil(t, summary)
+ require.Equal(t, "monthly", summary.PeriodType)
+ require.Nil(t, summary.UsagePercent)
+ require.InDelta(t, 50, *summary.UsedPercent, 1e-9)
+}
+
+func TestMergeBillingProbeResultRetainsFailedWindow(t *testing.T) {
+ t.Parallel()
+ previous := &BillingSummary{
+ PeriodType: "weekly",
+ UsagePercent: floatPointer(100),
+ PeriodEnd: "2026-07-16T00:00:00Z",
+ MonthlyLimitCents: floatPointer(15000),
+ UsedPercent: floatPointer(20),
+ BillingPeriodEnd: "2026-08-01T00:00:00Z",
+ WeeklyUpdatedAt: "2026-07-10T00:00:00Z",
+ MonthlyUpdatedAt: "2026-07-10T00:00:00Z",
+ FailedWindows: []string{"monthly"},
+ }
+ monthly := &BillingSummary{
+ PeriodType: "monthly",
+ MonthlyLimitCents: floatPointer(15000),
+ UsedPercent: floatPointer(30),
+ BillingPeriodEnd: "2026-08-01T00:00:00Z",
+ }
+
+ merged := MergeBillingProbeResult(previous, nil, monthly, false, true)
+ require.Equal(t, "weekly", merged.PeriodType)
+ require.InDelta(t, 100, *merged.UsagePercent, 1e-9)
+ require.Equal(t, previous.WeeklyUpdatedAt, merged.WeeklyUpdatedAt)
+ require.InDelta(t, 30, *merged.UsedPercent, 1e-9)
+ require.NotEqual(t, previous.MonthlyUpdatedAt, merged.MonthlyUpdatedAt)
+ require.True(t, merged.Partial)
+ require.Equal(t, []string{"weekly"}, merged.FailedWindows)
+ require.Equal(t, []string{"monthly"}, previous.FailedWindows)
+}
+
+func floatPointer(value float64) *float64 {
+ return &value
+}
diff --git a/backend/internal/repository/account_repo.go b/backend/internal/repository/account_repo.go
index 261be7c14b..6d26d474c6 100644
--- a/backend/internal/repository/account_repo.go
+++ b/backend/internal/repository/account_repo.go
@@ -61,6 +61,7 @@ var schedulerNeutralExtraKeyPrefixes = []string{
var schedulerNeutralExtraKeys = map[string]struct{}{
"codex_usage_updated_at": {},
+ "grok_billing_snapshot": {},
"session_window_utilization": {},
}
diff --git a/backend/internal/repository/account_repo_grok_billing_test.go b/backend/internal/repository/account_repo_grok_billing_test.go
new file mode 100644
index 0000000000..fb41ae5ffa
--- /dev/null
+++ b/backend/internal/repository/account_repo_grok_billing_test.go
@@ -0,0 +1,16 @@
+package repository
+
+import (
+ "testing"
+
+ "github.com/stretchr/testify/require"
+)
+
+func TestGrokBillingSnapshotIsSchedulerNeutral(t *testing.T) {
+ t.Parallel()
+
+ require.True(t, isSchedulerNeutralExtraKey("grok_billing_snapshot"))
+ require.False(t, shouldEnqueueSchedulerOutboxForExtraUpdates(map[string]any{
+ "grok_billing_snapshot": map[string]any{"usage_percent": 50},
+ }))
+}
diff --git a/backend/internal/repository/backup_s3_store.go b/backend/internal/repository/backup_s3_store.go
index 5d419f574b..2104e1e5d7 100644
--- a/backend/internal/repository/backup_s3_store.go
+++ b/backend/internal/repository/backup_s3_store.go
@@ -13,6 +13,7 @@ import (
"github.com/aws/aws-sdk-go-v2/credentials"
"github.com/aws/aws-sdk-go-v2/service/s3"
+ "github.com/Wei-Shaw/sub2api/internal/pkg/servertiming"
"github.com/Wei-Shaw/sub2api/internal/service"
)
@@ -63,12 +64,14 @@ func (s *S3BackupStore) Upload(ctx context.Context, key string, body io.Reader,
return 0, fmt.Errorf("read body: %w", err)
}
+ finish := servertiming.ObserveDependency(ctx, "s3")
_, err = s.client.PutObject(ctx, &s3.PutObjectInput{
Bucket: &s.bucket,
Key: &key,
Body: bytes.NewReader(data),
ContentType: &contentType,
})
+ finish()
if err != nil {
return 0, fmt.Errorf("S3 PutObject: %w", err)
}
@@ -76,10 +79,12 @@ func (s *S3BackupStore) Upload(ctx context.Context, key string, body io.Reader,
}
func (s *S3BackupStore) Download(ctx context.Context, key string) (io.ReadCloser, error) {
+ finish := servertiming.ObserveDependency(ctx, "s3")
result, err := s.client.GetObject(ctx, &s3.GetObjectInput{
Bucket: &s.bucket,
Key: &key,
})
+ finish()
if err != nil {
return nil, fmt.Errorf("S3 GetObject: %w", err)
}
@@ -87,10 +92,12 @@ func (s *S3BackupStore) Download(ctx context.Context, key string) (io.ReadCloser
}
func (s *S3BackupStore) Delete(ctx context.Context, key string) error {
+ finish := servertiming.ObserveDependency(ctx, "s3")
_, err := s.client.DeleteObject(ctx, &s3.DeleteObjectInput{
Bucket: &s.bucket,
Key: &key,
})
+ finish()
return err
}
@@ -107,9 +114,11 @@ func (s *S3BackupStore) PresignURL(ctx context.Context, key string, expiry time.
}
func (s *S3BackupStore) HeadBucket(ctx context.Context) error {
+ finish := servertiming.ObserveDependency(ctx, "s3")
_, err := s.client.HeadBucket(ctx, &s3.HeadBucketInput{
Bucket: &s.bucket,
})
+ finish()
if err != nil {
return fmt.Errorf("S3 HeadBucket failed: %w", err)
}
diff --git a/backend/internal/repository/claude_oauth_service.go b/backend/internal/repository/claude_oauth_service.go
index 5c5f27c86a..ec2d426ecb 100644
--- a/backend/internal/repository/claude_oauth_service.go
+++ b/backend/internal/repository/claude_oauth_service.go
@@ -276,5 +276,5 @@ func createReqClient(proxyURL string) (*req.Client, error) {
client.SetProxyURL(trimmed)
}
- return client, nil
+ return instrumentReqClient(client), nil
}
diff --git a/backend/internal/repository/ent.go b/backend/internal/repository/ent.go
index 64d321924d..3abb528e98 100644
--- a/backend/internal/repository/ent.go
+++ b/backend/internal/repository/ent.go
@@ -15,7 +15,7 @@ import (
"entgo.io/ent/dialect"
entsql "entgo.io/ent/dialect/sql"
- _ "github.com/lib/pq" // PostgreSQL 驱动,通过副作用导入注册驱动
+ "github.com/lib/pq"
)
// InitEnt 初始化 Ent ORM 客户端并返回客户端实例和底层的 *sql.DB。
@@ -48,9 +48,19 @@ func InitEnt(cfg *config.Config) (*ent.Client, *sql.DB, error) {
// 使用 Ent 的 SQL 驱动打开 PostgreSQL 连接。
// dialect.Postgres 指定使用 PostgreSQL 方言进行 SQL 生成。
- drv, err := entsql.Open(dialect.Postgres, dsn)
- if err != nil {
- return nil, nil, err
+ var drv *entsql.Driver
+ if cfg.Server.EnableServerTiming {
+ connector, err := pq.NewConnector(dsn)
+ if err != nil {
+ return nil, nil, err
+ }
+ drv = entsql.OpenDB(dialect.Postgres, sql.OpenDB(newServerTimingConnector(connector)))
+ } else {
+ var err error
+ drv, err = entsql.Open(dialect.Postgres, dsn)
+ if err != nil {
+ return nil, nil, err
+ }
}
applyDBPoolSettings(drv.DB(), cfg)
diff --git a/backend/internal/repository/http_upstream.go b/backend/internal/repository/http_upstream.go
index bb079b0789..0d5afd500c 100644
--- a/backend/internal/repository/http_upstream.go
+++ b/backend/internal/repository/http_upstream.go
@@ -25,6 +25,7 @@ import (
"github.com/Wei-Shaw/sub2api/internal/config"
"github.com/Wei-Shaw/sub2api/internal/pkg/proxyurl"
"github.com/Wei-Shaw/sub2api/internal/pkg/proxyutil"
+ "github.com/Wei-Shaw/sub2api/internal/pkg/servertiming"
"github.com/Wei-Shaw/sub2api/internal/pkg/tlsfingerprint"
"github.com/Wei-Shaw/sub2api/internal/service"
"github.com/Wei-Shaw/sub2api/internal/util/urlvalidator"
@@ -186,7 +187,7 @@ func (s *httpUpstreamService) Do(req *http.Request, proxyURL string, accountID i
}
// 执行请求
- resp, err := entry.client.Do(req)
+ resp, err := servertiming.Do(entry.client, req)
if err != nil {
s.recordOpenAIHTTP2Failure(profile, entry.protocolMode, entry.proxyKey, err)
// 请求失败,立即减少计数
@@ -243,7 +244,7 @@ func (s *httpUpstreamService) DoWithTLS(req *http.Request, proxyURL string, acco
return nil, err
}
- resp, err := entry.client.Do(req)
+ resp, err := servertiming.Do(entry.client, req)
if err != nil {
atomic.AddInt64(&entry.inFlight, -1)
atomic.StoreInt64(&entry.lastUsed, time.Now().UnixNano())
diff --git a/backend/internal/repository/openai_long_context_billing_migration_integration_test.go b/backend/internal/repository/openai_long_context_billing_migration_integration_test.go
new file mode 100644
index 0000000000..5f50ed0729
--- /dev/null
+++ b/backend/internal/repository/openai_long_context_billing_migration_integration_test.go
@@ -0,0 +1,159 @@
+//go:build integration
+
+package repository
+
+import (
+ "context"
+ "testing"
+
+ dbmigrations "github.com/Wei-Shaw/sub2api/migrations"
+ "github.com/stretchr/testify/require"
+)
+
+func TestMigration175EnforcesOpenAILongContextBillingWriteInvariant(t *testing.T) {
+ tx := testTx(t)
+ ctx := context.Background()
+ migrationSQL, err := dbmigrations.FS.ReadFile("175_default_openai_long_context_billing.sql")
+ require.NoError(t, err)
+ _, err = tx.ExecContext(ctx, `
+DROP TRIGGER IF EXISTS accounts_propagate_openai_long_context_billing_extra ON accounts;
+DROP TRIGGER IF EXISTS accounts_enforce_openai_long_context_billing_extra ON accounts;
+`)
+ require.NoError(t, err)
+
+ var ordinaryID int64
+ require.NoError(t, tx.QueryRowContext(ctx, `
+INSERT INTO accounts (name, platform, type, extra)
+VALUES ('migration-175-ordinary', 'openai', 'oauth', '{}'::jsonb)
+RETURNING id
+`).Scan(&ordinaryID))
+
+ var parentID int64
+ require.NoError(t, tx.QueryRowContext(ctx, `
+INSERT INTO accounts (name, platform, type, extra)
+VALUES ('migration-175-parent', 'openai', 'oauth', '{"openai_long_context_billing_enabled":false}'::jsonb)
+RETURNING id
+`).Scan(&parentID))
+
+ var shadowID int64
+ require.NoError(t, tx.QueryRowContext(ctx, `
+INSERT INTO accounts (name, platform, type, extra, parent_account_id, quota_dimension)
+VALUES ('migration-175-shadow', 'openai', 'oauth', '{}'::jsonb, $1, 'spark')
+RETURNING id
+`, parentID).Scan(&shadowID))
+
+ var malformedLegacyID int64
+ require.NoError(t, tx.QueryRowContext(ctx, `
+INSERT INTO accounts (name, platform, type, extra)
+VALUES ('migration-175-malformed-legacy', 'openai', 'oauth', '{"openai_long_context_billing_enabled":"false"}'::jsonb)
+RETURNING id
+`).Scan(&malformedLegacyID))
+
+ _, err = tx.ExecContext(ctx, string(migrationSQL))
+ require.NoError(t, err)
+ _, err = tx.ExecContext(ctx, string(migrationSQL))
+ require.NoError(t, err)
+
+ var ordinaryEnabled bool
+ require.NoError(t, tx.QueryRowContext(ctx, `
+SELECT (extra->>'openai_long_context_billing_enabled')::boolean
+FROM accounts
+WHERE id = $1
+`, ordinaryID).Scan(&ordinaryEnabled))
+ require.False(t, ordinaryEnabled)
+
+ var shadowEnabled bool
+ require.NoError(t, tx.QueryRowContext(ctx, `
+SELECT (extra->>'openai_long_context_billing_enabled')::boolean
+FROM accounts
+WHERE id = $1
+`, shadowID).Scan(&shadowEnabled))
+ require.False(t, shadowEnabled)
+
+ var initialShadowOutboxEvents int
+ require.NoError(t, tx.QueryRowContext(ctx, `
+SELECT COUNT(*)
+FROM scheduler_outbox
+WHERE event_type = 'account_changed' AND account_id = $1
+`, shadowID).Scan(&initialShadowOutboxEvents))
+ require.Equal(t, 1, initialShadowOutboxEvents)
+
+ var malformedLegacyEnabled bool
+ require.NoError(t, tx.QueryRowContext(ctx, `
+SELECT (extra->>'openai_long_context_billing_enabled')::boolean
+FROM accounts
+WHERE id = $1
+`, malformedLegacyID).Scan(&malformedLegacyEnabled))
+ require.False(t, malformedLegacyEnabled)
+ _, err = tx.ExecContext(ctx, `
+UPDATE accounts
+SET extra = extra || '{"migration_175_unrelated_update":true}'::jsonb
+WHERE id = $1
+`, malformedLegacyID)
+ require.NoError(t, err)
+
+ _, err = tx.ExecContext(ctx, "TRUNCATE scheduler_outbox")
+ require.NoError(t, err)
+ _, err = tx.ExecContext(ctx, `
+UPDATE accounts
+SET extra = '{"legacy_writer_replaced_extra":true}'::jsonb
+WHERE id = $1
+`, parentID)
+ require.NoError(t, err)
+ var parentEnabled bool
+ require.NoError(t, tx.QueryRowContext(ctx, `
+SELECT (extra->>'openai_long_context_billing_enabled')::boolean
+FROM accounts
+WHERE id = $1
+`, parentID).Scan(&parentEnabled))
+ require.False(t, parentEnabled)
+ require.NoError(t, tx.QueryRowContext(ctx, `
+SELECT (extra->>'openai_long_context_billing_enabled')::boolean
+FROM accounts
+WHERE id = $1
+`, shadowID).Scan(&shadowEnabled))
+ require.False(t, shadowEnabled)
+ var preservedOptOutEvents int
+ require.NoError(t, tx.QueryRowContext(ctx, `
+SELECT COUNT(*)
+FROM scheduler_outbox
+WHERE event_type = 'account_changed' AND account_id = $1
+`, shadowID).Scan(&preservedOptOutEvents))
+ require.Zero(t, preservedOptOutEvents)
+
+ require.NoError(t, tx.QueryRowContext(ctx, `
+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.False(t, ordinaryEnabled)
+
+ _, err = tx.ExecContext(ctx, "TRUNCATE scheduler_outbox")
+ require.NoError(t, err)
+ _, err = tx.ExecContext(ctx, `
+UPDATE accounts
+SET extra = jsonb_set(extra, '{openai_long_context_billing_enabled}', 'true'::jsonb, true)
+WHERE id = $1
+`, parentID)
+ require.NoError(t, err)
+ require.NoError(t, tx.QueryRowContext(ctx, `
+SELECT (extra->>'openai_long_context_billing_enabled')::boolean
+FROM accounts
+WHERE id = $1
+`, shadowID).Scan(&shadowEnabled))
+ require.True(t, shadowEnabled)
+
+ var shadowOutboxEvents int
+ require.NoError(t, tx.QueryRowContext(ctx, `
+SELECT COUNT(*)
+FROM scheduler_outbox
+WHERE event_type = 'account_changed' AND account_id = $1
+`, shadowID).Scan(&shadowOutboxEvents))
+ require.Equal(t, 1, shadowOutboxEvents)
+
+ _, err = tx.ExecContext(ctx, `
+INSERT INTO accounts (name, platform, type, extra)
+VALUES ('migration-175-malformed', 'openai', 'oauth', '{"openai_long_context_billing_enabled":"false"}'::jsonb)
+`)
+ require.ErrorContains(t, err, "openai_long_context_billing_enabled must be a boolean")
+}
diff --git a/backend/internal/repository/ops_repo.go b/backend/internal/repository/ops_repo.go
index 2129a451c4..900abcf212 100644
--- a/backend/internal/repository/ops_repo.go
+++ b/backend/internal/repository/ops_repo.go
@@ -718,6 +718,7 @@ func (r *opsRepository) BatchInsertSystemLogs(ctx context.Context, inputs []*ser
stmt, err := tx.PrepareContext(ctx, pq.CopyIn(
"ops_system_logs",
"created_at",
+ "host",
"level",
"component",
"message",
@@ -760,6 +761,7 @@ func (r *opsRepository) BatchInsertSystemLogs(ctx context.Context, inputs []*ser
if _, err := stmt.ExecContext(
ctx,
createdAt.UTC(),
+ opsNullString(input.Host),
level,
component,
message,
@@ -827,6 +829,7 @@ func (r *opsRepository) ListSystemLogs(ctx context.Context, filter *service.OpsS
SELECT
l.id,
l.created_at,
+ COALESCE(l.host, ''),
l.level,
COALESCE(l.component, ''),
COALESCE(l.message, ''),
@@ -859,6 +862,7 @@ LIMIT $` + itoa(len(args)+1) + ` OFFSET $` + itoa(len(args)+2)
if err := rows.Scan(
&item.ID,
&item.CreatedAt,
+ &item.Host,
&item.Level,
&item.Component,
&item.Message,
@@ -1130,6 +1134,11 @@ func buildOpsSystemLogsWhere(filter *service.OpsSystemLogFilter) (string, []any,
hasConstraint = true
}
if filter != nil {
+ if v := strings.TrimSpace(filter.Host); v != "" {
+ args = append(args, v)
+ clauses = append(clauses, "l.host = $"+itoa(len(args)))
+ hasConstraint = true
+ }
if v := strings.ToLower(strings.TrimSpace(filter.Level)); v != "" {
args = append(args, v)
clauses = append(clauses, "LOWER(COALESCE(l.level,'')) = $"+itoa(len(args)))
@@ -1194,6 +1203,7 @@ func buildOpsSystemLogsCleanupWhere(filter *service.OpsSystemLogCleanupFilter) (
listFilter := &service.OpsSystemLogFilter{
StartTime: filter.StartTime,
EndTime: filter.EndTime,
+ Host: filter.Host,
Level: filter.Level,
Component: filter.Component,
RequestID: filter.RequestID,
diff --git a/backend/internal/repository/ops_repo_system_logs_test.go b/backend/internal/repository/ops_repo_system_logs_test.go
index 98199f4828..48be3e7256 100644
--- a/backend/internal/repository/ops_repo_system_logs_test.go
+++ b/backend/internal/repository/ops_repo_system_logs_test.go
@@ -18,6 +18,7 @@ func TestBuildOpsSystemLogsWhere_WithClientRequestIDAndUserID(t *testing.T) {
filter := &service.OpsSystemLogFilter{
StartTime: &start,
EndTime: &end,
+ Host: "api-node-1",
Level: "warn",
Component: "http.access",
RequestID: "req-1",
@@ -37,8 +38,11 @@ func TestBuildOpsSystemLogsWhere_WithClientRequestIDAndUserID(t *testing.T) {
if where == "" {
t.Fatalf("where should not be empty")
}
- if len(args) != 12 {
- t.Fatalf("args len = %d, want 12", len(args))
+ if len(args) != 13 {
+ t.Fatalf("args len = %d, want 13", len(args))
+ }
+ if !contains(where, "l.host = $") {
+ t.Fatalf("where should include host condition: %s", where)
}
if !contains(where, "COALESCE(l.client_request_id,'') = $") {
t.Fatalf("where should include client_request_id condition: %s", where)
@@ -68,6 +72,7 @@ func TestBuildOpsSystemLogsCleanupWhere_WithClientRequestIDAndUserID(t *testing.
userID := int64(9)
apiKeyID := int64(10)
filter := &service.OpsSystemLogCleanupFilter{
+ Host: "api-node-2",
ClientRequestID: "creq-9",
UserID: &userID,
APIKeyID: &apiKeyID,
@@ -77,8 +82,11 @@ func TestBuildOpsSystemLogsCleanupWhere_WithClientRequestIDAndUserID(t *testing.
if !hasConstraint {
t.Fatalf("expected hasConstraint=true")
}
- if len(args) != 3 {
- t.Fatalf("args len = %d, want 3", len(args))
+ if len(args) != 4 {
+ t.Fatalf("args len = %d, want 4", len(args))
+ }
+ if !contains(where, "l.host = $") {
+ t.Fatalf("where should include host condition: %s", where)
}
if !contains(where, "COALESCE(l.client_request_id,'') = $") {
t.Fatalf("where should include client_request_id condition: %s", where)
diff --git a/backend/internal/repository/redis.go b/backend/internal/repository/redis.go
index 2b4ee4e636..0ead4644c1 100644
--- a/backend/internal/repository/redis.go
+++ b/backend/internal/repository/redis.go
@@ -21,7 +21,11 @@ import (
// 2. MinIdleConns: 保持最小空闲连接,减少冷启动延迟(默认 10)
// 3. DialTimeout/ReadTimeout/WriteTimeout: 精确控制各阶段超时
func InitRedis(cfg *config.Config) *redis.Client {
- return redis.NewClient(buildRedisOptions(cfg))
+ client := redis.NewClient(buildRedisOptions(cfg))
+ if cfg.Server.EnableServerTiming {
+ client.AddHook(serverTimingRedisHook{})
+ }
+ return client
}
// buildRedisOptions 构建 Redis 连接选项
diff --git a/backend/internal/repository/req_client_pool.go b/backend/internal/repository/req_client_pool.go
index 32501f7b19..95ab27ce32 100644
--- a/backend/internal/repository/req_client_pool.go
+++ b/backend/internal/repository/req_client_pool.go
@@ -2,11 +2,13 @@ package repository
import (
"fmt"
+ "net/http"
"strings"
"sync"
"time"
"github.com/Wei-Shaw/sub2api/internal/pkg/proxyurl"
+ "github.com/Wei-Shaw/sub2api/internal/pkg/servertiming"
"github.com/imroc/req/v3"
)
@@ -57,6 +59,7 @@ func getSharedReqClient(opts reqClientOptions) (*req.Client, error) {
if trimmed != "" {
client.SetProxyURL(trimmed)
}
+ client = instrumentReqClient(client)
actual, _ := sharedReqClients.LoadOrStore(key, client)
if c, ok := actual.(*req.Client); ok {
@@ -65,6 +68,17 @@ func getSharedReqClient(opts reqClientOptions) (*req.Client, error) {
return client, nil
}
+func instrumentReqClient(client *req.Client) *req.Client {
+ if client == nil {
+ return nil
+ }
+ client.GetTransport().WrapRoundTripFunc(func(rt http.RoundTripper) req.HttpRoundTripFunc {
+ timed := servertiming.WrapRoundTripper(rt)
+ return timed.RoundTrip
+ })
+ return client
+}
+
func buildReqClientKey(opts reqClientOptions) string {
return fmt.Sprintf("%s|%s|%t|%t",
strings.TrimSpace(opts.ProxyURL),
diff --git a/backend/internal/repository/req_client_pool_test.go b/backend/internal/repository/req_client_pool_test.go
index 9067d0129f..3a27841c5a 100644
--- a/backend/internal/repository/req_client_pool_test.go
+++ b/backend/internal/repository/req_client_pool_test.go
@@ -1,12 +1,17 @@
package repository
import (
+ "context"
+ "net/http"
+ "net/http/httptest"
"reflect"
+ "strings"
"sync"
"testing"
"time"
"unsafe"
+ "github.com/Wei-Shaw/sub2api/internal/pkg/servertiming"
"github.com/imroc/req/v3"
"github.com/stretchr/testify/require"
)
@@ -118,3 +123,20 @@ func TestCreateGeminiReqClient_ForceHTTP2Disabled(t *testing.T) {
require.NoError(t, err)
require.Equal(t, "", forceHTTPVersion(t, client))
}
+
+func TestInstrumentReqClientRecordsDependency(t *testing.T) {
+ server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
+ w.WriteHeader(http.StatusNoContent)
+ }))
+ defer server.Close()
+
+ collector := servertiming.New(time.Now())
+ ctx := servertiming.WithCollector(context.Background(), collector)
+ client := instrumentReqClient(req.C())
+ response, err := client.R().SetContext(ctx).Get(server.URL)
+ require.NoError(t, err)
+ require.Equal(t, http.StatusNoContent, response.StatusCode)
+
+ header := collector.HeaderValue(time.Now(), "bypass")
+ require.True(t, strings.Contains(header, "dep_http;dur="), header)
+}
diff --git a/backend/internal/repository/server_timing_redis.go b/backend/internal/repository/server_timing_redis.go
new file mode 100644
index 0000000000..dba35450de
--- /dev/null
+++ b/backend/internal/repository/server_timing_redis.go
@@ -0,0 +1,39 @@
+package repository
+
+import (
+ "context"
+ "time"
+
+ "github.com/Wei-Shaw/sub2api/internal/pkg/servertiming"
+ "github.com/redis/go-redis/v9"
+)
+
+type serverTimingRedisHook struct{}
+
+func (serverTimingRedisHook) DialHook(next redis.DialHook) redis.DialHook {
+ return next
+}
+
+func (serverTimingRedisHook) ProcessHook(next redis.ProcessHook) redis.ProcessHook {
+ return func(ctx context.Context, cmd redis.Cmder) error {
+ if !servertiming.Active(ctx) {
+ return next(ctx, cmd)
+ }
+ startedAt := time.Now()
+ err := next(ctx, cmd)
+ servertiming.Record(ctx, servertiming.MetricRedis, startedAt, time.Now(), 1)
+ return err
+ }
+}
+
+func (serverTimingRedisHook) ProcessPipelineHook(next redis.ProcessPipelineHook) redis.ProcessPipelineHook {
+ return func(ctx context.Context, cmds []redis.Cmder) error {
+ if !servertiming.Active(ctx) {
+ return next(ctx, cmds)
+ }
+ startedAt := time.Now()
+ err := next(ctx, cmds)
+ servertiming.Record(ctx, servertiming.MetricRedis, startedAt, time.Now(), len(cmds))
+ return err
+ }
+}
diff --git a/backend/internal/repository/server_timing_redis_test.go b/backend/internal/repository/server_timing_redis_test.go
new file mode 100644
index 0000000000..d1ae47e3b0
--- /dev/null
+++ b/backend/internal/repository/server_timing_redis_test.go
@@ -0,0 +1,63 @@
+package repository
+
+import (
+ "context"
+ "errors"
+ "strings"
+ "testing"
+ "time"
+
+ "github.com/Wei-Shaw/sub2api/internal/pkg/servertiming"
+ "github.com/redis/go-redis/v9"
+)
+
+func TestServerTimingRedisHookRecordsCommands(t *testing.T) {
+ collector := servertiming.New(time.Now())
+ ctx := servertiming.WithCollector(context.Background(), collector)
+ hook := serverTimingRedisHook{}
+
+ process := hook.ProcessHook(func(context.Context, redis.Cmder) error {
+ time.Sleep(time.Millisecond)
+ return errors.New("redis failure")
+ })
+ if err := process(ctx, redis.NewStringCmd(ctx, "get", "sensitive-key")); err == nil {
+ t.Fatal("ProcessHook did not return the underlying error")
+ }
+
+ pipeline := hook.ProcessPipelineHook(func(context.Context, []redis.Cmder) error {
+ time.Sleep(time.Millisecond)
+ return nil
+ })
+ commands := []redis.Cmder{
+ redis.NewStringCmd(ctx, "get", "first-secret"),
+ redis.NewStringCmd(ctx, "get", "second-secret"),
+ redis.NewStatusCmd(ctx, "set", "third-secret", "value"),
+ }
+ if err := pipeline(ctx, commands); err != nil {
+ t.Fatal(err)
+ }
+
+ header := collector.HeaderValue(time.Now(), "bypass")
+ if !strings.Contains(header, `commands=4`) {
+ t.Fatalf("header %q does not report one command and a three-command pipeline", header)
+ }
+ if strings.Contains(header, "secret") || strings.Contains(header, "get") {
+ t.Fatalf("Redis command details leaked into header: %q", header)
+ }
+}
+
+func TestServerTimingRedisHookSkipsInactiveContext(t *testing.T) {
+ called := false
+ hook := serverTimingRedisHook{}
+ process := hook.ProcessHook(func(context.Context, redis.Cmder) error {
+ called = true
+ return nil
+ })
+ ctx := context.Background()
+ if err := process(ctx, redis.NewStringCmd(ctx, "ping")); err != nil {
+ t.Fatal(err)
+ }
+ if !called {
+ t.Fatal("inactive Redis command did not reach the next hook")
+ }
+}
diff --git a/backend/internal/repository/server_timing_sql.go b/backend/internal/repository/server_timing_sql.go
new file mode 100644
index 0000000000..062663f08b
--- /dev/null
+++ b/backend/internal/repository/server_timing_sql.go
@@ -0,0 +1,311 @@
+package repository
+
+import (
+ "context"
+ "database/sql/driver"
+ "errors"
+ "io"
+ "reflect"
+ "time"
+
+ "github.com/Wei-Shaw/sub2api/internal/pkg/servertiming"
+)
+
+type serverTimingConnector struct {
+ base driver.Connector
+}
+
+func newServerTimingConnector(base driver.Connector) driver.Connector {
+ return &serverTimingConnector{base: base}
+}
+
+func (c *serverTimingConnector) Connect(ctx context.Context) (driver.Conn, error) {
+ startedAt := time.Now()
+ conn, err := c.base.Connect(ctx)
+ servertiming.RecordInterval(ctx, servertiming.MetricDatabase, startedAt, time.Now())
+ if err != nil {
+ return nil, err
+ }
+ return &serverTimingConn{Conn: conn}, nil
+}
+
+func (c *serverTimingConnector) Driver() driver.Driver {
+ return c.base.Driver()
+}
+
+type serverTimingConn struct {
+ driver.Conn
+}
+
+func (c *serverTimingConn) Prepare(query string) (driver.Stmt, error) {
+ stmt, err := c.Conn.Prepare(query)
+ if err != nil {
+ return nil, err
+ }
+ return &serverTimingStmt{Stmt: stmt}, nil
+}
+
+func (c *serverTimingConn) PrepareContext(ctx context.Context, query string) (driver.Stmt, error) {
+ startedAt := time.Now()
+ var (
+ stmt driver.Stmt
+ err error
+ )
+ if preparer, ok := c.Conn.(driver.ConnPrepareContext); ok {
+ stmt, err = preparer.PrepareContext(ctx, query)
+ } else {
+ stmt, err = c.Conn.Prepare(query)
+ }
+ servertiming.Record(ctx, servertiming.MetricDatabase, startedAt, time.Now(), 1)
+ if err != nil {
+ return nil, err
+ }
+ return &serverTimingStmt{Stmt: stmt}, nil
+}
+
+func (c *serverTimingConn) ExecContext(ctx context.Context, query string, args []driver.NamedValue) (driver.Result, error) {
+ execer, ok := c.Conn.(driver.ExecerContext)
+ if !ok {
+ return nil, driver.ErrSkip
+ }
+ startedAt := time.Now()
+ result, err := execer.ExecContext(ctx, query, args)
+ servertiming.Record(ctx, servertiming.MetricDatabase, startedAt, time.Now(), 1)
+ return result, err
+}
+
+func (c *serverTimingConn) QueryContext(ctx context.Context, query string, args []driver.NamedValue) (driver.Rows, error) {
+ queryer, ok := c.Conn.(driver.QueryerContext)
+ if !ok {
+ return nil, driver.ErrSkip
+ }
+ startedAt := time.Now()
+ rows, err := queryer.QueryContext(ctx, query, args)
+ servertiming.Record(ctx, servertiming.MetricDatabase, startedAt, time.Now(), 1)
+ if err != nil || rows == nil {
+ return rows, err
+ }
+ return newServerTimingRows(ctx, rows), nil
+}
+
+func (c *serverTimingConn) BeginTx(ctx context.Context, opts driver.TxOptions) (driver.Tx, error) {
+ startedAt := time.Now()
+ var (
+ tx driver.Tx
+ err error
+ )
+ if beginner, ok := c.Conn.(driver.ConnBeginTx); ok {
+ tx, err = beginner.BeginTx(ctx, opts)
+ } else {
+ if opts.Isolation != driver.IsolationLevel(0) {
+ return nil, errors.New("driver does not support non-default isolation")
+ }
+ if opts.ReadOnly {
+ return nil, errors.New("driver does not support read-only transactions")
+ }
+ // The wrapper exposes ConnBeginTx, so it must retain database/sql's
+ // legacy fallback for drivers that only implement Conn.Begin.
+ tx, err = c.Conn.Begin() //nolint:staticcheck // Required driver compatibility fallback.
+ }
+ servertiming.RecordInterval(ctx, servertiming.MetricDatabase, startedAt, time.Now())
+ if err != nil || tx == nil {
+ return tx, err
+ }
+ return &serverTimingTx{Tx: tx, ctx: ctx}, nil
+}
+
+func (c *serverTimingConn) Ping(ctx context.Context) error {
+ if pinger, ok := c.Conn.(driver.Pinger); ok {
+ startedAt := time.Now()
+ err := pinger.Ping(ctx)
+ servertiming.RecordInterval(ctx, servertiming.MetricDatabase, startedAt, time.Now())
+ return err
+ }
+ return nil
+}
+
+func (c *serverTimingConn) ResetSession(ctx context.Context) error {
+ if resetter, ok := c.Conn.(driver.SessionResetter); ok {
+ startedAt := time.Now()
+ err := resetter.ResetSession(ctx)
+ servertiming.RecordInterval(ctx, servertiming.MetricDatabase, startedAt, time.Now())
+ return err
+ }
+ return nil
+}
+
+func (c *serverTimingConn) IsValid() bool {
+ if validator, ok := c.Conn.(driver.Validator); ok {
+ return validator.IsValid()
+ }
+ return true
+}
+
+func (c *serverTimingConn) CheckNamedValue(value *driver.NamedValue) error {
+ if checker, ok := c.Conn.(driver.NamedValueChecker); ok {
+ return checker.CheckNamedValue(value)
+ }
+ return driver.ErrSkip
+}
+
+type serverTimingStmt struct {
+ driver.Stmt
+}
+
+func (s *serverTimingStmt) ExecContext(ctx context.Context, args []driver.NamedValue) (driver.Result, error) {
+ startedAt := time.Now()
+ var (
+ result driver.Result
+ err error
+ )
+ if execer, ok := s.Stmt.(driver.StmtExecContext); ok {
+ result, err = execer.ExecContext(ctx, args)
+ } else {
+ var values []driver.Value
+ values, err = namedValues(args)
+ if err == nil {
+ // The wrapper exposes StmtExecContext and must preserve the fallback
+ // database/sql would use for a legacy driver statement.
+ result, err = s.Stmt.Exec(values) //nolint:staticcheck // Required driver compatibility fallback.
+ }
+ }
+ servertiming.Record(ctx, servertiming.MetricDatabase, startedAt, time.Now(), 1)
+ return result, err
+}
+
+func (s *serverTimingStmt) QueryContext(ctx context.Context, args []driver.NamedValue) (driver.Rows, error) {
+ startedAt := time.Now()
+ var (
+ rows driver.Rows
+ err error
+ )
+ if queryer, ok := s.Stmt.(driver.StmtQueryContext); ok {
+ rows, err = queryer.QueryContext(ctx, args)
+ } else {
+ var values []driver.Value
+ values, err = namedValues(args)
+ if err == nil {
+ // The wrapper exposes StmtQueryContext and must preserve the fallback
+ // database/sql would use for a legacy driver statement.
+ rows, err = s.Stmt.Query(values) //nolint:staticcheck // Required driver compatibility fallback.
+ }
+ }
+ servertiming.Record(ctx, servertiming.MetricDatabase, startedAt, time.Now(), 1)
+ if err != nil || rows == nil {
+ return rows, err
+ }
+ return newServerTimingRows(ctx, rows), nil
+}
+
+func (s *serverTimingStmt) CheckNamedValue(value *driver.NamedValue) error {
+ if checker, ok := s.Stmt.(driver.NamedValueChecker); ok {
+ return checker.CheckNamedValue(value)
+ }
+ return driver.ErrSkip
+}
+
+func namedValues(args []driver.NamedValue) ([]driver.Value, error) {
+ values := make([]driver.Value, len(args))
+ for i, arg := range args {
+ if arg.Name != "" {
+ return nil, errors.New("named parameters are not supported")
+ }
+ values[i] = arg.Value
+ }
+ return values, nil
+}
+
+type serverTimingRows struct {
+ driver.Rows
+ ctx context.Context
+}
+
+func newServerTimingRows(ctx context.Context, rows driver.Rows) *serverTimingRows {
+ return &serverTimingRows{Rows: rows, ctx: ctx}
+}
+
+func (r *serverTimingRows) Close() error {
+ startedAt := time.Now()
+ err := r.Rows.Close()
+ servertiming.RecordInterval(r.ctx, servertiming.MetricDatabase, startedAt, time.Now())
+ return err
+}
+
+func (r *serverTimingRows) Next(dest []driver.Value) error {
+ startedAt := time.Now()
+ err := r.Rows.Next(dest)
+ servertiming.RecordInterval(r.ctx, servertiming.MetricDatabase, startedAt, time.Now())
+ return err
+}
+
+func (r *serverTimingRows) HasNextResultSet() bool {
+ if rows, ok := r.Rows.(driver.RowsNextResultSet); ok {
+ return rows.HasNextResultSet()
+ }
+ return false
+}
+
+func (r *serverTimingRows) NextResultSet() error {
+ rows, ok := r.Rows.(driver.RowsNextResultSet)
+ if !ok {
+ return io.EOF
+ }
+ startedAt := time.Now()
+ err := rows.NextResultSet()
+ servertiming.RecordInterval(r.ctx, servertiming.MetricDatabase, startedAt, time.Now())
+ return err
+}
+
+func (r *serverTimingRows) ColumnTypeScanType(index int) reflect.Type {
+ if rows, ok := r.Rows.(driver.RowsColumnTypeScanType); ok {
+ return rows.ColumnTypeScanType(index)
+ }
+ return reflect.TypeOf(new(any)).Elem()
+}
+
+func (r *serverTimingRows) ColumnTypeDatabaseTypeName(index int) string {
+ if rows, ok := r.Rows.(driver.RowsColumnTypeDatabaseTypeName); ok {
+ return rows.ColumnTypeDatabaseTypeName(index)
+ }
+ return ""
+}
+
+func (r *serverTimingRows) ColumnTypeLength(index int) (int64, bool) {
+ if rows, ok := r.Rows.(driver.RowsColumnTypeLength); ok {
+ return rows.ColumnTypeLength(index)
+ }
+ return 0, false
+}
+
+func (r *serverTimingRows) ColumnTypeNullable(index int) (bool, bool) {
+ if rows, ok := r.Rows.(driver.RowsColumnTypeNullable); ok {
+ return rows.ColumnTypeNullable(index)
+ }
+ return false, false
+}
+
+func (r *serverTimingRows) ColumnTypePrecisionScale(index int) (int64, int64, bool) {
+ if rows, ok := r.Rows.(driver.RowsColumnTypePrecisionScale); ok {
+ return rows.ColumnTypePrecisionScale(index)
+ }
+ return 0, 0, false
+}
+
+type serverTimingTx struct {
+ driver.Tx
+ ctx context.Context
+}
+
+func (t *serverTimingTx) Commit() error {
+ startedAt := time.Now()
+ err := t.Tx.Commit()
+ servertiming.RecordInterval(t.ctx, servertiming.MetricDatabase, startedAt, time.Now())
+ return err
+}
+
+func (t *serverTimingTx) Rollback() error {
+ startedAt := time.Now()
+ err := t.Tx.Rollback()
+ servertiming.RecordInterval(t.ctx, servertiming.MetricDatabase, startedAt, time.Now())
+ return err
+}
diff --git a/backend/internal/repository/server_timing_sql_test.go b/backend/internal/repository/server_timing_sql_test.go
new file mode 100644
index 0000000000..3a8bbbe03e
--- /dev/null
+++ b/backend/internal/repository/server_timing_sql_test.go
@@ -0,0 +1,258 @@
+package repository
+
+import (
+ "context"
+ "database/sql/driver"
+ "io"
+ "regexp"
+ "strconv"
+ "strings"
+ "testing"
+ "time"
+
+ "github.com/Wei-Shaw/sub2api/internal/pkg/servertiming"
+)
+
+const fakeDriverDelay = 2 * time.Millisecond
+
+type timingFakeDriver struct{}
+
+func (timingFakeDriver) Open(string) (driver.Conn, error) { return newTimingFakeConn(), nil }
+
+type timingFakeConnector struct {
+ conn driver.Conn
+}
+
+func (c timingFakeConnector) Connect(context.Context) (driver.Conn, error) {
+ time.Sleep(fakeDriverDelay)
+ return c.conn, nil
+}
+
+func (timingFakeConnector) Driver() driver.Driver { return timingFakeDriver{} }
+
+type timingFakeConn struct{}
+
+func newTimingFakeConn() *timingFakeConn { return &timingFakeConn{} }
+
+func (c *timingFakeConn) Prepare(string) (driver.Stmt, error) {
+ time.Sleep(fakeDriverDelay)
+ return &timingFakeStmt{}, nil
+}
+
+func (c *timingFakeConn) PrepareContext(context.Context, string) (driver.Stmt, error) {
+ time.Sleep(fakeDriverDelay)
+ return &timingFakeStmt{}, nil
+}
+
+func (c *timingFakeConn) Close() error { return nil }
+
+func (c *timingFakeConn) Begin() (driver.Tx, error) {
+ time.Sleep(fakeDriverDelay)
+ return &timingFakeTx{}, nil
+}
+
+func (c *timingFakeConn) BeginTx(context.Context, driver.TxOptions) (driver.Tx, error) {
+ time.Sleep(fakeDriverDelay)
+ return &timingFakeTx{}, nil
+}
+
+func (c *timingFakeConn) ExecContext(context.Context, string, []driver.NamedValue) (driver.Result, error) {
+ time.Sleep(fakeDriverDelay)
+ return driver.RowsAffected(1), nil
+}
+
+func (c *timingFakeConn) QueryContext(context.Context, string, []driver.NamedValue) (driver.Rows, error) {
+ time.Sleep(fakeDriverDelay)
+ return &timingFakeRows{values: [][]driver.Value{{"value"}}}, nil
+}
+
+func (c *timingFakeConn) Ping(context.Context) error {
+ time.Sleep(fakeDriverDelay)
+ return nil
+}
+
+func (c *timingFakeConn) ResetSession(context.Context) error {
+ time.Sleep(fakeDriverDelay)
+ return nil
+}
+
+type timingFakeStmt struct{}
+
+func (s *timingFakeStmt) Close() error { return nil }
+func (s *timingFakeStmt) NumInput() int { return -1 }
+
+func (s *timingFakeStmt) Exec([]driver.Value) (driver.Result, error) {
+ time.Sleep(fakeDriverDelay)
+ return driver.RowsAffected(1), nil
+}
+
+func (s *timingFakeStmt) Query([]driver.Value) (driver.Rows, error) {
+ time.Sleep(fakeDriverDelay)
+ return &timingFakeRows{values: [][]driver.Value{{"value"}}}, nil
+}
+
+func (s *timingFakeStmt) ExecContext(context.Context, []driver.NamedValue) (driver.Result, error) {
+ time.Sleep(fakeDriverDelay)
+ return driver.RowsAffected(1), nil
+}
+
+func (s *timingFakeStmt) QueryContext(context.Context, []driver.NamedValue) (driver.Rows, error) {
+ time.Sleep(fakeDriverDelay)
+ return &timingFakeRows{values: [][]driver.Value{{"value"}}}, nil
+}
+
+type timingFakeRows struct {
+ values [][]driver.Value
+ index int
+}
+
+func (r *timingFakeRows) Columns() []string { return []string{"value"} }
+
+func (r *timingFakeRows) Close() error {
+ time.Sleep(fakeDriverDelay)
+ return nil
+}
+
+func (r *timingFakeRows) Next(dest []driver.Value) error {
+ time.Sleep(fakeDriverDelay)
+ if r.index >= len(r.values) {
+ return io.EOF
+ }
+ copy(dest, r.values[r.index])
+ r.index++
+ return nil
+}
+
+type timingFakeTx struct{}
+
+func (t *timingFakeTx) Commit() error {
+ time.Sleep(fakeDriverDelay)
+ return nil
+}
+
+func (t *timingFakeTx) Rollback() error {
+ time.Sleep(fakeDriverDelay)
+ return nil
+}
+
+func metricDuration(t *testing.T, header, metric string) float64 {
+ t.Helper()
+ re := regexp.MustCompile(`(?:^|, )` + regexp.QuoteMeta(metric) + `;dur=([0-9]+(?:\.[0-9]+)?)`)
+ match := re.FindStringSubmatch(header)
+ if len(match) != 2 {
+ t.Fatalf("metric %q missing from header %q", metric, header)
+ }
+ value, err := strconv.ParseFloat(match[1], 64)
+ if err != nil {
+ t.Fatalf("parse %s duration: %v", metric, err)
+ }
+ return value
+}
+
+func TestServerTimingConnectorRecordsDriverCallsWithoutRowLifetime(t *testing.T) {
+ startedAt := time.Now()
+ collector := servertiming.New(startedAt)
+ ctx := servertiming.WithCollector(context.Background(), collector)
+
+ wrapped := newServerTimingConnector(timingFakeConnector{conn: newTimingFakeConn()})
+ rawConn, err := wrapped.Connect(ctx)
+ if err != nil {
+ t.Fatal(err)
+ }
+ conn, ok := rawConn.(*serverTimingConn)
+ if !ok {
+ t.Fatalf("Connect() returned %T, want *serverTimingConn", rawConn)
+ }
+
+ if _, err := conn.ExecContext(ctx, "sensitive update", nil); err != nil {
+ t.Fatal(err)
+ }
+ rows, err := conn.QueryContext(ctx, "sensitive select", nil)
+ if err != nil {
+ t.Fatal(err)
+ }
+ values := make([]driver.Value, 1)
+ if err := rows.Next(values); err != nil {
+ t.Fatal(err)
+ }
+
+ // Application work between row reads must remain app time.
+ time.Sleep(30 * time.Millisecond)
+ if err := rows.Next(values); err != io.EOF {
+ t.Fatalf("rows.Next() = %v, want EOF", err)
+ }
+ if err := rows.Close(); err != nil {
+ t.Fatal(err)
+ }
+
+ header := collector.HeaderValue(time.Now(), "bypass")
+ if !strings.Contains(header, `queries=2`) {
+ t.Fatalf("header %q does not report two SQL operations", header)
+ }
+ if strings.Contains(header, "sensitive") {
+ t.Fatalf("SQL text leaked into header: %q", header)
+ }
+ if app, db := metricDuration(t, header, "app"), metricDuration(t, header, "db"); app <= db {
+ t.Fatalf("row processing gap was counted as DB time: app=%.1fms db=%.1fms header=%q", app, db, header)
+ }
+}
+
+func TestServerTimingPreparedStatementsAndTransactions(t *testing.T) {
+ collector := servertiming.New(time.Now())
+ ctx := servertiming.WithCollector(context.Background(), collector)
+ conn := &serverTimingConn{Conn: newTimingFakeConn()}
+
+ stmt, err := conn.PrepareContext(ctx, "prepare sensitive statement")
+ if err != nil {
+ t.Fatal(err)
+ }
+ timedStmt, ok := stmt.(*serverTimingStmt)
+ if !ok {
+ t.Fatalf("PrepareContext() returned %T, want *serverTimingStmt", stmt)
+ }
+ if _, err := timedStmt.ExecContext(ctx, nil); err != nil {
+ t.Fatal(err)
+ }
+ rows, err := timedStmt.QueryContext(ctx, nil)
+ if err != nil {
+ t.Fatal(err)
+ }
+ if err := rows.Close(); err != nil {
+ t.Fatal(err)
+ }
+
+ tx, err := conn.BeginTx(ctx, driver.TxOptions{})
+ if err != nil {
+ t.Fatal(err)
+ }
+ if err := tx.Commit(); err != nil {
+ t.Fatal(err)
+ }
+ if err := conn.Ping(ctx); err != nil {
+ t.Fatal(err)
+ }
+ if err := conn.ResetSession(ctx); err != nil {
+ t.Fatal(err)
+ }
+
+ header := collector.HeaderValue(time.Now(), "bypass")
+ if !strings.Contains(header, `queries=3`) {
+ t.Fatalf("header %q does not report prepare, exec, and query operations", header)
+ }
+ if metricDuration(t, header, "db") <= 0 {
+ t.Fatalf("DB duration was not recorded: %q", header)
+ }
+}
+
+func TestNamedValuesRejectNamedParameters(t *testing.T) {
+ if _, err := namedValues([]driver.NamedValue{{Name: "secret", Value: 1}}); err == nil {
+ t.Fatal("namedValues accepted a named parameter")
+ }
+ values, err := namedValues([]driver.NamedValue{{Ordinal: 1, Value: "value"}})
+ if err != nil {
+ t.Fatal(err)
+ }
+ if len(values) != 1 || values[0] != "value" {
+ t.Fatalf("namedValues() = %#v", values)
+ }
+}
diff --git a/backend/internal/repository/usage_log_repo_insert.go b/backend/internal/repository/usage_log_repo_insert.go
index dfd8969512..ec09b308a0 100644
--- a/backend/internal/repository/usage_log_repo_insert.go
+++ b/backend/internal/repository/usage_log_repo_insert.go
@@ -71,6 +71,7 @@ var usageLogInsertArgTypes = [...]string{
"text", // inbound_endpoint
"text", // upstream_endpoint
"boolean", // cache_ttl_overridden
+ "boolean", // long_context_billing_applied
"bigint", // channel_id
"text", // model_mapping_chain
"text", // billing_tier
@@ -263,6 +264,7 @@ func (r *usageLogRepository) createSingle(ctx context.Context, sqlq sqlExecutor,
inbound_endpoint,
upstream_endpoint,
cache_ttl_overridden,
+ long_context_billing_applied,
channel_id,
model_mapping_chain,
billing_tier,
@@ -275,7 +277,7 @@ func (r *usageLogRepository) createSingle(ctx context.Context, sqlq sqlExecutor,
$10, $11, $12, $13,
$14, $15, $16, $17,
$18, $19, $20, $21, $22, $23,
- $24, $25, $26, $27, $28, $29, $30, $31, $32, $33, $34, $35, $36, $37, $38, $39, $40, $41, $42, $43, $44, $45, $46, $47, $48, $49, $50, $51, $52, $53
+ $24, $25, $26, $27, $28, $29, $30, $31, $32, $33, $34, $35, $36, $37, $38, $39, $40, $41, $42, $43, $44, $45, $46, $47, $48, $49, $50, $51, $52, $53, $54
)
ON CONFLICT (request_id, api_key_id) DO NOTHING
RETURNING id, created_at
@@ -714,6 +716,7 @@ func buildUsageLogBatchInsertQuery(keys []string, preparedByKey map[string]usage
inbound_endpoint,
upstream_endpoint,
cache_ttl_overridden,
+ long_context_billing_applied,
channel_id,
model_mapping_chain,
billing_tier,
@@ -722,7 +725,7 @@ func buildUsageLogBatchInsertQuery(keys []string, preparedByKey map[string]usage
created_at
) AS (VALUES `)
- args := make([]any, 0, len(keys)*53)
+ args := make([]any, 0, len(keys)*54)
argPos := 1
for idx, key := range keys {
if idx > 0 {
@@ -798,6 +801,7 @@ func buildUsageLogBatchInsertQuery(keys []string, preparedByKey map[string]usage
inbound_endpoint,
upstream_endpoint,
cache_ttl_overridden,
+ long_context_billing_applied,
channel_id,
model_mapping_chain,
billing_tier,
@@ -853,6 +857,7 @@ func buildUsageLogBatchInsertQuery(keys []string, preparedByKey map[string]usage
inbound_endpoint,
upstream_endpoint,
cache_ttl_overridden,
+ long_context_billing_applied,
channel_id,
model_mapping_chain,
billing_tier,
@@ -948,6 +953,7 @@ func buildUsageLogBestEffortInsertQuery(preparedList []usageLogInsertPrepared) (
inbound_endpoint,
upstream_endpoint,
cache_ttl_overridden,
+ long_context_billing_applied,
channel_id,
model_mapping_chain,
billing_tier,
@@ -956,7 +962,7 @@ func buildUsageLogBestEffortInsertQuery(preparedList []usageLogInsertPrepared) (
created_at
) AS (VALUES `)
- args := make([]any, 0, len(preparedList)*53)
+ args := make([]any, 0, len(preparedList)*54)
argPos := 1
for idx, prepared := range preparedList {
if idx > 0 {
@@ -1029,6 +1035,7 @@ func buildUsageLogBestEffortInsertQuery(preparedList []usageLogInsertPrepared) (
inbound_endpoint,
upstream_endpoint,
cache_ttl_overridden,
+ long_context_billing_applied,
channel_id,
model_mapping_chain,
billing_tier,
@@ -1084,6 +1091,7 @@ func buildUsageLogBestEffortInsertQuery(preparedList []usageLogInsertPrepared) (
inbound_endpoint,
upstream_endpoint,
cache_ttl_overridden,
+ long_context_billing_applied,
channel_id,
model_mapping_chain,
billing_tier,
@@ -1147,6 +1155,7 @@ func execUsageLogInsertNoResult(ctx context.Context, sqlq sqlExecutor, prepared
inbound_endpoint,
upstream_endpoint,
cache_ttl_overridden,
+ long_context_billing_applied,
channel_id,
model_mapping_chain,
billing_tier,
@@ -1159,7 +1168,7 @@ func execUsageLogInsertNoResult(ctx context.Context, sqlq sqlExecutor, prepared
$10, $11, $12, $13,
$14, $15, $16, $17,
$18, $19, $20, $21, $22, $23,
- $24, $25, $26, $27, $28, $29, $30, $31, $32, $33, $34, $35, $36, $37, $38, $39, $40, $41, $42, $43, $44, $45, $46, $47, $48, $49, $50, $51, $52, $53
+ $24, $25, $26, $27, $28, $29, $30, $31, $32, $33, $34, $35, $36, $37, $38, $39, $40, $41, $42, $43, $44, $45, $46, $47, $48, $49, $50, $51, $52, $53, $54
)
ON CONFLICT (request_id, api_key_id) DO NOTHING
`, prepared.args...)
@@ -1264,6 +1273,7 @@ func prepareUsageLogInsert(log *service.UsageLog) usageLogInsertPrepared {
inboundEndpoint,
upstreamEndpoint,
log.CacheTTLOverridden,
+ log.LongContextBillingApplied,
channelID,
modelMappingChain,
billingTier,
diff --git a/backend/internal/repository/usage_log_repo_query.go b/backend/internal/repository/usage_log_repo_query.go
index c178429bab..1fdedd8665 100644
--- a/backend/internal/repository/usage_log_repo_query.go
+++ b/backend/internal/repository/usage_log_repo_query.go
@@ -19,7 +19,7 @@ import (
"github.com/Wei-Shaw/sub2api/internal/service"
)
-const usageLogSelectColumns = "id, user_id, api_key_id, account_id, request_id, model, requested_model, upstream_model, group_id, subscription_id, input_tokens, output_tokens, cache_creation_tokens, cache_read_tokens, cache_creation_5m_tokens, cache_creation_1h_tokens, image_output_tokens, image_output_cost, input_cost, output_cost, cache_creation_cost, cache_read_cost, total_cost, actual_cost, rate_multiplier, account_rate_multiplier, billing_type, request_type, stream, openai_ws_mode, duration_ms, first_token_ms, user_agent, ip_address, image_count, image_size, image_input_size, image_output_size, image_size_source, image_size_breakdown, video_count, video_resolution, video_duration_seconds, service_tier, reasoning_effort, inbound_endpoint, upstream_endpoint, cache_ttl_overridden, channel_id, model_mapping_chain, billing_tier, billing_mode, account_stats_cost, created_at"
+const usageLogSelectColumns = "id, user_id, api_key_id, account_id, request_id, model, requested_model, upstream_model, group_id, subscription_id, input_tokens, output_tokens, cache_creation_tokens, cache_read_tokens, cache_creation_5m_tokens, cache_creation_1h_tokens, image_output_tokens, image_output_cost, input_cost, output_cost, cache_creation_cost, cache_read_cost, total_cost, actual_cost, rate_multiplier, account_rate_multiplier, billing_type, request_type, stream, openai_ws_mode, duration_ms, first_token_ms, user_agent, ip_address, image_count, image_size, image_input_size, image_output_size, image_size_source, image_size_breakdown, video_count, video_resolution, video_duration_seconds, service_tier, reasoning_effort, inbound_endpoint, upstream_endpoint, cache_ttl_overridden, long_context_billing_applied, channel_id, model_mapping_chain, billing_tier, billing_mode, account_stats_cost, created_at"
func (r *usageLogRepository) GetByID(ctx context.Context, id int64) (log *service.UsageLog, err error) {
query := "SELECT " + usageLogSelectColumns + " FROM usage_logs WHERE id = $1"
@@ -425,60 +425,61 @@ func (r *usageLogRepository) loadSubscriptions(ctx context.Context, ids []int64)
func scanUsageLog(scanner interface{ Scan(...any) error }) (*service.UsageLog, error) {
var (
- id int64
- userID int64
- apiKeyID int64
- accountID int64
- requestID sql.NullString
- model string
- requestedModel sql.NullString
- upstreamModel sql.NullString
- groupID sql.NullInt64
- subscriptionID sql.NullInt64
- inputTokens int
- outputTokens int
- cacheCreationTokens int
- cacheReadTokens int
- cacheCreation5m int
- cacheCreation1h int
- imageOutputTokens int
- imageOutputCost float64
- inputCost float64
- outputCost float64
- cacheCreationCost float64
- cacheReadCost float64
- totalCost float64
- actualCost float64
- rateMultiplier float64
- accountRateMultiplier sql.NullFloat64
- billingType int16
- requestTypeRaw int16
- stream bool
- openaiWSMode bool
- durationMs sql.NullInt64
- firstTokenMs sql.NullInt64
- userAgent sql.NullString
- ipAddress sql.NullString
- imageCount int
- imageSize sql.NullString
- imageInputSize sql.NullString
- imageOutputSize sql.NullString
- imageSizeSource sql.NullString
- imageSizeBreakdown sql.NullString
- videoCount int
- videoResolution sql.NullString
- videoDurationSeconds sql.NullInt64
- serviceTier sql.NullString
- reasoningEffort sql.NullString
- inboundEndpoint sql.NullString
- upstreamEndpoint sql.NullString
- cacheTTLOverridden bool
- channelID sql.NullInt64
- modelMappingChain sql.NullString
- billingTier sql.NullString
- billingMode sql.NullString
- accountStatsCost sql.NullFloat64
- createdAt time.Time
+ id int64
+ userID int64
+ apiKeyID int64
+ accountID int64
+ requestID sql.NullString
+ model string
+ requestedModel sql.NullString
+ upstreamModel sql.NullString
+ groupID sql.NullInt64
+ subscriptionID sql.NullInt64
+ inputTokens int
+ outputTokens int
+ cacheCreationTokens int
+ cacheReadTokens int
+ cacheCreation5m int
+ cacheCreation1h int
+ imageOutputTokens int
+ imageOutputCost float64
+ inputCost float64
+ outputCost float64
+ cacheCreationCost float64
+ cacheReadCost float64
+ totalCost float64
+ actualCost float64
+ rateMultiplier float64
+ accountRateMultiplier sql.NullFloat64
+ billingType int16
+ requestTypeRaw int16
+ stream bool
+ openaiWSMode bool
+ durationMs sql.NullInt64
+ firstTokenMs sql.NullInt64
+ userAgent sql.NullString
+ ipAddress sql.NullString
+ imageCount int
+ imageSize sql.NullString
+ imageInputSize sql.NullString
+ imageOutputSize sql.NullString
+ imageSizeSource sql.NullString
+ imageSizeBreakdown sql.NullString
+ videoCount int
+ videoResolution sql.NullString
+ videoDurationSeconds sql.NullInt64
+ serviceTier sql.NullString
+ reasoningEffort sql.NullString
+ inboundEndpoint sql.NullString
+ upstreamEndpoint sql.NullString
+ cacheTTLOverridden bool
+ longContextBillingApplied bool
+ channelID sql.NullInt64
+ modelMappingChain sql.NullString
+ billingTier sql.NullString
+ billingMode sql.NullString
+ accountStatsCost sql.NullFloat64
+ createdAt time.Time
)
if err := scanner.Scan(
@@ -530,6 +531,7 @@ func scanUsageLog(scanner interface{ Scan(...any) error }) (*service.UsageLog, e
&inboundEndpoint,
&upstreamEndpoint,
&cacheTTLOverridden,
+ &longContextBillingApplied,
&channelID,
&modelMappingChain,
&billingTier,
@@ -541,34 +543,35 @@ func scanUsageLog(scanner interface{ Scan(...any) error }) (*service.UsageLog, e
}
log := &service.UsageLog{
- ID: id,
- UserID: userID,
- APIKeyID: apiKeyID,
- AccountID: accountID,
- Model: model,
- RequestedModel: coalesceTrimmedString(requestedModel, model),
- InputTokens: inputTokens,
- OutputTokens: outputTokens,
- CacheCreationTokens: cacheCreationTokens,
- CacheReadTokens: cacheReadTokens,
- CacheCreation5mTokens: cacheCreation5m,
- CacheCreation1hTokens: cacheCreation1h,
- ImageOutputTokens: imageOutputTokens,
- ImageOutputCost: imageOutputCost,
- InputCost: inputCost,
- OutputCost: outputCost,
- CacheCreationCost: cacheCreationCost,
- CacheReadCost: cacheReadCost,
- TotalCost: totalCost,
- ActualCost: actualCost,
- RateMultiplier: rateMultiplier,
- AccountRateMultiplier: nullFloat64Ptr(accountRateMultiplier),
- BillingType: int8(billingType),
- RequestType: service.RequestTypeFromInt16(requestTypeRaw),
- ImageCount: imageCount,
- VideoCount: videoCount,
- CacheTTLOverridden: cacheTTLOverridden,
- CreatedAt: createdAt,
+ ID: id,
+ UserID: userID,
+ APIKeyID: apiKeyID,
+ AccountID: accountID,
+ Model: model,
+ RequestedModel: coalesceTrimmedString(requestedModel, model),
+ InputTokens: inputTokens,
+ OutputTokens: outputTokens,
+ CacheCreationTokens: cacheCreationTokens,
+ CacheReadTokens: cacheReadTokens,
+ CacheCreation5mTokens: cacheCreation5m,
+ CacheCreation1hTokens: cacheCreation1h,
+ ImageOutputTokens: imageOutputTokens,
+ ImageOutputCost: imageOutputCost,
+ InputCost: inputCost,
+ OutputCost: outputCost,
+ CacheCreationCost: cacheCreationCost,
+ CacheReadCost: cacheReadCost,
+ TotalCost: totalCost,
+ ActualCost: actualCost,
+ RateMultiplier: rateMultiplier,
+ AccountRateMultiplier: nullFloat64Ptr(accountRateMultiplier),
+ BillingType: int8(billingType),
+ RequestType: service.RequestTypeFromInt16(requestTypeRaw),
+ ImageCount: imageCount,
+ VideoCount: videoCount,
+ CacheTTLOverridden: cacheTTLOverridden,
+ LongContextBillingApplied: longContextBillingApplied,
+ CreatedAt: createdAt,
}
// 先回填 legacy 字段,再基于 legacy + request_type 计算最终请求类型,保证历史数据兼容。
log.Stream = stream
diff --git a/backend/internal/repository/usage_log_repo_request_type_test.go b/backend/internal/repository/usage_log_repo_request_type_test.go
index c32ad2b63f..052c319183 100644
--- a/backend/internal/repository/usage_log_repo_request_type_test.go
+++ b/backend/internal/repository/usage_log_repo_request_type_test.go
@@ -88,6 +88,7 @@ func TestUsageLogRepositoryCreateSyncRequestTypeAndLegacyFields(t *testing.T) {
sqlmock.AnyArg(), // inbound_endpoint
sqlmock.AnyArg(), // upstream_endpoint
log.CacheTTLOverridden,
+ log.LongContextBillingApplied,
sqlmock.AnyArg(), // channel_id
sqlmock.AnyArg(), // model_mapping_chain
sqlmock.AnyArg(), // billing_tier
@@ -174,6 +175,7 @@ func TestUsageLogRepositoryCreate_PersistsServiceTier(t *testing.T) {
sqlmock.AnyArg(),
sqlmock.AnyArg(),
log.CacheTTLOverridden,
+ log.LongContextBillingApplied,
sqlmock.AnyArg(), // channel_id
sqlmock.AnyArg(), // model_mapping_chain
sqlmock.AnyArg(), // billing_tier
@@ -813,6 +815,7 @@ func TestScanUsageLogRequestTypeAndLegacyFallback(t *testing.T) {
sql.NullString{},
sql.NullString{},
false,
+ false,
sql.NullInt64{},
sql.NullString{},
sql.NullString{},
@@ -884,6 +887,7 @@ func TestScanUsageLogRequestTypeAndLegacyFallback(t *testing.T) {
sql.NullString{},
sql.NullString{},
false,
+ false,
sql.NullInt64{}, // channel_id
sql.NullString{}, // model_mapping_chain
sql.NullString{}, // billing_tier
@@ -939,6 +943,7 @@ func TestScanUsageLogRequestTypeAndLegacyFallback(t *testing.T) {
sql.NullString{},
sql.NullString{},
false,
+ false,
sql.NullInt64{}, // channel_id
sql.NullString{}, // model_mapping_chain
sql.NullString{}, // billing_tier
@@ -994,6 +999,7 @@ func TestScanUsageLogRequestTypeAndLegacyFallback(t *testing.T) {
sql.NullString{},
sql.NullString{},
false,
+ false,
sql.NullInt64{}, // channel_id
sql.NullString{}, // model_mapping_chain
sql.NullString{}, // billing_tier
diff --git a/backend/internal/server/api_contract_test.go b/backend/internal/server/api_contract_test.go
index 372cc46bbf..a5e3fde155 100644
--- a/backend/internal/server/api_contract_test.go
+++ b/backend/internal/server/api_contract_test.go
@@ -594,6 +594,7 @@ func TestAPIContracts(t *testing.T) {
"total_cost": 0.5,
"actual_cost": 0.5,
"rate_multiplier": 1,
+ "long_context_billing_applied": false,
"billing_type": 0,
"stream": true,
"duration_ms": 100,
diff --git a/backend/internal/server/middleware/cors.go b/backend/internal/server/middleware/cors.go
index 03d5d025de..0283d53115 100644
--- a/backend/internal/server/middleware/cors.go
+++ b/backend/internal/server/middleware/cors.go
@@ -52,7 +52,7 @@ func CORS(cfg config.CORSConfig) gin.HandlerFunc {
}
allowHeaders := []string{
"Content-Type", "Content-Length", "Accept-Encoding", "X-CSRF-Token", "Authorization",
- "accept", "origin", "Cache-Control", "X-Requested-With", "X-API-Key",
+ "accept", "origin", "Cache-Control", "X-Requested-With", "X-API-Key", "X-Admin-UI-Request",
}
// OpenAI Node SDK 会发送 x-stainless-* 请求头,需在 CORS 中显式放行。
openAIProperties := []string{
@@ -83,7 +83,7 @@ func CORS(cfg config.CORSConfig) gin.HandlerFunc {
}
c.Writer.Header().Set("Access-Control-Allow-Headers", allowHeadersValue)
c.Writer.Header().Set("Access-Control-Allow-Methods", "POST, OPTIONS, GET, PUT, DELETE, PATCH")
- c.Writer.Header().Set("Access-Control-Expose-Headers", "ETag")
+ c.Writer.Header().Set("Access-Control-Expose-Headers", "ETag, Server-Timing")
c.Writer.Header().Set("Access-Control-Max-Age", "86400")
}
// 处理预检请求
diff --git a/backend/internal/server/middleware/cors_test.go b/backend/internal/server/middleware/cors_test.go
index 6d0bea3608..6a61f696df 100644
--- a/backend/internal/server/middleware/cors_test.go
+++ b/backend/internal/server/middleware/cors_test.go
@@ -103,8 +103,10 @@ func TestCORS_AllowedOrigin_HasAllowHeaders(t *testing.T) {
// 应设置 Allow-Headers、Allow-Methods 和 Max-Age
assert.NotEmpty(t, w.Header().Get("Access-Control-Allow-Headers"),
"允许的 origin 应收到 Allow-Headers")
+ assert.Contains(t, w.Header().Get("Access-Control-Allow-Headers"), "X-Admin-UI-Request")
assert.NotEmpty(t, w.Header().Get("Access-Control-Allow-Methods"),
"允许的 origin 应收到 Allow-Methods")
+ assert.Contains(t, w.Header().Get("Access-Control-Expose-Headers"), "Server-Timing")
assert.Equal(t, "86400", w.Header().Get("Access-Control-Max-Age"),
"允许的 origin 应收到 Max-Age=86400")
assert.Equal(t, "https://allowed.example.com", w.Header().Get("Access-Control-Allow-Origin"),
diff --git a/backend/internal/server/middleware/server_timing.go b/backend/internal/server/middleware/server_timing.go
new file mode 100644
index 0000000000..2bb21071e0
--- /dev/null
+++ b/backend/internal/server/middleware/server_timing.go
@@ -0,0 +1,132 @@
+package middleware
+
+import (
+ "net/http"
+ "strings"
+ "sync"
+ "time"
+
+ "github.com/Wei-Shaw/sub2api/internal/pkg/servertiming"
+ "github.com/gin-gonic/gin"
+)
+
+const (
+ snapshotCacheHeader = "X-Snapshot-Cache"
+ usageCacheHeader = "X-Usage-Stats-Cache"
+)
+
+type serverTimingResponseWriter struct {
+ gin.ResponseWriter
+ context *gin.Context
+ once sync.Once
+}
+
+func (w *serverTimingResponseWriter) Unwrap() http.ResponseWriter {
+ return w.ResponseWriter
+}
+
+// ServerTiming collects timing only for requests made by the Admin web UI.
+func ServerTiming(enabled bool) gin.HandlerFunc {
+ if !enabled {
+ return func(c *gin.Context) {
+ c.Next()
+ }
+ }
+ return func(c *gin.Context) {
+ if !isAdminUIRequest(c) || c.Request == nil {
+ c.Next()
+ return
+ }
+
+ collector := servertiming.New(time.Now())
+ c.Request = c.Request.WithContext(servertiming.WithCollector(c.Request.Context(), collector))
+ writer := &serverTimingResponseWriter{
+ ResponseWriter: c.Writer,
+ context: c,
+ }
+ c.Writer = writer
+ c.Next()
+ writer.finalize()
+ }
+}
+
+func (w *serverTimingResponseWriter) WriteHeader(statusCode int) {
+ w.ResponseWriter.WriteHeader(statusCode)
+}
+
+func (w *serverTimingResponseWriter) WriteHeaderNow() {
+ w.finalize()
+ w.ResponseWriter.WriteHeaderNow()
+}
+
+func (w *serverTimingResponseWriter) Write(data []byte) (int, error) {
+ w.finalize()
+ return w.ResponseWriter.Write(data)
+}
+
+func (w *serverTimingResponseWriter) WriteString(data string) (int, error) {
+ w.finalize()
+ return w.ResponseWriter.WriteString(data)
+}
+
+func (w *serverTimingResponseWriter) Flush() {
+ w.finalize()
+ w.ResponseWriter.Flush()
+}
+
+func (w *serverTimingResponseWriter) finalize() {
+ if w == nil {
+ return
+ }
+ w.once.Do(func() {
+ if value := ServerTimingHeaderValue(w.context); value != "" {
+ w.ResponseWriter.Header().Set(servertiming.HeaderName, value)
+ }
+ })
+}
+
+// ServerTimingHeaderValue returns a timing value only for an authenticated admin.
+func ServerTimingHeaderValue(c *gin.Context) string {
+ if c == nil || c.Request == nil {
+ return ""
+ }
+ role, ok := GetUserRoleFromContext(c)
+ if !ok || role != "admin" {
+ return ""
+ }
+ return servertiming.HeaderValue(c.Request.Context(), time.Now(), responseCacheStatus(c.Writer.Header()))
+}
+
+// ServerTimingResponseHeader builds the extra header map required by WebSocket upgrades.
+func ServerTimingResponseHeader(c *gin.Context) http.Header {
+ value := ServerTimingHeaderValue(c)
+ if value == "" {
+ return nil
+ }
+ return http.Header{servertiming.HeaderName: []string{value}}
+}
+
+func isAdminUIRequest(c *gin.Context) bool {
+ if c == nil || c.Request == nil || c.Request.URL == nil {
+ return false
+ }
+ if strings.TrimSpace(c.GetHeader(servertiming.AdminUIHeader)) == "1" {
+ return true
+ }
+ path := strings.TrimSpace(c.Request.URL.Path)
+ return path == "/api/v1/admin" || strings.HasPrefix(path, "/api/v1/admin/")
+}
+
+func responseCacheStatus(header http.Header) string {
+ for _, name := range []string{snapshotCacheHeader, usageCacheHeader} {
+ switch strings.ToLower(strings.TrimSpace(header.Get(name))) {
+ case "hit":
+ return "hit"
+ case "miss":
+ return "miss"
+ case "bypass":
+ return "bypass"
+ }
+ }
+ return "bypass"
+}
diff --git a/backend/internal/server/middleware/server_timing_test.go b/backend/internal/server/middleware/server_timing_test.go
new file mode 100644
index 0000000000..c064840ece
--- /dev/null
+++ b/backend/internal/server/middleware/server_timing_test.go
@@ -0,0 +1,188 @@
+package middleware
+
+import (
+ "net/http"
+ "net/http/httptest"
+ "strings"
+ "testing"
+ "time"
+
+ "github.com/Wei-Shaw/sub2api/internal/pkg/servertiming"
+ "github.com/gin-gonic/gin"
+)
+
+func runServerTimingRequest(
+ t *testing.T,
+ enabled bool,
+ path string,
+ marker string,
+ role string,
+ handler gin.HandlerFunc,
+) *httptest.ResponseRecorder {
+ t.Helper()
+ gin.SetMode(gin.TestMode)
+ engine := gin.New()
+ engine.Use(ServerTiming(enabled))
+ engine.Any("/*path", func(c *gin.Context) {
+ if role != "" {
+ c.Set(string(ContextKeyUserRole), role)
+ }
+ handler(c)
+ })
+
+ recorder := httptest.NewRecorder()
+ request := httptest.NewRequest(http.MethodGet, path, nil)
+ if marker != "" {
+ request.Header.Set(servertiming.AdminUIHeader, marker)
+ }
+ engine.ServeHTTP(recorder, request)
+ return recorder
+}
+
+func TestServerTimingScopesAndRoleGate(t *testing.T) {
+ tests := []struct {
+ name string
+ enabled bool
+ path string
+ marker string
+ role string
+ wantHeader bool
+ }{
+ {name: "disabled", enabled: false, path: "/api/v1/admin/users", role: "admin"},
+ {name: "admin API path", enabled: true, path: "/api/v1/admin/users", role: "admin", wantHeader: true},
+ {name: "shared API marked by admin UI", enabled: true, path: "/api/v1/groups/available", marker: "1", role: "admin", wantHeader: true},
+ {name: "non admin role", enabled: true, path: "/api/v1/groups/available", marker: "1", role: "user"},
+ {name: "unauthenticated public request", enabled: true, path: "/api/v1/settings/public", marker: "1"},
+ {name: "unmarked shared API", enabled: true, path: "/api/v1/groups/available", role: "admin"},
+ {name: "invalid marker", enabled: true, path: "/api/v1/groups/available", marker: "true", role: "admin"},
+ {name: "admin prefix boundary", enabled: true, path: "/api/v1/administrator", role: "admin"},
+ }
+
+ for _, tt := range tests {
+ t.Run(tt.name, func(t *testing.T) {
+ recorder := runServerTimingRequest(t, tt.enabled, tt.path, tt.marker, tt.role, func(c *gin.Context) {
+ c.JSON(http.StatusOK, gin.H{"ok": true})
+ })
+ header := recorder.Header().Get(servertiming.HeaderName)
+ if tt.wantHeader && header == "" {
+ t.Fatalf("%s header missing", servertiming.HeaderName)
+ }
+ if !tt.wantHeader && header != "" {
+ t.Fatalf("unexpected %s header: %q", servertiming.HeaderName, header)
+ }
+ if header != "" && (!strings.Contains(header, "total;dur=") || !strings.Contains(header, `cache;desc="bypass"`)) {
+ t.Fatalf("incomplete timing header: %q", header)
+ }
+ })
+ }
+}
+
+func TestServerTimingCollectorIsRequestScoped(t *testing.T) {
+ active := false
+ recorder := runServerTimingRequest(t, true, "/api/v1/keys", "1", "admin", func(c *gin.Context) {
+ active = servertiming.Active(c.Request.Context())
+ c.Status(http.StatusNoContent)
+ })
+ if !active {
+ t.Fatal("collector was not attached to marked request context")
+ }
+ if recorder.Header().Get(servertiming.HeaderName) == "" {
+ t.Fatal("timing header missing from status-only response")
+ }
+}
+
+func TestServerTimingFinalizesBeforeEarlyCommit(t *testing.T) {
+ recorder := runServerTimingRequest(t, true, "/api/v1/admin/stream", "", "admin", func(c *gin.Context) {
+ c.Status(http.StatusAccepted)
+ c.Writer.WriteHeaderNow()
+ })
+ if got := recorder.Header().Get(servertiming.HeaderName); got == "" {
+ t.Fatal("timing header was not written before response commit")
+ }
+}
+
+func TestServerTimingFinalizesOnFlush(t *testing.T) {
+ recorder := runServerTimingRequest(t, true, "/api/v1/admin/export", "", "admin", func(c *gin.Context) {
+ c.Writer.Flush()
+ })
+ if got := recorder.Header().Get(servertiming.HeaderName); got == "" {
+ t.Fatal("timing header was not written before stream flush")
+ }
+}
+
+func TestServerTimingStatusResponses(t *testing.T) {
+ tests := []struct {
+ name string
+ status int
+ }{
+ {name: "not modified", status: http.StatusNotModified},
+ {name: "internal error", status: http.StatusInternalServerError},
+ }
+ for _, tt := range tests {
+ t.Run(tt.name, func(t *testing.T) {
+ recorder := runServerTimingRequest(t, true, "/api/v1/admin/test", "", "admin", func(c *gin.Context) {
+ c.Status(tt.status)
+ })
+ if recorder.Code != tt.status {
+ t.Fatalf("status = %d, want %d", recorder.Code, tt.status)
+ }
+ if got := recorder.Header().Get(servertiming.HeaderName); got == "" {
+ t.Fatalf("timing header missing from status %d response", tt.status)
+ }
+ })
+ }
+}
+
+func TestServerTimingResponseWriterUnwraps(t *testing.T) {
+ recorder := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(recorder)
+ baseWriter := c.Writer
+ writer := &serverTimingResponseWriter{ResponseWriter: baseWriter}
+ if got := writer.Unwrap(); got != baseWriter {
+ t.Fatalf("Unwrap() = %T, want original Gin writer", got)
+ }
+}
+
+func TestServerTimingCacheOutcome(t *testing.T) {
+ tests := []struct {
+ name string
+ headerName string
+ value string
+ want string
+ }{
+ {name: "snapshot hit", headerName: snapshotCacheHeader, value: "hit", want: "hit"},
+ {name: "usage miss", headerName: usageCacheHeader, value: "MISS", want: "miss"},
+ {name: "invalid", headerName: snapshotCacheHeader, value: "stale", want: "bypass"},
+ }
+ for _, tt := range tests {
+ t.Run(tt.name, func(t *testing.T) {
+ recorder := runServerTimingRequest(t, true, "/api/v1/admin/dashboard", "", "admin", func(c *gin.Context) {
+ c.Header(tt.headerName, tt.value)
+ c.JSON(http.StatusOK, gin.H{"ok": true})
+ })
+ want := `cache;desc="` + tt.want + `"`
+ if got := recorder.Header().Get(servertiming.HeaderName); !strings.Contains(got, want) {
+ t.Fatalf("timing header %q does not contain %q", got, want)
+ }
+ })
+ }
+}
+
+func TestServerTimingResponseHeaderForWebSocket(t *testing.T) {
+ recorder := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(recorder)
+ c.Request = httptest.NewRequest(http.MethodGet, "/api/v1/admin/ops/ws/qps", nil)
+ collector := servertiming.New(time.Now())
+ c.Request = c.Request.WithContext(servertiming.WithCollector(c.Request.Context(), collector))
+ c.Set(string(ContextKeyUserRole), "admin")
+
+ header := ServerTimingResponseHeader(c)
+ if header.Get(servertiming.HeaderName) == "" {
+ t.Fatal("WebSocket response header missing timing value")
+ }
+
+ c.Set(string(ContextKeyUserRole), "user")
+ if got := ServerTimingResponseHeader(c); got != nil {
+ t.Fatalf("non-admin WebSocket received timing header: %#v", got)
+ }
+}
diff --git a/backend/internal/server/router.go b/backend/internal/server/router.go
index 3d86373779..5fc70149fe 100644
--- a/backend/internal/server/router.go
+++ b/backend/internal/server/router.go
@@ -60,6 +60,7 @@ func SetupRouter(
}
return nil
}))
+ r.Use(middleware2.ServerTiming(cfg.Server.EnableServerTiming))
// Serve embedded frontend with settings injection if available
if web.HasEmbeddedFrontend() {
diff --git a/backend/internal/service/account.go b/backend/internal/service/account.go
index 1ab2e21fdf..3e67fa5982 100644
--- a/backend/internal/service/account.go
+++ b/backend/internal/service/account.go
@@ -83,6 +83,8 @@ type Account struct {
type OpenAIEndpointCapability string
+const openAILongContextBillingEnabledKey = "openai_long_context_billing_enabled"
+
const (
OpenAIEndpointCapabilityChatCompletions OpenAIEndpointCapability = "chat_completions"
OpenAIEndpointCapabilityEmbeddings OpenAIEndpointCapability = "embeddings"
@@ -1192,6 +1194,14 @@ func (a *Account) IsOpenAI() bool {
return a.Platform == PlatformOpenAI
}
+func (a *Account) IsOpenAILongContextBillingEnabled() bool {
+ if a == nil || !a.IsOpenAI() || a.Extra == nil {
+ return false
+ }
+ enabled, ok := a.Extra[openAILongContextBillingEnabledKey].(bool)
+ return ok && enabled
+}
+
func (a *Account) IsAnthropic() bool {
return a.Platform == PlatformAnthropic
}
diff --git a/backend/internal/service/account_long_context_billing_test.go b/backend/internal/service/account_long_context_billing_test.go
new file mode 100644
index 0000000000..709559d932
--- /dev/null
+++ b/backend/internal/service/account_long_context_billing_test.go
@@ -0,0 +1,290 @@
+//go:build unit
+
+package service
+
+import (
+ "context"
+ "net/http"
+ "testing"
+
+ infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors"
+ "github.com/stretchr/testify/require"
+)
+
+func TestAccountIsOpenAILongContextBillingEnabled(t *testing.T) {
+ tests := []struct {
+ name string
+ account *Account
+ want bool
+ }{
+ {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 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},
+ }
+
+ for _, tt := range tests {
+ t.Run(tt.name, func(t *testing.T) {
+ require.Equal(t, tt.want, tt.account.IsOpenAILongContextBillingEnabled())
+ })
+ }
+}
+
+func TestNormalizeOpenAILongContextBillingExtra(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, false, extra["openai_long_context_billing_enabled"])
+ })
+
+ t.Run("OpenAI explicit false is preserved", func(t *testing.T) {
+ extra, err := normalizeOpenAILongContextBillingExtra(PlatformOpenAI, map[string]any{"openai_long_context_billing_enabled": false})
+
+ require.NoError(t, err)
+ require.Equal(t, false, extra["openai_long_context_billing_enabled"])
+ })
+
+ t.Run("OpenAI malformed value is rejected", func(t *testing.T) {
+ _, err := normalizeOpenAILongContextBillingExtra(PlatformOpenAI, map[string]any{"openai_long_context_billing_enabled": "false"})
+
+ require.Error(t, err)
+ require.Equal(t, http.StatusBadRequest, infraerrors.Code(err))
+ })
+
+ t.Run("non OpenAI extra is unchanged", func(t *testing.T) {
+ extra, err := normalizeOpenAILongContextBillingExtra(PlatformGrok, nil)
+
+ require.NoError(t, err)
+ require.Nil(t, extra)
+ })
+
+ t.Run("non OpenAI malformed value is ignored", func(t *testing.T) {
+ extra := map[string]any{openAILongContextBillingEnabledKey: "provider-owned"}
+ normalized, err := normalizeOpenAILongContextBillingExtra(PlatformAnthropic, extra)
+
+ require.NoError(t, err)
+ require.Equal(t, extra, normalized)
+ })
+}
+
+type longContextBillingRepoStub struct {
+ accountRepoStub
+ account *Account
+ accounts []*Account
+ createdAccount *Account
+ updateExtraCalls int
+ bulkUpdateCalls int
+}
+
+func (r *longContextBillingRepoStub) Create(_ context.Context, account *Account) error {
+ account.ID = 1
+ r.account = account
+ r.createdAccount = account
+ return nil
+}
+
+func (r *longContextBillingRepoStub) GetByID(_ context.Context, _ int64) (*Account, error) {
+ return r.account, nil
+}
+
+func (r *longContextBillingRepoStub) GetByIDs(_ context.Context, _ []int64) ([]*Account, error) {
+ if r.accounts != nil {
+ return r.accounts, nil
+ }
+ if r.account == nil {
+ return nil, nil
+ }
+ return []*Account{r.account}, nil
+}
+
+func (r *longContextBillingRepoStub) Update(_ context.Context, account *Account) error {
+ r.account = account
+ return nil
+}
+
+func (r *longContextBillingRepoStub) UpdateExtra(_ context.Context, _ int64, _ map[string]any) error {
+ r.updateExtraCalls++
+ return nil
+}
+
+func (r *longContextBillingRepoStub) BulkUpdate(_ context.Context, _ []int64, _ AccountBulkUpdate) (int64, error) {
+ r.bulkUpdateCalls++
+ return 1, nil
+}
+
+func TestAdminServiceCreateAccountDefaultsOpenAILongContextBillingDisabled(t *testing.T) {
+ repo := &longContextBillingRepoStub{}
+ svc := &adminServiceImpl{accountRepo: repo}
+
+ account, err := svc.CreateAccount(context.Background(), &CreateAccountInput{
+ Name: "openai-account",
+ Platform: PlatformOpenAI,
+ Type: AccountTypeAPIKey,
+ Credentials: map[string]any{"api_key": "test"},
+ SkipDefaultGroupBind: true,
+ })
+
+ require.NoError(t, err)
+ require.Same(t, account, repo.createdAccount)
+ require.Equal(t, false, account.Extra[openAILongContextBillingEnabledKey])
+}
+
+func TestAdminServiceCreateAccountRejectsMalformedOpenAILongContextBillingValue(t *testing.T) {
+ repo := &longContextBillingRepoStub{}
+ svc := &adminServiceImpl{accountRepo: repo}
+
+ account, err := svc.CreateAccount(context.Background(), &CreateAccountInput{
+ Platform: PlatformOpenAI,
+ Extra: map[string]any{openAILongContextBillingEnabledKey: "false"},
+ })
+
+ require.Nil(t, account)
+ require.Equal(t, http.StatusBadRequest, infraerrors.Code(err))
+ require.Nil(t, repo.createdAccount)
+}
+
+func TestAdminServiceUpdateAccountPreservesOpenAILongContextBillingOptOutWhenOmitted(t *testing.T) {
+ repo := &longContextBillingRepoStub{account: &Account{
+ ID: 1,
+ Platform: PlatformOpenAI,
+ Type: AccountTypeOAuth,
+ Extra: map[string]any{openAILongContextBillingEnabledKey: false},
+ }}
+ svc := &adminServiceImpl{accountRepo: repo}
+
+ account, err := svc.UpdateAccount(context.Background(), 1, &UpdateAccountInput{Extra: map[string]any{}})
+
+ require.NoError(t, err)
+ require.Equal(t, false, account.Extra[openAILongContextBillingEnabledKey])
+}
+
+func TestAdminServiceUpdateAccountAllowsExplicitCodexImportOptIn(t *testing.T) {
+ repo := &longContextBillingRepoStub{account: &Account{
+ ID: 1,
+ Platform: PlatformOpenAI,
+ Type: AccountTypeOAuth,
+ Credentials: map[string]any{"access_token": "old-token"},
+ Extra: map[string]any{
+ openAILongContextBillingEnabledKey: false,
+ "import_source": "codex_session",
+ },
+ }}
+ svc := &adminServiceImpl{accountRepo: repo}
+
+ account, err := svc.UpdateAccount(context.Background(), 1, &UpdateAccountInput{
+ Credentials: map[string]any{"access_token": "new-token"},
+ Extra: map[string]any{
+ openAILongContextBillingEnabledKey: true,
+ "import_source": "codex_session",
+ },
+ })
+
+ require.NoError(t, err)
+ require.Equal(t, true, account.Extra[openAILongContextBillingEnabledKey])
+}
+
+func TestAdminServiceUpdateAccountAllowsExplicitOptInOutsideCodexImport(t *testing.T) {
+ repo := &longContextBillingRepoStub{account: &Account{
+ ID: 1,
+ Platform: PlatformOpenAI,
+ Type: AccountTypeOAuth,
+ Extra: map[string]any{
+ openAILongContextBillingEnabledKey: false,
+ "import_source": "codex_session",
+ },
+ }}
+ svc := &adminServiceImpl{accountRepo: repo}
+
+ account, err := svc.UpdateAccount(context.Background(), 1, &UpdateAccountInput{Extra: map[string]any{
+ openAILongContextBillingEnabledKey: true,
+ "import_source": "codex_session",
+ }})
+
+ require.NoError(t, err)
+ require.Equal(t, true, account.Extra[openAILongContextBillingEnabledKey])
+}
+
+func TestAdminServiceUpdateAccountRejectsMalformedOpenAILongContextBillingValue(t *testing.T) {
+ repo := &longContextBillingRepoStub{account: &Account{ID: 1, Platform: PlatformOpenAI}}
+ svc := &adminServiceImpl{accountRepo: repo}
+
+ account, err := svc.UpdateAccount(context.Background(), 1, &UpdateAccountInput{Extra: map[string]any{
+ openAILongContextBillingEnabledKey: 1,
+ }})
+
+ require.Nil(t, account)
+ require.Equal(t, http.StatusBadRequest, infraerrors.Code(err))
+}
+
+func TestAdminServiceUpdateAccountExtraRejectsMalformedOpenAILongContextBillingValue(t *testing.T) {
+ repo := &longContextBillingRepoStub{account: &Account{ID: 1, Platform: PlatformOpenAI}}
+ svc := &adminServiceImpl{accountRepo: repo}
+
+ err := svc.UpdateAccountExtra(context.Background(), 1, map[string]any{
+ openAILongContextBillingEnabledKey: "true",
+ })
+
+ require.Equal(t, http.StatusBadRequest, infraerrors.Code(err))
+ require.Zero(t, repo.updateExtraCalls)
+}
+
+func TestAdminServiceUpdateAccountExtraAllowsProviderOwnedValueForNonOpenAIAccount(t *testing.T) {
+ repo := &longContextBillingRepoStub{account: &Account{ID: 1, Platform: PlatformAnthropic}}
+ svc := &adminServiceImpl{accountRepo: repo}
+
+ err := svc.UpdateAccountExtra(context.Background(), 1, map[string]any{
+ openAILongContextBillingEnabledKey: "provider-owned",
+ })
+
+ require.NoError(t, err)
+ require.Equal(t, 1, repo.updateExtraCalls)
+}
+
+func TestAdminServiceBulkUpdateAccountsRejectsMalformedOpenAILongContextBillingValue(t *testing.T) {
+ repo := &longContextBillingRepoStub{account: &Account{ID: 1, Platform: PlatformOpenAI}}
+ svc := &adminServiceImpl{accountRepo: repo}
+
+ result, err := svc.BulkUpdateAccounts(context.Background(), &BulkUpdateAccountsInput{
+ AccountIDs: []int64{1},
+ Extra: map[string]any{openAILongContextBillingEnabledKey: []bool{true}},
+ })
+
+ require.Nil(t, result)
+ require.Equal(t, http.StatusBadRequest, infraerrors.Code(err))
+ require.Zero(t, repo.bulkUpdateCalls)
+}
+
+func TestAdminServiceBulkUpdateAccountsAllowsProviderOwnedValueForNonOpenAIAccounts(t *testing.T) {
+ repo := &longContextBillingRepoStub{account: &Account{ID: 1, Platform: PlatformGrok}}
+ svc := &adminServiceImpl{accountRepo: repo}
+
+ result, err := svc.BulkUpdateAccounts(context.Background(), &BulkUpdateAccountsInput{
+ AccountIDs: []int64{1},
+ Extra: map[string]any{openAILongContextBillingEnabledKey: []string{"provider-owned"}},
+ })
+
+ require.NoError(t, err)
+ require.NotNil(t, result)
+ require.Equal(t, 1, repo.bulkUpdateCalls)
+}
+
+func TestAdminServiceBulkUpdateAccountsRejectsMalformedValueForMixedTargetsIncludingOpenAI(t *testing.T) {
+ repo := &longContextBillingRepoStub{accounts: []*Account{
+ {ID: 1, Platform: PlatformGrok},
+ {ID: 2, Platform: PlatformOpenAI},
+ }}
+ svc := &adminServiceImpl{accountRepo: repo}
+
+ result, err := svc.BulkUpdateAccounts(context.Background(), &BulkUpdateAccountsInput{
+ AccountIDs: []int64{1, 2},
+ Extra: map[string]any{openAILongContextBillingEnabledKey: "malformed"},
+ })
+
+ require.Nil(t, result)
+ require.Equal(t, http.StatusBadRequest, infraerrors.Code(err))
+ require.Zero(t, repo.bulkUpdateCalls)
+}
diff --git a/backend/internal/service/account_test_service.go b/backend/internal/service/account_test_service.go
index 31549a6b17..222b3f8a4d 100644
--- a/backend/internal/service/account_test_service.go
+++ b/backend/internal/service/account_test_service.go
@@ -669,7 +669,7 @@ func (s *AccountTestService) testGrokAccountConnection(c *gin.Context, account *
testModelID := strings.TrimSpace(modelID)
if testModelID == "" {
- testModelID = "grok-4.3"
+ testModelID = grokDefaultResponsesModel
}
if mapped := strings.TrimSpace(account.GetMappedModel(testModelID)); mapped != "" {
testModelID = mapped
diff --git a/backend/internal/service/account_test_service_grok_test.go b/backend/internal/service/account_test_service_grok_test.go
index 497224b713..4b0890ff44 100644
--- a/backend/internal/service/account_test_service_grok_test.go
+++ b/backend/internal/service/account_test_service_grok_test.go
@@ -80,6 +80,47 @@ func TestAccountTestService_TestAccountConnection_GrokUsesXAIResponses(t *testin
require.Contains(t, rec.Body.String(), `"type":"test_complete"`)
}
+func TestAccountTestService_TestAccountConnection_GrokDefaultsEmptyModelTo45(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+
+ account := &Account{
+ ID: 16,
+ Name: "grok-oauth-default-model",
+ Platform: PlatformGrok,
+ Type: AccountTypeOAuth,
+ Status: StatusActive,
+ Schedulable: true,
+ Concurrency: 1,
+ Credentials: map[string]any{
+ "access_token": "grok-access-token",
+ "expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339),
+ },
+ }
+ repo := &mockAccountRepoForGemini{accountsByID: map[int64]*Account{account.ID: account}}
+ upstream := &httpUpstreamRecorder{resp: &http.Response{
+ StatusCode: http.StatusOK,
+ Header: http.Header{"Content-Type": []string{"text/event-stream"}},
+ Body: io.NopCloser(strings.NewReader(
+ "data: {\"type\":\"response.output_text.delta\",\"delta\":\"ok\"}\n\n" +
+ "data: {\"type\":\"response.completed\"}\n\n",
+ )),
+ }}
+ svc := &AccountTestService{
+ accountRepo: repo,
+ grokTokenProvider: NewGrokTokenProvider(repo, nil),
+ httpUpstream: upstream,
+ }
+ recorder := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(recorder)
+ c.Request = httptest.NewRequest(http.MethodPost, "/api/v1/admin/accounts/16/test", nil)
+
+ err := svc.TestAccountConnection(c, account.ID, "", "", AccountTestModeDefault)
+
+ require.NoError(t, err)
+ require.Equal(t, grokDefaultResponsesModel, gjson.GetBytes(upstream.lastBody, "model").String())
+ require.Contains(t, recorder.Body.String(), `"model":"grok-4.5"`)
+}
+
func TestAccountTestService_Grok429PersistsRateLimitReset(t *testing.T) {
gin.SetMode(gin.TestMode)
diff --git a/backend/internal/service/account_usage_service.go b/backend/internal/service/account_usage_service.go
index 281122d4f0..9966c66515 100644
--- a/backend/internal/service/account_usage_service.go
+++ b/backend/internal/service/account_usage_service.go
@@ -111,6 +111,7 @@ const (
apiQueryMaxJitter = 800 * time.Millisecond // 用量查询最大随机延迟
windowStatsCacheTTL = 1 * time.Minute
openAIProbeCacheTTL = 10 * time.Minute
+ grokProbeRetryTTL = 1 * time.Minute
openAICodexProbeVersion = "0.144.1"
)
@@ -122,6 +123,7 @@ type UsageCache struct {
apiFlight singleflight.Group // 防止同一账号的并发请求击穿缓存(Anthropic)
antigravityFlight singleflight.Group // 防止同一 Antigravity 账号的并发请求击穿缓存
openAIProbeCache sync.Map // accountID -> time.Time
+ grokProbeCache sync.Map // accountID -> last billing probe attempt
}
// NewUsageCache 创建 UsageCache 实例
@@ -196,15 +198,18 @@ type UsageInfo struct {
AntigravityQuota map[string]*AntigravityModelQuota `json:"antigravity_quota,omitempty"`
// Grok / xAI 被动额度快照
- GrokRequestQuota *xai.QuotaWindow `json:"grok_request_quota,omitempty"`
- GrokTokenQuota *xai.QuotaWindow `json:"grok_token_quota,omitempty"`
- GrokRetryAfterSeconds *int `json:"grok_retry_after_seconds,omitempty"`
- GrokEntitlementStatus string `json:"grok_entitlement_status,omitempty"`
- GrokQuotaSnapshotState string `json:"grok_quota_snapshot_state,omitempty"`
- GrokLastQuotaProbeAt string `json:"grok_last_quota_probe_at,omitempty"`
- GrokLastHeadersSeenAt string `json:"grok_last_headers_seen_at,omitempty"`
- GrokLastStatusCode int `json:"grok_last_status_code,omitempty"`
- GrokLocalUsage *WindowStats `json:"grok_local_usage,omitempty"`
+ GrokRequestQuota *xai.QuotaWindow `json:"grok_request_quota,omitempty"`
+ GrokTokenQuota *xai.QuotaWindow `json:"grok_token_quota,omitempty"`
+ GrokRetryAfterSeconds *int `json:"grok_retry_after_seconds,omitempty"`
+ GrokEntitlementStatus string `json:"grok_entitlement_status,omitempty"`
+ GrokQuotaSnapshotState string `json:"grok_quota_snapshot_state,omitempty"`
+ GrokLastQuotaProbeAt string `json:"grok_last_quota_probe_at,omitempty"`
+ GrokLastHeadersSeenAt string `json:"grok_last_headers_seen_at,omitempty"`
+ GrokLastStatusCode int `json:"grok_last_status_code,omitempty"`
+ GrokLocalUsage *WindowStats `json:"grok_local_usage,omitempty"`
+ GrokLocalUsage7d *WindowStats `json:"grok_local_usage_7d,omitempty"`
+ GrokLocalUsageMonthly *WindowStats `json:"grok_local_usage_monthly,omitempty"`
+ GrokBilling *xai.BillingSummary `json:"grok_billing,omitempty"`
// Antigravity 账号级信息
SubscriptionTier string `json:"subscription_tier,omitempty"` // 归一化订阅等级: FREE/PRO/ULTRA/UNKNOWN
@@ -287,6 +292,7 @@ type AccountUsageService struct {
geminiQuotaService *GeminiQuotaService
antigravityQuotaFetcher *AntigravityQuotaFetcher
grokQuotaFetcher *GrokQuotaFetcher
+ grokQuotaService *GrokQuotaService
openAIQuotaService *OpenAIQuotaService
cache *UsageCache
identityCache IdentityCache
@@ -301,6 +307,7 @@ func NewAccountUsageService(
geminiQuotaService *GeminiQuotaService,
antigravityQuotaFetcher *AntigravityQuotaFetcher,
grokQuotaFetcher *GrokQuotaFetcher,
+ grokQuotaService *GrokQuotaService,
openAIQuotaService *OpenAIQuotaService,
cache *UsageCache,
identityCache IdentityCache,
@@ -313,6 +320,7 @@ func NewAccountUsageService(
geminiQuotaService: geminiQuotaService,
antigravityQuotaFetcher: antigravityQuotaFetcher,
grokQuotaFetcher: grokQuotaFetcher,
+ grokQuotaService: grokQuotaService,
openAIQuotaService: openAIQuotaService,
cache: cache,
identityCache: identityCache,
@@ -358,8 +366,8 @@ func (s *AccountUsageService) GetUsage(ctx context.Context, accountID int64, for
}
if account.Platform == PlatformGrok {
- usage, err := s.getGrokUsage(ctx, account)
- if err == nil {
+ usage, err := s.getGrokUsage(ctx, account, forceProbe)
+ if err == nil && usage != nil && usage.Error == "" {
s.tryClearRecoverableAccountError(ctx, account)
}
return usage, err
@@ -930,11 +938,19 @@ func (s *AccountUsageService) getAntigravityUsage(ctx context.Context, account *
return usage, nil
}
-func (s *AccountUsageService) getGrokUsage(ctx context.Context, account *Account) (*UsageInfo, error) {
+func (s *AccountUsageService) getGrokUsage(ctx context.Context, account *Account, force bool) (*UsageInfo, error) {
if s.grokQuotaFetcher == nil {
now := time.Now()
return &UsageInfo{UpdatedAt: &now}, nil
}
+ if account != nil && account.IsGrokOAuth() && s.grokQuotaService != nil && (force || grokBillingSnapshotNeedsRefresh(account, time.Now())) && s.shouldProbeGrokBilling(account.ID, time.Now(), force) {
+ result, err := s.grokQuotaService.ProbeBilling(ctx, account.ID)
+ if err == nil && result != nil && result.Billing != nil {
+ mergeAccountExtra(account, map[string]any{grokBillingExtraKey: result.Billing})
+ } else if err != nil && force {
+ return nil, err
+ }
+ }
usage := s.grokQuotaFetcher.BuildUsageInfo(account)
if usage.GrokQuotaSnapshotState == "" {
if usage.ErrorCode == "quota_unknown" {
@@ -948,12 +964,90 @@ func (s *AccountUsageService) getGrokUsage(ctx context.Context, account *Account
if stats, err := s.usageLogRepo.GetAccountTodayStats(ctx, account.ID); err == nil && stats != nil {
usage.GrokLocalUsage = windowStatsFromAccountStats(stats)
}
+ usage.GrokLocalUsage7d, usage.GrokLocalUsageMonthly = grokLocalUsageForBilling(ctx, s.usageLogRepo, account.ID, usage.GrokBilling, time.Now().UTC())
}
enrichUsageWithAccountError(usage, account)
return usage, nil
}
+func grokLocalUsageForBilling(
+ ctx context.Context,
+ repo UsageLogRepository,
+ accountID int64,
+ billing *xai.BillingSummary,
+ now time.Time,
+) (*WindowStats, *WindowStats) {
+ var weekly *WindowStats
+ var monthly *WindowStats
+ if repo == nil || accountID <= 0 {
+ return weekly, monthly
+ }
+ if start, ok := currentGrokBillingWindow(billing, true, now); ok {
+ if stats, err := repo.GetAccountWindowStats(ctx, accountID, start); err == nil {
+ weekly = windowStatsFromAccountStats(stats)
+ } else {
+ slog.Warn("grok_window_usage_query_failed", "account_id", accountID, "window_start", start, "error", err)
+ }
+ }
+ if start, ok := currentGrokBillingWindow(billing, false, now); ok {
+ if stats, err := repo.GetAccountWindowStats(ctx, accountID, start); err == nil {
+ monthly = windowStatsFromAccountStats(stats)
+ } else {
+ slog.Warn("grok_monthly_usage_query_failed", "account_id", accountID, "window_start", start, "error", err)
+ }
+ }
+ return weekly, monthly
+}
+
+func currentGrokBillingWindow(billing *xai.BillingSummary, weekly bool, now time.Time) (time.Time, bool) {
+ if billing == nil {
+ return time.Time{}, false
+ }
+ startRaw, endRaw := billing.BillingPeriodStart, billing.BillingPeriodEnd
+ if weekly {
+ if billing.PeriodType != "weekly" {
+ return time.Time{}, false
+ }
+ startRaw, endRaw = billing.PeriodStart, billing.PeriodEnd
+ }
+ start, startErr := parseTime(strings.TrimSpace(startRaw))
+ end, endErr := parseTime(strings.TrimSpace(endRaw))
+ if startErr != nil || endErr != nil || now.Before(start) || !now.Before(end) {
+ return time.Time{}, false
+ }
+ return start, true
+}
+
+func grokBillingSnapshotNeedsRefresh(account *Account, now time.Time) bool {
+ if account == nil {
+ return false
+ }
+ billing, err := grokBillingSnapshotFromExtra(account.Extra)
+ if err != nil || billing == nil || billing.Partial || len(billing.FailedWindows) > 0 {
+ return true
+ }
+ stamp := strings.TrimSpace(billing.UpdatedAt)
+ if stamp == "" {
+ stamp = strings.TrimSpace(billing.FetchedAt)
+ }
+ updatedAt, err := parseTime(stamp)
+ return err != nil || now.Sub(updatedAt) >= openAIProbeCacheTTL
+}
+
+func (s *AccountUsageService) shouldProbeGrokBilling(accountID int64, now time.Time, force bool) bool {
+ if force || s == nil || s.cache == nil || accountID <= 0 {
+ return true
+ }
+ if cached, ok := s.cache.grokProbeCache.Load(accountID); ok {
+ if ts, ok := cached.(time.Time); ok && now.Sub(ts) < grokProbeRetryTTL {
+ return false
+ }
+ }
+ s.cache.grokProbeCache.Store(accountID, now)
+ return true
+}
+
// recalcAntigravityRemainingSeconds 重新计算 Antigravity UsageInfo 中各窗口的 RemainingSeconds
// 用于从缓存取出时更新倒计时,避免返回过时的剩余秒数
func recalcAntigravityRemainingSeconds(info *UsageInfo) {
diff --git a/backend/internal/service/admin_account.go b/backend/internal/service/admin_account.go
index 52e5ce719b..8cb6d8e63b 100644
--- a/backend/internal/service/admin_account.go
+++ b/backend/internal/service/admin_account.go
@@ -5,6 +5,7 @@ import (
"errors"
"fmt"
"log/slog"
+ "maps"
"net/http"
"strconv"
"strings"
@@ -68,7 +69,65 @@ func normalizeAccountConcurrency(platform, accountType string, concurrency int)
return concurrency
}
+// ValidateOpenAILongContextBillingExtra validates the OpenAI account billing flag when present.
+func ValidateOpenAILongContextBillingExtra(platform string, extra map[string]any) error {
+ if platform != PlatformOpenAI {
+ return nil
+ }
+ raw, exists := extra[openAILongContextBillingEnabledKey]
+ if !exists {
+ return nil
+ }
+ if _, ok := raw.(bool); !ok {
+ return infraerrors.BadRequest(
+ "OPENAI_LONG_CONTEXT_BILLING_INVALID",
+ "openai_long_context_billing_enabled must be a boolean",
+ )
+ }
+ return nil
+}
+
+func normalizeOpenAILongContextBillingExtra(platform string, extra map[string]any) (map[string]any, error) {
+ if platform != PlatformOpenAI {
+ return extra, nil
+ }
+ if err := ValidateOpenAILongContextBillingExtra(platform, extra); err != nil {
+ return nil, err
+ }
+
+ normalized := maps.Clone(extra)
+ if normalized == nil {
+ normalized = make(map[string]any, 1)
+ }
+ _, exists := normalized[openAILongContextBillingEnabledKey]
+ if !exists {
+ normalized[openAILongContextBillingEnabledKey] = false
+ }
+ return normalized, nil
+}
+
+func normalizeOpenAILongContextBillingUpdateExtra(account *Account, input *UpdateAccountInput) (map[string]any, error) {
+ normalized, err := normalizeOpenAILongContextBillingExtra(account.Platform, input.Extra)
+ if err != nil || account.Platform != PlatformOpenAI {
+ return normalized, err
+ }
+
+ _, provided := input.Extra[openAILongContextBillingEnabledKey]
+ current, hasCurrent := account.Extra[openAILongContextBillingEnabledKey].(bool)
+ if !provided {
+ if hasCurrent {
+ normalized[openAILongContextBillingEnabledKey] = current
+ }
+ }
+ return normalized, nil
+}
+
func (s *adminServiceImpl) CreateAccount(ctx context.Context, input *CreateAccountInput) (*Account, error) {
+ accountExtra, err := normalizeOpenAILongContextBillingExtra(input.Platform, input.Extra)
+ if err != nil {
+ return nil, err
+ }
+
// 绑定分组
groupIDs := input.GroupIDs
// 如果没有指定分组,自动绑定对应平台的默认分组
@@ -103,7 +162,7 @@ func (s *adminServiceImpl) CreateAccount(ctx context.Context, input *CreateAccou
Platform: input.Platform,
Type: input.Type,
Credentials: input.Credentials,
- Extra: input.Extra,
+ Extra: accountExtra,
ProxyID: input.ProxyID,
Concurrency: normalizeAccountConcurrency(input.Platform, input.Type, input.Concurrency),
Priority: input.Priority,
@@ -183,6 +242,13 @@ func (s *adminServiceImpl) UpdateAccount(ctx context.Context, id int64, input *U
if err != nil {
return nil, err
}
+ var normalizedExtra map[string]any
+ if input.Extra != nil {
+ normalizedExtra, err = normalizeOpenAILongContextBillingUpdateExtra(account, input)
+ if err != nil {
+ return nil, err
+ }
+ }
// 安全/身份不变量(影子账号):通用更新路径被 edit/re-auth/refresh/batch 共用,
// 必须在此守住,否则仅在创建时的保证可被这些路径绕过。
if account.IsCredentialShadow() {
@@ -238,10 +304,10 @@ func (s *adminServiceImpl) UpdateAccount(ctx context.Context, id int64, input *U
// 保留配额用量字段,防止编辑账号时意外重置
for _, key := range []string{"quota_used", "quota_daily_used", "quota_daily_start", "quota_weekly_used", "quota_weekly_start"} {
if v, ok := account.Extra[key]; ok {
- input.Extra[key] = v
+ normalizedExtra[key] = v
}
}
- account.Extra = input.Extra
+ account.Extra = normalizedExtra
if account.Platform == PlatformAntigravity && wasOveragesEnabled && !account.IsOveragesEnabled() {
delete(account.Extra, "antigravity_credits_overages") // 清理旧版 overages 运行态
// 清除 AICredits 限流 key
@@ -353,6 +419,15 @@ func (s *adminServiceImpl) UpdateAccount(ctx context.Context, id int64, input *U
// UpdateAccountExtra 仅对 Extra JSONB 做 key 级合并,避免覆盖其它运行态键
// (如 model_rate_limits / passive_usage_* 等)。
func (s *adminServiceImpl) UpdateAccountExtra(ctx context.Context, id int64, updates map[string]any) error {
+ if _, exists := updates[openAILongContextBillingEnabledKey]; exists {
+ account, err := s.accountRepo.GetByID(ctx, id)
+ if err != nil {
+ return err
+ }
+ if err := ValidateOpenAILongContextBillingExtra(account.Platform, updates); err != nil {
+ return err
+ }
+ }
if len(updates) == 0 {
return nil
}
@@ -386,16 +461,28 @@ func (s *adminServiceImpl) BulkUpdateAccounts(ctx context.Context, input *BulkUp
}
needMixedChannelCheck := input.GroupIDs != nil && !input.SkipMixedChannelCheck
+ _, hasLongContextBillingUpdate := input.Extra[openAILongContextBillingEnabledKey]
// 预取所有目标账号,供凭据守卫/代理守卫/混合渠道检查共用,避免多次 DB 查询。
var cachedTargets []*Account
- if len(input.Credentials) > 0 || input.ProxyID != nil || needMixedChannelCheck {
+ if len(input.Credentials) > 0 || input.ProxyID != nil || needMixedChannelCheck || hasLongContextBillingUpdate {
loaded, err := s.accountRepo.GetByIDs(ctx, input.AccountIDs)
if err != nil {
return nil, err
}
cachedTargets = loaded
}
+ if hasLongContextBillingUpdate {
+ for _, account := range cachedTargets {
+ if account == nil || account.Platform != PlatformOpenAI {
+ continue
+ }
+ if err := ValidateOpenAILongContextBillingExtra(account.Platform, input.Extra); err != nil {
+ return nil, err
+ }
+ break
+ }
+ }
// 影子账号绝不持有凭据:批量更新携带凭据时,目标中不得含影子(外审 G5,与单账号
// UpdateAccount 守卫对齐)。覆盖显式 IDs 与 filter 解析出的 IDs(此处 AccountIDs 已解析完成)。
@@ -745,6 +832,9 @@ func (s *adminServiceImpl) CreateShadow(ctx context.Context, parentID int64, opt
Priority: priority,
Concurrency: concurrency,
Schedulable: true,
+ Extra: map[string]any{
+ openAILongContextBillingEnabledKey: parent.IsOpenAILongContextBillingEnabled(),
+ },
}
// 5. 持久化(Create 填充 shadow.ID)。并发竞态:预查(步骤2)放行后另一请求抢先建成,本次会撞
diff --git a/backend/internal/service/admin_service_spark_shadow_test.go b/backend/internal/service/admin_service_spark_shadow_test.go
index 6b4017207a..0eda0d93c7 100644
--- a/backend/internal/service/admin_service_spark_shadow_test.go
+++ b/backend/internal/service/admin_service_spark_shadow_test.go
@@ -157,6 +157,39 @@ func TestCreateShadow(t *testing.T) {
require.Error(t, err)
}
+func TestCreateShadowInheritsParentEffectiveOpenAILongContextBillingValue(t *testing.T) {
+ tests := []struct {
+ name string
+ parentExtra map[string]any
+ want bool
+ }{
+ {name: "missing parent value defaults disabled", want: false},
+ {name: "explicit parent opt-out is inherited", parentExtra: map[string]any{openAILongContextBillingEnabledKey: false}, want: false},
+ }
+
+ for _, tt := range tests {
+ t.Run(tt.name, func(t *testing.T) {
+ repo := newSparkShadowRepoStub()
+ svc := &adminServiceImpl{accountRepo: repo}
+ parent := &Account{
+ Name: "parent",
+ Platform: PlatformOpenAI,
+ Type: AccountTypeOAuth,
+ Status: StatusActive,
+ Credentials: map[string]any{"access_token": "token"},
+ Extra: tt.parentExtra,
+ }
+ require.NoError(t, repo.Create(context.Background(), parent))
+
+ shadow, err := svc.CreateShadow(context.Background(), parent.ID, ShadowOptions{Name: "shadow"})
+
+ require.NoError(t, err)
+ require.Equal(t, tt.want, shadow.Extra[openAILongContextBillingEnabledKey])
+ require.Equal(t, tt.want, shadow.IsOpenAILongContextBillingEnabled())
+ })
+ }
+}
+
// TestCreateShadow_BindGroups は BindGroups の後置呼び出しを検証する。
// 影子账号が指定グループに属し、ListSchedulableByGroupID で取得可能であること。
func TestCreateShadow_BindGroups(t *testing.T) {
diff --git a/backend/internal/service/billing_service.go b/backend/internal/service/billing_service.go
index 43d7f22d78..7fa69d41ea 100644
--- a/backend/internal/service/billing_service.go
+++ b/backend/internal/service/billing_service.go
@@ -153,14 +153,15 @@ type UsageTokens struct {
// CostBreakdown 费用明细
type CostBreakdown struct {
- InputCost float64
- OutputCost float64
- ImageOutputCost float64
- CacheCreationCost float64
- CacheReadCost float64
- TotalCost float64
- ActualCost float64 // 应用倍率后的实际费用
- BillingMode string // 计费模式("token"/"per_request"/"image"),由 CalculateCostUnified 填充
+ InputCost float64
+ OutputCost float64
+ ImageOutputCost float64
+ CacheCreationCost float64
+ CacheReadCost float64
+ TotalCost float64
+ ActualCost float64 // 应用倍率后的实际费用
+ BillingMode string // 计费模式("token"/"per_request"/"image"),由 CalculateCostUnified 填充
+ LongContextBillingApplied bool
}
// ErrModelPricingUnavailable indicates that none of the configured pricing
@@ -865,16 +866,17 @@ func (s *BillingService) GetModelPricingWithChannel(model string, channelPricing
// CostInput 统一计费输入
type CostInput struct {
- Ctx context.Context
- Model string
- GroupID *int64 // 用于渠道定价查找
- Tokens UsageTokens
- RequestCount int // 按次计费时使用
- SizeTier string // 按次/图片模式的层级标签("1K","2K","4K","HD" 等)
- RateMultiplier float64
- ServiceTier string // "priority","flex","" 等
- Resolver *ModelPricingResolver // 定价解析器
- Resolved *ResolvedPricing // 可选:预解析的定价结果(避免重复 Resolve 调用)
+ Ctx context.Context
+ Model string
+ GroupID *int64 // 用于渠道定价查找
+ Tokens UsageTokens
+ RequestCount int // 按次计费时使用
+ SizeTier string // 按次/图片模式的层级标签("1K","2K","4K","HD" 等)
+ RateMultiplier float64
+ ServiceTier string // "priority","flex","" 等
+ Resolver *ModelPricingResolver // 定价解析器
+ Resolved *ResolvedPricing // 可选:预解析的定价结果(避免重复 Resolve 调用)
+ LongContextBillingEnabled *bool
}
// CalculateCostUnified 统一计费入口,支持三种计费模式。
@@ -882,7 +884,18 @@ type CostInput struct {
func (s *BillingService) CalculateCostUnified(input CostInput) (*CostBreakdown, error) {
if input.Resolver == nil {
// 无 Resolver,回退到旧路径
- return s.calculateCostInternal(input.Model, input.Tokens, input.RateMultiplier, input.ServiceTier, nil)
+ applyLongContextBilling := true
+ if input.LongContextBillingEnabled != nil {
+ applyLongContextBilling = *input.LongContextBillingEnabled
+ }
+ return s.calculateCostInternalWithPolicy(
+ input.Model,
+ input.Tokens,
+ input.RateMultiplier,
+ input.ServiceTier,
+ nil,
+ applyLongContextBilling,
+ )
}
// 优先使用预解析结果,避免重复 Resolve 调用
@@ -929,6 +942,9 @@ func (s *BillingService) calculateTokenCost(resolved *ResolvedPricing, input Cos
// 长上下文定价仅在无区间定价时应用(区间定价已包含上下文分层)
applyLongCtx := len(resolved.Intervals) == 0
+ if input.LongContextBillingEnabled != nil {
+ applyLongCtx = applyLongCtx && *input.LongContextBillingEnabled
+ }
return s.computeTokenBreakdown(pricing, input.Tokens, input.RateMultiplier, input.ServiceTier, applyLongCtx), nil
}
@@ -969,7 +985,10 @@ func (s *BillingService) computeTokenBreakdown(
tierMultiplier = serviceTierCostMultiplier(serviceTier)
}
- if applyLongCtx && s.shouldApplySessionLongContextPricing(tokens, pricing) {
+ longContextPricingEligible := applyLongCtx && s.shouldApplySessionLongContextPricing(tokens, pricing)
+ var baselineCost *CostBreakdown
+ if longContextPricingEligible {
+ baselineCost = s.computeTokenBreakdown(pricing, tokens, rateMultiplier, serviceTier, false)
inputPrice *= pricing.LongContextInputMultiplier
outputPrice *= pricing.LongContextOutputMultiplier
// 缓存读取本质上是输入侧的复用,应与 input 一同应用长上下文倍率;
@@ -1033,6 +1052,7 @@ func (s *BillingService) computeTokenBreakdown(
bd.TotalCost = bd.InputCost + bd.OutputCost + bd.ImageOutputCost +
bd.CacheCreationCost + bd.CacheReadCost
bd.ActualCost = bd.TotalCost * rateMultiplier
+ bd.LongContextBillingApplied = baselineCost != nil && bd.ActualCost > baselineCost.ActualCost
return bd
}
@@ -1092,7 +1112,28 @@ func (s *BillingService) CalculateCostWithServiceTier(model string, tokens Usage
return s.calculateCostInternal(model, tokens, rateMultiplier, serviceTier, nil)
}
+func (s *BillingService) calculateCostWithServiceTierPolicy(
+ model string,
+ tokens UsageTokens,
+ rateMultiplier float64,
+ serviceTier string,
+ longContextBillingEnabled bool,
+) (*CostBreakdown, error) {
+ return s.calculateCostInternalWithPolicy(model, tokens, rateMultiplier, serviceTier, nil, longContextBillingEnabled)
+}
+
func (s *BillingService) calculateCostInternal(model string, tokens UsageTokens, rateMultiplier float64, serviceTier string, channelPricing *ChannelModelPricing) (*CostBreakdown, error) {
+ return s.calculateCostInternalWithPolicy(model, tokens, rateMultiplier, serviceTier, channelPricing, true)
+}
+
+func (s *BillingService) calculateCostInternalWithPolicy(
+ model string,
+ tokens UsageTokens,
+ rateMultiplier float64,
+ serviceTier string,
+ channelPricing *ChannelModelPricing,
+ longContextBillingEnabled bool,
+) (*CostBreakdown, error) {
var pricing *ModelPricing
var err error
if channelPricing != nil {
@@ -1104,8 +1145,7 @@ func (s *BillingService) calculateCostInternal(model string, tokens UsageTokens,
return nil, err
}
- // 旧路径始终检查长上下文定价(无区间定价概念)
- return s.computeTokenBreakdown(pricing, tokens, rateMultiplier, serviceTier, true), nil
+ return s.computeTokenBreakdown(pricing, tokens, rateMultiplier, serviceTier, longContextBillingEnabled), nil
}
func (s *BillingService) applyModelSpecificPricingPolicy(model string, pricing *ModelPricing) *ModelPricing {
@@ -1236,13 +1276,14 @@ func (s *BillingService) CalculateCostWithLongContext(model string, tokens Usage
// 合并成本
return &CostBreakdown{
- InputCost: inRangeCost.InputCost + outRangeCost.InputCost,
- OutputCost: inRangeCost.OutputCost,
- ImageOutputCost: inRangeCost.ImageOutputCost,
- CacheCreationCost: inRangeCost.CacheCreationCost,
- CacheReadCost: inRangeCost.CacheReadCost + outRangeCost.CacheReadCost,
- TotalCost: inRangeCost.TotalCost + outRangeCost.TotalCost,
- ActualCost: inRangeCost.ActualCost + outRangeCost.ActualCost,
+ InputCost: inRangeCost.InputCost + outRangeCost.InputCost,
+ OutputCost: inRangeCost.OutputCost,
+ ImageOutputCost: inRangeCost.ImageOutputCost,
+ CacheCreationCost: inRangeCost.CacheCreationCost,
+ CacheReadCost: inRangeCost.CacheReadCost + outRangeCost.CacheReadCost,
+ TotalCost: inRangeCost.TotalCost + outRangeCost.TotalCost,
+ ActualCost: inRangeCost.ActualCost + outRangeCost.ActualCost,
+ LongContextBillingApplied: outRangeCost.ActualCost > 0,
}, nil
}
diff --git a/backend/internal/service/billing_service_test.go b/backend/internal/service/billing_service_test.go
index 53014412fd..885da194e3 100644
--- a/backend/internal/service/billing_service_test.go
+++ b/backend/internal/service/billing_service_test.go
@@ -261,6 +261,23 @@ func TestCalculateCost_OpenAIGPT54LongContextAppliesWholeSessionMultipliers(t *t
require.InDelta(t, expectedOutput, cost.OutputCost, 1e-10)
require.InDelta(t, expectedInput+expectedOutput, cost.TotalCost, 1e-10)
require.InDelta(t, expectedInput+expectedOutput, cost.ActualCost, 1e-10)
+ require.True(t, cost.LongContextBillingApplied)
+}
+
+func TestCalculateCost_OpenAIGPT54LongContextMarkerRequiresActualCostIncrease(t *testing.T) {
+ svc := newTestBillingService()
+
+ cost, err := svc.calculateCostWithServiceTierPolicy(
+ "gpt-5.4-2026-03-05",
+ UsageTokens{InputTokens: 300000},
+ 0,
+ "",
+ true,
+ )
+
+ require.NoError(t, err)
+ require.Zero(t, cost.ActualCost)
+ require.False(t, cost.LongContextBillingApplied)
}
func TestCalculateCost_OpenAIGPT55ProUsesGPT55PricingPolicy(t *testing.T) {
@@ -831,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()
diff --git a/backend/internal/service/channel_monitor_checker.go b/backend/internal/service/channel_monitor_checker.go
index 7fb829a3cb..889b2bbed7 100644
--- a/backend/internal/service/channel_monitor_checker.go
+++ b/backend/internal/service/channel_monitor_checker.go
@@ -13,6 +13,7 @@ import (
"strings"
"time"
+ "github.com/Wei-Shaw/sub2api/internal/pkg/servertiming"
"github.com/tidwall/gjson"
)
@@ -34,7 +35,7 @@ func newSSRFSafeHTTPClient(timeout time.Duration) *http.Client {
TLSHandshakeTimeout: monitorTLSHandshakeTimeout,
ResponseHeaderTimeout: monitorResponseHeaderTimeout,
}
- return &http.Client{Timeout: timeout, Transport: tr}
+ return &http.Client{Timeout: timeout, Transport: servertiming.WrapRoundTripper(tr)}
}
// CheckOptions 承载一次检测的自定义入参。
diff --git a/backend/internal/service/content_moderation.go b/backend/internal/service/content_moderation.go
index 6d3b91d205..f633c8ad17 100644
--- a/backend/internal/service/content_moderation.go
+++ b/backend/internal/service/content_moderation.go
@@ -22,6 +22,7 @@ import (
infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors"
"github.com/Wei-Shaw/sub2api/internal/pkg/pagination"
+ "github.com/Wei-Shaw/sub2api/internal/pkg/servertiming"
)
const (
@@ -561,7 +562,7 @@ func NewContentModerationService(
userRepo: userRepo,
authCacheInvalidator: authCacheInvalidator,
emailService: emailService,
- httpClient: &http.Client{},
+ httpClient: servertiming.InstrumentClient(nil),
workerCount: maxContentModerationWorkerCount,
asyncQueue: make(chan contentModerationTask, maxContentModerationQueueSize),
keyHealth: make(map[string]*contentModerationKeyHealth),
diff --git a/backend/internal/service/crs_sync_long_context_billing_test.go b/backend/internal/service/crs_sync_long_context_billing_test.go
new file mode 100644
index 0000000000..6439f08190
--- /dev/null
+++ b/backend/internal/service/crs_sync_long_context_billing_test.go
@@ -0,0 +1,169 @@
+//go:build unit
+
+package service
+
+import (
+ "context"
+ "encoding/json"
+ "net/http"
+ "net/http/httptest"
+ "testing"
+
+ "github.com/Wei-Shaw/sub2api/internal/config"
+ "github.com/stretchr/testify/require"
+)
+
+type crsLongContextAccountRepo struct {
+ AccountRepository
+ accounts map[string]*Account
+ nextID int64
+}
+
+type crsOpenAILongContextSource struct {
+ collection string
+ credentials map[string]any
+ extra map[string]any
+}
+
+func newCRSLongContextAccountRepo(existing ...*Account) *crsLongContextAccountRepo {
+ repo := &crsLongContextAccountRepo{accounts: make(map[string]*Account)}
+ for _, account := range existing {
+ if account == nil {
+ continue
+ }
+ crsID, _ := account.Extra["crs_account_id"].(string)
+ repo.accounts[crsID] = account
+ if account.ID > repo.nextID {
+ repo.nextID = account.ID
+ }
+ }
+ return repo
+}
+
+func (r *crsLongContextAccountRepo) Create(_ context.Context, account *Account) error {
+ r.nextID++
+ account.ID = r.nextID
+ crsID, _ := account.Extra["crs_account_id"].(string)
+ r.accounts[crsID] = account
+ return nil
+}
+
+func (r *crsLongContextAccountRepo) Update(_ context.Context, account *Account) error {
+ crsID, _ := account.Extra["crs_account_id"].(string)
+ r.accounts[crsID] = account
+ return nil
+}
+
+func (r *crsLongContextAccountRepo) GetByCRSAccountID(_ context.Context, crsID string) (*Account, error) {
+ return r.accounts[crsID], nil
+}
+
+func (r *crsLongContextAccountRepo) ListShadowsByParent(_ context.Context, _ int64) ([]*Account, error) {
+ return nil, nil
+}
+
+func TestCRSSyncOpenAILongContextBilling(t *testing.T) {
+ tests := []struct {
+ name string
+ collection string
+ credentials map[string]any
+ sourceExtra map[string]any
+ existingExtra map[string]any
+ wantAction string
+ wantEnabled bool
+ }{
+ {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 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 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 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"},
+ {name: "API key update rejects malformed source value", collection: "openaiResponsesAccounts", credentials: map[string]any{"api_key": "sk-test"}, sourceExtra: map[string]any{openAILongContextBillingEnabledKey: "false"}, existingExtra: map[string]any{openAILongContextBillingEnabledKey: true}, wantAction: "failed"},
+ }
+
+ for _, tt := range tests {
+ t.Run(tt.name, func(t *testing.T) {
+ const crsID = "crs-openai-1"
+ var existing *Account
+ if tt.existingExtra != nil {
+ existingExtra := mergeMap(tt.existingExtra, map[string]any{"crs_account_id": crsID})
+ accountType := AccountTypeOAuth
+ if tt.collection == "openaiResponsesAccounts" {
+ accountType = AccountTypeAPIKey
+ }
+ existing = &Account{ID: 41, Platform: PlatformOpenAI, Type: accountType, Extra: existingExtra}
+ }
+ repo := newCRSLongContextAccountRepo(existing)
+ result := runCRSOpenAILongContextSync(t, repo, crsOpenAILongContextSource{
+ collection: tt.collection,
+ credentials: tt.credentials,
+ extra: tt.sourceExtra,
+ })
+
+ require.Len(t, result.Items, 1)
+ require.Equal(t, tt.wantAction, result.Items[0].Action)
+ if tt.wantAction == "failed" {
+ require.Contains(t, result.Items[0].Error, "openai_long_context_billing_enabled must be a boolean")
+ return
+ }
+ stored, ok := repo.accounts[crsID].Extra[openAILongContextBillingEnabledKey]
+ require.True(t, ok)
+ require.Equal(t, tt.wantEnabled, stored)
+ })
+ }
+}
+
+func runCRSOpenAILongContextSync(t *testing.T, repo AccountRepository, source crsOpenAILongContextSource) *SyncFromCRSResult {
+ t.Helper()
+ account := map[string]any{
+ "kind": "openai",
+ "id": "crs-openai-1",
+ "name": "OpenAI CRS",
+ "isActive": true,
+ "schedulable": true,
+ "credentials": source.credentials,
+ }
+ if source.extra != nil {
+ account["extra"] = source.extra
+ }
+
+ server := httptest.NewServer(http.HandlerFunc(func(response http.ResponseWriter, request *http.Request) {
+ response.Header().Set("Content-Type", "application/json")
+ if request.URL.Path == "/web/auth/login" {
+ _, _ = response.Write([]byte(`{"success":true,"token":"admin-token"}`))
+ return
+ }
+ require.Equal(t, "/admin/sync/export-accounts", request.URL.Path)
+ require.NoError(t, json.NewEncoder(response).Encode(map[string]any{
+ "success": true,
+ "data": map[string]any{source.collection: []any{account}},
+ }))
+ }))
+ t.Cleanup(server.Close)
+
+ cfg := &config.Config{}
+ cfg.Security.URLAllowlist.AllowInsecureHTTP = true
+ service := NewCRSSyncService(repo, nil, nil, nil, nil, cfg)
+ result, err := service.SyncFromCRS(context.Background(), SyncFromCRSInput{
+ BaseURL: server.URL,
+ Username: "admin",
+ Password: "password",
+ })
+ require.NoError(t, err)
+ return result
+}
diff --git a/backend/internal/service/crs_sync_service.go b/backend/internal/service/crs_sync_service.go
index edf3cd43d2..d0abc74038 100644
--- a/backend/internal/service/crs_sync_service.go
+++ b/backend/internal/service/crs_sync_service.go
@@ -168,6 +168,7 @@ type crsOpenAIResponsesAccount struct {
Status string `json:"status"`
Proxy *crsProxy `json:"proxy"`
Credentials map[string]any `json:"credentials"`
+ Extra map[string]any `json:"extra"`
}
type crsOpenAIOAuthAccount struct {
@@ -632,6 +633,18 @@ func (s *CRSSyncService) SyncFromCRS(ctx context.Context, input SyncFromCRSInput
result.Items = append(result.Items, item)
continue
}
+ var existingExtra map[string]any
+ if existing != nil {
+ existingExtra = existing.Extra
+ }
+ extra, err = mergeCRSOpenAILongContextBillingExtra(existingExtra, extra)
+ if err != nil {
+ item.Action = "failed"
+ item.Error = err.Error()
+ result.Failed++
+ result.Items = append(result.Items, item)
+ continue
+ }
if existing == nil {
if !shouldCreateAccount(src.ID, selectedSet) {
@@ -670,7 +683,7 @@ func (s *CRSSyncService) SyncFromCRS(ctx context.Context, input SyncFromCRSInput
continue
}
- existing.Extra = mergeMap(existing.Extra, extra)
+ existing.Extra = extra
existing.Name = defaultName(src.Name, src.ID)
existing.Platform = PlatformOpenAI
existing.Type = AccountTypeOAuth
@@ -751,11 +764,13 @@ func (s *CRSSyncService) SyncFromCRS(ctx context.Context, input SyncFromCRSInput
concurrency := 3
status := mapCRSStatus(src.IsActive, src.Status)
- extra := map[string]any{
- "crs_account_id": src.ID,
- "crs_kind": src.Kind,
- "crs_synced_at": now,
+ extra := make(map[string]any, len(src.Extra)+3)
+ for key, value := range src.Extra {
+ extra[key] = value
}
+ extra["crs_account_id"] = src.ID
+ extra["crs_kind"] = src.Kind
+ extra["crs_synced_at"] = now
existing, err := s.accountRepo.GetByCRSAccountID(ctx, src.ID)
if err != nil {
@@ -765,6 +780,18 @@ func (s *CRSSyncService) SyncFromCRS(ctx context.Context, input SyncFromCRSInput
result.Items = append(result.Items, item)
continue
}
+ var existingExtra map[string]any
+ if existing != nil {
+ existingExtra = existing.Extra
+ }
+ extra, err = mergeCRSOpenAILongContextBillingExtra(existingExtra, extra)
+ if err != nil {
+ item.Action = "failed"
+ item.Error = err.Error()
+ result.Failed++
+ result.Items = append(result.Items, item)
+ continue
+ }
if existing == nil {
if !shouldCreateAccount(src.ID, selectedSet) {
@@ -809,7 +836,7 @@ func (s *CRSSyncService) SyncFromCRS(ctx context.Context, input SyncFromCRSInput
continue
}
- existing.Extra = mergeMap(existing.Extra, extra)
+ existing.Extra = extra
existing.Name = defaultName(src.Name, src.ID)
existing.Platform = PlatformOpenAI
existing.Type = AccountTypeAPIKey
@@ -1098,6 +1125,10 @@ func mergeMap(existing map[string]any, updates map[string]any) map[string]any {
return out
}
+func mergeCRSOpenAILongContextBillingExtra(existing, updates map[string]any) (map[string]any, error) {
+ return normalizeOpenAILongContextBillingExtra(PlatformOpenAI, mergeMap(existing, updates))
+}
+
func (s *CRSSyncService) mapOrCreateProxy(ctx context.Context, enabled bool, cached *[]Proxy, src *crsProxy, defaultName string) (*int64, error) {
if !enabled || src == nil {
return nil, nil
diff --git a/backend/internal/service/gateway_usage_billing.go b/backend/internal/service/gateway_usage_billing.go
index 8a95915981..61ab3abd2f 100644
--- a/backend/internal/service/gateway_usage_billing.go
+++ b/backend/internal/service/gateway_usage_billing.go
@@ -947,6 +947,7 @@ func (s *GatewayService) buildRecordUsageLog(
usageLog.CacheReadCost = cost.CacheReadCost
usageLog.TotalCost = cost.TotalCost
usageLog.ActualCost = cost.ActualCost
+ usageLog.LongContextBillingApplied = cost.LongContextBillingApplied
}
return usageLog
diff --git a/backend/internal/service/grok_quota_fetcher.go b/backend/internal/service/grok_quota_fetcher.go
index 0939b78e20..f220fe33b9 100644
--- a/backend/internal/service/grok_quota_fetcher.go
+++ b/backend/internal/service/grok_quota_fetcher.go
@@ -3,6 +3,8 @@ package service
import (
"encoding/json"
"fmt"
+ "net/http"
+ "strings"
"time"
"github.com/Wei-Shaw/sub2api/internal/pkg/xai"
@@ -24,54 +26,150 @@ func (f *GrokQuotaFetcher) BuildUsageInfo(account *Account) *UsageInfo {
}
if account == nil {
usage.ErrorCode = "quota_unknown"
- usage.Error = "Grok quota is unknown until the first upstream response includes xAI rate-limit headers"
+ usage.Error = "Grok quota is unknown until billing is probed or an upstream response includes xAI rate-limit headers"
return usage
}
+ billing, _ := grokBillingSnapshotFromExtra(account.Extra)
snapshot, err := grokQuotaSnapshotFromExtra(account.Extra)
+ if billing != nil {
+ usage.GrokBilling = billing
+ if billing.Plan != "" {
+ usage.SubscriptionTier = billing.Plan
+ usage.SubscriptionTierRaw = billing.Plan
+ }
+ if parsedAt, parseErr := time.Parse(time.RFC3339, billing.UpdatedAt); parseErr == nil {
+ usage.UpdatedAt = &parsedAt
+ }
+ if billing.FetchedAt != "" {
+ usage.GrokLastQuotaProbeAt = billing.FetchedAt
+ }
+ usage.GrokQuotaSnapshotState = "billing_observed"
+ usage.GrokLastStatusCode = billing.StatusCode
+ switch billing.StatusCode {
+ case 401:
+ usage.NeedsReauth = true
+ usage.ErrorCode = "unauthenticated"
+ case 403:
+ usage.IsForbidden = true
+ usage.ForbiddenType = "forbidden"
+ usage.ErrorCode = "forbidden"
+ case 429:
+ usage.ErrorCode = "rate_limited"
+ }
+ }
+
if err != nil || snapshot == nil {
- usage.ErrorCode = "quota_unknown"
- usage.Error = "Grok quota is unknown until the first upstream response includes xAI rate-limit headers"
+ applyGrokCredentialUsageFallback(usage, account)
+ if billing == nil {
+ usage.ErrorCode = "quota_unknown"
+ usage.Error = "Grok quota is unknown until billing is probed or an upstream response includes xAI rate-limit headers"
+ }
return usage
}
- if parsedAt, err := time.Parse(time.RFC3339, snapshot.UpdatedAt); err == nil {
- usage.UpdatedAt = &parsedAt
+ if parsedAt, parseErr := time.Parse(time.RFC3339, snapshot.UpdatedAt); parseErr == nil {
+ if billing == nil || usage.UpdatedAt == nil || parsedAt.After(*usage.UpdatedAt) {
+ usage.UpdatedAt = &parsedAt
+ }
}
usage.GrokRequestQuota = snapshot.Requests
usage.GrokTokenQuota = snapshot.Tokens
usage.GrokRetryAfterSeconds = snapshot.RetryAfterSeconds
- usage.SubscriptionTier = snapshot.SubscriptionTier
- usage.SubscriptionTierRaw = snapshot.SubscriptionTier
- usage.GrokEntitlementStatus = snapshot.EntitlementStatus
- usage.GrokLastQuotaProbeAt = snapshot.LastProbeAt
+ if usage.SubscriptionTier == "" {
+ usage.SubscriptionTier = snapshot.SubscriptionTier
+ usage.SubscriptionTierRaw = snapshot.SubscriptionTier
+ }
+ if usage.GrokEntitlementStatus == "" {
+ usage.GrokEntitlementStatus = snapshot.EntitlementStatus
+ }
+ if usage.GrokLastQuotaProbeAt == "" {
+ usage.GrokLastQuotaProbeAt = snapshot.LastProbeAt
+ }
usage.GrokLastHeadersSeenAt = snapshot.LastHeadersSeenAt
- usage.GrokLastStatusCode = snapshot.StatusCode
+ if snapshot.StatusCode >= http.StatusBadRequest || usage.GrokLastStatusCode == 0 {
+ usage.GrokLastStatusCode = snapshot.StatusCode
+ }
if snapshot.HasObservedHeaders() {
- usage.GrokQuotaSnapshotState = "observed"
- } else {
+ if usage.GrokQuotaSnapshotState == "" {
+ usage.GrokQuotaSnapshotState = "observed"
+ }
+ } else if billing == nil {
usage.GrokQuotaSnapshotState = "no_headers"
usage.ErrorCode = "quota_unknown"
usage.Error = "No xAI quota headers observed on the latest Grok probe"
}
- switch snapshot.StatusCode {
- case 401:
- usage.NeedsReauth = true
- usage.ErrorCode = "unauthenticated"
- case 403:
- usage.IsForbidden = true
- usage.ForbiddenType = "forbidden"
- usage.ErrorCode = "forbidden"
- if usage.GrokEntitlementStatus == "" {
- usage.GrokEntitlementStatus = "forbidden"
+ if usage.ErrorCode == "" {
+ switch snapshot.StatusCode {
+ case 401:
+ usage.NeedsReauth = true
+ usage.ErrorCode = "unauthenticated"
+ case 403:
+ usage.IsForbidden = true
+ usage.ForbiddenType = "forbidden"
+ usage.ErrorCode = "forbidden"
+ if usage.GrokEntitlementStatus == "" {
+ usage.GrokEntitlementStatus = "forbidden"
+ }
+ case 429:
+ usage.ErrorCode = "rate_limited"
}
- case 429:
- usage.ErrorCode = "rate_limited"
}
+ applyGrokCredentialUsageFallback(usage, account)
return usage
}
+func applyGrokCredentialUsageFallback(usage *UsageInfo, account *Account) {
+ if usage == nil || account == nil {
+ return
+ }
+ if usage.SubscriptionTier == "" {
+ tier := strings.TrimSpace(account.GetCredential("subscription_tier"))
+ usage.SubscriptionTier = tier
+ usage.SubscriptionTierRaw = tier
+ }
+ if usage.GrokEntitlementStatus == "" {
+ usage.GrokEntitlementStatus = strings.TrimSpace(account.GetCredential("entitlement_status"))
+ }
+}
+
+func grokBillingSnapshotFromExtra(extra map[string]any) (*xai.BillingSummary, error) {
+ if extra == nil {
+ return nil, nil
+ }
+ raw, ok := extra[grokBillingExtraKey]
+ if !ok || raw == nil {
+ return nil, nil
+ }
+ switch snapshot := raw.(type) {
+ case *xai.BillingSummary:
+ return snapshot, nil
+ case xai.BillingSummary:
+ return &snapshot, nil
+ case map[string]any:
+ data, err := json.Marshal(snapshot)
+ if err != nil {
+ return nil, err
+ }
+ var out xai.BillingSummary
+ if err := json.Unmarshal(data, &out); err != nil {
+ return nil, err
+ }
+ return &out, nil
+ default:
+ data, err := json.Marshal(raw)
+ if err != nil {
+ return nil, fmt.Errorf("marshal grok billing snapshot: %w", err)
+ }
+ var out xai.BillingSummary
+ if err := json.Unmarshal(data, &out); err != nil {
+ return nil, err
+ }
+ return &out, nil
+ }
+}
+
func grokQuotaSnapshotFromExtra(extra map[string]any) (*xai.QuotaSnapshot, error) {
if extra == nil {
return nil, nil
diff --git a/backend/internal/service/grok_quota_fetcher_test.go b/backend/internal/service/grok_quota_fetcher_test.go
index d2d9c14993..1de9b51c9e 100644
--- a/backend/internal/service/grok_quota_fetcher_test.go
+++ b/backend/internal/service/grok_quota_fetcher_test.go
@@ -20,7 +20,34 @@ func TestGrokQuotaFetcherBuildUsageInfoUnknownUntilFirstSnapshot(t *testing.T) {
usage := NewGrokQuotaFetcher().BuildUsageInfo(&Account{Platform: PlatformGrok, Type: AccountTypeOAuth})
require.Equal(t, "passive", usage.Source)
require.Equal(t, "quota_unknown", usage.ErrorCode)
- require.Contains(t, usage.Error, "unknown until the first upstream response")
+ require.Contains(t, usage.Error, "unknown until billing is probed")
+}
+
+func TestGrokQuotaFetcherUsesCredentialTierWhenBillingHasNoPlan(t *testing.T) {
+ t.Parallel()
+
+ account := &Account{
+ Platform: PlatformGrok,
+ Type: AccountTypeOAuth,
+ Credentials: map[string]any{
+ "subscription_tier": " FREE ",
+ "entitlement_status": " active ",
+ },
+ Extra: map[string]any{
+ grokBillingExtraKey: &xai.BillingSummary{
+ PeriodType: "weekly",
+ StatusCode: http.StatusOK,
+ UpdatedAt: "2030-01-01T00:00:00Z",
+ },
+ },
+ }
+
+ usage := NewGrokQuotaFetcher().BuildUsageInfo(account)
+
+ require.NotNil(t, usage.GrokBilling)
+ require.Equal(t, "FREE", usage.SubscriptionTier)
+ require.Equal(t, "FREE", usage.SubscriptionTierRaw)
+ require.Equal(t, "active", usage.GrokEntitlementStatus)
}
func TestGrokQuotaFetcherBuildUsageInfoFromSnapshot(t *testing.T) {
@@ -68,6 +95,32 @@ func TestGrokQuotaFetcherBuildUsageInfoFromSnapshot(t *testing.T) {
require.True(t, usage.UpdatedAt.Equal(time.Date(2030, 1, 1, 0, 0, 0, 0, time.UTC)))
}
+func TestGrokQuotaFetcherSnapshotErrorOverridesSuccessfulBillingStatus(t *testing.T) {
+ t.Parallel()
+
+ updatedAt := "2030-01-01T00:00:00Z"
+ account := &Account{
+ Platform: PlatformGrok,
+ Type: AccountTypeOAuth,
+ Extra: map[string]any{
+ grokBillingExtraKey: &xai.BillingSummary{
+ PeriodType: "weekly",
+ StatusCode: http.StatusOK,
+ UpdatedAt: updatedAt,
+ },
+ grokQuotaSnapshotExtraKey: &xai.QuotaSnapshot{
+ StatusCode: http.StatusTooManyRequests,
+ UpdatedAt: updatedAt,
+ },
+ },
+ }
+
+ usage := NewGrokQuotaFetcher().BuildUsageInfo(account)
+
+ require.Equal(t, "rate_limited", usage.ErrorCode)
+ require.Equal(t, http.StatusTooManyRequests, usage.GrokLastStatusCode)
+}
+
func TestGrokQuotaFetcherBuildUsageInfoFromNoHeadersProbe(t *testing.T) {
t.Parallel()
diff --git a/backend/internal/service/grok_quota_service.go b/backend/internal/service/grok_quota_service.go
index 17e91dac6f..aa838cb491 100644
--- a/backend/internal/service/grok_quota_service.go
+++ b/backend/internal/service/grok_quota_service.go
@@ -7,27 +7,36 @@ import (
"io"
"log/slog"
"net/http"
+ "strconv"
"strings"
+ "sync"
"time"
infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors"
"github.com/Wei-Shaw/sub2api/internal/pkg/xai"
+ "golang.org/x/sync/singleflight"
)
const (
grokQuotaUpstreamTimeout = 20 * time.Second
grokQuotaProbeInput = "."
- grokQuotaDefaultModel = "grok-4.3"
+ grokQuotaDefaultModel = grokDefaultResponsesModel
+ grokBillingExtraKey = "grok_billing_snapshot"
)
type GrokQuotaProbeResult struct {
- Source string `json:"source"`
- Model string `json:"model"`
- Snapshot *xai.QuotaSnapshot `json:"snapshot,omitempty"`
- StatusCode int `json:"status_code,omitempty"`
- HeadersObserved bool `json:"headers_observed"`
- ResetSupported bool `json:"reset_supported"`
- FetchedAt int64 `json:"fetched_at"`
+ Source string `json:"source"`
+ Model string `json:"model,omitempty"`
+ Billing *xai.BillingSummary `json:"billing,omitempty"`
+ Snapshot *xai.QuotaSnapshot `json:"snapshot,omitempty"`
+ LocalUsage7d *WindowStats `json:"local_usage_7d,omitempty"`
+ LocalUsageMonthly *WindowStats `json:"local_usage_monthly,omitempty"`
+ StatusCode int `json:"status_code,omitempty"`
+ HeadersObserved bool `json:"headers_observed"`
+ ResetSupported bool `json:"reset_supported"`
+ FetchedAt int64 `json:"fetched_at"`
+ Persisted bool `json:"persisted"`
+ ProbeError string `json:"probe_error,omitempty"`
}
type GrokQuotaResetResult struct {
@@ -41,6 +50,8 @@ type GrokQuotaService struct {
proxyRepo ProxyRepository
tokenProvider *GrokTokenProvider
httpUpstream HTTPUpstream
+ usageLogRepo UsageLogRepository
+ probeFlight singleflight.Group
}
func NewGrokQuotaService(
@@ -48,16 +59,70 @@ func NewGrokQuotaService(
proxyRepo ProxyRepository,
tokenProvider *GrokTokenProvider,
httpUpstream HTTPUpstream,
+ usageLogRepos ...UsageLogRepository,
) *GrokQuotaService {
+ var usageLogRepo UsageLogRepository
+ if len(usageLogRepos) > 0 {
+ usageLogRepo = usageLogRepos[0]
+ }
return &GrokQuotaService{
accountRepo: accountRepo,
proxyRepo: proxyRepo,
tokenProvider: tokenProvider,
httpUpstream: httpUpstream,
+ usageLogRepo: usageLogRepo,
}
}
+// QueryQuota combines xAI billing data with an active quota-header probe for
+// Free accounts, whose billing response does not include usage_percent.
+func (s *GrokQuotaService) QueryQuota(ctx context.Context, accountID int64) (*GrokQuotaProbeResult, error) {
+ billingResult, billingErr := s.ProbeBilling(ctx, accountID)
+ if billingErr == nil && billingResult != nil && grokBillingHasAuthoritativeQuota(billingResult.Billing) {
+ return billingResult, nil
+ }
+
+ probeResult, probeErr := s.ProbeUsage(ctx, accountID)
+ if probeErr != nil {
+ if billingResult != nil && billingResult.Billing != nil {
+ billingResult.ProbeError = probeErr.Error()
+ return billingResult, nil
+ }
+ return nil, probeErr
+ }
+ if probeResult == nil {
+ if billingErr != nil {
+ return nil, billingErr
+ }
+ return nil, infraerrors.New(http.StatusBadGateway, "GROK_QUOTA_PROBE_EMPTY", "Grok quota probe returned no result")
+ }
+ if billingResult != nil {
+ probeResult.Source = "hybrid_probe"
+ probeResult.Billing = billingResult.Billing
+ probeResult.LocalUsage7d = billingResult.LocalUsage7d
+ probeResult.LocalUsageMonthly = billingResult.LocalUsageMonthly
+ probeResult.Persisted = probeResult.Persisted || billingResult.Persisted
+ }
+ return probeResult, nil
+}
+
+func grokBillingHasAuthoritativeQuota(billing *xai.BillingSummary) bool {
+ if billing == nil {
+ return false
+ }
+ return billing.UsagePercent != nil ||
+ billing.UsedPercent != nil ||
+ (billing.MonthlyLimitCents != nil && *billing.MonthlyLimitCents > 0) ||
+ strings.TrimSpace(billing.Plan) != ""
+}
+
func (s *GrokQuotaService) ProbeUsage(ctx context.Context, accountID int64) (*GrokQuotaProbeResult, error) {
+ return s.runProbeFlight(ctx, "active:"+strconv.FormatInt(accountID, 10), func(sharedCtx context.Context) (*GrokQuotaProbeResult, error) {
+ return s.probeUsage(sharedCtx, accountID)
+ })
+}
+
+func (s *GrokQuotaService) probeUsage(ctx context.Context, accountID int64) (*GrokQuotaProbeResult, error) {
account, token, proxyURL, err := s.prepareProbe(ctx, accountID)
if err != nil {
return nil, err
@@ -95,7 +160,7 @@ func (s *GrokQuotaService) ProbeUsage(ctx context.Context, accountID int64) (*Gr
if limited {
normalizeGrokExhaustedWindowResets(snapshot, resetAt, time.Now())
}
- _ = s.accountRepo.UpdateExtra(ctx, account.ID, map[string]any{
+ persistErr := s.accountRepo.UpdateExtra(ctx, account.ID, map[string]any{
grokQuotaSnapshotExtraKey: snapshot,
})
if limited {
@@ -110,6 +175,7 @@ func (s *GrokQuotaService) ProbeUsage(ctx context.Context, accountID int64) (*Gr
HeadersObserved: snapshot.HeadersObserved,
ResetSupported: false,
FetchedAt: time.Now().Unix(),
+ Persisted: persistErr == nil,
}
if resp.StatusCode == http.StatusTooManyRequests {
return result, nil
@@ -123,6 +189,173 @@ func (s *GrokQuotaService) ProbeUsage(ctx context.Context, accountID int64) (*Gr
return result, nil
}
+// ProbeBilling only calls the xAI billing endpoints. Account usage refreshes
+// use this method so opening the account list never consumes model quota.
+func (s *GrokQuotaService) ProbeBilling(ctx context.Context, accountID int64) (*GrokQuotaProbeResult, error) {
+ return s.runProbeFlight(ctx, "billing:"+strconv.FormatInt(accountID, 10), func(sharedCtx context.Context) (*GrokQuotaProbeResult, error) {
+ return s.probeBilling(sharedCtx, accountID)
+ })
+}
+
+func (s *GrokQuotaService) probeBilling(ctx context.Context, accountID int64) (*GrokQuotaProbeResult, error) {
+ account, token, proxyURL, err := s.prepareProbe(ctx, accountID)
+ if err != nil {
+ return nil, err
+ }
+
+ probeCtx, cancel := context.WithTimeout(ctx, grokQuotaUpstreamTimeout)
+ defer cancel()
+ type billingResult struct {
+ summary *xai.BillingSummary
+ status int
+ err error
+ }
+ var weekly, monthly billingResult
+ var wg sync.WaitGroup
+ wg.Add(2)
+ go func() {
+ defer wg.Done()
+ weekly.summary, weekly.status, weekly.err = s.fetchBilling(probeCtx, account, token, proxyURL, true)
+ }()
+ go func() {
+ defer wg.Done()
+ monthly.summary, monthly.status, monthly.err = s.fetchBilling(probeCtx, account, token, proxyURL, false)
+ }()
+ wg.Wait()
+
+ weeklyOK := weekly.summary != nil
+ monthlyOK := monthly.summary != nil
+ if !weeklyOK && !monthlyOK {
+ return nil, mergeGrokBillingProbeErrors(weekly.status, monthly.status, weekly.err, monthly.err)
+ }
+ statusCode := preferSuccessfulBillingStatus(weekly.status, monthly.status, weeklyOK, monthlyOK)
+ previous, _ := grokBillingSnapshotFromExtra(account.Extra)
+ billing := xai.MergeBillingProbeResult(previous, weekly.summary, monthly.summary, weeklyOK, monthlyOK)
+ billing = xai.StampBillingSummary(billing, statusCode, "billing_probe")
+ persistErr := s.accountRepo.UpdateExtra(ctx, account.ID, map[string]any{
+ grokBillingExtraKey: billing,
+ })
+ if persistErr != nil {
+ slog.Warn("grok_billing_persist_failed", "account_id", account.ID, "error", persistErr)
+ }
+ localUsage7d, localUsageMonthly := grokLocalUsageForBilling(ctx, s.usageLogRepo, account.ID, billing, time.Now().UTC())
+ return &GrokQuotaProbeResult{
+ Source: "billing_probe",
+ Billing: billing,
+ LocalUsage7d: localUsage7d,
+ LocalUsageMonthly: localUsageMonthly,
+ StatusCode: statusCode,
+ FetchedAt: time.Now().Unix(),
+ Persisted: persistErr == nil,
+ }, nil
+}
+
+func (s *GrokQuotaService) runProbeFlight(
+ ctx context.Context,
+ key string,
+ probe func(context.Context) (*GrokQuotaProbeResult, error),
+) (*GrokQuotaProbeResult, error) {
+ if s == nil {
+ return nil, infraerrors.New(http.StatusInternalServerError, "GROK_QUOTA_NOT_CONFIGURED", "grok quota service is not configured")
+ }
+ resultCh := s.probeFlight.DoChan(key, func() (any, error) {
+ sharedCtx, cancel := context.WithTimeout(context.Background(), grokQuotaUpstreamTimeout+5*time.Second)
+ defer cancel()
+ return probe(sharedCtx)
+ })
+ select {
+ case <-ctx.Done():
+ return nil, ctx.Err()
+ case flightResult := <-resultCh:
+ if flightResult.Err != nil {
+ return nil, flightResult.Err
+ }
+ result, ok := flightResult.Val.(*GrokQuotaProbeResult)
+ if !ok || result == nil {
+ return nil, infraerrors.New(http.StatusInternalServerError, "GROK_QUOTA_PROBE_RESULT_INVALID", "invalid Grok quota probe result")
+ }
+ cloned := *result
+ return &cloned, nil
+ }
+}
+
+func (s *GrokQuotaService) fetchBilling(
+ ctx context.Context,
+ account *Account,
+ token string,
+ proxyURL string,
+ weekly bool,
+) (*xai.BillingSummary, int, error) {
+ req, err := http.NewRequestWithContext(ctx, http.MethodGet, xai.BuildBillingURL(weekly), nil)
+ if err != nil {
+ return nil, 0, infraerrors.Newf(http.StatusInternalServerError, "GROK_QUOTA_PROBE_REQUEST_BUILD_FAILED", "failed to build billing request: %v", err)
+ }
+ xai.ApplyCLIBillingHeaders(req, token)
+ resp, err := s.httpUpstream.Do(req, proxyURL, account.ID, maxInt(account.Concurrency, 2))
+ if err != nil {
+ return nil, 0, infraerrors.Newf(http.StatusBadGateway, "GROK_QUOTA_PROBE_REQUEST_FAILED", "billing request failed: %v", err)
+ }
+ defer func() { _ = resp.Body.Close() }()
+
+ bodyBytes, _ := io.ReadAll(io.LimitReader(resp.Body, 1<<20))
+ if resp.StatusCode == http.StatusTooManyRequests {
+ return nil, resp.StatusCode, nil
+ }
+ if resp.StatusCode >= 400 {
+ bodyText := truncate(strings.TrimSpace(string(bodyBytes)), 240)
+ slog.Warn("grok_quota_billing_failed", "account_id", account.ID, "weekly", weekly, "status", resp.StatusCode, "body", bodyText)
+ return nil, resp.StatusCode, infraerrors.Newf(mapUpstreamStatus(resp.StatusCode), "GROK_QUOTA_PROBE_UPSTREAM_ERROR", "billing returned %d: %s", resp.StatusCode, bodyText)
+ }
+ payload, err := xai.ParseBillingPayload(bodyBytes)
+ if err != nil {
+ return nil, resp.StatusCode, infraerrors.Newf(http.StatusBadGateway, "GROK_QUOTA_BILLING_PARSE_ERROR", "failed to parse billing body: %v", err)
+ }
+ return xai.BuildBillingSummary(payload.Config), resp.StatusCode, nil
+}
+
+func mergeGrokBillingProbeErrors(weeklyStatus, monthlyStatus int, weeklyErr, monthlyErr error) error {
+ weeklyKey := grokBillingProbeErrorKey(weeklyStatus, weeklyErr)
+ monthlyKey := grokBillingProbeErrorKey(monthlyStatus, monthlyErr)
+ if weeklyKey == monthlyKey {
+ switch {
+ case weeklyErr != nil:
+ return weeklyErr
+ case monthlyErr != nil:
+ return monthlyErr
+ case weeklyStatus == http.StatusTooManyRequests:
+ return infraerrors.New(http.StatusTooManyRequests, "GROK_QUOTA_PROBE_UPSTREAM_ERROR", "billing rate limited")
+ case weeklyStatus != 0 && weeklyStatus != http.StatusOK:
+ return infraerrors.New(mapUpstreamStatus(weeklyStatus), "GROK_QUOTA_PROBE_UPSTREAM_ERROR", "xAI billing endpoints returned the same upstream error")
+ default:
+ return infraerrors.New(http.StatusBadGateway, "GROK_QUOTA_BILLING_EMPTY", "xAI billing endpoints returned no quota data")
+ }
+ }
+ slog.Warn("grok_quota_probe_parts_failed", "weekly_status", weeklyStatus, "weekly_error", weeklyErr, "monthly_status", monthlyStatus, "monthly_error", monthlyErr)
+ return infraerrors.New(http.StatusBadGateway, "GROK_QUOTA_PROBE_PARTS_FAILED", "weekly and monthly billing probes failed differently").WithMetadata(map[string]string{
+ "weekly_status": strconv.Itoa(weeklyStatus), "monthly_status": strconv.Itoa(monthlyStatus),
+ })
+}
+
+func grokBillingProbeErrorKey(status int, err error) string {
+ if err != nil {
+ return strconv.Itoa(status) + ":" + strconv.Itoa(infraerrors.Code(err)) + ":" + infraerrors.Reason(err)
+ }
+ return strconv.Itoa(status) + ":empty"
+}
+
+func preferSuccessfulBillingStatus(weeklyStatus, monthlyStatus int, weeklyOK, monthlyOK bool) int {
+ if weeklyOK && weeklyStatus >= 200 && weeklyStatus < 300 {
+ return weeklyStatus
+ }
+ if monthlyOK && monthlyStatus >= 200 && monthlyStatus < 300 {
+ return monthlyStatus
+ }
+ if weeklyStatus != 0 {
+ return weeklyStatus
+ }
+ return monthlyStatus
+}
+
func (s *GrokQuotaService) ResetQuota(ctx context.Context, accountID int64) (*GrokQuotaResetResult, error) {
if _, err := s.loadGrokOAuthAccount(ctx, accountID); err != nil {
return nil, err
diff --git a/backend/internal/service/grok_quota_service_test.go b/backend/internal/service/grok_quota_service_test.go
index 2248674899..d49497e69b 100644
--- a/backend/internal/service/grok_quota_service_test.go
+++ b/backend/internal/service/grok_quota_service_test.go
@@ -6,11 +6,14 @@ import (
"context"
"io"
"net/http"
+ "strconv"
"strings"
+ "sync"
"testing"
"time"
infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors"
+ "github.com/Wei-Shaw/sub2api/internal/pkg/usagestats"
"github.com/Wei-Shaw/sub2api/internal/pkg/xai"
"github.com/stretchr/testify/require"
"github.com/tidwall/gjson"
@@ -63,6 +66,107 @@ type grokQuotaProxyRepo struct {
calls int
}
+type grokQuotaUsageLogRepo struct {
+ UsageLogRepository
+ stats *usagestats.AccountStats
+ err error
+ calls int
+}
+
+func (r *grokQuotaUsageLogRepo) GetAccountWindowStats(context.Context, int64, time.Time) (*usagestats.AccountStats, error) {
+ r.calls++
+ return r.stats, r.err
+}
+
+type grokHybridUpstream struct {
+ httpUpstreamRecorder
+ mu sync.Mutex
+ requests []*http.Request
+ bodies [][]byte
+ weeklyUsagePercent *float64
+ monthlyLimitCents *float64
+ activeStatus int
+ activeHeaders http.Header
+ billingStarted chan struct{}
+ billingRelease <-chan struct{}
+ billingStartOnce sync.Once
+ billingStatus int
+ billingHeaders http.Header
+}
+
+func (u *grokHybridUpstream) Do(req *http.Request, _ string, _ int64, _ int) (*http.Response, error) {
+ var body []byte
+ if req != nil && req.Body != nil {
+ body, _ = io.ReadAll(req.Body)
+ }
+ u.mu.Lock()
+ u.requests = append(u.requests, req)
+ u.bodies = append(u.bodies, body)
+ u.mu.Unlock()
+
+ if req.URL.Path == "/v1/responses" {
+ status := u.activeStatus
+ if status == 0 {
+ status = http.StatusOK
+ }
+ headers := u.activeHeaders
+ if headers == nil {
+ headers = http.Header{
+ "X-Ratelimit-Limit-Tokens": []string{"2000000"},
+ "X-Ratelimit-Remaining-Tokens": []string{"1500000"},
+ }
+ }
+ return &http.Response{StatusCode: status, Header: headers, Body: io.NopCloser(strings.NewReader(`{"id":"resp_probe"}`))}, nil
+ }
+ if u.billingStarted != nil {
+ u.billingStartOnce.Do(func() { close(u.billingStarted) })
+ }
+ if u.billingRelease != nil {
+ select {
+ case <-u.billingRelease:
+ case <-req.Context().Done():
+ return nil, req.Context().Err()
+ }
+ }
+ if u.billingStatus != 0 && u.billingStatus != http.StatusOK {
+ return &http.Response{
+ StatusCode: u.billingStatus,
+ Header: u.billingHeaders,
+ Body: io.NopCloser(strings.NewReader(`{"error":{"message":"billing limited"}}`)),
+ }, nil
+ }
+
+ if req.URL.RawQuery == "format=credits" {
+ usage := ""
+ if u.weeklyUsagePercent != nil {
+ usage = `,"creditUsagePercent":` + strconv.FormatFloat(*u.weeklyUsagePercent, 'f', -1, 64)
+ }
+ payload := `{"config":{"currentPeriod":{"type":"WEEKLY","start":"2026-07-09T03:25:00Z","end":"2026-07-16T03:25:00Z"}` + usage + `}}`
+ return &http.Response{StatusCode: http.StatusOK, Header: make(http.Header), Body: io.NopCloser(strings.NewReader(payload))}, nil
+ }
+ monthlyLimit := ""
+ if u.monthlyLimitCents != nil {
+ monthlyLimit = `,"monthlyLimit":{"val":` + strconv.FormatFloat(*u.monthlyLimitCents, 'f', -1, 64) + `}`
+ }
+ monthlyPayload := `{"config":{"billingPeriodStart":"2026-07-01T00:00:00Z","billingPeriodEnd":"2026-08-01T00:00:00Z"` + monthlyLimit + `}}`
+ return &http.Response{
+ StatusCode: http.StatusOK,
+ Header: make(http.Header),
+ Body: io.NopCloser(strings.NewReader(monthlyPayload)),
+ }, nil
+}
+
+func (u *grokHybridUpstream) snapshot() ([]*http.Request, [][]byte) {
+ u.mu.Lock()
+ defer u.mu.Unlock()
+ requests := append([]*http.Request(nil), u.requests...)
+ bodies := make([][]byte, len(u.bodies))
+ for i := range u.bodies {
+ bodies[i] = append([]byte(nil), u.bodies[i]...)
+ }
+ return requests, bodies
+}
+
func (r *grokQuotaProxyRepo) GetByID(_ context.Context, id int64) (*Proxy, error) {
r.calls++
return r.proxies[id], nil
@@ -102,7 +206,7 @@ func TestGrokQuotaServiceProbeUsageStoresHeaders(t *testing.T) {
result, err := svc.ProbeUsage(context.Background(), 42)
require.NoError(t, err)
require.Equal(t, http.StatusOK, result.StatusCode)
- require.Equal(t, "grok-4.3", result.Model)
+ require.Equal(t, "grok-4.5", result.Model)
require.True(t, result.HeadersObserved)
require.NotNil(t, result.Snapshot)
require.True(t, result.Snapshot.HeadersObserved)
@@ -115,7 +219,7 @@ func TestGrokQuotaServiceProbeUsageStoresHeaders(t *testing.T) {
require.Equal(t, "https://cli-chat-proxy.grok.com/v1/responses", upstream.lastReq.URL.String())
require.Equal(t, "Bearer access-token", upstream.lastReq.Header.Get("Authorization"))
require.Equal(t, grokCLIVersion, upstream.lastReq.Header.Get("X-Grok-Client-Version"))
- require.Equal(t, "grok-4.3", gjson.GetBytes(upstream.lastBody, "model").String())
+ require.Equal(t, "grok-4.5", gjson.GetBytes(upstream.lastBody, "model").String())
require.Contains(t, string(upstream.lastBody), `"max_output_tokens":1`)
require.Contains(t, string(upstream.lastBody), `"store":false`)
require.NotNil(t, repo.updates[42][grokQuotaSnapshotExtraKey])
@@ -152,8 +256,8 @@ func TestGrokQuotaServiceProbeUsageIgnoresAccountGrokMapping(t *testing.T) {
result, err := svc.ProbeUsage(context.Background(), 47)
require.NoError(t, err)
- require.Equal(t, "grok-4.3", result.Model)
- require.Equal(t, "grok-4.3", gjson.GetBytes(upstream.lastBody, "model").String())
+ require.Equal(t, "grok-4.5", result.Model)
+ require.Equal(t, "grok-4.5", gjson.GetBytes(upstream.lastBody, "model").String())
require.NotContains(t, string(upstream.lastBody), "grok-composer")
}
@@ -185,7 +289,7 @@ func TestGrokQuotaServiceProbeUsageReportsProbeModelOnUpstreamError(t *testing.T
_, err := svc.ProbeUsage(context.Background(), 48)
require.Error(t, err)
require.Equal(t, "GROK_QUOTA_PROBE_UPSTREAM_ERROR", infraerrors.Reason(err))
- require.Contains(t, infraerrors.Message(err), `probe model "grok-4.3"`)
+ require.Contains(t, infraerrors.Message(err), `probe model "grok-4.5"`)
}
func TestGrokQuotaServiceProbeUsageLoadsProxyWhenAccountEdgeMissing(t *testing.T) {
@@ -308,6 +412,299 @@ func TestGrokQuotaServiceProbeUsageReturnsRateLimitedSnapshot(t *testing.T) {
require.Zero(t, repo.tempUnschedCalls)
}
+func TestGrokQuotaServiceQueryQuotaFreeFallsBackToGrok45(t *testing.T) {
+ t.Parallel()
+
+ account := &Account{
+ ID: 51, Platform: PlatformGrok, Type: AccountTypeOAuth, Concurrency: 1,
+ Credentials: map[string]any{
+ "access_token": "access-token",
+ "expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339),
+ },
+ }
+ repo := &grokQuotaAccountRepo{mockAccountRepoForPlatform: &mockAccountRepoForPlatform{
+ accountsByID: map[int64]*Account{account.ID: account},
+ }}
+ upstream := &grokHybridUpstream{}
+ svc := NewGrokQuotaService(repo, nil, NewGrokTokenProvider(repo, nil), upstream)
+
+ result, err := svc.QueryQuota(context.Background(), account.ID)
+ require.NoError(t, err)
+ require.Equal(t, "hybrid_probe", result.Source)
+ require.Equal(t, "grok-4.5", result.Model)
+ require.NotNil(t, result.Billing)
+ require.Nil(t, result.Billing.UsagePercent)
+ require.NotNil(t, result.Snapshot)
+ require.NotNil(t, result.Snapshot.Tokens)
+ require.EqualValues(t, 2_000_000, *result.Snapshot.Tokens.Limit)
+ require.True(t, result.HeadersObserved)
+
+ requests, bodies := upstream.snapshot()
+ require.Len(t, requests, 3)
+ responseCalls := 0
+ for i, req := range requests {
+ if req.URL.Path != "/v1/responses" {
+ continue
+ }
+ responseCalls++
+ require.Equal(t, http.MethodPost, req.Method)
+ require.Equal(t, "grok-4.5", gjson.GetBytes(bodies[i], "model").String())
+ require.EqualValues(t, 1, gjson.GetBytes(bodies[i], "max_output_tokens").Int())
+ }
+ require.Equal(t, 1, responseCalls)
+}
+
+func TestGrokQuotaServiceQueryQuotaPaidBillingSkipsActiveProbe(t *testing.T) {
+ t.Parallel()
+
+ account := &Account{
+ ID: 52, Platform: PlatformGrok, Type: AccountTypeOAuth, Concurrency: 1,
+ Credentials: map[string]any{
+ "access_token": "access-token",
+ "expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339),
+ },
+ }
+ repo := &grokQuotaAccountRepo{mockAccountRepoForPlatform: &mockAccountRepoForPlatform{
+ accountsByID: map[int64]*Account{account.ID: account},
+ }}
+ usagePercent := 25.0
+ upstream := &grokHybridUpstream{weeklyUsagePercent: &usagePercent}
+ svc := NewGrokQuotaService(repo, nil, NewGrokTokenProvider(repo, nil), upstream)
+
+ result, err := svc.QueryQuota(context.Background(), account.ID)
+ require.NoError(t, err)
+ require.Equal(t, "billing_probe", result.Source)
+ require.NotNil(t, result.Billing)
+ require.InDelta(t, usagePercent, *result.Billing.UsagePercent, 1e-9)
+ require.Nil(t, result.Snapshot)
+ require.Empty(t, result.Model)
+
+ requests, _ := upstream.snapshot()
+ require.Len(t, requests, 2)
+ for _, req := range requests {
+ require.Equal(t, "/v1/billing", req.URL.Path)
+ }
+}
+
+func TestGrokQuotaServiceQueryQuotaCustomPaidMonthlyLimitSkipsActiveProbe(t *testing.T) {
+ t.Parallel()
+
+ account := &Account{
+ ID: 57, Platform: PlatformGrok, Type: AccountTypeOAuth, Concurrency: 1,
+ Credentials: map[string]any{
+ "access_token": "access-token",
+ "expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339),
+ },
+ }
+ repo := &grokQuotaAccountRepo{mockAccountRepoForPlatform: &mockAccountRepoForPlatform{
+ accountsByID: map[int64]*Account{account.ID: account},
+ }}
+ monthlyLimit := 25_000.0
+ upstream := &grokHybridUpstream{monthlyLimitCents: &monthlyLimit}
+ svc := NewGrokQuotaService(repo, nil, NewGrokTokenProvider(repo, nil), upstream)
+
+ result, err := svc.QueryQuota(context.Background(), account.ID)
+ require.NoError(t, err)
+ require.Equal(t, "billing_probe", result.Source)
+ require.NotNil(t, result.Billing)
+ require.InDelta(t, monthlyLimit, *result.Billing.MonthlyLimitCents, 1e-9)
+ require.Nil(t, result.Snapshot)
+
+ requests, _ := upstream.snapshot()
+ require.Len(t, requests, 2)
+ for _, req := range requests {
+ require.Equal(t, "/v1/billing", req.URL.Path)
+ }
+}
+
+func TestGrokLocalUsageForBillingOnlyReturnsAvailableWindows(t *testing.T) {
+ t.Parallel()
+
+ now := time.Date(2026, 7, 13, 12, 0, 0, 0, time.UTC)
+ billing := &xai.BillingSummary{
+ PeriodType: "weekly",
+ PeriodStart: now.Add(-4 * 24 * time.Hour).Format(time.RFC3339),
+ PeriodEnd: now.Add(3 * 24 * time.Hour).Format(time.RFC3339),
+ }
+
+ t.Run("valid weekly window", func(t *testing.T) {
+ repo := &grokQuotaUsageLogRepo{stats: &usagestats.AccountStats{Tokens: 1_500_000}}
+ weekly, monthly := grokLocalUsageForBilling(context.Background(), repo, 57, billing, now)
+ require.NotNil(t, weekly)
+ require.EqualValues(t, 1_500_000, weekly.Tokens)
+ require.Nil(t, monthly)
+ require.Equal(t, 1, repo.calls)
+ })
+
+ t.Run("query failure", func(t *testing.T) {
+ repo := &grokQuotaUsageLogRepo{err: context.DeadlineExceeded}
+ weekly, monthly := grokLocalUsageForBilling(context.Background(), repo, 57, billing, now)
+ require.Nil(t, weekly)
+ require.Nil(t, monthly)
+ require.Equal(t, 1, repo.calls)
+ })
+
+ t.Run("missing billing window", func(t *testing.T) {
+ repo := &grokQuotaUsageLogRepo{}
+ weekly, monthly := grokLocalUsageForBilling(context.Background(), repo, 57, nil, now)
+ require.Nil(t, weekly)
+ require.Nil(t, monthly)
+ require.Zero(t, repo.calls)
+ })
+}
+
+func TestAccountUsageServiceGrokRefreshUsesBillingOnly(t *testing.T) {
+ t.Parallel()
+
+ account := &Account{
+ ID: 54, Platform: PlatformGrok, Type: AccountTypeOAuth, Concurrency: 1,
+ Credentials: map[string]any{
+ "access_token": "access-token",
+ "expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339),
+ },
+ }
+ repo := &grokQuotaAccountRepo{mockAccountRepoForPlatform: &mockAccountRepoForPlatform{
+ accountsByID: map[int64]*Account{account.ID: account},
+ }}
+ upstream := &grokHybridUpstream{}
+ quotaService := NewGrokQuotaService(repo, nil, NewGrokTokenProvider(repo, nil), upstream)
+ usageService := &AccountUsageService{
+ grokQuotaFetcher: NewGrokQuotaFetcher(),
+ grokQuotaService: quotaService,
+ cache: NewUsageCache(),
+ }
+
+ usage, err := usageService.getGrokUsage(context.Background(), account, false)
+ require.NoError(t, err)
+ require.NotNil(t, usage.GrokBilling)
+ require.Nil(t, usage.GrokBilling.UsagePercent)
+
+ requests, _ := upstream.snapshot()
+ require.Len(t, requests, 2)
+ for _, req := range requests {
+ require.Equal(t, http.MethodGet, req.Method)
+ require.Equal(t, "/v1/billing", req.URL.Path)
+ }
+}
+
+func TestGrokQuotaServiceProbeFlightsDeduplicateBillingAndSeparateActive(t *testing.T) {
+ t.Parallel()
+
+ account := &Account{
+ ID: 55, Platform: PlatformGrok, Type: AccountTypeOAuth, Concurrency: 1,
+ Credentials: map[string]any{
+ "access_token": "access-token",
+ "expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339),
+ },
+ }
+ repo := &grokQuotaAccountRepo{mockAccountRepoForPlatform: &mockAccountRepoForPlatform{
+ accountsByID: map[int64]*Account{account.ID: account},
+ }}
+ billingStarted := make(chan struct{})
+ billingRelease := make(chan struct{})
+ upstream := &grokHybridUpstream{billingStarted: billingStarted, billingRelease: billingRelease}
+ svc := NewGrokQuotaService(repo, nil, NewGrokTokenProvider(repo, nil), upstream)
+
+ type probeOutcome struct {
+ result *GrokQuotaProbeResult
+ err error
+ }
+ billingOutcomes := make(chan probeOutcome, 2)
+ go func() {
+ result, err := svc.ProbeBilling(context.Background(), account.ID)
+ billingOutcomes <- probeOutcome{result: result, err: err}
+ }()
+ <-billingStarted
+ secondStarted := make(chan struct{})
+ go func() {
+ close(secondStarted)
+ result, err := svc.ProbeBilling(context.Background(), account.ID)
+ billingOutcomes <- probeOutcome{result: result, err: err}
+ }()
+ <-secondStarted
+ time.Sleep(25 * time.Millisecond)
+
+ activeResult, err := svc.ProbeUsage(context.Background(), account.ID)
+ require.NoError(t, err)
+ require.NotNil(t, activeResult.Snapshot)
+ close(billingRelease)
+ for range 2 {
+ outcome := <-billingOutcomes
+ require.NoError(t, outcome.err)
+ require.NotNil(t, outcome.result.Billing)
+ }
+
+ requests, _ := upstream.snapshot()
+ billingCalls := 0
+ activeCalls := 0
+ for _, req := range requests {
+ switch req.URL.Path {
+ case "/v1/billing":
+ billingCalls++
+ case "/v1/responses":
+ activeCalls++
+ }
+ }
+ require.Equal(t, 2, billingCalls)
+ require.Equal(t, 1, activeCalls)
+}
+
+func TestGrokQuotaServiceBilling429DoesNotPauseModelScheduling(t *testing.T) {
+ t.Parallel()
+
+ account := &Account{
+ ID: 56, Platform: PlatformGrok, Type: AccountTypeOAuth, Concurrency: 1,
+ Credentials: map[string]any{
+ "access_token": "access-token",
+ "expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339),
+ },
+ }
+ repo := &grokQuotaAccountRepo{mockAccountRepoForPlatform: &mockAccountRepoForPlatform{
+ accountsByID: map[int64]*Account{account.ID: account},
+ }}
+ upstream := &grokHybridUpstream{
+ billingStatus: http.StatusTooManyRequests,
+ billingHeaders: http.Header{"Retry-After": []string{"45"}},
+ }
+ svc := NewGrokQuotaService(repo, nil, NewGrokTokenProvider(repo, nil), upstream)
+
+ result, err := svc.ProbeBilling(context.Background(), account.ID)
+
+ require.Error(t, err)
+ require.Nil(t, result)
+ require.Zero(t, repo.rateLimitedCalls)
+}
+
+func TestGrokQuotaServiceQueryQuotaFree429PersistsLimitAndKeepsBilling(t *testing.T) {
+ t.Parallel()
+
+ account := &Account{
+ ID: 53, Platform: PlatformGrok, Type: AccountTypeOAuth, Concurrency: 1,
+ Credentials: map[string]any{
+ "access_token": "access-token",
+ "expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339),
+ },
+ }
+ repo := &grokQuotaAccountRepo{mockAccountRepoForPlatform: &mockAccountRepoForPlatform{
+ accountsByID: map[int64]*Account{account.ID: account},
+ }}
+ upstream := &grokHybridUpstream{
+ activeStatus: http.StatusTooManyRequests,
+ activeHeaders: http.Header{"Retry-After": []string{"45"}},
+ }
+ svc := NewGrokQuotaService(repo, nil, NewGrokTokenProvider(repo, nil), upstream)
+
+ result, err := svc.QueryQuota(context.Background(), account.ID)
+ require.NoError(t, err)
+ require.Equal(t, http.StatusTooManyRequests, result.StatusCode)
+ require.NotNil(t, result.Billing)
+ require.NotNil(t, result.Snapshot)
+ require.Equal(t, 45, *result.Snapshot.RetryAfterSeconds)
+ require.Equal(t, 1, repo.rateLimitedCalls)
+ require.Equal(t, account.ID, repo.lastRateLimitedID)
+ require.WithinDuration(t, time.Now().Add(45*time.Second), repo.lastRateLimitResetAt, time.Second)
+}
+
func TestGrokQuotaServiceResetQuotaUnsupported(t *testing.T) {
t.Parallel()
diff --git a/backend/internal/service/openai_alpha_search_billing_test.go b/backend/internal/service/openai_alpha_search_billing_test.go
index 1251ee43f9..7151725763 100644
--- a/backend/internal/service/openai_alpha_search_billing_test.go
+++ b/backend/internal/service/openai_alpha_search_billing_test.go
@@ -50,7 +50,7 @@ func TestCalculateOpenAIRecordUsageCostWebSearchPerCall(t *testing.T) {
// 即使 token 倍率(含高峰,3.0)更高也不采用。
apiKey := &APIKey{ID: 1, GroupID: &groupID, Group: &Group{ID: groupID, Platform: PlatformOpenAI}}
result := &OpenAIForwardResult{Model: "gpt-5.6-sol", UpstreamModel: "gpt-5.6-sol", WebSearchCalls: 1}
- cost, err := svc.calculateOpenAIRecordUsageCost(context.Background(), result, apiKey, []string{"gpt-5.6-sol"}, 3.0, 1.0, 1.0, 2.0, UsageTokens{}, "")
+ cost, err := svc.calculateOpenAIRecordUsageCost(context.Background(), result, apiKey, []string{"gpt-5.6-sol"}, 3.0, 1.0, 1.0, 2.0, UsageTokens{}, "", false)
require.NoError(t, err)
require.Equal(t, string(BillingModePerRequest), cost.BillingMode)
require.InDelta(t, 0.01, cost.TotalCost, 1e-12)
@@ -58,7 +58,7 @@ func TestCalculateOpenAIRecordUsageCostWebSearchPerCall(t *testing.T) {
// 分组配置单价 0.005
apiKey.Group.WebSearchPricePerCall = float64Ptr(0.005)
- cost, err = svc.calculateOpenAIRecordUsageCost(context.Background(), result, apiKey, []string{"gpt-5.6-sol"}, 1.0, 1.0, 1.0, 1.0, UsageTokens{}, "")
+ cost, err = svc.calculateOpenAIRecordUsageCost(context.Background(), result, apiKey, []string{"gpt-5.6-sol"}, 1.0, 1.0, 1.0, 1.0, UsageTokens{}, "", false)
require.NoError(t, err)
require.InDelta(t, 0.005, cost.TotalCost, 1e-12)
require.InDelta(t, 0.005, cost.ActualCost, 1e-12)
@@ -66,7 +66,7 @@ func TestCalculateOpenAIRecordUsageCostWebSearchPerCall(t *testing.T) {
// WebSearchCalls = 0 时不得走按次分支(无定价数据会返回 pricing 错误,
// 证明回落到了 token 路径而不是被按次分支吞掉)。
result.WebSearchCalls = 0
- _, err = svc.calculateOpenAIRecordUsageCost(context.Background(), result, apiKey, []string{"gpt-5.6-sol"}, 1.0, 1.0, 1.0, 1.0, UsageTokens{InputTokens: 10}, "")
+ _, err = svc.calculateOpenAIRecordUsageCost(context.Background(), result, apiKey, []string{"gpt-5.6-sol"}, 1.0, 1.0, 1.0, 1.0, UsageTokens{InputTokens: 10}, "", false)
require.Error(t, err)
}
diff --git a/backend/internal/service/openai_codex_models_service.go b/backend/internal/service/openai_codex_models_service.go
index 8a919fa2b0..a7331f64c9 100644
--- a/backend/internal/service/openai_codex_models_service.go
+++ b/backend/internal/service/openai_codex_models_service.go
@@ -2,21 +2,36 @@ package service
import (
"context"
+ "crypto/sha256"
+ "errors"
+ "fmt"
"io"
+ "net"
"net/http"
"net/url"
+ "sort"
"strings"
+ "sync"
"time"
infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors"
"github.com/Wei-Shaw/sub2api/internal/pkg/httpclient"
+ "golang.org/x/net/http2"
+ "golang.org/x/sync/singleflight"
)
// chatgptCodexModelsURL is the ChatGPT Codex models manifest endpoint.
// Package-level variable so tests can point it at a stub server.
var chatgptCodexModelsURL = "https://chatgpt.com/backend-api/codex/models"
-const codexModelsManifestBodyLimit int64 = 8 << 20
+const (
+ codexModelsManifestBodyLimit int64 = 8 << 20
+ codexModelsManifestCacheBodyLimit = 1 << 20
+ codexModelsManifestCacheMaxEntries = 64
+ codexModelsManifestCacheTTL = 30 * time.Second
+ codexModelsManifestCacheStaleTTL = 5 * time.Minute
+ codexModelsManifestRequestTimeout = 15 * time.Second
+)
// CodexModelsManifest carries the raw upstream manifest payload plus caching
// metadata so handlers can pass both through to the client untouched.
@@ -26,8 +41,180 @@ type CodexModelsManifest struct {
NotModified bool
}
-// FetchCodexModelsManifest fetches the live Codex models manifest from the
-// ChatGPT backend using the account's OAuth credentials.
+type codexModelsManifestUpstreamError struct {
+ err error
+ retryable bool
+}
+
+func (e *codexModelsManifestUpstreamError) Error() string { return e.err.Error() }
+
+func (e *codexModelsManifestUpstreamError) Unwrap() error { return e.err }
+
+// IsRetryableCodexModelsManifestError reports whether another selected account
+// may succeed without changing the request. Configuration and upstream 4xx
+// responses, except 429, are intentionally not retried.
+func IsRetryableCodexModelsManifestError(err error) bool {
+ var upstreamErr *codexModelsManifestUpstreamError
+ return errors.As(err, &upstreamErr) && upstreamErr.retryable
+}
+
+func isRetryableCodexModelsManifestTransportError(err error) bool {
+ if err == nil || errors.Is(err, context.Canceled) {
+ return false
+ }
+ if errors.Is(err, context.DeadlineExceeded) ||
+ errors.Is(err, io.EOF) ||
+ errors.Is(err, io.ErrUnexpectedEOF) ||
+ errors.Is(err, net.ErrClosed) {
+ return true
+ }
+
+ var opErr *net.OpError
+ if errors.As(err, &opErr) {
+ return true
+ }
+ var dnsErr *net.DNSError
+ if errors.As(err, &dnsErr) {
+ return true
+ }
+ var goAwayErr http2.GoAwayError
+ if errors.As(err, &goAwayErr) {
+ return true
+ }
+ var streamErr http2.StreamError
+ if errors.As(err, &streamErr) {
+ return true
+ }
+ var connectionErr http2.ConnectionError
+ if errors.As(err, &connectionErr) {
+ return true
+ }
+ var netErr net.Error
+ if errors.As(err, &netErr) && netErr.Timeout() {
+ return true
+ }
+
+ // net/http uses unexported HTTP/2 error types, so typed matching is not
+ // possible for errors produced by the standard library transport.
+ message := strings.ToLower(err.Error())
+ if strings.Contains(message, "http2:") &&
+ (strings.Contains(message, "goaway") ||
+ strings.Contains(message, "refused_stream") ||
+ strings.Contains(message, "frame too large")) {
+ return true
+ }
+ if strings.Contains(message, "stream error: stream id ") {
+ return true
+ }
+ for _, code := range []http2.ErrCode{
+ http2.ErrCodeNo,
+ http2.ErrCodeProtocol,
+ http2.ErrCodeInternal,
+ http2.ErrCodeFlowControl,
+ http2.ErrCodeSettingsTimeout,
+ http2.ErrCodeStreamClosed,
+ http2.ErrCodeFrameSize,
+ http2.ErrCodeRefusedStream,
+ http2.ErrCodeCancel,
+ http2.ErrCodeCompression,
+ http2.ErrCodeConnect,
+ http2.ErrCodeEnhanceYourCalm,
+ http2.ErrCodeInadequateSecurity,
+ http2.ErrCodeHTTP11Required,
+ } {
+ if strings.Contains(message, "connection error: "+strings.ToLower(code.String())) {
+ return true
+ }
+ }
+ return false
+}
+
+type codexModelsManifestRequest struct {
+ url string
+ headers http.Header
+ proxyURL string
+ accountID int64
+ credentialAccountID int64
+ accountConcurrency int
+ useAPIKeyUpstream bool
+}
+
+type codexModelsManifestCacheEntry struct {
+ manifest *CodexModelsManifest
+ order uint64
+ expiresAt time.Time
+ staleUntil time.Time
+}
+
+type codexModelsManifestCacheState uint8
+
+const (
+ codexModelsManifestCacheMiss codexModelsManifestCacheState = iota
+ codexModelsManifestCacheFresh
+ codexModelsManifestCacheStale
+)
+
+type codexModelsManifestCache struct {
+ mu sync.Mutex
+ entries map[string]codexModelsManifestCacheEntry
+ nextOrder uint64
+ refresh singleflight.Group
+}
+
+func (c *codexModelsManifestCache) get(key string, now time.Time) (*CodexModelsManifest, codexModelsManifestCacheState) {
+ c.mu.Lock()
+ defer c.mu.Unlock()
+ entry, ok := c.entries[key]
+ if !ok {
+ return nil, codexModelsManifestCacheMiss
+ }
+ if !now.Before(entry.staleUntil) {
+ delete(c.entries, key)
+ return nil, codexModelsManifestCacheMiss
+ }
+ if now.Before(entry.expiresAt) {
+ return entry.manifest, codexModelsManifestCacheFresh
+ }
+ return entry.manifest, codexModelsManifestCacheStale
+}
+
+func (c *codexModelsManifestCache) set(key string, manifest *CodexModelsManifest, now time.Time) {
+ if manifest == nil || len(manifest.Body) > codexModelsManifestCacheBodyLimit {
+ return
+ }
+ c.mu.Lock()
+ defer c.mu.Unlock()
+ if c.entries == nil {
+ c.entries = make(map[string]codexModelsManifestCacheEntry)
+ }
+ if _, exists := c.entries[key]; !exists && len(c.entries) >= codexModelsManifestCacheMaxEntries {
+ oldestKey := ""
+ var oldestOrder uint64
+ for candidateKey, entry := range c.entries {
+ if !now.Before(entry.staleUntil) {
+ delete(c.entries, candidateKey)
+ continue
+ }
+ if oldestKey == "" || entry.order < oldestOrder {
+ oldestKey = candidateKey
+ oldestOrder = entry.order
+ }
+ }
+ if len(c.entries) >= codexModelsManifestCacheMaxEntries && oldestKey != "" {
+ delete(c.entries, oldestKey)
+ }
+ }
+ c.nextOrder++
+ c.entries[key] = codexModelsManifestCacheEntry{
+ manifest: manifest,
+ order: c.nextOrder,
+ expiresAt: now.Add(codexModelsManifestCacheTTL),
+ staleUntil: now.Add(codexModelsManifestCacheStaleTTL),
+ }
+}
+
+// FetchCodexModelsManifest fetches the live Codex models manifest from either
+// the ChatGPT backend for OAuth accounts or a custom upstream for API key accounts.
//
// The response body is passed through verbatim: the manifest schema evolves
// with Codex client releases, and interpreting it here would force the gateway
@@ -41,49 +228,171 @@ func (s *OpenAIGatewayService) FetchCodexModelsManifest(ctx context.Context, acc
if err != nil {
return nil, infraerrors.Newf(http.StatusInternalServerError, "OPENAI_CODEX_MODELS_CREDENTIALS_FAILED", "resolve credential account: %v", err)
}
- accessToken := credAccount.GetOpenAIAccessToken()
- if accessToken == "" {
- return nil, infraerrors.New(http.StatusBadGateway, "OPENAI_CODEX_MODELS_TOKEN_MISSING", "account has no Codex backend access token")
- }
clientVersion = strings.TrimSpace(clientVersion)
if clientVersion == "" {
clientVersion = openAICodexProbeVersion
}
- requestURL := chatgptCodexModelsURL + "?client_version=" + url.QueryEscape(clientVersion)
- reqCtx, cancel := context.WithTimeout(ctx, 15*time.Second)
- defer cancel()
- req, err := http.NewRequestWithContext(reqCtx, http.MethodGet, requestURL, nil)
+ requestEndpoint := chatgptCodexModelsURL
+ authToken := ""
+ useAPIKeyUpstream := false
+ appendModelsPath := false
+ switch {
+ case credAccount.IsOpenAIOAuth():
+ authToken = strings.TrimSpace(credAccount.GetOpenAIAccessToken())
+ if authToken == "" {
+ return nil, infraerrors.New(http.StatusBadGateway, "OPENAI_CODEX_MODELS_TOKEN_MISSING", "account has no Codex backend access token")
+ }
+ case credAccount.IsOpenAIApiKey():
+ baseURL := strings.TrimSpace(credAccount.GetCredential("base_url"))
+ if baseURL == "" || isOfficialOpenAIModelsBaseURL(baseURL) {
+ return nil, infraerrors.New(
+ http.StatusBadGateway,
+ "OPENAI_CODEX_MODELS_API_KEY_UPSTREAM_UNSUPPORTED",
+ "Codex models manifest requires a custom API key upstream base URL",
+ )
+ }
+ authToken = strings.TrimSpace(credAccount.GetOpenAIApiKey())
+ if authToken == "" {
+ return nil, infraerrors.New(http.StatusBadGateway, "OPENAI_CODEX_MODELS_API_KEY_MISSING", "account has no API key for the Codex models upstream")
+ }
+ normalizedBaseURL, validateErr := s.validateUpstreamBaseURL(baseURL)
+ if validateErr != nil {
+ return nil, infraerrors.Newf(http.StatusBadGateway, "OPENAI_CODEX_MODELS_API_KEY_UPSTREAM_INVALID", "invalid Codex models upstream base URL: %v", validateErr)
+ }
+ requestEndpoint = normalizedBaseURL
+ useAPIKeyUpstream = true
+ appendModelsPath = true
+ default:
+ return nil, infraerrors.Newf(http.StatusBadGateway, "OPENAI_CODEX_MODELS_ACCOUNT_TYPE_UNSUPPORTED", "account type %q cannot fetch the Codex models manifest", credAccount.Type)
+ }
+
+ requestURL, err := buildCodexModelsManifestURL(requestEndpoint, appendModelsPath, clientVersion)
if err != nil {
- return nil, infraerrors.Newf(http.StatusInternalServerError, "OPENAI_CODEX_MODELS_REQUEST_FAILED", "create codex models request: %v", err)
+ if useAPIKeyUpstream {
+ return nil, infraerrors.Newf(http.StatusBadGateway, "OPENAI_CODEX_MODELS_API_KEY_UPSTREAM_INVALID", "invalid Codex models upstream base URL: %v", err)
+ }
+ return nil, infraerrors.Newf(http.StatusInternalServerError, "OPENAI_CODEX_MODELS_REQUEST_FAILED", "parse codex models request URL: %v", err)
}
- req.Header.Set("Authorization", "Bearer "+accessToken)
- req.Header.Set("Accept", "application/json")
- req.Header.Set("Originator", "codex_cli_rs")
- req.Header.Set("Version", clientVersion)
- req.Header.Set("User-Agent", codexCLIUserAgent)
- if ifNoneMatch = strings.TrimSpace(ifNoneMatch); ifNoneMatch != "" {
- req.Header.Set("If-None-Match", ifNoneMatch)
+
+ headers := make(http.Header)
+ headers.Set("Authorization", "Bearer "+authToken)
+ headers.Set("Accept", "application/json")
+ headers.Set("Originator", "codex_cli_rs")
+ headers.Set("Version", clientVersion)
+ headers.Set("User-Agent", codexCLIUserAgent)
+ if useAPIKeyUpstream {
+ credAccount.ApplyHeaderOverrides(headers)
+ } else {
+ setOpenAIChatGPTAccountHeaders(headers, credAccount)
}
- setOpenAIChatGPTAccountHeaders(req.Header, credAccount)
proxyURL := ""
if account.ProxyID != nil && account.Proxy != nil {
proxyURL = account.Proxy.URL()
}
- client, err := httpclient.GetClient(httpclient.Options{
- ProxyURL: proxyURL,
- Timeout: 15 * time.Second,
- ResponseHeaderTimeout: 10 * time.Second,
+
+ request := codexModelsManifestRequest{
+ url: requestURL.String(),
+ headers: headers,
+ proxyURL: proxyURL,
+ accountID: account.ID,
+ credentialAccountID: credAccount.ID,
+ accountConcurrency: account.Concurrency,
+ useAPIKeyUpstream: useAPIKeyUpstream,
+ }
+ if useAPIKeyUpstream {
+ return s.fetchCachedAPIKeyCodexModelsManifest(ctx, request, ifNoneMatch)
+ }
+ return s.fetchCodexModelsManifestUpstream(ctx, request, ifNoneMatch)
+}
+
+func (s *OpenAIGatewayService) fetchCachedAPIKeyCodexModelsManifest(ctx context.Context, request codexModelsManifestRequest, ifNoneMatch string) (*CodexModelsManifest, error) {
+ if err := ctx.Err(); err != nil {
+ return nil, err
+ }
+ cacheKey := buildCodexModelsManifestCacheKey(request)
+ manifest, state := s.codexModelsManifestCache.get(cacheKey, time.Now())
+ if state == codexModelsManifestCacheFresh {
+ return codexModelsManifestForClient(manifest, ifNoneMatch), nil
+ }
+ resultCh := s.refreshCachedAPIKeyCodexModelsManifest(cacheKey, request)
+ if state == codexModelsManifestCacheStale {
+ return codexModelsManifestForClient(manifest, ifNoneMatch), nil
+ }
+ select {
+ case <-ctx.Done():
+ return nil, ctx.Err()
+ case result := <-resultCh:
+ if result.Err != nil {
+ return nil, result.Err
+ }
+ manifest, ok := result.Val.(*CodexModelsManifest)
+ if !ok || manifest == nil {
+ return nil, infraerrors.New(http.StatusInternalServerError, "OPENAI_CODEX_MODELS_REQUEST_FAILED", "invalid shared Codex models manifest result")
+ }
+ return codexModelsManifestForClient(manifest, ifNoneMatch), nil
+ }
+}
+
+func (s *OpenAIGatewayService) refreshCachedAPIKeyCodexModelsManifest(cacheKey string, request codexModelsManifestRequest) <-chan singleflight.Result {
+ return s.codexModelsManifestCache.refresh.DoChan(cacheKey, func() (any, error) {
+ cached, _ := s.codexModelsManifestCache.get(cacheKey, time.Now())
+ ifNoneMatch := ""
+ if cached != nil {
+ ifNoneMatch = cached.ETag
+ }
+ manifest, err := s.fetchCodexModelsManifestUpstream(context.Background(), request, ifNoneMatch)
+ if err != nil {
+ return nil, err
+ }
+ if manifest.NotModified && cached != nil {
+ s.codexModelsManifestCache.set(cacheKey, cached, time.Now())
+ return cached, nil
+ }
+ if !manifest.NotModified {
+ s.codexModelsManifestCache.set(cacheKey, manifest, time.Now())
+ }
+ return manifest, nil
})
+}
+
+func (s *OpenAIGatewayService) fetchCodexModelsManifestUpstream(ctx context.Context, request codexModelsManifestRequest, ifNoneMatch string) (*CodexModelsManifest, error) {
+ reqCtx, cancel := context.WithTimeout(ctx, codexModelsManifestRequestTimeout)
+ defer cancel()
+ req, err := http.NewRequestWithContext(reqCtx, http.MethodGet, request.url, nil)
if err != nil {
- return nil, infraerrors.Newf(http.StatusInternalServerError, "OPENAI_CODEX_MODELS_PROXY_INVALID", "invalid proxy configuration: %v", err)
+ return nil, infraerrors.Newf(http.StatusInternalServerError, "OPENAI_CODEX_MODELS_REQUEST_FAILED", "create codex models request: %v", err)
+ }
+ req.Header = request.headers.Clone()
+ if ifNoneMatch = strings.TrimSpace(ifNoneMatch); ifNoneMatch != "" {
+ req.Header.Set("If-None-Match", ifNoneMatch)
}
- resp, err := client.Do(req)
+ var resp *http.Response
+ if request.useAPIKeyUpstream {
+ if s.httpUpstream == nil {
+ return nil, infraerrors.New(http.StatusInternalServerError, "OPENAI_CODEX_MODELS_UPSTREAM_NOT_CONFIGURED", "Codex models upstream HTTP client is not configured")
+ }
+ req = req.WithContext(WithHTTPUpstreamProfile(req.Context(), HTTPUpstreamProfileOpenAI))
+ resp, err = s.httpUpstream.Do(req, request.proxyURL, request.accountID, request.accountConcurrency)
+ } else {
+ client, clientErr := httpclient.GetClient(httpclient.Options{
+ ProxyURL: request.proxyURL,
+ Timeout: codexModelsManifestRequestTimeout,
+ ResponseHeaderTimeout: 10 * time.Second,
+ })
+ if clientErr != nil {
+ return nil, infraerrors.Newf(http.StatusInternalServerError, "OPENAI_CODEX_MODELS_PROXY_INVALID", "invalid proxy configuration: %v", clientErr)
+ }
+ resp, err = client.Do(req)
+ }
if err != nil {
- return nil, infraerrors.Newf(http.StatusBadGateway, "OPENAI_CODEX_MODELS_UPSTREAM_FAILED", "codex models manifest request failed: %v", err)
+ return nil, &codexModelsManifestUpstreamError{
+ err: infraerrors.Newf(http.StatusBadGateway, "OPENAI_CODEX_MODELS_UPSTREAM_FAILED", "codex models manifest request failed: %v", err),
+ retryable: isRetryableCodexModelsManifestTransportError(err),
+ }
}
defer func() { _ = resp.Body.Close() }()
@@ -96,12 +405,100 @@ func (s *OpenAIGatewayService) FetchCodexModelsManifest(ctx context.Context, acc
if message == "" {
message = resp.Status
}
- return nil, infraerrors.Newf(http.StatusBadGateway, "OPENAI_CODEX_MODELS_UPSTREAM_FAILED", "codex models manifest upstream error %d: %s", resp.StatusCode, message)
+ return nil, &codexModelsManifestUpstreamError{
+ err: infraerrors.Newf(http.StatusBadGateway, "OPENAI_CODEX_MODELS_UPSTREAM_FAILED", "codex models manifest upstream error %d: %s", resp.StatusCode, message),
+ retryable: resp.StatusCode == http.StatusTooManyRequests ||
+ (resp.StatusCode >= http.StatusInternalServerError && resp.StatusCode < 600),
+ }
}
body, err := io.ReadAll(io.LimitReader(resp.Body, codexModelsManifestBodyLimit))
if err != nil {
- return nil, infraerrors.Newf(http.StatusBadGateway, "OPENAI_CODEX_MODELS_UPSTREAM_FAILED", "read codex models manifest response: %v", err)
+ return nil, &codexModelsManifestUpstreamError{
+ err: infraerrors.Newf(http.StatusBadGateway, "OPENAI_CODEX_MODELS_UPSTREAM_FAILED", "read codex models manifest response: %v", err),
+ retryable: isRetryableCodexModelsManifestTransportError(err),
+ }
}
return &CodexModelsManifest{Body: body, ETag: resp.Header.Get("ETag")}, nil
}
+
+func buildCodexModelsManifestCacheKey(request codexModelsManifestRequest) string {
+ hasher := sha256.New()
+ _, _ = fmt.Fprintf(hasher, "%d\n%d\n%s\n%s\n", request.accountID, request.credentialAccountID, request.proxyURL, request.url)
+ headerNames := make([]string, 0, len(request.headers))
+ for name := range request.headers {
+ headerNames = append(headerNames, name)
+ }
+ sort.Strings(headerNames)
+ for _, name := range headerNames {
+ _, _ = fmt.Fprintf(hasher, "%s\n", strings.ToLower(name))
+ for _, value := range request.headers[name] {
+ _, _ = fmt.Fprintf(hasher, "%s\n", value)
+ }
+ }
+ return fmt.Sprintf("%x", hasher.Sum(nil))
+}
+
+func codexModelsManifestForClient(manifest *CodexModelsManifest, ifNoneMatch string) *CodexModelsManifest {
+ if manifest == nil {
+ return nil
+ }
+ if codexModelsManifestETagMatches(ifNoneMatch, manifest.ETag) {
+ return &CodexModelsManifest{ETag: manifest.ETag, NotModified: true}
+ }
+ return manifest
+}
+
+func codexModelsManifestETagMatches(ifNoneMatch, etag string) bool {
+ etag = strings.TrimSpace(etag)
+ if etag == "" {
+ return false
+ }
+ normalize := func(value string) string {
+ value = strings.TrimSpace(value)
+ if len(value) >= 2 && strings.EqualFold(value[:2], "W/") {
+ value = strings.TrimSpace(value[2:])
+ }
+ return value
+ }
+ want := normalize(etag)
+ for _, candidate := range strings.Split(ifNoneMatch, ",") {
+ candidate = strings.TrimSpace(candidate)
+ if candidate == "*" || normalize(candidate) == want {
+ return true
+ }
+ }
+ return false
+}
+
+func isOfficialOpenAIModelsBaseURL(raw string) bool {
+ parsed, err := url.Parse(strings.TrimSpace(raw))
+ if err != nil {
+ return false
+ }
+ hostname := strings.TrimSuffix(parsed.Hostname(), ".")
+ return strings.EqualFold(hostname, "api.openai.com")
+}
+
+func buildCodexModelsManifestURL(endpoint string, appendModelsPath bool, clientVersion string) (*url.URL, error) {
+ requestURL, err := url.Parse(endpoint)
+ if err != nil {
+ return nil, err
+ }
+ if requestURL.Fragment != "" {
+ return nil, fmt.Errorf("URL fragments are not supported")
+ }
+
+ query := requestURL.Query()
+ requestURL.RawQuery = ""
+ requestURL.ForceQuery = false
+ if appendModelsPath {
+ requestURL, err = url.Parse(buildOpenAIModelsURL(requestURL.String()))
+ if err != nil {
+ return nil, err
+ }
+ }
+ query.Set("client_version", clientVersion)
+ requestURL.RawQuery = query.Encode()
+ return requestURL, nil
+}
diff --git a/backend/internal/service/openai_codex_models_service_test.go b/backend/internal/service/openai_codex_models_service_test.go
index c9eae35629..2c87861f60 100644
--- a/backend/internal/service/openai_codex_models_service_test.go
+++ b/backend/internal/service/openai_codex_models_service_test.go
@@ -2,11 +2,146 @@ package service
import (
"context"
+ "errors"
+ "io"
+ "net"
"net/http"
"net/http/httptest"
+ "net/url"
+ "strings"
+ "sync"
+ "sync/atomic"
"testing"
+ "time"
+
+ "github.com/Wei-Shaw/sub2api/internal/config"
+ infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors"
+ "github.com/Wei-Shaw/sub2api/internal/pkg/tlsfingerprint"
+ "golang.org/x/net/http2"
)
+type codexModelsHTTPUpstreamStub struct {
+ do func(req *http.Request, proxyURL string, accountID int64, accountConcurrency int) (*http.Response, error)
+}
+
+type codexModelsBlockingBody struct {
+ ctx context.Context
+ readStarted chan struct{}
+ startedOnce *sync.Once
+ release <-chan struct{}
+ body *strings.Reader
+}
+
+func (b *codexModelsBlockingBody) Read(p []byte) (int, error) {
+ b.startedOnce.Do(func() { close(b.readStarted) })
+ select {
+ case <-b.release:
+ return b.body.Read(p)
+ case <-b.ctx.Done():
+ return 0, b.ctx.Err()
+ }
+}
+
+func (b *codexModelsBlockingBody) Close() error { return nil }
+
+func (s *codexModelsHTTPUpstreamStub) Do(req *http.Request, proxyURL string, accountID int64, accountConcurrency int) (*http.Response, error) {
+ return s.do(req, proxyURL, accountID, accountConcurrency)
+}
+
+func (s *codexModelsHTTPUpstreamStub) DoWithTLS(req *http.Request, proxyURL string, accountID int64, accountConcurrency int, _ *tlsfingerprint.Profile) (*http.Response, error) {
+ return s.Do(req, proxyURL, accountID, accountConcurrency)
+}
+
+func TestIsRetryableCodexModelsManifestTransportError(t *testing.T) {
+ tests := []struct {
+ name string
+ err error
+ retryable bool
+ }{
+ {name: "nil", err: nil},
+ {name: "configuration error", err: errors.New("invalid proxy URL")},
+ {name: "upstream configuration error", err: errors.New("upstream error: invalid proxy")},
+ {name: "proxy connection configuration error", err: errors.New("proxy connection error: invalid configuration")},
+ {name: "canceled request", err: context.Canceled},
+ {
+ name: "redirect policy error",
+ err: &url.Error{
+ Op: "Get",
+ URL: "https://upstream.example/v1/models",
+ Err: errors.New("stopped after 10 redirects"),
+ },
+ },
+ {name: "deadline exceeded", err: context.DeadlineExceeded, retryable: true},
+ {name: "unexpected EOF", err: io.ErrUnexpectedEOF, retryable: true},
+ {name: "closed connection", err: net.ErrClosed, retryable: true},
+ {
+ name: "network operation",
+ err: &net.OpError{
+ Op: "read",
+ Net: "tcp",
+ Err: errors.New("connection reset"),
+ },
+ retryable: true,
+ },
+ {
+ name: "DNS error",
+ err: &net.DNSError{Err: "temporary failure", Name: "upstream.example"},
+ retryable: true,
+ },
+ {
+ name: "typed HTTP2 GOAWAY",
+ err: http2.GoAwayError{ErrCode: http2.ErrCodeNo},
+ retryable: true,
+ },
+ {
+ name: "stdlib HTTP2 GOAWAY",
+ err: errors.New("http2: server sent GOAWAY and closed the connection; LastStreamID=1, ErrCode=NO_ERROR"),
+ retryable: true,
+ },
+ {
+ name: "stdlib HTTP2 refused stream",
+ err: errors.New("stream error: stream ID 3; REFUSED_STREAM"),
+ retryable: true,
+ },
+ {
+ name: "stdlib HTTP2 connection error",
+ err: errors.New(`Get "https://upstream.example/v1/models": connection error: PROTOCOL_ERROR`),
+ retryable: true,
+ },
+ }
+
+ for _, tt := range tests {
+ t.Run(tt.name, func(t *testing.T) {
+ if got := isRetryableCodexModelsManifestTransportError(tt.err); got != tt.retryable {
+ t.Fatalf("retryable = %v, want %v", got, tt.retryable)
+ }
+ })
+ }
+}
+
+func newCodexModelsAPIKeyTestService(upstream HTTPUpstream) *OpenAIGatewayService {
+ return &OpenAIGatewayService{
+ cfg: &config.Config{Security: config.SecurityConfig{URLAllowlist: config.URLAllowlistConfig{
+ Enabled: false,
+ }}},
+ httpUpstream: upstream,
+ }
+}
+
+func newCodexModelsAPIKeyTestAccount(baseURL string) *Account {
+ credentials := map[string]any{"api_key": "sk-upstream"}
+ if baseURL != "" {
+ credentials["base_url"] = baseURL
+ }
+ return &Account{
+ ID: 2,
+ Platform: PlatformOpenAI,
+ Type: AccountTypeAPIKey,
+ Credentials: credentials,
+ Concurrency: 3,
+ }
+}
+
func newCodexModelsTestAccount() *Account {
return &Account{
ID: 1,
@@ -136,3 +271,679 @@ func TestFetchCodexModelsManifestMissingToken(t *testing.T) {
t.Fatal("expected error for missing access token, got nil")
}
}
+
+func TestFetchCodexModelsManifestAPIKeyCustomUpstream(t *testing.T) {
+ manifestBody := `{"models":[{"slug":"gpt-5.6"}]}`
+ var gotRequest *http.Request
+ var gotProxyURL string
+ var gotAccountID int64
+ var gotConcurrency int
+ upstream := &codexModelsHTTPUpstreamStub{do: func(req *http.Request, proxyURL string, accountID int64, accountConcurrency int) (*http.Response, error) {
+ gotRequest = req
+ gotProxyURL = proxyURL
+ gotAccountID = accountID
+ gotConcurrency = accountConcurrency
+ header := make(http.Header)
+ header.Set("ETag", `W/"api-key-manifest"`)
+ return &http.Response{
+ StatusCode: http.StatusOK,
+ Header: header,
+ Body: io.NopCloser(strings.NewReader(manifestBody)),
+ }, nil
+ }}
+
+ s := newCodexModelsAPIKeyTestService(upstream)
+ manifest, err := s.FetchCodexModelsManifest(
+ context.Background(),
+ newCodexModelsAPIKeyTestAccount("https://upstream.example/v1"),
+ "0.144.0",
+ "",
+ )
+ if err != nil {
+ t.Fatalf("FetchCodexModelsManifest returned error: %v", err)
+ }
+
+ if gotRequest == nil {
+ t.Fatal("expected request to custom API key upstream")
+ }
+ if gotRequest.Method != http.MethodGet {
+ t.Errorf("method: got %q", gotRequest.Method)
+ }
+ if gotRequest.URL.String() != "https://upstream.example/v1/models?client_version=0.144.0" {
+ t.Errorf("request URL: got %q", gotRequest.URL.String())
+ }
+ if gotRequest.Header.Get("Authorization") != "Bearer sk-upstream" {
+ t.Errorf("authorization header: got %q", gotRequest.Header.Get("Authorization"))
+ }
+ if gotRequest.Header.Get("Originator") != "codex_cli_rs" {
+ t.Errorf("originator header: got %q", gotRequest.Header.Get("Originator"))
+ }
+ if gotRequest.Header.Get("Version") != "0.144.0" {
+ t.Errorf("version header: got %q", gotRequest.Header.Get("Version"))
+ }
+ if gotRequest.Header.Get("User-Agent") != codexCLIUserAgent {
+ t.Errorf("user-agent header: got %q", gotRequest.Header.Get("User-Agent"))
+ }
+ if gotRequest.Header.Get("chatgpt-account-id") != "" {
+ t.Errorf("chatgpt-account-id must not be sent to API key upstream: got %q", gotRequest.Header.Get("chatgpt-account-id"))
+ }
+ if gotProxyURL != "" || gotAccountID != 2 || gotConcurrency != 3 {
+ t.Errorf("upstream routing metadata: proxy=%q account_id=%d concurrency=%d", gotProxyURL, gotAccountID, gotConcurrency)
+ }
+ if string(manifest.Body) != manifestBody {
+ t.Errorf("body not passed through verbatim: got %q", manifest.Body)
+ }
+ if manifest.ETag != `W/"api-key-manifest"` {
+ t.Errorf("etag not passed through: got %q", manifest.ETag)
+ }
+}
+
+func TestFetchCodexModelsManifestAPIKeySharedRefreshSurvivesCallerCancellation(t *testing.T) {
+ const manifestBody = `{"models":[{"slug":"gpt-5.6"}]}`
+ var calls atomic.Int32
+ var readStartedOnce sync.Once
+ readStarted := make(chan struct{})
+ deadlineRemaining := make(chan time.Duration, 1)
+ release := make(chan struct{})
+ upstream := &codexModelsHTTPUpstreamStub{do: func(req *http.Request, _ string, _ int64, _ int) (*http.Response, error) {
+ calls.Add(1)
+ deadline, ok := req.Context().Deadline()
+ if !ok {
+ deadlineRemaining <- 0
+ } else {
+ deadlineRemaining <- time.Until(deadline)
+ }
+ return &http.Response{
+ StatusCode: http.StatusOK,
+ Header: http.Header{"Etag": []string{`W/"shared"`}},
+ Body: &codexModelsBlockingBody{
+ ctx: req.Context(),
+ readStarted: readStarted,
+ startedOnce: &readStartedOnce,
+ release: release,
+ body: strings.NewReader(manifestBody),
+ },
+ }, nil
+ }}
+
+ s := newCodexModelsAPIKeyTestService(upstream)
+ account := newCodexModelsAPIKeyTestAccount("https://upstream.example")
+ firstCtx, cancelFirst := context.WithCancel(context.Background())
+ firstErr := make(chan error, 1)
+ go func() {
+ _, err := s.FetchCodexModelsManifest(firstCtx, account, "0.144.0", "")
+ firstErr <- err
+ }()
+
+ select {
+ case <-readStarted:
+ case <-time.After(time.Second):
+ t.Fatal("upstream body read did not start")
+ }
+ remaining := <-deadlineRemaining
+ if remaining < 14*time.Second || remaining > codexModelsManifestRequestTimeout {
+ t.Errorf("detached refresh deadline: got %s, want approximately %s", remaining, codexModelsManifestRequestTimeout)
+ }
+ cancelFirst()
+ select {
+ case err := <-firstErr:
+ if !errors.Is(err, context.Canceled) {
+ t.Fatalf("first caller error: got %v, want context.Canceled", err)
+ }
+ case <-time.After(time.Second):
+ t.Fatal("canceled caller did not return promptly")
+ }
+
+ secondResult := make(chan struct {
+ manifest *CodexModelsManifest
+ err error
+ }, 1)
+ go func() {
+ manifest, err := s.FetchCodexModelsManifest(context.Background(), account, "0.144.0", "")
+ secondResult <- struct {
+ manifest *CodexModelsManifest
+ err error
+ }{manifest: manifest, err: err}
+ }()
+
+ time.Sleep(50 * time.Millisecond)
+ if got := calls.Load(); got != 1 {
+ t.Errorf("upstream calls before shared refresh completed: got %d, want 1", got)
+ }
+ close(release)
+ select {
+ case result := <-secondResult:
+ if result.err != nil {
+ t.Fatalf("second caller returned error: %v", result.err)
+ }
+ if string(result.manifest.Body) != manifestBody {
+ t.Errorf("second caller body: got %q", result.manifest.Body)
+ }
+ case <-time.After(time.Second):
+ t.Fatal("second caller did not receive shared refresh result")
+ }
+ if got := calls.Load(); got != 1 {
+ t.Errorf("total upstream calls: got %d, want 1", got)
+ }
+}
+
+func TestFetchCodexModelsManifestAPIKeyConcurrentRequestsShareRefresh(t *testing.T) {
+ const callers = 8
+ var calls atomic.Int32
+ started := make(chan struct{})
+ var startedOnce sync.Once
+ release := make(chan struct{})
+ upstream := &codexModelsHTTPUpstreamStub{do: func(_ *http.Request, _ string, _ int64, _ int) (*http.Response, error) {
+ calls.Add(1)
+ startedOnce.Do(func() { close(started) })
+ <-release
+ return &http.Response{
+ StatusCode: http.StatusOK,
+ Header: make(http.Header),
+ Body: io.NopCloser(strings.NewReader(`{"models":[]}`)),
+ }, nil
+ }}
+
+ s := newCodexModelsAPIKeyTestService(upstream)
+ account := newCodexModelsAPIKeyTestAccount("https://upstream.example")
+ begin := make(chan struct{})
+ errs := make(chan error, callers)
+ for i := 0; i < callers; i++ {
+ go func() {
+ <-begin
+ _, err := s.FetchCodexModelsManifest(context.Background(), account, "0.144.0", "")
+ errs <- err
+ }()
+ }
+ close(begin)
+ select {
+ case <-started:
+ case <-time.After(time.Second):
+ t.Fatal("upstream request did not start")
+ }
+ time.Sleep(50 * time.Millisecond)
+ if got := calls.Load(); got != 1 {
+ t.Errorf("concurrent upstream calls: got %d, want 1", got)
+ }
+ close(release)
+ for i := 0; i < callers; i++ {
+ if err := <-errs; err != nil {
+ t.Errorf("caller %d returned error: %v", i, err)
+ }
+ }
+}
+
+func TestFetchCodexModelsManifestAPIKeyFreshCacheHandlesETagLocally(t *testing.T) {
+ var calls atomic.Int32
+ upstream := &codexModelsHTTPUpstreamStub{do: func(req *http.Request, _ string, _ int64, _ int) (*http.Response, error) {
+ calls.Add(1)
+ if got := req.Header.Get("If-None-Match"); got != "" {
+ t.Errorf("cache refresh must not inherit a caller's If-None-Match: got %q", got)
+ }
+ return &http.Response{
+ StatusCode: http.StatusOK,
+ Header: http.Header{"Etag": []string{`W/"cached"`}},
+ Body: io.NopCloser(strings.NewReader(`{"models":[]}`)),
+ }, nil
+ }}
+
+ s := newCodexModelsAPIKeyTestService(upstream)
+ account := newCodexModelsAPIKeyTestAccount("https://upstream.example")
+ if _, err := s.FetchCodexModelsManifest(context.Background(), account, "0.144.0", ""); err != nil {
+ t.Fatalf("initial fetch returned error: %v", err)
+ }
+ manifest, err := s.FetchCodexModelsManifest(context.Background(), account, "0.144.0", `W/"cached"`)
+ if err != nil {
+ t.Fatalf("cached fetch returned error: %v", err)
+ }
+ if !manifest.NotModified {
+ t.Fatal("matching cached ETag must return NotModified")
+ }
+ if got := calls.Load(); got != 1 {
+ t.Errorf("upstream calls: got %d, want 1", got)
+ }
+}
+
+func TestFetchCodexModelsManifestAPIKeyCacheKeyIsolatesRequestIdentity(t *testing.T) {
+ var calls atomic.Int32
+ upstream := &codexModelsHTTPUpstreamStub{do: func(_ *http.Request, _ string, _ int64, _ int) (*http.Response, error) {
+ calls.Add(1)
+ return &http.Response{
+ StatusCode: http.StatusOK,
+ Header: make(http.Header),
+ Body: io.NopCloser(strings.NewReader(`{"models":[]}`)),
+ }, nil
+ }}
+ s := newCodexModelsAPIKeyTestService(upstream)
+
+ base := newCodexModelsAPIKeyTestAccount("https://upstream.example")
+ fetch := func(account *Account, version string) {
+ t.Helper()
+ if _, err := s.FetchCodexModelsManifest(context.Background(), account, version, ""); err != nil {
+ t.Fatalf("fetch returned error: %v", err)
+ }
+ }
+ fetch(base, "0.144.0")
+ fetch(base, "0.144.0")
+
+ differentAccount := newCodexModelsAPIKeyTestAccount("https://upstream.example")
+ differentAccount.ID = 3
+ fetch(differentAccount, "0.144.0")
+
+ differentToken := newCodexModelsAPIKeyTestAccount("https://upstream.example")
+ differentToken.Credentials["api_key"] = "sk-other"
+ fetch(differentToken, "0.144.0")
+
+ differentUpstream := newCodexModelsAPIKeyTestAccount("https://other-upstream.example")
+ fetch(differentUpstream, "0.144.0")
+ fetch(base, "0.145.0")
+
+ differentHeaders := newCodexModelsAPIKeyTestAccount("https://upstream.example")
+ differentHeaders.Credentials[credKeyHeaderOverrideEnabled] = true
+ differentHeaders.Credentials[credKeyHeaderOverrides] = map[string]any{"x-tenant": "other"}
+ fetch(differentHeaders, "0.144.0")
+
+ proxyID := int64(9)
+ differentProxy := newCodexModelsAPIKeyTestAccount("https://upstream.example")
+ differentProxy.ProxyID = &proxyID
+ differentProxy.Proxy = &Proxy{Protocol: "http", Host: "127.0.0.1", Port: 8080}
+ fetch(differentProxy, "0.144.0")
+ fetch(differentProxy, "0.144.0")
+
+ if got := calls.Load(); got != 7 {
+ t.Errorf("isolated upstream calls: got %d, want 7", got)
+ }
+}
+
+func TestFetchCodexModelsManifestAPIKeyCacheBoundsEntriesAndBodySize(t *testing.T) {
+ var calls atomic.Int32
+ upstream := &codexModelsHTTPUpstreamStub{do: func(req *http.Request, _ string, _ int64, _ int) (*http.Response, error) {
+ calls.Add(1)
+ body := `{"models":[]}`
+ if strings.Contains(req.URL.Host, "large") {
+ body = strings.Repeat("x", (1<<20)+1)
+ }
+ return &http.Response{
+ StatusCode: http.StatusOK,
+ Header: make(http.Header),
+ Body: io.NopCloser(strings.NewReader(body)),
+ }, nil
+ }}
+ s := newCodexModelsAPIKeyTestService(upstream)
+ fetch := func(account *Account) {
+ t.Helper()
+ if _, err := s.FetchCodexModelsManifest(context.Background(), account, "0.144.0", ""); err != nil {
+ t.Fatalf("fetch returned error: %v", err)
+ }
+ }
+
+ small := newCodexModelsAPIKeyTestAccount("https://small.example")
+ fetch(small)
+ fetch(small)
+ large := newCodexModelsAPIKeyTestAccount("https://large.example")
+ large.ID = 3
+ fetch(large)
+ fetch(large)
+ if got := calls.Load(); got != 3 {
+ t.Fatalf("body-size bounded cache calls: got %d, want 3", got)
+ }
+
+ for i := int64(10); i < 75; i++ {
+ account := newCodexModelsAPIKeyTestAccount("https://bounded.example")
+ account.ID = i
+ fetch(account)
+ }
+ last := newCodexModelsAPIKeyTestAccount("https://bounded.example")
+ last.ID = 74
+ fetch(last)
+ if got := calls.Load(); got != 68 {
+ t.Fatalf("most recent cache entry was not retained: calls=%d, want 68", got)
+ }
+ first := newCodexModelsAPIKeyTestAccount("https://bounded.example")
+ first.ID = 10
+ fetch(first)
+ if got := calls.Load(); got != 69 {
+ t.Errorf("oldest cache entry was not evicted: calls=%d, want 69", got)
+ }
+}
+
+func TestFetchCodexModelsManifestAPIKeyServesStaleWhileRefreshing(t *testing.T) {
+ var calls atomic.Int32
+ refreshStarted := make(chan struct{})
+ releaseRefresh := make(chan struct{})
+ upstream := &codexModelsHTTPUpstreamStub{do: func(_ *http.Request, _ string, _ int64, _ int) (*http.Response, error) {
+ call := calls.Add(1)
+ body := `{"models":[{"slug":"old"}]}`
+ if call > 1 {
+ if call == 2 {
+ close(refreshStarted)
+ }
+ <-releaseRefresh
+ body = `{"models":[{"slug":"new"}]}`
+ }
+ return &http.Response{
+ StatusCode: http.StatusOK,
+ Header: make(http.Header),
+ Body: io.NopCloser(strings.NewReader(body)),
+ }, nil
+ }}
+ s := newCodexModelsAPIKeyTestService(upstream)
+ account := newCodexModelsAPIKeyTestAccount("https://upstream.example")
+ if _, err := s.FetchCodexModelsManifest(context.Background(), account, "0.144.0", ""); err != nil {
+ t.Fatalf("initial fetch returned error: %v", err)
+ }
+
+ s.codexModelsManifestCache.mu.Lock()
+ for key, entry := range s.codexModelsManifestCache.entries {
+ entry.expiresAt = time.Now().Add(-time.Second)
+ s.codexModelsManifestCache.entries[key] = entry
+ }
+ s.codexModelsManifestCache.mu.Unlock()
+
+ resultCh := make(chan struct {
+ manifest *CodexModelsManifest
+ err error
+ }, 1)
+ go func() {
+ manifest, err := s.FetchCodexModelsManifest(context.Background(), account, "0.144.0", "")
+ resultCh <- struct {
+ manifest *CodexModelsManifest
+ err error
+ }{manifest: manifest, err: err}
+ }()
+ select {
+ case <-refreshStarted:
+ case <-time.After(time.Second):
+ t.Fatal("background refresh did not start")
+ }
+
+ var staleResult struct {
+ manifest *CodexModelsManifest
+ err error
+ }
+ select {
+ case staleResult = <-resultCh:
+ case <-time.After(100 * time.Millisecond):
+ t.Error("stale manifest was not returned while refresh was blocked")
+ close(releaseRefresh)
+ staleResult = <-resultCh
+ }
+ if staleResult.err != nil {
+ t.Fatalf("stale fetch returned error: %v", staleResult.err)
+ }
+ if got := string(staleResult.manifest.Body); got != `{"models":[{"slug":"old"}]}` {
+ t.Errorf("stale body: got %q", got)
+ }
+ if got := calls.Load(); got != 2 {
+ t.Errorf("upstream calls during stale refresh: got %d, want 2", got)
+ }
+
+ select {
+ case <-releaseRefresh:
+ default:
+ close(releaseRefresh)
+ }
+ deadline := time.Now().Add(time.Second)
+ for {
+ manifest, err := s.FetchCodexModelsManifest(context.Background(), account, "0.144.0", "")
+ if err == nil && string(manifest.Body) == `{"models":[{"slug":"new"}]}` {
+ break
+ }
+ if time.Now().After(deadline) {
+ t.Fatalf("refreshed manifest was not cached: manifest=%v err=%v", manifest, err)
+ }
+ time.Sleep(10 * time.Millisecond)
+ }
+ if got := calls.Load(); got != 2 {
+ t.Errorf("stale refresh was not deduplicated: calls=%d, want 2", got)
+ }
+}
+
+func TestFetchCodexModelsManifestAPIKeyRevalidatesStaleETag(t *testing.T) {
+ var calls atomic.Int32
+ refreshDone := make(chan struct{})
+ upstream := &codexModelsHTTPUpstreamStub{do: func(req *http.Request, _ string, _ int64, _ int) (*http.Response, error) {
+ call := calls.Add(1)
+ if call == 1 {
+ header := make(http.Header)
+ header.Set("ETag", `W/"cached"`)
+ return &http.Response{
+ StatusCode: http.StatusOK,
+ Header: header,
+ Body: io.NopCloser(strings.NewReader(`{"models":[{"slug":"cached"}]}`)),
+ }, nil
+ }
+ if got := req.Header.Get("If-None-Match"); got != `W/"cached"` {
+ t.Errorf("background revalidation If-None-Match: got %q", got)
+ }
+ close(refreshDone)
+ header := make(http.Header)
+ header.Set("ETag", `W/"cached"`)
+ return &http.Response{StatusCode: http.StatusNotModified, Header: header, Body: http.NoBody}, nil
+ }}
+ s := newCodexModelsAPIKeyTestService(upstream)
+ account := newCodexModelsAPIKeyTestAccount("https://upstream.example")
+ if _, err := s.FetchCodexModelsManifest(context.Background(), account, "0.144.0", ""); err != nil {
+ t.Fatalf("initial fetch returned error: %v", err)
+ }
+ s.codexModelsManifestCache.mu.Lock()
+ for key, entry := range s.codexModelsManifestCache.entries {
+ entry.expiresAt = time.Now().Add(-time.Second)
+ s.codexModelsManifestCache.entries[key] = entry
+ }
+ s.codexModelsManifestCache.mu.Unlock()
+
+ manifest, err := s.FetchCodexModelsManifest(context.Background(), account, "0.144.0", "")
+ if err != nil {
+ t.Fatalf("stale fetch returned error: %v", err)
+ }
+ if got := string(manifest.Body); got != `{"models":[{"slug":"cached"}]}` {
+ t.Fatalf("stale body: got %q", got)
+ }
+ select {
+ case <-refreshDone:
+ case <-time.After(time.Second):
+ t.Fatal("ETag revalidation did not complete")
+ }
+
+ deadline := time.Now().Add(time.Second)
+ for {
+ s.codexModelsManifestCache.mu.Lock()
+ fresh := false
+ for _, entry := range s.codexModelsManifestCache.entries {
+ fresh = time.Now().Before(entry.expiresAt)
+ }
+ s.codexModelsManifestCache.mu.Unlock()
+ if fresh {
+ break
+ }
+ if time.Now().After(deadline) {
+ t.Fatal("304 revalidation did not renew the cached manifest")
+ }
+ time.Sleep(10 * time.Millisecond)
+ }
+ manifest, err = s.FetchCodexModelsManifest(context.Background(), account, "0.144.0", "")
+ if err != nil || string(manifest.Body) != `{"models":[{"slug":"cached"}]}` {
+ t.Fatalf("renewed cached manifest: body=%q err=%v", manifest.Body, err)
+ }
+ if got := calls.Load(); got != 2 {
+ t.Errorf("upstream calls: got %d, want 2", got)
+ }
+}
+
+func TestFetchCodexModelsManifestAPIKeyColdCacheHandlesNotModifiedLocally(t *testing.T) {
+ var gotIfNoneMatch string
+ upstream := &codexModelsHTTPUpstreamStub{do: func(req *http.Request, _ string, _ int64, _ int) (*http.Response, error) {
+ gotIfNoneMatch = req.Header.Get("If-None-Match")
+ header := make(http.Header)
+ header.Set("ETag", `W/"api-key-manifest"`)
+ return &http.Response{
+ StatusCode: http.StatusOK,
+ Header: header,
+ Body: io.NopCloser(strings.NewReader(`{"models":[]}`)),
+ }, nil
+ }}
+
+ s := newCodexModelsAPIKeyTestService(upstream)
+ manifest, err := s.FetchCodexModelsManifest(
+ context.Background(),
+ newCodexModelsAPIKeyTestAccount("https://upstream.example"),
+ "0.144.0",
+ `W/"api-key-manifest"`,
+ )
+ if err != nil {
+ t.Fatalf("FetchCodexModelsManifest returned error: %v", err)
+ }
+ if !manifest.NotModified {
+ t.Error("expected NotModified to be true")
+ }
+ if manifest.ETag != `W/"api-key-manifest"` {
+ t.Errorf("etag not passed through: got %q", manifest.ETag)
+ }
+ if gotIfNoneMatch != "" {
+ t.Errorf("cold shared refresh must not inherit caller if-none-match: got %q", gotIfNoneMatch)
+ }
+}
+
+func TestFetchCodexModelsManifestAPIKeyDoesNotCacheUnexpectedColdNotModified(t *testing.T) {
+ var calls atomic.Int32
+ upstream := &codexModelsHTTPUpstreamStub{do: func(req *http.Request, _ string, _ int64, _ int) (*http.Response, error) {
+ calls.Add(1)
+ if got := req.Header.Get("If-None-Match"); got != "" {
+ t.Errorf("cold shared refresh If-None-Match: got %q", got)
+ }
+ header := make(http.Header)
+ header.Set("ETag", `W/"unexpected"`)
+ return &http.Response{StatusCode: http.StatusNotModified, Header: header, Body: http.NoBody}, nil
+ }}
+ s := newCodexModelsAPIKeyTestService(upstream)
+ account := newCodexModelsAPIKeyTestAccount("https://upstream.example")
+ for i := 0; i < 2; i++ {
+ manifest, err := s.FetchCodexModelsManifest(context.Background(), account, "0.144.0", "")
+ if err != nil {
+ t.Fatalf("fetch %d returned error: %v", i, err)
+ }
+ if !manifest.NotModified {
+ t.Fatalf("fetch %d: expected upstream NotModified response", i)
+ }
+ }
+ if got := calls.Load(); got != 2 {
+ t.Errorf("unexpected cold 304 was cached: upstream calls=%d, want 2", got)
+ }
+}
+
+func TestFetchCodexModelsManifestAPIKeyPreservesBaseURLQuery(t *testing.T) {
+ var gotURL string
+ upstream := &codexModelsHTTPUpstreamStub{do: func(req *http.Request, _ string, _ int64, _ int) (*http.Response, error) {
+ gotURL = req.URL.String()
+ return &http.Response{
+ StatusCode: http.StatusOK,
+ Header: make(http.Header),
+ Body: io.NopCloser(strings.NewReader(`{"models":[]}`)),
+ }, nil
+ }}
+
+ s := newCodexModelsAPIKeyTestService(upstream)
+ _, err := s.FetchCodexModelsManifest(
+ context.Background(),
+ newCodexModelsAPIKeyTestAccount("https://upstream.example/v1?tenant=acme"),
+ "0.144.0",
+ "",
+ )
+ if err != nil {
+ t.Fatalf("FetchCodexModelsManifest returned error: %v", err)
+ }
+ if gotURL != "https://upstream.example/v1/models?client_version=0.144.0&tenant=acme" {
+ t.Errorf("request URL: got %q", gotURL)
+ }
+}
+
+func TestFetchCodexModelsManifestAPIKeyRejectsBaseURLFragment(t *testing.T) {
+ called := false
+ upstream := &codexModelsHTTPUpstreamStub{do: func(_ *http.Request, _ string, _ int64, _ int) (*http.Response, error) {
+ called = true
+ return &http.Response{
+ StatusCode: http.StatusOK,
+ Header: make(http.Header),
+ Body: io.NopCloser(strings.NewReader(`{"models":[]}`)),
+ }, nil
+ }}
+
+ s := newCodexModelsAPIKeyTestService(upstream)
+ _, err := s.FetchCodexModelsManifest(
+ context.Background(),
+ newCodexModelsAPIKeyTestAccount("https://upstream.example/v1#models"),
+ "0.144.0",
+ "",
+ )
+ if err == nil {
+ t.Fatal("expected invalid upstream base URL error, got nil")
+ }
+ if infraerrors.Reason(err) != "OPENAI_CODEX_MODELS_API_KEY_UPSTREAM_INVALID" {
+ t.Errorf("error reason: got %q", infraerrors.Reason(err))
+ }
+ if called {
+ t.Fatal("fragment-bearing base URL must be rejected before the upstream request")
+ }
+}
+
+func TestFetchCodexModelsManifestAPIKeyUpstreamError(t *testing.T) {
+ upstream := &codexModelsHTTPUpstreamStub{do: func(_ *http.Request, _ string, _ int64, _ int) (*http.Response, error) {
+ return &http.Response{
+ StatusCode: http.StatusTooManyRequests,
+ Status: "429 Too Many Requests",
+ Header: make(http.Header),
+ Body: io.NopCloser(strings.NewReader(`{"error":"rate limited"}`)),
+ }, nil
+ }}
+
+ s := newCodexModelsAPIKeyTestService(upstream)
+ _, err := s.FetchCodexModelsManifest(
+ context.Background(),
+ newCodexModelsAPIKeyTestAccount("https://upstream.example"),
+ "0.144.0",
+ "",
+ )
+ if err == nil {
+ t.Fatal("expected error for upstream 429, got nil")
+ }
+ if infraerrors.Code(err) != http.StatusBadGateway {
+ t.Errorf("error status: got %d, want %d", infraerrors.Code(err), http.StatusBadGateway)
+ }
+ if infraerrors.Reason(err) != "OPENAI_CODEX_MODELS_UPSTREAM_FAILED" {
+ t.Errorf("error reason: got %q", infraerrors.Reason(err))
+ }
+}
+
+func TestFetchCodexModelsManifestAPIKeyRejectsOfficialOpenAIBaseURL(t *testing.T) {
+ tests := []struct {
+ name string
+ baseURL string
+ }{
+ {name: "missing base URL"},
+ {name: "official host", baseURL: "https://api.openai.com"},
+ {name: "official versioned URL", baseURL: "https://API.OPENAI.COM:443/v1/"},
+ }
+
+ for _, tt := range tests {
+ t.Run(tt.name, func(t *testing.T) {
+ s := newCodexModelsAPIKeyTestService(&codexModelsHTTPUpstreamStub{do: func(_ *http.Request, _ string, _ int64, _ int) (*http.Response, error) {
+ t.Fatal("official OpenAI API key must not be used as a Codex manifest upstream")
+ return nil, nil
+ }})
+
+ _, err := s.FetchCodexModelsManifest(
+ context.Background(),
+ newCodexModelsAPIKeyTestAccount(tt.baseURL),
+ "0.144.0",
+ "",
+ )
+ if err == nil {
+ t.Fatal("expected unsupported API key upstream error, got nil")
+ }
+ if infraerrors.Reason(err) != "OPENAI_CODEX_MODELS_API_KEY_UPSTREAM_UNSUPPORTED" {
+ t.Errorf("error reason: got %q", infraerrors.Reason(err))
+ }
+ })
+ }
+}
diff --git a/backend/internal/service/openai_codex_transform.go b/backend/internal/service/openai_codex_transform.go
index 3869e97c99..3bde4abcfb 100644
--- a/backend/internal/service/openai_codex_transform.go
+++ b/backend/internal/service/openai_codex_transform.go
@@ -838,6 +838,9 @@ func ensureOpenAIResponsesImageGenerationTool(reqBody map[string]any) bool {
if isCodexSparkModel(firstNonEmptyString(reqBody["model"])) {
return false
}
+ if hasOpenAIImageGenerationTool(reqBody) {
+ return false
+ }
tool := map[string]any{
"type": "image_generation",
@@ -855,16 +858,6 @@ func ensureOpenAIResponsesImageGenerationTool(reqBody map[string]any) bool {
reqBody["tools"] = []any{tool}
return true
}
- for _, rawTool := range tools {
- toolMap, ok := rawTool.(map[string]any)
- if !ok {
- continue
- }
- if strings.TrimSpace(firstNonEmptyString(toolMap["type"])) == "image_generation" {
- return false
- }
- }
-
reqBody["tools"] = append(tools, tool)
return true
}
diff --git a/backend/internal/service/openai_codex_transform_test.go b/backend/internal/service/openai_codex_transform_test.go
index b226655eeb..456740136b 100644
--- a/backend/internal/service/openai_codex_transform_test.go
+++ b/backend/internal/service/openai_codex_transform_test.go
@@ -617,6 +617,65 @@ func TestEnsureOpenAIResponsesImageGenerationTool_PreservesExistingImageTool(t *
require.Equal(t, "webp", tool["output_format"])
}
+func TestEnsureOpenAIResponsesImageGenerationTool_PreservesImageGenNamespace(t *testing.T) {
+ tests := []struct {
+ name string
+ reqBody map[string]any
+ }{
+ {
+ name: "top-level tools",
+ reqBody: map[string]any{
+ "model": "gpt-5.5",
+ "tools": []any{
+ map[string]any{
+ "type": "namespace",
+ "name": "image_gen",
+ "tools": []any{
+ map[string]any{"type": "function", "name": "imagegen"},
+ },
+ },
+ },
+ },
+ },
+ {
+ name: "responses lite additional_tools",
+ reqBody: map[string]any{
+ "model": "gpt-5.5",
+ "input": []any{
+ map[string]any{
+ "type": "additional_tools",
+ "tools": []any{
+ map[string]any{
+ "type": "namespace",
+ "name": "image_gen",
+ "tools": []any{
+ map[string]any{"type": "function", "name": "imagegen"},
+ },
+ },
+ },
+ },
+ },
+ },
+ },
+ }
+
+ for _, tt := range tests {
+ t.Run(tt.name, func(t *testing.T) {
+ require.True(t, hasOpenAIImageGenerationTool(tt.reqBody))
+
+ modified := ensureOpenAIResponsesImageGenerationTool(tt.reqBody)
+
+ require.False(t, modified)
+ tools, _ := tt.reqBody["tools"].([]any)
+ for _, rawTool := range tools {
+ tool, ok := rawTool.(map[string]any)
+ require.True(t, ok)
+ require.NotEqual(t, "image_generation", firstNonEmptyString(tool["type"]))
+ }
+ })
+ }
+}
+
func TestApplyCodexImageGenerationBridgeInstructions_AppendsBridgeOnce(t *testing.T) {
reqBody := map[string]any{
"model": "gpt-5.4",
diff --git a/backend/internal/service/openai_compat_model_test.go b/backend/internal/service/openai_compat_model_test.go
index 69b6ddbca2..e1007c507a 100644
--- a/backend/internal/service/openai_compat_model_test.go
+++ b/backend/internal/service/openai_compat_model_test.go
@@ -124,6 +124,55 @@ func TestApplyOpenAICompatModelNormalization(t *testing.T) {
})
}
+func TestForwardAsAnthropic_UsesExactFableMessagesDispatchModel(t *testing.T) {
+ t.Parallel()
+ gin.SetMode(gin.TestMode)
+
+ rec := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(rec)
+ body := []byte(`{"model":"claude-fable-5","max_tokens":16,"messages":[{"role":"user","content":"hello"}],"stream":false}`)
+ c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages", bytes.NewReader(body))
+ c.Request.Header.Set("Content-Type", "application/json")
+
+ upstreamBody := strings.Join([]string{
+ `data: {"type":"response.completed","response":{"id":"resp_fable","object":"response","model":"gpt-5.6-sol","status":"completed","output":[{"type":"message","id":"msg_fable","role":"assistant","status":"completed","content":[{"type":"output_text","text":"ok"}]}],"usage":{"input_tokens":5,"output_tokens":2,"total_tokens":7}}}`,
+ "",
+ "data: [DONE]",
+ "",
+ }, "\n")
+ upstream := &httpUpstreamRecorder{resp: &http.Response{
+ StatusCode: http.StatusOK,
+ Header: http.Header{"Content-Type": []string{"text/event-stream"}, "x-request-id": []string{"rid_fable"}},
+ Body: io.NopCloser(strings.NewReader(upstreamBody)),
+ }}
+
+ svc := &OpenAIGatewayService{
+ httpUpstream: upstream,
+ cfg: &config.Config{Security: config.SecurityConfig{URLAllowlist: config.URLAllowlistConfig{Enabled: false}}},
+ }
+ account := &Account{
+ ID: 1,
+ Name: "openai-oauth",
+ Platform: PlatformOpenAI,
+ Type: AccountTypeOAuth,
+ Concurrency: 1,
+ Credentials: map[string]any{
+ "access_token": "oauth-token",
+ "chatgpt_account_id": "chatgpt-acc",
+ },
+ }
+
+ result, err := svc.ForwardAsAnthropic(context.Background(), c, account, body, "", "gpt-5.6-sol")
+ require.NoError(t, err)
+ require.NotNil(t, result)
+ require.Equal(t, "claude-fable-5", result.Model)
+ require.Equal(t, "gpt-5.6-sol", result.BillingModel)
+ require.Equal(t, "gpt-5.6-sol", result.UpstreamModel)
+ require.Equal(t, "gpt-5.6-sol", gjson.GetBytes(upstream.lastBody, "model").String())
+ require.NotContains(t, string(upstream.lastBody), "claude-fable-5")
+ require.Equal(t, "claude-fable-5", gjson.GetBytes(rec.Body.Bytes(), "model").String())
+}
+
func TestForwardAsAnthropic_NormalizesRoutingAndEffortForGpt54XHigh(t *testing.T) {
t.Parallel()
gin.SetMode(gin.TestMode)
diff --git a/backend/internal/service/openai_gateway_chat_completions_test.go b/backend/internal/service/openai_gateway_chat_completions_test.go
index b85ee33947..5186598a70 100644
--- a/backend/internal/service/openai_gateway_chat_completions_test.go
+++ b/backend/internal/service/openai_gateway_chat_completions_test.go
@@ -98,7 +98,7 @@ func TestNormalizeResponsesBodyServiceTier(t *testing.T) {
require.False(t, gjson.GetBytes(body, "service_tier").Exists())
}
-func TestForwardAsChatCompletions_UnknownModelDoesNotUseDefaultMappedModel(t *testing.T) {
+func TestForwardAsChatCompletions_UnknownModelWithoutMessagesDispatchKeepsRequestedModel(t *testing.T) {
gin.SetMode(gin.TestMode)
rec := httptest.NewRecorder()
@@ -129,7 +129,7 @@ func TestForwardAsChatCompletions_UnknownModelDoesNotUseDefaultMappedModel(t *te
},
}
- result, err := svc.ForwardAsChatCompletions(context.Background(), c, account, body, "", "gpt-5.4")
+ result, err := svc.ForwardAsChatCompletions(context.Background(), c, account, body, "", "")
require.Error(t, err)
require.Nil(t, result)
require.Equal(t, "gpt6", gjson.GetBytes(upstream.lastBody, "model").String())
diff --git a/backend/internal/service/openai_gateway_grok.go b/backend/internal/service/openai_gateway_grok.go
index 379b586136..1906b83780 100644
--- a/backend/internal/service/openai_gateway_grok.go
+++ b/backend/internal/service/openai_gateway_grok.go
@@ -23,6 +23,7 @@ const (
grokComposerImageBridgeMaxOutputTokens = 512
grokUpstreamUserAgent = "sub2api-grok/1.0"
grokCLIVersion = "0.2.93"
+ grokDefaultResponsesModel = "grok-4.5"
grokRateLimitFallbackCooldown = 2 * time.Minute
)
@@ -41,7 +42,7 @@ func (s *OpenAIGatewayService) forwardGrokResponses(
upstreamModel := account.GetMappedModel(originalModel)
if strings.TrimSpace(upstreamModel) == "" {
- upstreamModel = "grok-4.3"
+ upstreamModel = grokDefaultResponsesModel
}
cacheIdentity := resolveGrokCacheIdentity(c, body, "", upstreamModel)
patchedBody, err := patchGrokResponsesBody(body, upstreamModel)
diff --git a/backend/internal/service/openai_gateway_grok_test.go b/backend/internal/service/openai_gateway_grok_test.go
index a3edfacf9d..b45e0bccc0 100644
--- a/backend/internal/service/openai_gateway_grok_test.go
+++ b/backend/internal/service/openai_gateway_grok_test.go
@@ -853,12 +853,12 @@ func TestForwardAsChatCompletionsForGrokStopFallsBackToXAIChatCompletions(t *tes
require.Equal(t, http.StatusOK, recorder.Code)
}
-func TestForwardGrokResponsesStreamingUsesXAIResponsesAndSnapshots(t *testing.T) {
+func TestForwardGrokResponsesStreamingDefaultsEmptyModelTo45AndSnapshots(t *testing.T) {
gin.SetMode(gin.TestMode)
recorder := httptest.NewRecorder()
c, _ := gin.CreateTestContext(recorder)
- body := []byte(`{"model":"grok","input":"hi","stream":true,"reasoning_effort":"high"}`)
+ body := []byte(`{"input":"hi","stream":true,"reasoning_effort":"high"}`)
c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", bytes.NewReader(body))
c.Request.Header.Set("Content-Type", "application/json")
c.Request.Header.Set("OpenAI-Beta", "responses=experimental")
@@ -905,7 +905,7 @@ func TestForwardGrokResponsesStreamingUsesXAIResponsesAndSnapshots(t *testing.T)
accountRepo: repo,
}
- result, err := svc.forwardGrokResponses(context.Background(), c, account, body, "grok", true, time.Now())
+ result, err := svc.forwardGrokResponses(context.Background(), c, account, body, "", true, time.Now())
require.NoError(t, err)
require.Equal(t, xai.DefaultCLIBaseURL+"/responses", upstream.lastReq.URL.String())
require.Equal(t, "Bearer access-token", upstream.lastReq.Header.Get("Authorization"))
diff --git a/backend/internal/service/openai_gateway_record_usage_test.go b/backend/internal/service/openai_gateway_record_usage_test.go
index 81578f630f..d2eca05d74 100644
--- a/backend/internal/service/openai_gateway_record_usage_test.go
+++ b/backend/internal/service/openai_gateway_record_usage_test.go
@@ -39,6 +39,17 @@ type openAIRecordUsageBillingRepoStub struct {
lastCtxErr error
}
+type openAIRecordUsageAccountRepoStub struct {
+ AccountRepository
+ account *Account
+ calls int
+}
+
+func (s *openAIRecordUsageAccountRepoStub) GetByID(_ context.Context, _ int64) (*Account, error) {
+ s.calls++
+ return s.account, nil
+}
+
func (s *openAIRecordUsageBillingRepoStub) Apply(ctx context.Context, cmd *UsageBillingCommand) (*UsageBillingApplyResult, error) {
s.calls++
s.lastCmd = cmd
@@ -1045,7 +1056,7 @@ func TestOpenAIGatewayServiceRecordUsage_GPT56SeparatesCacheWriteForBillingAndSt
require.InDelta(t, usageRepo.lastLog.TotalCost*1.1, usageRepo.lastLog.ActualCost, 1e-12)
}
-func TestOpenAIGatewayServiceRecordUsage_Gpt54LongContextBillsWholeSession(t *testing.T) {
+func TestOpenAIGatewayServiceRecordUsage_Gpt54LongContextBillingDisabledByDefault(t *testing.T) {
usageRepo := &openAIRecordUsageLogRepoStub{inserted: true}
userRepo := &openAIRecordUsageUserRepoStub{}
subRepo := &openAIRecordUsageSubRepoStub{}
@@ -1063,7 +1074,45 @@ func TestOpenAIGatewayServiceRecordUsage_Gpt54LongContextBillsWholeSession(t *te
},
APIKey: &APIKey{ID: 1014},
User: &User{ID: 2014},
- Account: &Account{ID: 3014},
+ Account: &Account{ID: 3014, Platform: PlatformOpenAI},
+ })
+
+ require.NoError(t, err)
+ require.NotNil(t, usageRepo.lastLog)
+
+ 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.False(t, usageRepo.lastLog.LongContextBillingApplied)
+ require.Equal(t, 1, userRepo.deductCalls)
+}
+
+func TestOpenAIGatewayServiceRecordUsage_Gpt54LongContextBillingEnabledPerAccount(t *testing.T) {
+ usageRepo := &openAIRecordUsageLogRepoStub{inserted: true}
+ userRepo := &openAIRecordUsageUserRepoStub{}
+ subRepo := &openAIRecordUsageSubRepoStub{}
+ svc := newOpenAIRecordUsageServiceForTest(usageRepo, userRepo, subRepo, nil)
+
+ err := svc.RecordUsage(context.Background(), &OpenAIRecordUsageInput{
+ Result: &OpenAIForwardResult{
+ RequestID: "resp_gpt54_long_context_disabled",
+ Usage: OpenAIUsage{
+ InputTokens: 300000,
+ OutputTokens: 2000,
+ },
+ Model: "gpt-5.4-2026-03-05",
+ Duration: time.Second,
+ },
+ APIKey: &APIKey{ID: 1015},
+ User: &User{ID: 2015},
+ Account: &Account{
+ ID: 3015,
+ Platform: PlatformOpenAI,
+ Extra: map[string]any{"openai_long_context_billing_enabled": true},
+ },
})
require.NoError(t, err)
@@ -1075,7 +1124,62 @@ func TestOpenAIGatewayServiceRecordUsage_Gpt54LongContextBillsWholeSession(t *te
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.Equal(t, 1, userRepo.deductCalls)
+ require.True(t, usageRepo.lastLog.LongContextBillingApplied)
+}
+
+func TestOpenAIGatewayServiceRecordUsage_SparkShadowUsesCurrentParentBillingSetting(t *testing.T) {
+ tests := []struct {
+ name string
+ parentEnabled bool
+ }{
+ {name: "parent opt out overrides stale enabled shadow", parentEnabled: false},
+ {name: "parent opt in overrides stale disabled shadow", parentEnabled: true},
+ }
+
+ for _, tt := range tests {
+ t.Run(tt.name, func(t *testing.T) {
+ usageRepo := &openAIRecordUsageLogRepoStub{inserted: true}
+ accountRepo := &openAIRecordUsageAccountRepoStub{account: &Account{
+ ID: 4016,
+ Platform: PlatformOpenAI,
+ Type: AccountTypeOAuth,
+ Extra: map[string]any{openAILongContextBillingEnabledKey: tt.parentEnabled},
+ }}
+ svc := newOpenAIRecordUsageServiceForTest(
+ usageRepo,
+ &openAIRecordUsageUserRepoStub{},
+ &openAIRecordUsageSubRepoStub{},
+ nil,
+ )
+ svc.accountRepo = accountRepo
+ parentID := int64(4016)
+
+ err := svc.RecordUsage(context.Background(), &OpenAIRecordUsageInput{
+ Result: &OpenAIForwardResult{
+ RequestID: "resp_gpt54_shadow_parent_setting",
+ Usage: OpenAIUsage{InputTokens: 300000, OutputTokens: 2000},
+ Model: "gpt-5.4-2026-03-05",
+ Duration: time.Second,
+ },
+ APIKey: &APIKey{ID: 1016},
+ User: &User{ID: 2016},
+ Account: &Account{
+ ID: 3016,
+ Platform: PlatformOpenAI,
+ Type: AccountTypeOAuth,
+ ParentAccountID: &parentID,
+ QuotaDimension: QuotaDimensionSpark,
+ Extra: map[string]any{
+ openAILongContextBillingEnabledKey: !tt.parentEnabled,
+ },
+ },
+ })
+
+ require.NoError(t, err)
+ require.Equal(t, 1, accountRepo.calls)
+ require.Equal(t, tt.parentEnabled, usageRepo.lastLog.LongContextBillingApplied)
+ })
+ }
}
func TestOpenAIGatewayServiceRecordUsage_ServiceTierPriorityUsesFastPricing(t *testing.T) {
diff --git a/backend/internal/service/openai_gateway_request_body.go b/backend/internal/service/openai_gateway_request_body.go
index 935a32f58b..d56b11fdcc 100644
--- a/backend/internal/service/openai_gateway_request_body.go
+++ b/backend/internal/service/openai_gateway_request_body.go
@@ -365,15 +365,57 @@ func newOpenAIRequestView(body []byte) openAIRequestView {
if len(body) == 0 {
return openAIRequestView{}
}
- return openAIRequestView{
- body: body,
- Model: strings.TrimSpace(gjson.GetBytes(body, "model").String()),
- Stream: gjson.GetBytes(body, "stream").Bool(),
- PromptCacheKey: strings.TrimSpace(gjson.GetBytes(body, "prompt_cache_key").String()),
- PreviousResponseID: strings.TrimSpace(gjson.GetBytes(body, "previous_response_id").String()),
- ServiceTier: strings.TrimSpace(gjson.GetBytes(body, "service_tier").String()),
- ReasoningEffort: strings.TrimSpace(gjson.GetBytes(body, "reasoning.effort").String()),
- }
+
+ const (
+ modelField uint8 = 1 << iota
+ streamField
+ promptCacheKeyField
+ previousResponseIDField
+ serviceTierField
+ reasoningField
+ allRequestViewFields = modelField | streamField | promptCacheKeyField |
+ previousResponseIDField | serviceTierField | reasoningField
+ )
+
+ view := openAIRequestView{body: body}
+ var seen uint8
+ // parseRawJSONView reads body without copying; view keeps body alive for extracted strings.
+ parseRawJSONView(body).ForEach(func(key, value gjson.Result) bool {
+ switch key.Str {
+ case "model":
+ if seen&modelField == 0 {
+ view.Model = strings.TrimSpace(value.String())
+ seen |= modelField
+ }
+ case "stream":
+ if seen&streamField == 0 {
+ view.Stream = value.Bool()
+ seen |= streamField
+ }
+ case "prompt_cache_key":
+ if seen&promptCacheKeyField == 0 {
+ view.PromptCacheKey = strings.TrimSpace(value.String())
+ seen |= promptCacheKeyField
+ }
+ case "previous_response_id":
+ if seen&previousResponseIDField == 0 {
+ view.PreviousResponseID = strings.TrimSpace(value.String())
+ seen |= previousResponseIDField
+ }
+ case "service_tier":
+ if seen&serviceTierField == 0 {
+ view.ServiceTier = strings.TrimSpace(value.String())
+ seen |= serviceTierField
+ }
+ case "reasoning":
+ if seen&reasoningField == 0 {
+ view.ReasoningEffort = strings.TrimSpace(value.Get("effort").String())
+ seen |= reasoningField
+ }
+ }
+ return seen != allRequestViewFields
+ })
+ return view
}
// Decode 保留阶段一既有 full-map 行为;后续阶段会把调用点下沉到复杂分支。
diff --git a/backend/internal/service/openai_gateway_service.go b/backend/internal/service/openai_gateway_service.go
index 29c7d968a2..6817aa7e98 100644
--- a/backend/internal/service/openai_gateway_service.go
+++ b/backend/internal/service/openai_gateway_service.go
@@ -404,6 +404,7 @@ type OpenAIGatewayService struct {
openaiWSRetryMetrics openAIWSRetryMetrics
responseHeaderFilter *responseheaders.CompiledHeaderFilter
codexSnapshotThrottle *accountWriteThrottle
+ codexModelsManifestCache codexModelsManifestCache
openaiCompatSessionResponses sync.Map
openaiCompatAnthropicDigestSessions sync.Map
}
diff --git a/backend/internal/service/openai_gateway_service_hotpath_test.go b/backend/internal/service/openai_gateway_service_hotpath_test.go
index 1dde60c9f0..326fde534d 100644
--- a/backend/internal/service/openai_gateway_service_hotpath_test.go
+++ b/backend/internal/service/openai_gateway_service_hotpath_test.go
@@ -27,6 +27,33 @@ func TestOpenAIRequestView_ExtractsRawScalars(t *testing.T) {
require.Equal(t, "medium", view.ReasoningEffort)
}
+func TestOpenAIRequestView_ExtractsFieldsAfterLargeInput(t *testing.T) {
+ body := []byte(`{"model":"gpt-5","input":[{"content":"` + strings.Repeat("payload", 1024) + `"}],"stream":true,"prompt_cache_key":"session-1","previous_response_id":"resp-1","service_tier":"flex","reasoning":{"effort":"high"}}`)
+
+ view := newOpenAIRequestView(body)
+
+ require.Equal(t, "gpt-5", view.Model)
+ require.True(t, view.Stream)
+ require.Equal(t, "session-1", view.PromptCacheKey)
+ require.Equal(t, "resp-1", view.PreviousResponseID)
+ require.Equal(t, "flex", view.ServiceTier)
+ require.Equal(t, "high", view.ReasoningEffort)
+}
+
+func TestOpenAIRequestView_KeepsFirstDuplicateField(t *testing.T) {
+ view := newOpenAIRequestView([]byte(`{"model":"gpt-5","model":"gpt-5.1","reasoning":{"effort":"low"},"reasoning":{"effort":"high"}}`))
+
+ require.Equal(t, "gpt-5", view.Model)
+ require.Equal(t, "low", view.ReasoningEffort)
+}
+
+func TestOpenAIRequestView_KeepsLenientPrefixExtraction(t *testing.T) {
+ view := newOpenAIRequestView([]byte(`{"model":"gpt-5","stream":true,"input":[`))
+
+ require.Equal(t, "gpt-5", view.Model)
+ require.True(t, view.Stream)
+}
+
func TestOpenAIRequestView_DecodeKeepsFullMapBehavior(t *testing.T) {
view := newOpenAIRequestView([]byte(`{"model":"gpt-5","stream":true,"input":[{"type":"message","content":"hi"}]}`))
diff --git a/backend/internal/service/openai_gateway_usage.go b/backend/internal/service/openai_gateway_usage.go
index 9431b73ecb..410b4d944d 100644
--- a/backend/internal/service/openai_gateway_usage.go
+++ b/backend/internal/service/openai_gateway_usage.go
@@ -178,7 +178,27 @@ func (s *OpenAIGatewayService) RecordUsage(ctx context.Context, input *OpenAIRec
if result.ServiceTier != nil {
serviceTier = strings.TrimSpace(*result.ServiceTier)
}
- cost, err = s.calculateOpenAIRecordUsageCost(ctx, result, apiKey, billingModels, multiplier, imageMultiplier, videoMultiplier, baseMultiplier, tokens, serviceTier)
+ billingAccount := account
+ if account.IsShadow() {
+ billingAccount, err = resolveCredentialAccount(ctx, s.accountRepo, account)
+ if err != nil {
+ return err
+ }
+ }
+ longContextBillingEnabled := billingAccount.IsOpenAILongContextBillingEnabled()
+ cost, err = s.calculateOpenAIRecordUsageCost(
+ ctx,
+ result,
+ apiKey,
+ billingModels,
+ multiplier,
+ imageMultiplier,
+ videoMultiplier,
+ baseMultiplier,
+ tokens,
+ serviceTier,
+ longContextBillingEnabled,
+ )
if err != nil {
if !isUsagePricingUnavailableError(err) {
return err
@@ -257,6 +277,7 @@ func (s *OpenAIGatewayService) RecordUsage(ctx context.Context, input *OpenAIRec
usageLog.CacheReadCost = cost.CacheReadCost
usageLog.TotalCost = cost.TotalCost
usageLog.ActualCost = cost.ActualCost
+ usageLog.LongContextBillingApplied = cost.LongContextBillingApplied
}
if isVideoUsage && (cost == nil || cost.BillingMode != string(BillingModeToken)) {
usageLog.RateMultiplier = videoMultiplier
@@ -366,6 +387,7 @@ func (s *OpenAIGatewayService) calculateOpenAIRecordUsageCost(
webSearchMultiplier float64,
tokens UsageTokens,
serviceTier string,
+ longContextBillingEnabled bool,
) (*CostBreakdown, error) {
billingModel := firstUsageBillingModel(billingModels)
if result != nil && result.WebSearchCalls > 0 {
@@ -395,7 +417,15 @@ func (s *OpenAIGatewayService) calculateOpenAIRecordUsageCost(
if candidate == "" {
continue
}
- cost, err := s.calculateOpenAIRecordUsageTokenCost(ctx, apiKey, candidate, multiplier, tokens, serviceTier)
+ cost, err := s.calculateOpenAIRecordUsageTokenCost(
+ ctx,
+ apiKey,
+ candidate,
+ multiplier,
+ tokens,
+ serviceTier,
+ longContextBillingEnabled,
+ )
if err == nil {
return cost, nil
}
@@ -443,21 +473,29 @@ func (s *OpenAIGatewayService) calculateOpenAIRecordUsageTokenCost(
multiplier float64,
tokens UsageTokens,
serviceTier string,
+ longContextBillingEnabled bool,
) (*CostBreakdown, error) {
if s.resolver != nil && apiKey.Group != nil {
gid := apiKey.Group.ID
return s.billingService.CalculateCostUnified(CostInput{
- Ctx: ctx,
- Model: billingModel,
- GroupID: &gid,
- Tokens: tokens,
- RequestCount: 1,
- RateMultiplier: multiplier,
- ServiceTier: serviceTier,
- Resolver: s.resolver,
+ Ctx: ctx,
+ Model: billingModel,
+ GroupID: &gid,
+ Tokens: tokens,
+ RequestCount: 1,
+ RateMultiplier: multiplier,
+ ServiceTier: serviceTier,
+ Resolver: s.resolver,
+ LongContextBillingEnabled: &longContextBillingEnabled,
})
}
- return s.billingService.CalculateCostWithServiceTier(billingModel, tokens, multiplier, serviceTier)
+ return s.billingService.calculateCostWithServiceTierPolicy(
+ billingModel,
+ tokens,
+ multiplier,
+ serviceTier,
+ longContextBillingEnabled,
+ )
}
func (s *OpenAIGatewayService) calculateOpenAIImageCost(
diff --git a/backend/internal/service/openai_image_generation_controls_test.go b/backend/internal/service/openai_image_generation_controls_test.go
index af0cdf669c..49ba51fcb0 100644
--- a/backend/internal/service/openai_image_generation_controls_test.go
+++ b/backend/internal/service/openai_image_generation_controls_test.go
@@ -283,6 +283,44 @@ func TestOpenAIGatewayServiceForward_ChannelBridgeOverrideEnablesCodexInjection(
require.Contains(t, instructions, "image_generation")
}
+func TestOpenAIGatewayServiceForward_CodexBridgeDoesNotInjectHostedToolAlongsideImageGenNamespace(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+
+ upstream := &httpUpstreamRecorder{
+ resp: &http.Response{
+ StatusCode: http.StatusOK,
+ Header: http.Header{"Content-Type": []string{"application/json"}},
+ Body: io.NopCloser(strings.NewReader(`{"id":"resp_namespace_image","model":"gpt-5.5","usage":{"input_tokens":1,"output_tokens":1}}`)),
+ },
+ }
+ svc := newOpenAIImageGenerationControlTestService(upstream)
+ svc.cfg.Gateway.CodexImageGenerationBridgeEnabled = true
+ c, _ := newOpenAIImageGenerationControlTestContext(true, "codex_cli_rs/0.144.1")
+ account := newOpenAIImageGenerationControlTestAccount()
+ body := []byte(`{
+ "model":"gpt-5.5",
+ "stream":false,
+ "tools":[
+ {"type":"function","name":"shell","parameters":{"type":"object"}},
+ {"type":"namespace","name":"image_gen","tools":[{"type":"function","name":"imagegen"}]}
+ ],
+ "input":[
+ {"type":"message","role":"user","content":[{"type":"input_text","text":"draw a cat"}]},
+ {"type":"additional_tools","tools":[{"type":"namespace","name":"image_gen","tools":[{"type":"function","name":"imagegen"}]}]}
+ ],
+ "tool_choice":"auto"
+ }`)
+
+ result, err := svc.Forward(context.Background(), c, account, body)
+
+ require.NoError(t, err)
+ require.NotNil(t, result)
+ require.NotNil(t, upstream.lastReq)
+ require.False(t, gjson.GetBytes(upstream.lastBody, `tools.#(type=="image_generation")`).Exists())
+ require.Equal(t, "namespace", gjson.GetBytes(upstream.lastBody, `tools.#(name=="image_gen").type`).String())
+ require.Equal(t, "namespace", gjson.GetBytes(upstream.lastBody, `input.#(type=="additional_tools").tools.#(name=="image_gen").type`).String())
+}
+
func TestOpenAIGatewayServiceForward_CodexBridgePreservesExistingToolChoice(t *testing.T) {
gin.SetMode(gin.TestMode)
diff --git a/backend/internal/service/openai_messages_dispatch_test.go b/backend/internal/service/openai_messages_dispatch_test.go
index bafd36449b..db7804a4f3 100644
--- a/backend/internal/service/openai_messages_dispatch_test.go
+++ b/backend/internal/service/openai_messages_dispatch_test.go
@@ -37,3 +37,25 @@ func TestGroupResolveMessagesDispatchModel_GrokMapsClaudeFamilyToGrok(t *testing
require.Empty(t, group.ResolveMessagesDispatchModel("grok"))
require.Empty(t, group.ResolveMessagesDispatchModel("gpt-5.3-codex"))
}
+
+func TestSanitizeGroupMessagesDispatchFields_ClearsNonOpenAIPlatform(t *testing.T) {
+ t.Parallel()
+
+ group := &Group{
+ Platform: PlatformAnthropic,
+ AllowMessagesDispatch: true,
+ DefaultMappedModel: "gpt-5.6-sol",
+ MessagesDispatchModelConfig: OpenAIMessagesDispatchModelConfig{
+ SonnetMappedModel: "gpt-5.3-codex",
+ ExactModelMappings: map[string]string{
+ "claude-fable-5": "gpt-5.6-sol",
+ },
+ },
+ }
+
+ sanitizeGroupMessagesDispatchFields(group)
+
+ require.False(t, group.AllowMessagesDispatch)
+ require.Empty(t, group.DefaultMappedModel)
+ require.Equal(t, OpenAIMessagesDispatchModelConfig{}, group.MessagesDispatchModelConfig)
+}
diff --git a/backend/internal/service/openai_model_mapping.go b/backend/internal/service/openai_model_mapping.go
index cb7a8ca84b..8ba1d6fe1b 100644
--- a/backend/internal/service/openai_model_mapping.go
+++ b/backend/internal/service/openai_model_mapping.go
@@ -3,19 +3,20 @@ package service
import "strings"
// resolveOpenAIForwardModel 解析 OpenAI 兼容转发使用的模型。
-// defaultMappedModel 只服务于 /v1/messages 的 Claude 系列显式调度映射,
-// 不作为普通 OpenAI 请求的未知模型兜底。
-func resolveOpenAIForwardModel(account *Account, requestedModel, defaultMappedModel string) string {
+// messagesDispatchMappedModel 是调用方已为 /v1/messages 解析的显式调度结果;
+// 普通 OpenAI 请求必须传空,避免将分组配置作为通用模型兜底。
+func resolveOpenAIForwardModel(account *Account, requestedModel, messagesDispatchMappedModel string) string {
+ messagesDispatchMappedModel = strings.TrimSpace(messagesDispatchMappedModel)
if account == nil {
- if defaultMappedModel != "" && claudeMessagesDispatchFamily(requestedModel) != "" {
- return defaultMappedModel
+ if messagesDispatchMappedModel != "" {
+ return messagesDispatchMappedModel
}
return requestedModel
}
mappedModel, matched := account.ResolveMappedModel(requestedModel)
- if !matched && defaultMappedModel != "" && claudeMessagesDispatchFamily(requestedModel) != "" {
- return defaultMappedModel
+ if !matched && messagesDispatchMappedModel != "" {
+ return messagesDispatchMappedModel
}
return mappedModel
}
diff --git a/backend/internal/service/openai_model_mapping_test.go b/backend/internal/service/openai_model_mapping_test.go
index f2ceb3551c..7107a706ad 100644
--- a/backend/internal/service/openai_model_mapping_test.go
+++ b/backend/internal/service/openai_model_mapping_test.go
@@ -4,159 +4,156 @@ import "testing"
func TestResolveOpenAIForwardModel(t *testing.T) {
tests := []struct {
- name string
- account *Account
- requestedModel string
- defaultMappedModel string
- expectedModel string
+ name string
+ account *Account
+ requestedModel string
+ messagesDispatchMappedModel string
+ expectedModel string
}{
{
- name: "uses messages dispatch default for claude model",
+ name: "uses messages dispatch model for known claude family",
account: &Account{
Credentials: map[string]any{},
},
- requestedModel: "claude-opus-4-6",
- defaultMappedModel: "gpt-4o-mini",
- expectedModel: "gpt-4o-mini",
+ requestedModel: "claude-opus-4-6",
+ messagesDispatchMappedModel: "gpt-4o-mini",
+ expectedModel: "gpt-4o-mini",
},
{
- name: "does not fall back to group default for invalid gpt model",
+ name: "uses exact messages dispatch model for unknown claude family",
account: &Account{
Credentials: map[string]any{},
},
- requestedModel: "gpt6",
- defaultMappedModel: "gpt-5.4",
- expectedModel: "gpt6",
+ requestedModel: "claude-fable-5",
+ messagesDispatchMappedModel: " gpt-5.6-sol ",
+ expectedModel: "gpt-5.6-sol",
},
{
- name: "preserves explicit gpt-5.4 instead of group default",
+ name: "nil account uses messages dispatch model",
+ requestedModel: "claude-fable-5",
+ messagesDispatchMappedModel: "gpt-5.6-sol",
+ expectedModel: "gpt-5.6-sol",
+ },
+ {
+ name: "nil account without messages dispatch keeps requested model",
+ requestedModel: "claude-fable-5",
+ expectedModel: "claude-fable-5",
+ },
+ {
+ name: "ordinary unknown gpt model has no messages dispatch fallback",
account: &Account{
Credentials: map[string]any{},
},
- requestedModel: "gpt-5.4",
- defaultMappedModel: "gpt-4o-mini",
- expectedModel: "gpt-5.4",
+ requestedModel: "gpt6",
+ expectedModel: "gpt6",
},
{
- name: "preserves exact passthrough mapping instead of group default",
+ name: "account exact mapping overrides messages dispatch model",
account: &Account{
Credentials: map[string]any{
"model_mapping": map[string]any{
- "gpt-5.4": "gpt-5.4",
+ "claude-fable-5": "gpt-5.5",
},
},
},
- requestedModel: "gpt-5.4",
- defaultMappedModel: "gpt-4o-mini",
- expectedModel: "gpt-5.4",
+ requestedModel: "claude-fable-5",
+ messagesDispatchMappedModel: "gpt-5.6-sol",
+ expectedModel: "gpt-5.5",
},
{
- name: "preserves wildcard passthrough mapping instead of group default",
+ name: "account wildcard mapping overrides messages dispatch model",
account: &Account{
Credentials: map[string]any{
"model_mapping": map[string]any{
- "gpt-*": "gpt-5.4",
+ "claude-*": "gpt-5.4",
},
},
},
- requestedModel: "gpt-5.4",
- defaultMappedModel: "gpt-4o-mini",
- expectedModel: "gpt-5.4",
+ requestedModel: "claude-fable-5",
+ messagesDispatchMappedModel: "gpt-5.6-sol",
+ expectedModel: "gpt-5.4",
},
{
- name: "uses account remap when explicit target differs",
+ name: "account passthrough mapping overrides messages dispatch model",
account: &Account{
Credentials: map[string]any{
"model_mapping": map[string]any{
- "gpt-5": "gpt-5.4",
+ "claude-fable-5": "claude-fable-5",
},
},
},
- requestedModel: "gpt-5",
- defaultMappedModel: "gpt-4o-mini",
- expectedModel: "gpt-5.4",
+ requestedModel: "claude-fable-5",
+ messagesDispatchMappedModel: "gpt-5.6-sol",
+ expectedModel: "claude-fable-5",
},
{
- name: "preserves codex spark instead of group default",
+ name: "ordinary codex spark request keeps requested model",
account: &Account{
Credentials: map[string]any{},
},
- requestedModel: "gpt-5.3-codex-spark",
- defaultMappedModel: "gpt-5.4",
- expectedModel: "gpt-5.3-codex-spark",
+ requestedModel: "gpt-5.3-codex-spark",
+ expectedModel: "gpt-5.3-codex-spark",
},
{
- name: "preserves gpt-5.5 instead of group default",
+ name: "ordinary gpt-5.5 request keeps requested model",
account: &Account{
Credentials: map[string]any{},
},
- requestedModel: "gpt-5.5",
- defaultMappedModel: "gpt-5.4",
- expectedModel: "gpt-5.5",
+ requestedModel: "gpt-5.5",
+ expectedModel: "gpt-5.5",
},
{
- name: "preserves gpt-5.5-pro instead of group default",
+ name: "ordinary gpt-5.5-pro request keeps requested model",
account: &Account{
Credentials: map[string]any{},
},
- requestedModel: "gpt-5.5-pro",
- defaultMappedModel: "gpt-5.5",
- expectedModel: "gpt-5.5-pro",
+ requestedModel: "gpt-5.5-pro",
+ expectedModel: "gpt-5.5-pro",
},
{
- name: "preserves compact-spelled gpt5.5 instead of group default",
+ name: "ordinary compact-spelled gpt5.5 request keeps requested model",
account: &Account{
Credentials: map[string]any{},
},
- requestedModel: "gpt5.5",
- defaultMappedModel: "gpt-5.4",
- expectedModel: "gpt5.5",
+ requestedModel: "gpt5.5",
+ expectedModel: "gpt5.5",
},
{
- name: "preserves openai namespaced gpt-5.5 instead of group default",
+ name: "ordinary namespaced gpt-5.5 request keeps requested model",
account: &Account{
Credentials: map[string]any{},
},
- requestedModel: "openai/gpt-5.5",
- defaultMappedModel: "gpt-5.4",
- expectedModel: "openai/gpt-5.5",
+ requestedModel: "openai/gpt-5.5",
+ expectedModel: "openai/gpt-5.5",
},
{
- name: "preserves compact gpt-5.5 instead of group default",
+ name: "ordinary compact gpt-5.5 request keeps requested model",
account: &Account{
Credentials: map[string]any{},
},
- requestedModel: "gpt-5.5-openai-compact",
- defaultMappedModel: "gpt-5.4",
- expectedModel: "gpt-5.5-openai-compact",
+ requestedModel: "gpt-5.5-openai-compact",
+ expectedModel: "gpt-5.5-openai-compact",
+ },
+ {
+ name: "whitespace-only messages dispatch model is ignored",
+ account: &Account{
+ Credentials: map[string]any{},
+ },
+ requestedModel: "gpt-5.5",
+ messagesDispatchMappedModel: " ",
+ expectedModel: "gpt-5.5",
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
- if got := resolveOpenAIForwardModel(tt.account, tt.requestedModel, tt.defaultMappedModel); got != tt.expectedModel {
+ if got := resolveOpenAIForwardModel(tt.account, tt.requestedModel, tt.messagesDispatchMappedModel); got != tt.expectedModel {
t.Fatalf("resolveOpenAIForwardModel(...) = %q, want %q", got, tt.expectedModel)
}
})
}
}
-func TestResolveOpenAIForwardModel_PreventsClaudeModelFromFallingBackToGpt54(t *testing.T) {
- account := &Account{
- Credentials: map[string]any{},
- }
-
- withoutDefault := resolveOpenAIForwardModel(account, "claude-opus-4-6", "")
- if withoutDefault != "claude-opus-4-6" {
- t.Fatalf("resolveOpenAIForwardModel(...) = %q, want %q", withoutDefault, "claude-opus-4-6")
- }
-
- withDefault := resolveOpenAIForwardModel(account, "claude-opus-4-6", "gpt-5.4")
- if withDefault != "gpt-5.4" {
- t.Fatalf("resolveOpenAIForwardModel(...) = %q, want %q", withDefault, "gpt-5.4")
- }
-}
-
func TestResolveOpenAICompactForwardModel(t *testing.T) {
tests := []struct {
name string
diff --git a/backend/internal/service/openai_ws_http_bridge.go b/backend/internal/service/openai_ws_http_bridge.go
index 0afb9181a1..a5aba67c6d 100644
--- a/backend/internal/service/openai_ws_http_bridge.go
+++ b/backend/internal/service/openai_ws_http_bridge.go
@@ -431,7 +431,7 @@ func resolveGrokWSUpstreamModel(account *Account, body []byte, originalModel str
}
}
if upstreamModel == "" {
- upstreamModel = "grok-4.3"
+ upstreamModel = grokDefaultResponsesModel
}
return upstreamModel
}
diff --git a/backend/internal/service/openai_ws_http_bridge_test.go b/backend/internal/service/openai_ws_http_bridge_test.go
index 0105c7a331..d2046b9006 100644
--- a/backend/internal/service/openai_ws_http_bridge_test.go
+++ b/backend/internal/service/openai_ws_http_bridge_test.go
@@ -178,6 +178,51 @@ func TestOpenAIWSHTTPBridgeRelaysSSEFramesAsWebSocketMessages(t *testing.T) {
require.True(t, gjson.GetBytes(upstream.lastBody, "stream").Bool())
}
+func TestProxyOpenAIWSHTTPBridgeTurnForGrokDefaultsEmptyModelTo45(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+
+ upstream := &httpUpstreamRecorder{resp: &http.Response{
+ StatusCode: http.StatusOK,
+ Header: http.Header{"Content-Type": []string{"text/event-stream"}},
+ Body: io.NopCloser(strings.NewReader(strings.Join([]string{
+ `data: {"type":"response.created","response":{"id":"resp_grok_default","model":"grok-4.5"}}`,
+ "",
+ `data: {"type":"response.completed","response":{"id":"resp_grok_default","model":"grok-4.5","usage":{"input_tokens":1,"output_tokens":1}}}`,
+ "",
+ }, "\n"))),
+ }}
+ svc := &OpenAIGatewayService{
+ cfg: &config.Config{Gateway: config.GatewayConfig{MaxLineSize: defaultMaxLineSize}},
+ httpUpstream: upstream,
+ }
+ account := &Account{
+ ID: 72,
+ Platform: PlatformGrok,
+ Type: AccountTypeOAuth,
+ Concurrency: 1,
+ Credentials: map[string]any{"base_url": xai.DefaultCLIBaseURL},
+ }
+ payload := []byte(`{"type":"response.create","generate":true,"stream":true,"input":"hi"}`)
+ recorder := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(recorder)
+ c.Request = httptest.NewRequest(http.MethodGet, "/v1/responses", nil)
+ var events [][]byte
+
+ result, err := svc.proxyOpenAIWSHTTPBridgeTurn(
+ context.Background(), c, account, "access-token", payload, len(payload),
+ "", "", "", "", "", 1,
+ func(message []byte) error {
+ events = append(events, append([]byte(nil), message...))
+ return nil
+ },
+ )
+
+ require.NoError(t, err)
+ require.NotNil(t, result)
+ require.Equal(t, grokDefaultResponsesModel, gjson.GetBytes(upstream.lastBody, "model").String())
+ require.Len(t, events, 2)
+}
+
func TestProxyResponsesWebSocketFromClientForGrokUsesXAIHTTPBridge(t *testing.T) {
gin.SetMode(gin.TestMode)
diff --git a/backend/internal/service/ops_models.go b/backend/internal/service/ops_models.go
index e33dcf82a8..d95bbadbb8 100644
--- a/backend/internal/service/ops_models.go
+++ b/backend/internal/service/ops_models.go
@@ -8,6 +8,7 @@ import (
type OpsSystemLog struct {
ID int64 `json:"id"`
CreatedAt time.Time `json:"created_at"`
+ Host string `json:"host"`
Level string `json:"level"`
Component string `json:"component"`
Message string `json:"message"`
diff --git a/backend/internal/service/ops_port.go b/backend/internal/service/ops_port.go
index 46d171c7c3..2b73d4a694 100644
--- a/backend/internal/service/ops_port.go
+++ b/backend/internal/service/ops_port.go
@@ -194,6 +194,7 @@ type OpsInsertSystemMetricsInput struct {
type OpsInsertSystemLogInput struct {
CreatedAt time.Time
+ Host string
Level string
Component string
Message string
@@ -210,6 +211,7 @@ type OpsInsertSystemLogInput struct {
type OpsSystemLogFilter struct {
StartTime *time.Time
EndTime *time.Time
+ Host string
Level string
Component string
@@ -230,6 +232,7 @@ type OpsSystemLogFilter struct {
type OpsSystemLogCleanupFilter struct {
StartTime *time.Time
EndTime *time.Time
+ Host string
Level string
Component string
diff --git a/backend/internal/service/ops_system_log_service.go b/backend/internal/service/ops_system_log_service.go
index b3be37e8ae..b96ae89d92 100644
--- a/backend/internal/service/ops_system_log_service.go
+++ b/backend/internal/service/ops_system_log_service.go
@@ -89,6 +89,7 @@ func marshalSystemLogCleanupConditions(filter *OpsSystemLogCleanupFilter) string
return "{}"
}
payload := map[string]any{
+ "host": strings.TrimSpace(filter.Host),
"level": strings.TrimSpace(filter.Level),
"component": strings.TrimSpace(filter.Component),
"request_id": strings.TrimSpace(filter.RequestID),
diff --git a/backend/internal/service/ops_system_log_service_test.go b/backend/internal/service/ops_system_log_service_test.go
index 8b5a84c1f0..e8c6199f17 100644
--- a/backend/internal/service/ops_system_log_service_test.go
+++ b/backend/internal/service/ops_system_log_service_test.go
@@ -101,6 +101,7 @@ func TestOpsServiceCleanupSystemLogs_SuccessAndAudit(t *testing.T) {
now := time.Now().UTC()
filter := &OpsSystemLogCleanupFilter{
StartTime: &now,
+ Host: "api-node-1",
Level: "warn",
RequestID: "req-1",
ClientRequestID: "creq-1",
@@ -119,6 +120,9 @@ func TestOpsServiceCleanupSystemLogs_SuccessAndAudit(t *testing.T) {
if audit == nil {
t.Fatalf("expected cleanup audit")
}
+ if !strings.Contains(audit.Conditions, `"host":"api-node-1"`) {
+ t.Fatalf("audit conditions should include host: %s", audit.Conditions)
+ }
if !strings.Contains(audit.Conditions, `"client_request_id":"creq-1"`) {
t.Fatalf("audit conditions should include client_request_id: %s", audit.Conditions)
}
diff --git a/backend/internal/service/ops_system_log_sink.go b/backend/internal/service/ops_system_log_sink.go
index 2ff273be53..2e6f5515c8 100644
--- a/backend/internal/service/ops_system_log_sink.go
+++ b/backend/internal/service/ops_system_log_sink.go
@@ -27,6 +27,7 @@ type OpsSystemLogSinkHealth struct {
type OpsSystemLogSink struct {
opsRepo OpsRepository
+ host string
queue chan *logger.LogEvent
@@ -45,10 +46,14 @@ type OpsSystemLogSink struct {
lastError atomic.Value
}
+const maxSystemLogHostLength = 255
+
func NewOpsSystemLogSink(opsRepo OpsRepository) *OpsSystemLogSink {
ctx, cancel := context.WithCancel(context.Background())
+ rawHost, err := os.Hostname()
s := &OpsSystemLogSink{
opsRepo: opsRepo,
+ host: normalizeSystemLogHost(rawHost, err),
queue: make(chan *logger.LogEvent, 5000),
batchSize: 200,
flushInterval: time.Second,
@@ -59,6 +64,18 @@ func NewOpsSystemLogSink(opsRepo OpsRepository) *OpsSystemLogSink {
return s
}
+func normalizeSystemLogHost(host string, err error) string {
+ host = strings.TrimSpace(host)
+ if err != nil || host == "" {
+ return "unknown"
+ }
+ runes := []rune(host)
+ if len(runes) > maxSystemLogHostLength {
+ return string(runes[:maxSystemLogHostLength])
+ }
+ return host
+}
+
func (s *OpsSystemLogSink) Start() {
if s == nil || s.opsRepo == nil {
return
@@ -220,6 +237,7 @@ func (s *OpsSystemLogSink) flushBatch(baseCtx context.Context, batch []*logger.L
inputs = append(inputs, &OpsInsertSystemLogInput{
CreatedAt: createdAt,
+ Host: s.host,
Level: strings.ToLower(strings.TrimSpace(event.Level)),
Component: component,
Message: message,
diff --git a/backend/internal/service/ops_system_log_sink_test.go b/backend/internal/service/ops_system_log_sink_test.go
index b43d44c32e..0d15f1a662 100644
--- a/backend/internal/service/ops_system_log_sink_test.go
+++ b/backend/internal/service/ops_system_log_sink_test.go
@@ -140,6 +140,7 @@ func TestOpsSystemLogSink_StartStopAndFlushSuccess(t *testing.T) {
}
sink := NewOpsSystemLogSink(repo)
+ sink.host = "api-node-1"
sink.batchSize = 1
sink.flushInterval = 10 * time.Millisecond
sink.Start()
@@ -172,6 +173,9 @@ func TestOpsSystemLogSink_StartStopAndFlushSuccess(t *testing.T) {
t.Fatalf("captured len = %d, want 1", len(captured))
}
item := captured[0]
+ if item.Host != "api-node-1" {
+ t.Fatalf("host = %q, want api-node-1", item.Host)
+ }
if item.RequestID != "req-1" || item.ClientRequestID != "creq-1" {
t.Fatalf("unexpected request ids: %+v", item)
}
@@ -324,3 +328,20 @@ func TestOpsSystemLogSink_HelperFunctions(t *testing.T) {
}
}
}
+
+func TestNormalizeSystemLogHost(t *testing.T) {
+ if got := normalizeSystemLogHost(" api-node-1 ", nil); got != "api-node-1" {
+ t.Fatalf("trimmed host = %q, want api-node-1", got)
+ }
+ if got := normalizeSystemLogHost("", nil); got != "unknown" {
+ t.Fatalf("empty host = %q, want unknown", got)
+ }
+ if got := normalizeSystemLogHost("api-node-1", errors.New("hostname unavailable")); got != "unknown" {
+ t.Fatalf("errored host = %q, want unknown", got)
+ }
+ longHost := strings.Repeat("节", maxSystemLogHostLength+1)
+ got := normalizeSystemLogHost(longHost, nil)
+ if runeCount := len([]rune(got)); runeCount != maxSystemLogHostLength {
+ t.Fatalf("truncated host rune count = %d, want %d", runeCount, maxSystemLogHostLength)
+ }
+}
diff --git a/backend/internal/service/payment_order.go b/backend/internal/service/payment_order.go
index 04feb8002a..9b7ee08990 100644
--- a/backend/internal/service/payment_order.go
+++ b/backend/internal/service/payment_order.go
@@ -16,6 +16,7 @@ import (
"github.com/Wei-Shaw/sub2api/internal/payment"
"github.com/Wei-Shaw/sub2api/internal/payment/provider"
infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors"
+ "github.com/Wei-Shaw/sub2api/internal/pkg/servertiming"
"github.com/shopspring/decimal"
)
@@ -445,7 +446,9 @@ func (s *PaymentService) invokeProvider(ctx context.Context, order *dbent.Paymen
IsMobile: req.IsMobile,
ReturnURL: providerReturnURL,
}, sel, outTradeNo, payAmountStr, subject)
+ finishProviderCall := servertiming.ObserveDependency(ctx, "payment")
pr, err := prov.CreatePayment(ctx, providerReq)
+ finishProviderCall()
if err != nil {
slog.Error("[PaymentService] CreatePayment failed", "provider", sel.ProviderKey, "instance", sel.InstanceID, "error", err)
if appErr := new(infraerrors.ApplicationError); errors.As(err, &appErr) {
diff --git a/backend/internal/service/payment_order_lifecycle.go b/backend/internal/service/payment_order_lifecycle.go
index 8ed18797dd..46a2e00605 100644
--- a/backend/internal/service/payment_order_lifecycle.go
+++ b/backend/internal/service/payment_order_lifecycle.go
@@ -13,6 +13,7 @@ import (
"github.com/Wei-Shaw/sub2api/ent/paymentorder"
"github.com/Wei-Shaw/sub2api/internal/payment"
infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors"
+ "github.com/Wei-Shaw/sub2api/internal/pkg/servertiming"
)
// --- Cancel & Expire ---
@@ -157,7 +158,9 @@ func (s *PaymentService) checkPaidWithOptions(ctx context.Context, o *dbent.Paym
if queryRef == "" {
return ""
}
+ finishProviderCall := servertiming.ObserveDependency(ctx, "payment")
resp, err := prov.QueryOrder(ctx, queryRef)
+ finishProviderCall()
if err != nil {
slog.Warn("query upstream failed", "orderID", o.ID, "error", err)
return ""
@@ -199,7 +202,9 @@ func (s *PaymentService) checkPaidWithOptions(ctx context.Context, o *dbent.Paym
return ""
}
if cp, ok := prov.(payment.CancelableProvider); ok {
+ finishProviderCall := servertiming.ObserveDependency(ctx, "payment")
_ = cp.CancelPayment(ctx, queryRef)
+ finishProviderCall()
}
return ""
}
@@ -208,7 +213,9 @@ func requeryPaidOrderOnce(ctx context.Context, prov payment.Provider, queryRef s
if prov == nil || strings.TrimSpace(queryRef) == "" {
return nil, false
}
+ finishProviderCall := servertiming.ObserveDependency(ctx, "payment")
resp, err := prov.QueryOrder(ctx, queryRef)
+ finishProviderCall()
if err != nil {
slog.Warn("query upstream retry failed", "queryRef", queryRef, "error", err)
return nil, false
diff --git a/backend/internal/service/payment_refund.go b/backend/internal/service/payment_refund.go
index 91822680ed..bc073a2c34 100644
--- a/backend/internal/service/payment_refund.go
+++ b/backend/internal/service/payment_refund.go
@@ -19,6 +19,7 @@ import (
"github.com/Wei-Shaw/sub2api/internal/payment"
"github.com/Wei-Shaw/sub2api/internal/payment/provider"
infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors"
+ "github.com/Wei-Shaw/sub2api/internal/pkg/servertiming"
)
// --- Refund Flow ---
@@ -347,12 +348,14 @@ func (s *PaymentService) gwRefund(ctx context.Context, p *RefundPlan) (*payment.
})
return nil, err
}
+ finishProviderCall := servertiming.ObserveDependency(ctx, "payment")
resp, err := prov.Refund(ctx, payment.RefundRequest{
TradeNo: p.Order.PaymentTradeNo,
OrderID: p.Order.OutTradeNo,
Amount: formatGatewayRefundAmount(p.GatewayAmount, p.Order),
Reason: p.Reason,
})
+ finishProviderCall()
if err != nil {
if resp != nil && strings.TrimSpace(resp.Status) == payment.ProviderStatusPending {
return resp, nil
@@ -417,12 +420,14 @@ func (s *PaymentService) QueryAndFinalizeRefund(ctx context.Context, oid int64)
}
pendingDetail := s.latestRefundPendingDetail(ctx, oid)
+ finishProviderCall := servertiming.ObserveDependency(ctx, "payment")
resp, err := queryProvider.QueryRefund(ctx, payment.RefundQueryRequest{
TradeNo: o.PaymentTradeNo,
OrderID: o.OutTradeNo,
RefundID: pendingDetail.RefundID,
Amount: formatGatewayRefundAmount(o.RefundAmount, o),
})
+ finishProviderCall()
if err != nil {
return nil, fmt.Errorf("query refund: %w", err)
}
diff --git a/backend/internal/service/usage_log.go b/backend/internal/service/usage_log.go
index 62e48fc8f9..0adcc04a94 100644
--- a/backend/internal/service/usage_log.go
+++ b/backend/internal/service/usage_log.go
@@ -142,13 +142,14 @@ type UsageLog struct {
ImageOutputTokens int
ImageOutputCost float64
- InputCost float64
- OutputCost float64
- CacheCreationCost float64
- CacheReadCost float64
- TotalCost float64
- ActualCost float64
- RateMultiplier float64
+ InputCost float64
+ OutputCost float64
+ CacheCreationCost float64
+ CacheReadCost float64
+ TotalCost float64
+ ActualCost float64
+ RateMultiplier float64
+ LongContextBillingApplied bool
// AccountRateMultiplier 账号计费倍率快照(nil 表示历史数据,按 1.0 处理)
AccountRateMultiplier *float64
// AccountStatsCost 账号统计定价预计算费用(nil = 使用默认公式 total_cost × account_rate_multiplier)
diff --git a/backend/internal/service/vertex_service_account.go b/backend/internal/service/vertex_service_account.go
index 256695ded5..7ccbeee43c 100644
--- a/backend/internal/service/vertex_service_account.go
+++ b/backend/internal/service/vertex_service_account.go
@@ -18,6 +18,7 @@ import (
"github.com/Wei-Shaw/sub2api/internal/pkg/proxyurl"
"github.com/Wei-Shaw/sub2api/internal/pkg/proxyutil"
+ "github.com/Wei-Shaw/sub2api/internal/pkg/servertiming"
"github.com/golang-jwt/jwt/v5"
)
@@ -195,7 +196,7 @@ func vertexServiceAccountProxyURL(account *Account) string {
func newVertexServiceAccountHTTPClient(proxyURL string) (*http.Client, error) {
proxyURL = strings.TrimSpace(proxyURL)
if proxyURL == "" {
- return &http.Client{Timeout: 15 * time.Second}, nil
+ return servertiming.InstrumentClient(&http.Client{Timeout: 15 * time.Second}), nil
}
_, parsedProxy, err := proxyurl.Parse(proxyURL)
@@ -211,7 +212,7 @@ func newVertexServiceAccountHTTPClient(proxyURL string) (*http.Client, error) {
if err := proxyutil.ConfigureTransportProxy(transport, parsedProxy); err != nil {
return nil, err
}
- return &http.Client{Timeout: 15 * time.Second, Transport: transport}, nil
+ return servertiming.InstrumentClient(&http.Client{Timeout: 15 * time.Second, Transport: transport}), nil
}
func exchangeVertexServiceAccountToken(ctx context.Context, key *vertexServiceAccountKey, proxyURL string) (string, time.Duration, error) {
diff --git a/backend/internal/service/vertex_service_account_test.go b/backend/internal/service/vertex_service_account_test.go
index d77a1988e9..68c756eaa2 100644
--- a/backend/internal/service/vertex_service_account_test.go
+++ b/backend/internal/service/vertex_service_account_test.go
@@ -13,6 +13,7 @@ import (
"testing"
"time"
+ "github.com/Wei-Shaw/sub2api/internal/pkg/servertiming"
"github.com/stretchr/testify/require"
"github.com/tidwall/gjson"
)
@@ -101,6 +102,24 @@ func TestVertexServiceAccountProxyURL(t *testing.T) {
require.Empty(t, vertexServiceAccountProxyURL(&Account{ProxyID: &proxyID}))
}
+func TestVertexServiceAccountHTTPClientRecordsDependency(t *testing.T) {
+ server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
+ w.WriteHeader(http.StatusNoContent)
+ }))
+ defer server.Close()
+
+ client, err := newVertexServiceAccountHTTPClient("")
+ require.NoError(t, err)
+ collector := servertiming.New(time.Now())
+ ctx := servertiming.WithCollector(context.Background(), collector)
+ request, err := http.NewRequestWithContext(ctx, http.MethodGet, server.URL, nil)
+ require.NoError(t, err)
+ response, err := client.Do(request)
+ require.NoError(t, err)
+ require.NoError(t, response.Body.Close())
+ require.Contains(t, collector.HeaderValue(time.Now(), "bypass"), "dep_http;dur=")
+}
+
func TestExchangeVertexServiceAccountTokenUsesProxy(t *testing.T) {
privateKey, err := rsa.GenerateKey(rand.Reader, 2048)
require.NoError(t, err)
diff --git a/backend/internal/service/wire.go b/backend/internal/service/wire.go
index 7258ff05a3..d5d9124ec2 100644
--- a/backend/internal/service/wire.go
+++ b/backend/internal/service/wire.go
@@ -140,8 +140,9 @@ func ProvideGrokQuotaService(
proxyRepo ProxyRepository,
tokenProvider *GrokTokenProvider,
httpUpstream HTTPUpstream,
+ usageLogRepo UsageLogRepository,
) *GrokQuotaService {
- return NewGrokQuotaService(accountRepo, proxyRepo, tokenProvider, httpUpstream)
+ return NewGrokQuotaService(accountRepo, proxyRepo, tokenProvider, httpUpstream, usageLogRepo)
}
// ProvideGeminiTokenProvider creates GeminiTokenProvider with OAuthRefreshAPI injection
diff --git a/backend/migrations/174_add_usage_log_long_context_billing.sql b/backend/migrations/174_add_usage_log_long_context_billing.sql
new file mode 100644
index 0000000000..090403c310
--- /dev/null
+++ b/backend/migrations/174_add_usage_log_long_context_billing.sql
@@ -0,0 +1,4 @@
+-- Snapshot whether long-context pricing changed token prices for a request so
+-- usage history can explain the applied charge without inferring from totals.
+ALTER TABLE usage_logs
+ ADD COLUMN IF NOT EXISTS long_context_billing_applied BOOLEAN NOT NULL DEFAULT FALSE;
diff --git a/backend/migrations/175_add_ops_system_logs_host.sql b/backend/migrations/175_add_ops_system_logs_host.sql
new file mode 100644
index 0000000000..e5f9f7299c
--- /dev/null
+++ b/backend/migrations/175_add_ops_system_logs_host.sql
@@ -0,0 +1,3 @@
+-- Track the application host that emitted each indexed system log.
+ALTER TABLE ops_system_logs
+ ADD COLUMN IF NOT EXISTS host VARCHAR(255);
diff --git a/backend/migrations/175_default_openai_long_context_billing.sql b/backend/migrations/175_default_openai_long_context_billing.sql
new file mode 100644
index 0000000000..cccbea4108
--- /dev/null
+++ b/backend/migrations/175_default_openai_long_context_billing.sql
@@ -0,0 +1,162 @@
+-- Keep mixed-version writers consistent before backfilling rows that already exist.
+CREATE OR REPLACE FUNCTION public.enforce_openai_long_context_billing_extra()
+RETURNS TRIGGER
+LANGUAGE plpgsql
+AS $$
+DECLARE
+ parent_effective_value JSONB;
+BEGIN
+ IF NEW.platform IS DISTINCT FROM 'openai' THEN
+ RETURN NEW;
+ END IF;
+
+ NEW.extra := COALESCE(NEW.extra, '{}'::jsonb);
+ 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 'false'::jsonb
+ WHEN jsonb_typeof(parent.extra->'openai_long_context_billing_enabled') = 'boolean'
+ THEN parent.extra->'openai_long_context_billing_enabled'
+ ELSE 'false'::jsonb
+ END
+ INTO parent_effective_value
+ FROM accounts AS parent
+ WHERE parent.id = NEW.parent_account_id;
+
+ NEW.extra := jsonb_set(
+ NEW.extra,
+ '{openai_long_context_billing_enabled}',
+ COALESCE(parent_effective_value, 'false'::jsonb),
+ true
+ );
+ ELSIF NOT (NEW.extra ? 'openai_long_context_billing_enabled')
+ AND TG_OP = 'UPDATE'
+ AND OLD.platform = 'openai'
+ AND jsonb_typeof(OLD.extra->'openai_long_context_billing_enabled') = 'boolean' THEN
+ NEW.extra := jsonb_set(
+ NEW.extra,
+ '{openai_long_context_billing_enabled}',
+ OLD.extra->'openai_long_context_billing_enabled',
+ true
+ );
+ ELSIF NOT (NEW.extra ? 'openai_long_context_billing_enabled') THEN
+ NEW.extra := jsonb_set(
+ NEW.extra,
+ '{openai_long_context_billing_enabled}',
+ 'false'::jsonb,
+ true
+ );
+ END IF;
+
+ IF jsonb_typeof(NEW.extra->'openai_long_context_billing_enabled') IS DISTINCT FROM 'boolean' THEN
+ RAISE EXCEPTION 'openai_long_context_billing_enabled must be a boolean'
+ USING ERRCODE = '22023';
+ END IF;
+ RETURN NEW;
+END;
+$$;
+
+CREATE OR REPLACE FUNCTION public.propagate_openai_long_context_billing_extra_to_shadows()
+RETURNS TRIGGER
+LANGUAGE plpgsql
+AS $$
+BEGIN
+ WITH updated_shadows AS (
+ UPDATE accounts AS shadow
+ SET extra = jsonb_set(
+ COALESCE(shadow.extra, '{}'::jsonb),
+ '{openai_long_context_billing_enabled}',
+ NEW.extra->'openai_long_context_billing_enabled',
+ true
+ )
+ WHERE shadow.parent_account_id = NEW.id
+ AND shadow.platform = 'openai'
+ AND shadow.quota_dimension = 'spark'
+ AND shadow.extra->'openai_long_context_billing_enabled'
+ IS DISTINCT FROM NEW.extra->'openai_long_context_billing_enabled'
+ RETURNING shadow.id
+ )
+ INSERT INTO scheduler_outbox (event_type, account_id)
+ SELECT 'account_changed', id
+ FROM updated_shadows;
+ RETURN NULL;
+END;
+$$;
+
+DROP TRIGGER IF EXISTS accounts_enforce_openai_long_context_billing_extra ON accounts;
+CREATE TRIGGER accounts_enforce_openai_long_context_billing_extra
+BEFORE INSERT OR UPDATE OF platform, extra, parent_account_id, quota_dimension
+ON accounts
+FOR EACH ROW
+EXECUTE FUNCTION public.enforce_openai_long_context_billing_extra();
+
+DROP TRIGGER IF EXISTS accounts_propagate_openai_long_context_billing_extra ON accounts;
+CREATE TRIGGER accounts_propagate_openai_long_context_billing_extra
+AFTER UPDATE OF platform, extra
+ON accounts
+FOR EACH ROW
+WHEN (
+ NEW.platform = 'openai'
+ AND NEW.parent_account_id IS NULL
+ AND (
+ OLD.platform IS DISTINCT FROM NEW.platform
+ OR OLD.extra->'openai_long_context_billing_enabled'
+ IS DISTINCT FROM NEW.extra->'openai_long_context_billing_enabled'
+ )
+)
+EXECUTE FUNCTION public.propagate_openai_long_context_billing_extra_to_shadows();
+
+UPDATE accounts
+SET extra = jsonb_set(
+ COALESCE(extra, '{}'::jsonb),
+ '{openai_long_context_billing_enabled}',
+ 'false'::jsonb,
+ true
+)
+WHERE platform = 'openai'
+ AND COALESCE(extra, '{}'::jsonb) ? 'openai_long_context_billing_enabled'
+ AND jsonb_typeof(extra->'openai_long_context_billing_enabled') IS DISTINCT FROM 'boolean';
+
+UPDATE accounts
+SET extra = jsonb_set(
+ COALESCE(extra, '{}'::jsonb),
+ '{openai_long_context_billing_enabled}',
+ 'false'::jsonb,
+ true
+)
+WHERE platform = 'openai'
+ AND parent_account_id IS NULL
+ AND NOT (COALESCE(extra, '{}'::jsonb) ? 'openai_long_context_billing_enabled');
+
+WITH shadow_values AS (
+ SELECT
+ 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 'false'::jsonb
+ WHEN jsonb_typeof(parent.extra->'openai_long_context_billing_enabled') = 'boolean'
+ THEN parent.extra->'openai_long_context_billing_enabled'
+ ELSE 'false'::jsonb
+ END AS effective_value
+ FROM accounts AS shadow
+ JOIN accounts AS parent ON parent.id = shadow.parent_account_id
+ WHERE shadow.platform = 'openai'
+ AND shadow.quota_dimension = 'spark'
+),
+updated_shadows AS (
+ UPDATE accounts AS shadow
+ SET extra = jsonb_set(
+ COALESCE(shadow.extra, '{}'::jsonb),
+ '{openai_long_context_billing_enabled}',
+ shadow_values.effective_value,
+ true
+ )
+ FROM shadow_values
+ WHERE shadow.id = shadow_values.id
+ AND shadow.extra->'openai_long_context_billing_enabled'
+ IS DISTINCT FROM shadow_values.effective_value
+ RETURNING shadow.id
+)
+INSERT INTO scheduler_outbox (event_type, account_id)
+SELECT 'account_changed', id
+FROM updated_shadows;
diff --git a/backend/migrations/175a_add_ops_system_logs_host_index_notx.sql b/backend/migrations/175a_add_ops_system_logs_host_index_notx.sql
new file mode 100644
index 0000000000..ec2705e49b
--- /dev/null
+++ b/backend/migrations/175a_add_ops_system_logs_host_index_notx.sql
@@ -0,0 +1,2 @@
+CREATE INDEX CONCURRENTLY IF NOT EXISTS idx_ops_system_logs_host_created_at
+ ON ops_system_logs (host, created_at DESC);
diff --git a/backend/migrations/openai_long_context_billing_migration_test.go b/backend/migrations/openai_long_context_billing_migration_test.go
new file mode 100644
index 0000000000..212ac15d9d
--- /dev/null
+++ b/backend/migrations/openai_long_context_billing_migration_test.go
@@ -0,0 +1,36 @@
+package migrations
+
+import (
+ "testing"
+
+ "github.com/stretchr/testify/require"
+)
+
+func TestMigration175DefaultsOrdinaryOpenAIAndInheritsForSparkShadows(t *testing.T) {
+ content, err := FS.ReadFile("175_default_openai_long_context_billing.sql")
+ require.NoError(t, err)
+
+ sql := string(content)
+ require.Contains(t, sql, "parent_account_id IS NULL")
+ require.Contains(t, sql, "quota_dimension = 'spark'")
+ require.Contains(t, sql, "parent.extra")
+ require.Contains(t, sql, "jsonb_typeof")
+ require.Contains(t, sql, "openai_long_context_billing_enabled")
+}
+
+func TestMigration175GuardsMixedVersionAccountWrites(t *testing.T) {
+ content, err := FS.ReadFile("175_default_openai_long_context_billing.sql")
+ require.NoError(t, err)
+
+ sql := string(content)
+ require.Contains(t, sql, "RETURNS TRIGGER")
+ require.Contains(t, sql, "BEFORE INSERT OR UPDATE")
+ require.Contains(t, sql, "CREATE TRIGGER")
+ require.Contains(t, sql, "must be a boolean")
+ require.Contains(t, sql, "INSERT INTO scheduler_outbox")
+ require.Contains(t, sql, "'account_changed'")
+ require.Contains(t, sql, "jsonb_typeof(extra->'openai_long_context_billing_enabled') IS DISTINCT FROM 'boolean'")
+ require.Contains(t, sql, "WITH shadow_values AS")
+ require.Contains(t, sql, "TG_OP = 'UPDATE'")
+ require.Contains(t, sql, "OLD.extra->'openai_long_context_billing_enabled'")
+}
diff --git a/deploy/.env.example b/deploy/.env.example
index f68257df9f..1d1a6ec881 100644
--- a/deploy/.env.example
+++ b/deploy/.env.example
@@ -23,6 +23,9 @@ SERVER_PORT=8080
# Server mode: release or debug
SERVER_MODE=release
+# Return Server-Timing for authenticated requests made by the Admin web UI
+ENABLE_SERVER_TIMING=false
+
# Apple container image overrides (ignored by Docker Compose). Pin release tags
# or digests for repeatable operator-managed deployments.
APPLE_CONTAINER_SUB2API_IMAGE=weishaw/sub2api:latest
diff --git a/deploy/config.example.yaml b/deploy/config.example.yaml
index 954263d9c4..dfc584a59a 100644
--- a/deploy/config.example.yaml
+++ b/deploy/config.example.yaml
@@ -20,6 +20,9 @@ server:
# Mode: "debug" for development, "release" for production
# 运行模式:"debug" 用于开发,"release" 用于生产环境
mode: "release"
+ # Return Server-Timing for authenticated requests made by the Admin web UI
+ # 为管理端 Web 页面发出的已认证请求返回 Server-Timing
+ enable_server_timing: false
# Frontend base URL used to generate external links in emails (e.g. password reset)
# 用于生成邮件中的外部链接(例如:重置密码链接)的前端基础地址
# Example: "https://example.com"
diff --git a/deploy/docker-compose.dev.yml b/deploy/docker-compose.dev.yml
index 6f5b3f56f3..43f5dd3f60 100644
--- a/deploy/docker-compose.dev.yml
+++ b/deploy/docker-compose.dev.yml
@@ -26,6 +26,7 @@ services:
- SERVER_HOST=0.0.0.0
- SERVER_PORT=8080
- SERVER_MODE=debug
+ - ENABLE_SERVER_TIMING=${ENABLE_SERVER_TIMING:-false}
- RUN_MODE=${RUN_MODE:-standard}
- DATABASE_HOST=postgres
- DATABASE_PORT=5432
diff --git a/deploy/docker-compose.local.yml b/deploy/docker-compose.local.yml
index 042752e857..5fb161603b 100644
--- a/deploy/docker-compose.local.yml
+++ b/deploy/docker-compose.local.yml
@@ -51,6 +51,7 @@ services:
- SERVER_HOST=0.0.0.0
- SERVER_PORT=8080
- SERVER_MODE=${SERVER_MODE:-release}
+ - ENABLE_SERVER_TIMING=${ENABLE_SERVER_TIMING:-false}
- RUN_MODE=${RUN_MODE:-standard}
# =======================================================================
diff --git a/deploy/docker-compose.standalone.yml b/deploy/docker-compose.standalone.yml
index 2e1d335624..40ed4751d6 100644
--- a/deploy/docker-compose.standalone.yml
+++ b/deploy/docker-compose.standalone.yml
@@ -37,6 +37,7 @@ services:
- SERVER_HOST=0.0.0.0
- SERVER_PORT=8080
- SERVER_MODE=${SERVER_MODE:-release}
+ - ENABLE_SERVER_TIMING=${ENABLE_SERVER_TIMING:-false}
- RUN_MODE=${RUN_MODE:-standard}
# =======================================================================
diff --git a/deploy/docker-compose.yml b/deploy/docker-compose.yml
index 22713c59aa..6aecdcfa5a 100644
--- a/deploy/docker-compose.yml
+++ b/deploy/docker-compose.yml
@@ -47,6 +47,7 @@ services:
- SERVER_HOST=0.0.0.0
- SERVER_PORT=8080
- SERVER_MODE=${SERVER_MODE:-release}
+ - ENABLE_SERVER_TIMING=${ENABLE_SERVER_TIMING:-false}
- RUN_MODE=${RUN_MODE:-standard}
# =======================================================================
diff --git a/frontend/src/api/__tests__/adminUIRequest.spec.ts b/frontend/src/api/__tests__/adminUIRequest.spec.ts
new file mode 100644
index 0000000000..9064a52f1a
--- /dev/null
+++ b/frontend/src/api/__tests__/adminUIRequest.spec.ts
@@ -0,0 +1,38 @@
+import { describe, expect, it } from 'vitest'
+
+import {
+ ADMIN_UI_REQUEST_HEADER,
+ shouldMarkAdminUIRequest,
+} from '@/api/adminUIRequest'
+
+describe('Admin UI request marker', () => {
+ it('uses the stable request header name', () => {
+ expect(ADMIN_UI_REQUEST_HEADER).toBe('X-Admin-UI-Request')
+ })
+
+ it.each([
+ '/admin',
+ '/admin/users',
+ '/api/v1/admin',
+ '/api/v1/admin/accounts?status=active',
+ 'https://api.example.test/api/v1/admin/dashboard',
+ ])('marks Admin API request %s before page navigation', (requestURL) => {
+ expect(shouldMarkAdminUIRequest(requestURL, '/login')).toBe(true)
+ })
+
+ it.each(['/keys', '/groups/available', '/auth/me', '/announcements'])(
+ 'marks shared request %s while an Admin page is active',
+ (requestURL) => {
+ expect(shouldMarkAdminUIRequest(requestURL, '/admin/dashboard')).toBe(true)
+ }
+ )
+
+ it.each([
+ ['/keys', '/dashboard'],
+ ['/api/v1/administer', '/dashboard'],
+ ['/keys', '/administrator'],
+ ['', '/'],
+ ])('does not mark request %s on page %s', (requestURL, pagePath) => {
+ expect(shouldMarkAdminUIRequest(requestURL, pagePath)).toBe(false)
+ })
+})
diff --git a/frontend/src/api/__tests__/client.spec.ts b/frontend/src/api/__tests__/client.spec.ts
index a0a05410d4..b275cca34b 100644
--- a/frontend/src/api/__tests__/client.spec.ts
+++ b/frontend/src/api/__tests__/client.spec.ts
@@ -12,6 +12,7 @@ describe('API Client', () => {
beforeEach(async () => {
localStorage.clear()
+ window.history.replaceState({}, '', '/')
// 每次测试重新导入以获取干净的模块状态
vi.resetModules()
const mod = await import('@/api/client')
@@ -120,6 +121,55 @@ describe('API Client', () => {
const config = adapter.mock.calls[0][0]
expect(config.withCredentials).toBe(true)
})
+
+ it('Admin API 在进入管理页面前也带 Admin UI 标记', async () => {
+ const adapter = vi.fn().mockResolvedValue({
+ status: 200,
+ data: { code: 0, data: {} },
+ headers: {},
+ config: {},
+ statusText: 'OK',
+ })
+ apiClient.defaults.adapter = adapter
+
+ await apiClient.get('/admin/users')
+
+ const config = adapter.mock.calls[0][0]
+ expect(config.headers.get('X-Admin-UI-Request')).toBe('1')
+ })
+
+ it('管理页面调用共享 API 时带 Admin UI 标记', async () => {
+ window.history.replaceState({}, '', '/admin/dashboard')
+ const adapter = vi.fn().mockResolvedValue({
+ status: 200,
+ data: { code: 0, data: {} },
+ headers: {},
+ config: {},
+ statusText: 'OK',
+ })
+ apiClient.defaults.adapter = adapter
+
+ await apiClient.get('/groups/available')
+
+ const config = adapter.mock.calls[0][0]
+ expect(config.headers.get('X-Admin-UI-Request')).toBe('1')
+ })
+
+ it('普通用户页面调用共享 API 时不带 Admin UI 标记', async () => {
+ const adapter = vi.fn().mockResolvedValue({
+ status: 200,
+ data: { code: 0, data: {} },
+ headers: {},
+ config: {},
+ statusText: 'OK',
+ })
+ apiClient.defaults.adapter = adapter
+
+ await apiClient.get('/groups/available')
+
+ const config = adapter.mock.calls[0][0]
+ expect(config.headers.get('X-Admin-UI-Request')).toBeFalsy()
+ })
})
// --- 响应拦截器 ---
diff --git a/frontend/src/api/admin/grok.ts b/frontend/src/api/admin/grok.ts
index 6f8f945f2c..a61169551c 100644
--- a/frontend/src/api/admin/grok.ts
+++ b/frontend/src/api/admin/grok.ts
@@ -4,6 +4,9 @@
*/
import { apiClient } from '../client'
+import type { GrokBillingSummary, GrokQuotaWindow, WindowStats } from '@/types'
+
+export type { GrokBillingSummary, GrokQuotaWindow } from '@/types'
export interface GrokAuthUrlResponse {
auth_url: string
@@ -79,13 +82,6 @@ export function getGrokSSOImportTimeout(keyCount: number): number {
return batches * GROK_SSO_IMPORT_TIMEOUT_PER_BATCH_MS + GROK_SSO_IMPORT_TIMEOUT_BUFFER_MS
}
-export interface GrokQuotaWindow {
- limit?: number | null
- remaining?: number | null
- reset_unix?: number | null
- reset_at?: string | null
-}
-
export interface GrokQuotaSnapshot {
requests?: GrokQuotaWindow | null
tokens?: GrokQuotaWindow | null
@@ -102,13 +98,18 @@ export interface GrokQuotaSnapshot {
}
export interface GrokQuotaProbeResult {
- source: 'active_probe'
- model: string
+ source: 'active_probe' | 'billing_probe' | 'hybrid_probe'
+ model?: string
+ billing?: GrokBillingSummary | null
snapshot?: GrokQuotaSnapshot | null
+ local_usage_7d?: WindowStats | null
+ local_usage_monthly?: WindowStats | null
status_code?: number
headers_observed: boolean
reset_supported: boolean
fetched_at: number
+ persisted?: boolean
+ probe_error?: string
}
export interface GrokQuotaResetResult {
diff --git a/frontend/src/api/admin/ops.ts b/frontend/src/api/admin/ops.ts
index c7cbc64a4b..3284ef7e15 100644
--- a/frontend/src/api/admin/ops.ts
+++ b/frontend/src/api/admin/ops.ts
@@ -828,6 +828,7 @@ export interface OpsRuntimeLogConfig {
export interface OpsSystemLog {
id: number
created_at: string
+ host: string
level: string
component: string
message: string
@@ -849,6 +850,7 @@ export interface OpsSystemLogQuery {
time_range?: '5m' | '30m' | '1h' | '6h' | '24h' | '7d' | '30d'
start_time?: string
end_time?: string
+ host?: string
level?: string
component?: string
request_id?: string
@@ -864,6 +866,7 @@ export interface OpsSystemLogQuery {
export interface OpsSystemLogCleanupRequest {
start_time?: string
end_time?: string
+ host?: string
level?: string
component?: string
request_id?: string
diff --git a/frontend/src/api/adminUIRequest.ts b/frontend/src/api/adminUIRequest.ts
new file mode 100644
index 0000000000..2d60e2987d
--- /dev/null
+++ b/frontend/src/api/adminUIRequest.ts
@@ -0,0 +1,27 @@
+export const ADMIN_UI_REQUEST_HEADER = 'X-Admin-UI-Request'
+
+function isAdminPath(path: string): boolean {
+ return (
+ path === '/admin' ||
+ path.startsWith('/admin/') ||
+ path === '/api/v1/admin' ||
+ path.startsWith('/api/v1/admin/')
+ )
+}
+
+function requestPath(rawURL: string): string {
+ const value = rawURL.trim()
+ if (!value) return ''
+ try {
+ const origin = typeof window !== 'undefined' ? window.location.origin : 'http://localhost'
+ return new URL(value, origin).pathname
+ } catch {
+ return value.split(/[?#]/, 1)[0]
+ }
+}
+
+export function shouldMarkAdminUIRequest(requestURL: string, pagePath?: string): boolean {
+ const currentPath =
+ pagePath ?? (typeof window !== 'undefined' ? window.location.pathname : '')
+ return isAdminPath(requestPath(requestURL)) || isAdminPath(currentPath)
+}
diff --git a/frontend/src/api/client.ts b/frontend/src/api/client.ts
index 5df969f188..a2b4d2f650 100644
--- a/frontend/src/api/client.ts
+++ b/frontend/src/api/client.ts
@@ -6,6 +6,7 @@
import axios, { AxiosInstance, AxiosError, InternalAxiosRequestConfig, AxiosResponse } from 'axios'
import type { ApiResponse } from '@/types'
import { getLocale } from '@/i18n'
+import { ADMIN_UI_REQUEST_HEADER, shouldMarkAdminUIRequest } from './adminUIRequest'
import { getAPIBaseURL } from './url'
export { buildApiUrl, buildGatewayUrl } from './url'
@@ -74,6 +75,10 @@ apiClient.interceptors.request.use(
config.params.timezone = getUserTimezone()
}
+ if (config.headers && shouldMarkAdminUIRequest(String(config.url || ''))) {
+ config.headers[ADMIN_UI_REQUEST_HEADER] = '1'
+ }
+
return config
},
(error) => {
diff --git a/frontend/src/components/account/AccountUsageCell.vue b/frontend/src/components/account/AccountUsageCell.vue
index 5d7af74fc6..d751bf0429 100644
--- a/frontend/src/components/account/AccountUsageCell.vue
+++ b/frontend/src/components/account/AccountUsageCell.vue
@@ -382,7 +382,15 @@
+ {{ t('admin.accounts.openai.longContextBillingDesc') }} +
++ {{ t('admin.accounts.openai.longContextBillingDesc') }} +
+