mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-09-24 16:05:44 +08:00
Merge remote-tracking branch 'origin/main' into feat/grok-sso-device-oauth
# Conflicts: # frontend/src/api/admin/grok.ts
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
@@ -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))
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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))
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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,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
|
||||
|
||||
@@ -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),
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,104 @@
|
||||
package servertiming
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
type dependencyModuleKey struct{}
|
||||
|
||||
type timingRoundTripper struct {
|
||||
base http.RoundTripper
|
||||
}
|
||||
|
||||
// WithDependencyModule overrides the safe module name used for an outbound call.
|
||||
func WithDependencyModule(ctx context.Context, module string) context.Context {
|
||||
if ctx == nil {
|
||||
ctx = context.Background()
|
||||
}
|
||||
module = strings.TrimPrefix(normalizeMetricName(module), dependencyPrefix)
|
||||
if module == "" {
|
||||
return ctx
|
||||
}
|
||||
return context.WithValue(ctx, dependencyModuleKey{}, module)
|
||||
}
|
||||
|
||||
// WrapRoundTripper records outbound response-header latency for active requests.
|
||||
func WrapRoundTripper(base http.RoundTripper) http.RoundTripper {
|
||||
if base == nil {
|
||||
base = http.DefaultTransport
|
||||
}
|
||||
if _, ok := base.(*timingRoundTripper); ok {
|
||||
return base
|
||||
}
|
||||
return &timingRoundTripper{base: base}
|
||||
}
|
||||
|
||||
// InstrumentClient returns a shallow client copy with an instrumented transport.
|
||||
func InstrumentClient(client *http.Client) *http.Client {
|
||||
if client == nil {
|
||||
client = &http.Client{}
|
||||
}
|
||||
copyClient := *client
|
||||
copyClient.Transport = WrapRoundTripper(copyClient.Transport)
|
||||
return ©Client
|
||||
}
|
||||
|
||||
// Do records response-header latency without changing the client's transport
|
||||
// type. Use it for clients whose callers inspect or configure *http.Transport.
|
||||
func Do(client *http.Client, req *http.Request) (*http.Response, error) {
|
||||
if client == nil {
|
||||
client = http.DefaultClient
|
||||
}
|
||||
if req == nil || !Active(req.Context()) {
|
||||
return client.Do(req)
|
||||
}
|
||||
startedAt := time.Now()
|
||||
response, err := client.Do(req)
|
||||
RecordDependency(req.Context(), dependencyModule(req), startedAt, time.Now())
|
||||
return response, err
|
||||
}
|
||||
|
||||
func (t *timingRoundTripper) RoundTrip(req *http.Request) (*http.Response, error) {
|
||||
if req == nil || !Active(req.Context()) {
|
||||
return t.base.RoundTrip(req)
|
||||
}
|
||||
startedAt := time.Now()
|
||||
response, err := t.base.RoundTrip(req)
|
||||
RecordDependency(req.Context(), dependencyModule(req), startedAt, time.Now())
|
||||
return response, err
|
||||
}
|
||||
|
||||
func dependencyModule(req *http.Request) string {
|
||||
if req != nil {
|
||||
if module, ok := req.Context().Value(dependencyModuleKey{}).(string); ok && module != "" {
|
||||
return module
|
||||
}
|
||||
}
|
||||
if req == nil || req.URL == nil {
|
||||
return "http"
|
||||
}
|
||||
host := strings.ToLower(req.URL.Hostname())
|
||||
switch {
|
||||
case strings.Contains(host, "github"):
|
||||
return "github"
|
||||
case strings.Contains(host, "openai"):
|
||||
return "openai"
|
||||
case strings.Contains(host, "anthropic"):
|
||||
return "anthropic"
|
||||
case strings.Contains(host, "generativelanguage") || strings.Contains(host, "gemini"):
|
||||
return "gemini"
|
||||
case strings.Contains(host, "cloudcode") || strings.Contains(host, "antigravity"):
|
||||
return "antigravity"
|
||||
case strings.Contains(host, "googleapis") || strings.Contains(host, "google"):
|
||||
return "google"
|
||||
case strings.Contains(host, "amazonaws") || strings.Contains(host, "cloudflarestorage") || strings.Contains(host, "s3"):
|
||||
return "s3"
|
||||
case strings.Contains(host, "stripe") || strings.Contains(host, "airwallex") || strings.Contains(host, "alipay") || strings.Contains(host, "wechatpay") || strings.Contains(host, "paypal"):
|
||||
return "payment"
|
||||
default:
|
||||
return "http"
|
||||
}
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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())
|
||||
|
||||
+159
@@ -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")
|
||||
}
|
||||
@@ -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)
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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() {
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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
Reference in New Issue
Block a user