Merge remote-tracking branch 'origin/main' into feat/grok-sso-device-oauth

# Conflicts:
#	frontend/src/api/admin/grok.ts
This commit is contained in:
shaw
2026-07-14 10:19:16 +08:00
147 changed files with 8735 additions and 597 deletions
+2 -2
View File
@@ -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)
+15 -14
View File
@@ -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]},
},
},
}
+132 -78
View File
@@ -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
+17 -13
View File
@@ -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()
+3
View File
@@ -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").
+12 -1
View File
@@ -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))
+10
View File
@@ -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()
+15
View File
@@ -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))
+65
View File
@@ -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) {
+34
View File
@@ -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)
}
+5
View File
@@ -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秒空闲超时
+17
View File
@@ -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", "")
@@ -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
@@ -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) {
@@ -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
@@ -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)
}
@@ -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
}
@@ -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
@@ -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])
}
@@ -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
@@ -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),
@@ -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)
@@ -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
+49 -48
View File
@@ -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),
}
}
+8 -7
View File
@@ -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"`
@@ -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)
}
@@ -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
}
@@ -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) {
+9 -9
View File
@@ -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)
}
+2
View File
@@ -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,
@@ -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)
}
@@ -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)
}
}
+104
View File
@@ -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 &copyClient
}
// 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"
}
}
@@ -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")
}
}
+372
View File
@@ -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
}
+127
View File
@@ -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
}
@@ -61,6 +61,7 @@ var schedulerNeutralExtraKeyPrefixes = []string{
var schedulerNeutralExtraKeys = map[string]struct{}{
"codex_usage_updated_at": {},
"grok_billing_snapshot": {},
"session_window_utilization": {},
}
@@ -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},
}))
}
@@ -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)
}
@@ -276,5 +276,5 @@ func createReqClient(proxyURL string) (*req.Client, error) {
client.SetProxyURL(trimmed)
}
return client, nil
return instrumentReqClient(client), nil
}
+14 -4
View File
@@ -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)
+3 -2
View File
@@ -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())
@@ -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")
}
+10
View File
@@ -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,
@@ -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)
+5 -1
View File
@@ -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 连接选项
@@ -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),
@@ -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)
}
@@ -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
}
}
@@ -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")
}
}
@@ -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
}
@@ -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)
}
}
@@ -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,
@@ -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
@@ -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
@@ -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,
+2 -2
View File
@@ -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")
}
// 处理预检请求
@@ -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"),
@@ -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"
}
@@ -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)
}
}
+1
View File
@@ -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() {
+10
View File
@@ -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
}
@@ -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)
}
@@ -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
@@ -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)
+106 -12
View File
@@ -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) {
+94 -4
View File
@@ -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)放行后另一请求抢先建成,本次会撞
@@ -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) {
+70 -29
View File
@@ -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
}
@@ -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()
@@ -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 承载一次检测的自定义入参。
@@ -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),
@@ -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
}
+37 -6
View File
@@ -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
@@ -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
+122 -24
View File
@@ -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
@@ -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()
+242 -9
View File
@@ -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
@@ -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()
@@ -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)
}
@@ -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
}
@@ -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))
}
})
}
}
@@ -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
}
@@ -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",
@@ -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)
@@ -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())
@@ -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)
@@ -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"))
@@ -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) {
@@ -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 行为;后续阶段会把调用点下沉到复杂分支。
@@ -404,6 +404,7 @@ type OpenAIGatewayService struct {
openaiWSRetryMetrics openAIWSRetryMetrics
responseHeaderFilter *responseheaders.CompiledHeaderFilter
codexSnapshotThrottle *accountWriteThrottle
codexModelsManifestCache codexModelsManifestCache
openaiCompatSessionResponses sync.Map
openaiCompatAnthropicDigestSessions sync.Map
}
@@ -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"}]}`))
@@ -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(
@@ -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)
@@ -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)
}
@@ -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
}
@@ -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
@@ -431,7 +431,7 @@ func resolveGrokWSUpstreamModel(account *Account, body []byte, originalModel str
}
}
if upstreamModel == "" {
upstreamModel = "grok-4.3"
upstreamModel = grokDefaultResponsesModel
}
return upstreamModel
}
@@ -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)
+1
View File
@@ -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"`

Some files were not shown because too many files have changed in this diff Show More