diff --git a/backend/cmd/server/wire_gen.go b/backend/cmd/server/wire_gen.go index e4ef733b32..ba5b78a72f 100644 --- a/backend/cmd/server/wire_gen.go +++ b/backend/cmd/server/wire_gen.go @@ -190,9 +190,10 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) { claudeUsageFetcher := repository.NewClaudeUsageFetcher(httpUpstream) antigravityQuotaFetcher := service.NewAntigravityQuotaFetcher(proxyRepository) grokQuotaFetcher := service.NewGrokQuotaFetcher() + grokQuotaService := service.ProvideGrokQuotaService(accountRepository, proxyRepository, grokTokenProvider, httpUpstream, usageLogRepository) openAIQuotaService := service.ProvideOpenAIQuotaService(accountRepository, proxyRepository, openAITokenProvider, privacyClientFactory) usageCache := service.NewUsageCache() - accountUsageService := service.NewAccountUsageService(accountRepository, usageLogRepository, claudeUsageFetcher, geminiQuotaService, antigravityQuotaFetcher, grokQuotaFetcher, openAIQuotaService, usageCache, identityCache, tlsFingerprintProfileService) + accountUsageService := service.NewAccountUsageService(accountRepository, usageLogRepository, claudeUsageFetcher, geminiQuotaService, antigravityQuotaFetcher, grokQuotaFetcher, grokQuotaService, openAIQuotaService, usageCache, identityCache, tlsFingerprintProfileService) accountTestService := service.NewAccountTestService(accountRepository, geminiTokenProvider, claudeTokenProvider, grokTokenProvider, antigravityGatewayService, httpUpstream, configConfig, tlsFingerprintProfileService) crsSyncService := service.NewCRSSyncService(accountRepository, proxyRepository, oAuthService, openAIOAuthService, geminiOAuthService, configConfig) accountHandler := admin.NewAccountHandler(adminService, oAuthService, openAIOAuthService, geminiOAuthService, antigravityOAuthService, rateLimitService, accountUsageService, accountTestService, concurrencyService, crsSyncService, sessionLimitCache, rpmCache, compositeTokenCacheInvalidator) @@ -207,7 +208,6 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) { openAIOAuthHandler := admin.NewOpenAIOAuthHandler(openAIOAuthService, adminService, openAIQuotaService) geminiOAuthHandler := admin.NewGeminiOAuthHandler(geminiOAuthService) antigravityOAuthHandler := admin.NewAntigravityOAuthHandler(antigravityOAuthService) - grokQuotaService := service.ProvideGrokQuotaService(accountRepository, proxyRepository, grokTokenProvider, httpUpstream) grokOAuthHandler := admin.NewGrokOAuthHandler(grokOAuthService, adminService, grokQuotaService) proxyHandler := admin.NewProxyHandler(adminService) adminRedeemHandler := admin.NewRedeemHandler(adminService, redeemService) diff --git a/backend/ent/migrate/schema.go b/backend/ent/migrate/schema.go index 3441afec04..cac559d535 100644 --- a/backend/ent/migrate/schema.go +++ b/backend/ent/migrate/schema.go @@ -1560,6 +1560,7 @@ var ( {Name: "total_cost", Type: field.TypeFloat64, Default: 0, SchemaType: map[string]string{"postgres": "decimal(20,10)"}}, {Name: "actual_cost", Type: field.TypeFloat64, Default: 0, SchemaType: map[string]string{"postgres": "decimal(20,10)"}}, {Name: "rate_multiplier", Type: field.TypeFloat64, Default: 1, SchemaType: map[string]string{"postgres": "decimal(10,4)"}}, + {Name: "long_context_billing_applied", Type: field.TypeBool, Default: false}, {Name: "account_rate_multiplier", Type: field.TypeFloat64, Nullable: true, SchemaType: map[string]string{"postgres": "decimal(10,4)"}}, {Name: "billing_type", Type: field.TypeInt8, Default: 0}, {Name: "stream", Type: field.TypeBool, Default: false}, @@ -1592,31 +1593,31 @@ var ( ForeignKeys: []*schema.ForeignKey{ { Symbol: "usage_logs_api_keys_usage_logs", - Columns: []*schema.Column{UsageLogsColumns[40]}, + Columns: []*schema.Column{UsageLogsColumns[41]}, RefColumns: []*schema.Column{APIKeysColumns[0]}, OnDelete: schema.NoAction, }, { Symbol: "usage_logs_accounts_usage_logs", - Columns: []*schema.Column{UsageLogsColumns[41]}, + Columns: []*schema.Column{UsageLogsColumns[42]}, RefColumns: []*schema.Column{AccountsColumns[0]}, OnDelete: schema.NoAction, }, { Symbol: "usage_logs_groups_usage_logs", - Columns: []*schema.Column{UsageLogsColumns[42]}, + Columns: []*schema.Column{UsageLogsColumns[43]}, RefColumns: []*schema.Column{GroupsColumns[0]}, OnDelete: schema.SetNull, }, { Symbol: "usage_logs_users_usage_logs", - Columns: []*schema.Column{UsageLogsColumns[43]}, + Columns: []*schema.Column{UsageLogsColumns[44]}, RefColumns: []*schema.Column{UsersColumns[0]}, OnDelete: schema.NoAction, }, { Symbol: "usage_logs_user_subscriptions_usage_logs", - Columns: []*schema.Column{UsageLogsColumns[44]}, + Columns: []*schema.Column{UsageLogsColumns[45]}, RefColumns: []*schema.Column{UserSubscriptionsColumns[0]}, OnDelete: schema.SetNull, }, @@ -1625,32 +1626,32 @@ var ( { Name: "usagelog_user_id", Unique: false, - Columns: []*schema.Column{UsageLogsColumns[43]}, + Columns: []*schema.Column{UsageLogsColumns[44]}, }, { Name: "usagelog_api_key_id", Unique: false, - Columns: []*schema.Column{UsageLogsColumns[40]}, + Columns: []*schema.Column{UsageLogsColumns[41]}, }, { Name: "usagelog_account_id", Unique: false, - Columns: []*schema.Column{UsageLogsColumns[41]}, + Columns: []*schema.Column{UsageLogsColumns[42]}, }, { Name: "usagelog_group_id", Unique: false, - Columns: []*schema.Column{UsageLogsColumns[42]}, + Columns: []*schema.Column{UsageLogsColumns[43]}, }, { Name: "usagelog_subscription_id", Unique: false, - Columns: []*schema.Column{UsageLogsColumns[44]}, + Columns: []*schema.Column{UsageLogsColumns[45]}, }, { Name: "usagelog_created_at", Unique: false, - Columns: []*schema.Column{UsageLogsColumns[39]}, + Columns: []*schema.Column{UsageLogsColumns[40]}, }, { Name: "usagelog_model", @@ -1670,17 +1671,17 @@ var ( { Name: "usagelog_user_id_created_at", Unique: false, - Columns: []*schema.Column{UsageLogsColumns[43], UsageLogsColumns[39]}, + Columns: []*schema.Column{UsageLogsColumns[44], UsageLogsColumns[40]}, }, { Name: "usagelog_api_key_id_created_at", Unique: false, - Columns: []*schema.Column{UsageLogsColumns[40], UsageLogsColumns[39]}, + Columns: []*schema.Column{UsageLogsColumns[41], UsageLogsColumns[40]}, }, { Name: "usagelog_group_id_created_at", Unique: false, - Columns: []*schema.Column{UsageLogsColumns[42], UsageLogsColumns[39]}, + Columns: []*schema.Column{UsageLogsColumns[43], UsageLogsColumns[40]}, }, }, } diff --git a/backend/ent/mutation.go b/backend/ent/mutation.go index ab7c424a47..fb35531878 100644 --- a/backend/ent/mutation.go +++ b/backend/ent/mutation.go @@ -41763,83 +41763,84 @@ func (m *UsageCleanupTaskMutation) ResetEdge(name string) error { // UsageLogMutation represents an operation that mutates the UsageLog nodes in the graph. type UsageLogMutation struct { config - op Op - typ string - id *int64 - request_id *string - model *string - requested_model *string - upstream_model *string - channel_id *int64 - addchannel_id *int64 - model_mapping_chain *string - billing_tier *string - billing_mode *string - input_tokens *int - addinput_tokens *int - output_tokens *int - addoutput_tokens *int - cache_creation_tokens *int - addcache_creation_tokens *int - cache_read_tokens *int - addcache_read_tokens *int - cache_creation_5m_tokens *int - addcache_creation_5m_tokens *int - cache_creation_1h_tokens *int - addcache_creation_1h_tokens *int - input_cost *float64 - addinput_cost *float64 - output_cost *float64 - addoutput_cost *float64 - cache_creation_cost *float64 - addcache_creation_cost *float64 - cache_read_cost *float64 - addcache_read_cost *float64 - total_cost *float64 - addtotal_cost *float64 - actual_cost *float64 - addactual_cost *float64 - rate_multiplier *float64 - addrate_multiplier *float64 - account_rate_multiplier *float64 - addaccount_rate_multiplier *float64 - billing_type *int8 - addbilling_type *int8 - stream *bool - duration_ms *int - addduration_ms *int - first_token_ms *int - addfirst_token_ms *int - user_agent *string - ip_address *string - image_count *int - addimage_count *int - image_size *string - image_input_size *string - image_output_size *string - image_size_source *string - image_size_breakdown *map[string]int - video_count *int - addvideo_count *int - video_resolution *string - video_duration_seconds *int - addvideo_duration_seconds *int - cache_ttl_overridden *bool - created_at *time.Time - clearedFields map[string]struct{} - user *int64 - cleareduser bool - api_key *int64 - clearedapi_key bool - account *int64 - clearedaccount bool - group *int64 - clearedgroup bool - subscription *int64 - clearedsubscription bool - done bool - oldValue func(context.Context) (*UsageLog, error) - predicates []predicate.UsageLog + op Op + typ string + id *int64 + request_id *string + model *string + requested_model *string + upstream_model *string + channel_id *int64 + addchannel_id *int64 + model_mapping_chain *string + billing_tier *string + billing_mode *string + input_tokens *int + addinput_tokens *int + output_tokens *int + addoutput_tokens *int + cache_creation_tokens *int + addcache_creation_tokens *int + cache_read_tokens *int + addcache_read_tokens *int + cache_creation_5m_tokens *int + addcache_creation_5m_tokens *int + cache_creation_1h_tokens *int + addcache_creation_1h_tokens *int + input_cost *float64 + addinput_cost *float64 + output_cost *float64 + addoutput_cost *float64 + cache_creation_cost *float64 + addcache_creation_cost *float64 + cache_read_cost *float64 + addcache_read_cost *float64 + total_cost *float64 + addtotal_cost *float64 + actual_cost *float64 + addactual_cost *float64 + rate_multiplier *float64 + addrate_multiplier *float64 + long_context_billing_applied *bool + account_rate_multiplier *float64 + addaccount_rate_multiplier *float64 + billing_type *int8 + addbilling_type *int8 + stream *bool + duration_ms *int + addduration_ms *int + first_token_ms *int + addfirst_token_ms *int + user_agent *string + ip_address *string + image_count *int + addimage_count *int + image_size *string + image_input_size *string + image_output_size *string + image_size_source *string + image_size_breakdown *map[string]int + video_count *int + addvideo_count *int + video_resolution *string + video_duration_seconds *int + addvideo_duration_seconds *int + cache_ttl_overridden *bool + created_at *time.Time + clearedFields map[string]struct{} + user *int64 + cleareduser bool + api_key *int64 + clearedapi_key bool + account *int64 + clearedaccount bool + group *int64 + clearedgroup bool + subscription *int64 + clearedsubscription bool + done bool + oldValue func(context.Context) (*UsageLog, error) + predicates []predicate.UsageLog } var _ ent.Mutation = (*UsageLogMutation)(nil) @@ -43261,6 +43262,42 @@ func (m *UsageLogMutation) ResetRateMultiplier() { m.addrate_multiplier = nil } +// SetLongContextBillingApplied sets the "long_context_billing_applied" field. +func (m *UsageLogMutation) SetLongContextBillingApplied(b bool) { + m.long_context_billing_applied = &b +} + +// LongContextBillingApplied returns the value of the "long_context_billing_applied" field in the mutation. +func (m *UsageLogMutation) LongContextBillingApplied() (r bool, exists bool) { + v := m.long_context_billing_applied + if v == nil { + return + } + return *v, true +} + +// OldLongContextBillingApplied returns the old "long_context_billing_applied" field's value of the UsageLog entity. +// If the UsageLog object wasn't provided to the builder, the object is fetched from the database. +// An error is returned if the mutation operation is not UpdateOne, or the database query fails. +func (m *UsageLogMutation) OldLongContextBillingApplied(ctx context.Context) (v bool, err error) { + if !m.op.Is(OpUpdateOne) { + return v, errors.New("OldLongContextBillingApplied is only allowed on UpdateOne operations") + } + if m.id == nil || m.oldValue == nil { + return v, errors.New("OldLongContextBillingApplied requires an ID field in the mutation") + } + oldValue, err := m.oldValue(ctx) + if err != nil { + return v, fmt.Errorf("querying old value for OldLongContextBillingApplied: %w", err) + } + return oldValue.LongContextBillingApplied, nil +} + +// ResetLongContextBillingApplied resets all changes to the "long_context_billing_applied" field. +func (m *UsageLogMutation) ResetLongContextBillingApplied() { + m.long_context_billing_applied = nil +} + // SetAccountRateMultiplier sets the "account_rate_multiplier" field. func (m *UsageLogMutation) SetAccountRateMultiplier(f float64) { m.account_rate_multiplier = &f @@ -44378,7 +44415,7 @@ func (m *UsageLogMutation) Type() string { // order to get all numeric fields that were incremented/decremented, call // AddedFields(). func (m *UsageLogMutation) Fields() []string { - fields := make([]string, 0, 44) + fields := make([]string, 0, 45) if m.user != nil { fields = append(fields, usagelog.FieldUserID) } @@ -44457,6 +44494,9 @@ func (m *UsageLogMutation) Fields() []string { if m.rate_multiplier != nil { fields = append(fields, usagelog.FieldRateMultiplier) } + if m.long_context_billing_applied != nil { + fields = append(fields, usagelog.FieldLongContextBillingApplied) + } if m.account_rate_multiplier != nil { fields = append(fields, usagelog.FieldAccountRateMultiplier) } @@ -44571,6 +44611,8 @@ func (m *UsageLogMutation) Field(name string) (ent.Value, bool) { return m.ActualCost() case usagelog.FieldRateMultiplier: return m.RateMultiplier() + case usagelog.FieldLongContextBillingApplied: + return m.LongContextBillingApplied() case usagelog.FieldAccountRateMultiplier: return m.AccountRateMultiplier() case usagelog.FieldBillingType: @@ -44668,6 +44710,8 @@ func (m *UsageLogMutation) OldField(ctx context.Context, name string) (ent.Value return m.OldActualCost(ctx) case usagelog.FieldRateMultiplier: return m.OldRateMultiplier(ctx) + case usagelog.FieldLongContextBillingApplied: + return m.OldLongContextBillingApplied(ctx) case usagelog.FieldAccountRateMultiplier: return m.OldAccountRateMultiplier(ctx) case usagelog.FieldBillingType: @@ -44895,6 +44939,13 @@ func (m *UsageLogMutation) SetField(name string, value ent.Value) error { } m.SetRateMultiplier(v) return nil + case usagelog.FieldLongContextBillingApplied: + v, ok := value.(bool) + if !ok { + return fmt.Errorf("unexpected type %T for field %s", value, name) + } + m.SetLongContextBillingApplied(v) + return nil case usagelog.FieldAccountRateMultiplier: v, ok := value.(float64) if !ok { @@ -45526,6 +45577,9 @@ func (m *UsageLogMutation) ResetField(name string) error { case usagelog.FieldRateMultiplier: m.ResetRateMultiplier() return nil + case usagelog.FieldLongContextBillingApplied: + m.ResetLongContextBillingApplied() + return nil case usagelog.FieldAccountRateMultiplier: m.ResetAccountRateMultiplier() return nil diff --git a/backend/ent/runtime/runtime.go b/backend/ent/runtime/runtime.go index 4cb3f800f8..867f1cbdbd 100644 --- a/backend/ent/runtime/runtime.go +++ b/backend/ent/runtime/runtime.go @@ -1940,56 +1940,60 @@ func init() { usagelogDescRateMultiplier := usagelogFields[25].Descriptor() // usagelog.DefaultRateMultiplier holds the default value on creation for the rate_multiplier field. usagelog.DefaultRateMultiplier = usagelogDescRateMultiplier.Default.(float64) + // usagelogDescLongContextBillingApplied is the schema descriptor for long_context_billing_applied field. + usagelogDescLongContextBillingApplied := usagelogFields[26].Descriptor() + // usagelog.DefaultLongContextBillingApplied holds the default value on creation for the long_context_billing_applied field. + usagelog.DefaultLongContextBillingApplied = usagelogDescLongContextBillingApplied.Default.(bool) // usagelogDescBillingType is the schema descriptor for billing_type field. - usagelogDescBillingType := usagelogFields[27].Descriptor() + usagelogDescBillingType := usagelogFields[28].Descriptor() // usagelog.DefaultBillingType holds the default value on creation for the billing_type field. usagelog.DefaultBillingType = usagelogDescBillingType.Default.(int8) // usagelogDescStream is the schema descriptor for stream field. - usagelogDescStream := usagelogFields[28].Descriptor() + usagelogDescStream := usagelogFields[29].Descriptor() // usagelog.DefaultStream holds the default value on creation for the stream field. usagelog.DefaultStream = usagelogDescStream.Default.(bool) // usagelogDescUserAgent is the schema descriptor for user_agent field. - usagelogDescUserAgent := usagelogFields[31].Descriptor() + usagelogDescUserAgent := usagelogFields[32].Descriptor() // usagelog.UserAgentValidator is a validator for the "user_agent" field. It is called by the builders before save. usagelog.UserAgentValidator = usagelogDescUserAgent.Validators[0].(func(string) error) // usagelogDescIPAddress is the schema descriptor for ip_address field. - usagelogDescIPAddress := usagelogFields[32].Descriptor() + usagelogDescIPAddress := usagelogFields[33].Descriptor() // usagelog.IPAddressValidator is a validator for the "ip_address" field. It is called by the builders before save. usagelog.IPAddressValidator = usagelogDescIPAddress.Validators[0].(func(string) error) // usagelogDescImageCount is the schema descriptor for image_count field. - usagelogDescImageCount := usagelogFields[33].Descriptor() + usagelogDescImageCount := usagelogFields[34].Descriptor() // usagelog.DefaultImageCount holds the default value on creation for the image_count field. usagelog.DefaultImageCount = usagelogDescImageCount.Default.(int) // usagelogDescImageSize is the schema descriptor for image_size field. - usagelogDescImageSize := usagelogFields[34].Descriptor() + usagelogDescImageSize := usagelogFields[35].Descriptor() // usagelog.ImageSizeValidator is a validator for the "image_size" field. It is called by the builders before save. usagelog.ImageSizeValidator = usagelogDescImageSize.Validators[0].(func(string) error) // usagelogDescImageInputSize is the schema descriptor for image_input_size field. - usagelogDescImageInputSize := usagelogFields[35].Descriptor() + usagelogDescImageInputSize := usagelogFields[36].Descriptor() // usagelog.ImageInputSizeValidator is a validator for the "image_input_size" field. It is called by the builders before save. usagelog.ImageInputSizeValidator = usagelogDescImageInputSize.Validators[0].(func(string) error) // usagelogDescImageOutputSize is the schema descriptor for image_output_size field. - usagelogDescImageOutputSize := usagelogFields[36].Descriptor() + usagelogDescImageOutputSize := usagelogFields[37].Descriptor() // usagelog.ImageOutputSizeValidator is a validator for the "image_output_size" field. It is called by the builders before save. usagelog.ImageOutputSizeValidator = usagelogDescImageOutputSize.Validators[0].(func(string) error) // usagelogDescImageSizeSource is the schema descriptor for image_size_source field. - usagelogDescImageSizeSource := usagelogFields[37].Descriptor() + usagelogDescImageSizeSource := usagelogFields[38].Descriptor() // usagelog.ImageSizeSourceValidator is a validator for the "image_size_source" field. It is called by the builders before save. usagelog.ImageSizeSourceValidator = usagelogDescImageSizeSource.Validators[0].(func(string) error) // usagelogDescVideoCount is the schema descriptor for video_count field. - usagelogDescVideoCount := usagelogFields[39].Descriptor() + usagelogDescVideoCount := usagelogFields[40].Descriptor() // usagelog.DefaultVideoCount holds the default value on creation for the video_count field. usagelog.DefaultVideoCount = usagelogDescVideoCount.Default.(int) // usagelogDescVideoResolution is the schema descriptor for video_resolution field. - usagelogDescVideoResolution := usagelogFields[40].Descriptor() + usagelogDescVideoResolution := usagelogFields[41].Descriptor() // usagelog.VideoResolutionValidator is a validator for the "video_resolution" field. It is called by the builders before save. usagelog.VideoResolutionValidator = usagelogDescVideoResolution.Validators[0].(func(string) error) // usagelogDescCacheTTLOverridden is the schema descriptor for cache_ttl_overridden field. - usagelogDescCacheTTLOverridden := usagelogFields[42].Descriptor() + usagelogDescCacheTTLOverridden := usagelogFields[43].Descriptor() // usagelog.DefaultCacheTTLOverridden holds the default value on creation for the cache_ttl_overridden field. usagelog.DefaultCacheTTLOverridden = usagelogDescCacheTTLOverridden.Default.(bool) // usagelogDescCreatedAt is the schema descriptor for created_at field. - usagelogDescCreatedAt := usagelogFields[43].Descriptor() + usagelogDescCreatedAt := usagelogFields[44].Descriptor() // usagelog.DefaultCreatedAt holds the default value on creation for the created_at field. usagelog.DefaultCreatedAt = usagelogDescCreatedAt.Default.(func() time.Time) userMixin := schema.User{}.Mixin() diff --git a/backend/ent/schema/usage_log.go b/backend/ent/schema/usage_log.go index e84cc1c140..6d8c2d4191 100644 --- a/backend/ent/schema/usage_log.go +++ b/backend/ent/schema/usage_log.go @@ -100,6 +100,9 @@ func (UsageLog) Fields() []ent.Field { field.Float("rate_multiplier"). Default(1). SchemaType(map[string]string{dialect.Postgres: "decimal(10,4)"}), + field.Bool("long_context_billing_applied"). + Default(false). + Comment("Whether long-context pricing changed token prices for this request"), // account_rate_multiplier: 账号计费倍率快照(NULL 表示按 1.0 处理) field.Float("account_rate_multiplier"). diff --git a/backend/ent/usagelog.go b/backend/ent/usagelog.go index 4d374a8495..b13e29b2f7 100644 --- a/backend/ent/usagelog.go +++ b/backend/ent/usagelog.go @@ -75,6 +75,8 @@ type UsageLog struct { ActualCost float64 `json:"actual_cost,omitempty"` // RateMultiplier holds the value of the "rate_multiplier" field. RateMultiplier float64 `json:"rate_multiplier,omitempty"` + // Whether long-context pricing changed token prices for this request + LongContextBillingApplied bool `json:"long_context_billing_applied,omitempty"` // AccountRateMultiplier holds the value of the "account_rate_multiplier" field. AccountRateMultiplier *float64 `json:"account_rate_multiplier,omitempty"` // BillingType holds the value of the "billing_type" field. @@ -196,7 +198,7 @@ func (*UsageLog) scanValues(columns []string) ([]any, error) { switch columns[i] { case usagelog.FieldImageSizeBreakdown: values[i] = new([]byte) - case usagelog.FieldStream, usagelog.FieldCacheTTLOverridden: + case usagelog.FieldLongContextBillingApplied, usagelog.FieldStream, usagelog.FieldCacheTTLOverridden: values[i] = new(sql.NullBool) case usagelog.FieldInputCost, usagelog.FieldOutputCost, usagelog.FieldCacheCreationCost, usagelog.FieldCacheReadCost, usagelog.FieldTotalCost, usagelog.FieldActualCost, usagelog.FieldRateMultiplier, usagelog.FieldAccountRateMultiplier: values[i] = new(sql.NullFloat64) @@ -391,6 +393,12 @@ func (_m *UsageLog) assignValues(columns []string, values []any) error { } else if value.Valid { _m.RateMultiplier = value.Float64 } + case usagelog.FieldLongContextBillingApplied: + if value, ok := values[i].(*sql.NullBool); !ok { + return fmt.Errorf("unexpected type %T for field long_context_billing_applied", values[i]) + } else if value.Valid { + _m.LongContextBillingApplied = value.Bool + } case usagelog.FieldAccountRateMultiplier: if value, ok := values[i].(*sql.NullFloat64); !ok { return fmt.Errorf("unexpected type %T for field account_rate_multiplier", values[i]) @@ -667,6 +675,9 @@ func (_m *UsageLog) String() string { builder.WriteString("rate_multiplier=") builder.WriteString(fmt.Sprintf("%v", _m.RateMultiplier)) builder.WriteString(", ") + builder.WriteString("long_context_billing_applied=") + builder.WriteString(fmt.Sprintf("%v", _m.LongContextBillingApplied)) + builder.WriteString(", ") if v := _m.AccountRateMultiplier; v != nil { builder.WriteString("account_rate_multiplier=") builder.WriteString(fmt.Sprintf("%v", *v)) diff --git a/backend/ent/usagelog/usagelog.go b/backend/ent/usagelog/usagelog.go index a74a92c40f..a87d937195 100644 --- a/backend/ent/usagelog/usagelog.go +++ b/backend/ent/usagelog/usagelog.go @@ -66,6 +66,8 @@ const ( FieldActualCost = "actual_cost" // FieldRateMultiplier holds the string denoting the rate_multiplier field in the database. FieldRateMultiplier = "rate_multiplier" + // FieldLongContextBillingApplied holds the string denoting the long_context_billing_applied field in the database. + FieldLongContextBillingApplied = "long_context_billing_applied" // FieldAccountRateMultiplier holds the string denoting the account_rate_multiplier field in the database. FieldAccountRateMultiplier = "account_rate_multiplier" // FieldBillingType holds the string denoting the billing_type field in the database. @@ -180,6 +182,7 @@ var Columns = []string{ FieldTotalCost, FieldActualCost, FieldRateMultiplier, + FieldLongContextBillingApplied, FieldAccountRateMultiplier, FieldBillingType, FieldStream, @@ -251,6 +254,8 @@ var ( DefaultActualCost float64 // DefaultRateMultiplier holds the default value on creation for the "rate_multiplier" field. DefaultRateMultiplier float64 + // DefaultLongContextBillingApplied holds the default value on creation for the "long_context_billing_applied" field. + DefaultLongContextBillingApplied bool // DefaultBillingType holds the default value on creation for the "billing_type" field. DefaultBillingType int8 // DefaultStream holds the default value on creation for the "stream" field. @@ -417,6 +422,11 @@ func ByRateMultiplier(opts ...sql.OrderTermOption) OrderOption { return sql.OrderByField(FieldRateMultiplier, opts...).ToFunc() } +// ByLongContextBillingApplied orders the results by the long_context_billing_applied field. +func ByLongContextBillingApplied(opts ...sql.OrderTermOption) OrderOption { + return sql.OrderByField(FieldLongContextBillingApplied, opts...).ToFunc() +} + // ByAccountRateMultiplier orders the results by the account_rate_multiplier field. func ByAccountRateMultiplier(opts ...sql.OrderTermOption) OrderOption { return sql.OrderByField(FieldAccountRateMultiplier, opts...).ToFunc() diff --git a/backend/ent/usagelog/where.go b/backend/ent/usagelog/where.go index 4b08cc3425..a9462e0d0e 100644 --- a/backend/ent/usagelog/where.go +++ b/backend/ent/usagelog/where.go @@ -185,6 +185,11 @@ func RateMultiplier(v float64) predicate.UsageLog { return predicate.UsageLog(sql.FieldEQ(FieldRateMultiplier, v)) } +// LongContextBillingApplied applies equality check predicate on the "long_context_billing_applied" field. It's identical to LongContextBillingAppliedEQ. +func LongContextBillingApplied(v bool) predicate.UsageLog { + return predicate.UsageLog(sql.FieldEQ(FieldLongContextBillingApplied, v)) +} + // AccountRateMultiplier applies equality check predicate on the "account_rate_multiplier" field. It's identical to AccountRateMultiplierEQ. func AccountRateMultiplier(v float64) predicate.UsageLog { return predicate.UsageLog(sql.FieldEQ(FieldAccountRateMultiplier, v)) @@ -1465,6 +1470,16 @@ func RateMultiplierLTE(v float64) predicate.UsageLog { return predicate.UsageLog(sql.FieldLTE(FieldRateMultiplier, v)) } +// LongContextBillingAppliedEQ applies the EQ predicate on the "long_context_billing_applied" field. +func LongContextBillingAppliedEQ(v bool) predicate.UsageLog { + return predicate.UsageLog(sql.FieldEQ(FieldLongContextBillingApplied, v)) +} + +// LongContextBillingAppliedNEQ applies the NEQ predicate on the "long_context_billing_applied" field. +func LongContextBillingAppliedNEQ(v bool) predicate.UsageLog { + return predicate.UsageLog(sql.FieldNEQ(FieldLongContextBillingApplied, v)) +} + // AccountRateMultiplierEQ applies the EQ predicate on the "account_rate_multiplier" field. func AccountRateMultiplierEQ(v float64) predicate.UsageLog { return predicate.UsageLog(sql.FieldEQ(FieldAccountRateMultiplier, v)) diff --git a/backend/ent/usagelog_create.go b/backend/ent/usagelog_create.go index 3326f72fc0..31cf45328e 100644 --- a/backend/ent/usagelog_create.go +++ b/backend/ent/usagelog_create.go @@ -351,6 +351,20 @@ func (_c *UsageLogCreate) SetNillableRateMultiplier(v *float64) *UsageLogCreate return _c } +// SetLongContextBillingApplied sets the "long_context_billing_applied" field. +func (_c *UsageLogCreate) SetLongContextBillingApplied(v bool) *UsageLogCreate { + _c.mutation.SetLongContextBillingApplied(v) + return _c +} + +// SetNillableLongContextBillingApplied sets the "long_context_billing_applied" field if the given value is not nil. +func (_c *UsageLogCreate) SetNillableLongContextBillingApplied(v *bool) *UsageLogCreate { + if v != nil { + _c.SetLongContextBillingApplied(*v) + } + return _c +} + // SetAccountRateMultiplier sets the "account_rate_multiplier" field. func (_c *UsageLogCreate) SetAccountRateMultiplier(v float64) *UsageLogCreate { _c.mutation.SetAccountRateMultiplier(v) @@ -707,6 +721,10 @@ func (_c *UsageLogCreate) defaults() { v := usagelog.DefaultRateMultiplier _c.mutation.SetRateMultiplier(v) } + if _, ok := _c.mutation.LongContextBillingApplied(); !ok { + v := usagelog.DefaultLongContextBillingApplied + _c.mutation.SetLongContextBillingApplied(v) + } if _, ok := _c.mutation.BillingType(); !ok { v := usagelog.DefaultBillingType _c.mutation.SetBillingType(v) @@ -824,6 +842,9 @@ func (_c *UsageLogCreate) check() error { if _, ok := _c.mutation.RateMultiplier(); !ok { return &ValidationError{Name: "rate_multiplier", err: errors.New(`ent: missing required field "UsageLog.rate_multiplier"`)} } + if _, ok := _c.mutation.LongContextBillingApplied(); !ok { + return &ValidationError{Name: "long_context_billing_applied", err: errors.New(`ent: missing required field "UsageLog.long_context_billing_applied"`)} + } if _, ok := _c.mutation.BillingType(); !ok { return &ValidationError{Name: "billing_type", err: errors.New(`ent: missing required field "UsageLog.billing_type"`)} } @@ -997,6 +1018,10 @@ func (_c *UsageLogCreate) createSpec() (*UsageLog, *sqlgraph.CreateSpec) { _spec.SetField(usagelog.FieldRateMultiplier, field.TypeFloat64, value) _node.RateMultiplier = value } + if value, ok := _c.mutation.LongContextBillingApplied(); ok { + _spec.SetField(usagelog.FieldLongContextBillingApplied, field.TypeBool, value) + _node.LongContextBillingApplied = value + } if value, ok := _c.mutation.AccountRateMultiplier(); ok { _spec.SetField(usagelog.FieldAccountRateMultiplier, field.TypeFloat64, value) _node.AccountRateMultiplier = &value @@ -1650,6 +1675,18 @@ func (u *UsageLogUpsert) AddRateMultiplier(v float64) *UsageLogUpsert { return u } +// SetLongContextBillingApplied sets the "long_context_billing_applied" field. +func (u *UsageLogUpsert) SetLongContextBillingApplied(v bool) *UsageLogUpsert { + u.Set(usagelog.FieldLongContextBillingApplied, v) + return u +} + +// UpdateLongContextBillingApplied sets the "long_context_billing_applied" field to the value that was provided on create. +func (u *UsageLogUpsert) UpdateLongContextBillingApplied() *UsageLogUpsert { + u.SetExcluded(usagelog.FieldLongContextBillingApplied) + return u +} + // SetAccountRateMultiplier sets the "account_rate_multiplier" field. func (u *UsageLogUpsert) SetAccountRateMultiplier(v float64) *UsageLogUpsert { u.Set(usagelog.FieldAccountRateMultiplier, v) @@ -2531,6 +2568,20 @@ func (u *UsageLogUpsertOne) UpdateRateMultiplier() *UsageLogUpsertOne { }) } +// SetLongContextBillingApplied sets the "long_context_billing_applied" field. +func (u *UsageLogUpsertOne) SetLongContextBillingApplied(v bool) *UsageLogUpsertOne { + return u.Update(func(s *UsageLogUpsert) { + s.SetLongContextBillingApplied(v) + }) +} + +// UpdateLongContextBillingApplied sets the "long_context_billing_applied" field to the value that was provided on create. +func (u *UsageLogUpsertOne) UpdateLongContextBillingApplied() *UsageLogUpsertOne { + return u.Update(func(s *UsageLogUpsert) { + s.UpdateLongContextBillingApplied() + }) +} + // SetAccountRateMultiplier sets the "account_rate_multiplier" field. func (u *UsageLogUpsertOne) SetAccountRateMultiplier(v float64) *UsageLogUpsertOne { return u.Update(func(s *UsageLogUpsert) { @@ -3631,6 +3682,20 @@ func (u *UsageLogUpsertBulk) UpdateRateMultiplier() *UsageLogUpsertBulk { }) } +// SetLongContextBillingApplied sets the "long_context_billing_applied" field. +func (u *UsageLogUpsertBulk) SetLongContextBillingApplied(v bool) *UsageLogUpsertBulk { + return u.Update(func(s *UsageLogUpsert) { + s.SetLongContextBillingApplied(v) + }) +} + +// UpdateLongContextBillingApplied sets the "long_context_billing_applied" field to the value that was provided on create. +func (u *UsageLogUpsertBulk) UpdateLongContextBillingApplied() *UsageLogUpsertBulk { + return u.Update(func(s *UsageLogUpsert) { + s.UpdateLongContextBillingApplied() + }) +} + // SetAccountRateMultiplier sets the "account_rate_multiplier" field. func (u *UsageLogUpsertBulk) SetAccountRateMultiplier(v float64) *UsageLogUpsertBulk { return u.Update(func(s *UsageLogUpsert) { diff --git a/backend/ent/usagelog_update.go b/backend/ent/usagelog_update.go index 00a65ccff1..2a60d6f44d 100644 --- a/backend/ent/usagelog_update.go +++ b/backend/ent/usagelog_update.go @@ -542,6 +542,20 @@ func (_u *UsageLogUpdate) AddRateMultiplier(v float64) *UsageLogUpdate { return _u } +// SetLongContextBillingApplied sets the "long_context_billing_applied" field. +func (_u *UsageLogUpdate) SetLongContextBillingApplied(v bool) *UsageLogUpdate { + _u.mutation.SetLongContextBillingApplied(v) + return _u +} + +// SetNillableLongContextBillingApplied sets the "long_context_billing_applied" field if the given value is not nil. +func (_u *UsageLogUpdate) SetNillableLongContextBillingApplied(v *bool) *UsageLogUpdate { + if v != nil { + _u.SetLongContextBillingApplied(*v) + } + return _u +} + // SetAccountRateMultiplier sets the "account_rate_multiplier" field. func (_u *UsageLogUpdate) SetAccountRateMultiplier(v float64) *UsageLogUpdate { _u.mutation.ResetAccountRateMultiplier() @@ -1199,6 +1213,9 @@ func (_u *UsageLogUpdate) sqlSave(ctx context.Context) (_node int, err error) { if value, ok := _u.mutation.AddedRateMultiplier(); ok { _spec.AddField(usagelog.FieldRateMultiplier, field.TypeFloat64, value) } + if value, ok := _u.mutation.LongContextBillingApplied(); ok { + _spec.SetField(usagelog.FieldLongContextBillingApplied, field.TypeBool, value) + } if value, ok := _u.mutation.AccountRateMultiplier(); ok { _spec.SetField(usagelog.FieldAccountRateMultiplier, field.TypeFloat64, value) } @@ -1982,6 +1999,20 @@ func (_u *UsageLogUpdateOne) AddRateMultiplier(v float64) *UsageLogUpdateOne { return _u } +// SetLongContextBillingApplied sets the "long_context_billing_applied" field. +func (_u *UsageLogUpdateOne) SetLongContextBillingApplied(v bool) *UsageLogUpdateOne { + _u.mutation.SetLongContextBillingApplied(v) + return _u +} + +// SetNillableLongContextBillingApplied sets the "long_context_billing_applied" field if the given value is not nil. +func (_u *UsageLogUpdateOne) SetNillableLongContextBillingApplied(v *bool) *UsageLogUpdateOne { + if v != nil { + _u.SetLongContextBillingApplied(*v) + } + return _u +} + // SetAccountRateMultiplier sets the "account_rate_multiplier" field. func (_u *UsageLogUpdateOne) SetAccountRateMultiplier(v float64) *UsageLogUpdateOne { _u.mutation.ResetAccountRateMultiplier() @@ -2669,6 +2700,9 @@ func (_u *UsageLogUpdateOne) sqlSave(ctx context.Context) (_node *UsageLog, err if value, ok := _u.mutation.AddedRateMultiplier(); ok { _spec.AddField(usagelog.FieldRateMultiplier, field.TypeFloat64, value) } + if value, ok := _u.mutation.LongContextBillingApplied(); ok { + _spec.SetField(usagelog.FieldLongContextBillingApplied, field.TypeBool, value) + } if value, ok := _u.mutation.AccountRateMultiplier(); ok { _spec.SetField(usagelog.FieldAccountRateMultiplier, field.TypeFloat64, value) } diff --git a/backend/internal/config/config.go b/backend/internal/config/config.go index 8e081bac34..ea9169e0b7 100644 --- a/backend/internal/config/config.go +++ b/backend/internal/config/config.go @@ -601,6 +601,7 @@ type ServerConfig struct { Host string `mapstructure:"host"` Port int `mapstructure:"port"` Mode string `mapstructure:"mode"` // debug/release + EnableServerTiming bool `mapstructure:"enable_server_timing"` // Admin UI Server-Timing response header FrontendURL string `mapstructure:"frontend_url"` // 前端基础 URL,用于生成邮件中的外部链接 ReadHeaderTimeout int `mapstructure:"read_header_timeout"` // 读取请求头超时(秒) IdleTimeout int `mapstructure:"idle_timeout"` // 空闲连接超时(秒) @@ -1459,6 +1460,9 @@ func load(allowMissingJWTSecret bool) (*Config, error) { // 环境变量支持 viper.AutomaticEnv() viper.SetEnvKeyReplacer(strings.NewReplacer(".", "_")) + if err := viper.BindEnv("server.enable_server_timing", "ENABLE_SERVER_TIMING"); err != nil { + return nil, fmt.Errorf("bind ENABLE_SERVER_TIMING: %w", err) + } // 默认值 setDefaults() @@ -1614,6 +1618,7 @@ func setDefaults() { viper.SetDefault("server.host", "0.0.0.0") viper.SetDefault("server.port", 8080) viper.SetDefault("server.mode", "release") + viper.SetDefault("server.enable_server_timing", false) viper.SetDefault("server.frontend_url", "") viper.SetDefault("server.read_header_timeout", 30) // 30秒读取请求头 viper.SetDefault("server.idle_timeout", 120) // 120秒空闲超时 diff --git a/backend/internal/config/config_test.go b/backend/internal/config/config_test.go index 489bc346b8..4eea3a2840 100644 --- a/backend/internal/config/config_test.go +++ b/backend/internal/config/config_test.go @@ -17,6 +17,23 @@ func resetViperWithJWTSecret(t *testing.T) { t.Setenv("JWT_SECRET", strings.Repeat("x", 32)) } +func TestLoadServerTimingConfig(t *testing.T) { + t.Run("disabled by default", func(t *testing.T) { + resetViperWithJWTSecret(t) + cfg, err := Load() + require.NoError(t, err) + require.False(t, cfg.Server.EnableServerTiming) + }) + + t.Run("enabled by exact environment variable", func(t *testing.T) { + resetViperWithJWTSecret(t) + t.Setenv("ENABLE_SERVER_TIMING", "true") + cfg, err := Load() + require.NoError(t, err) + require.True(t, cfg.Server.EnableServerTiming) + }) +} + func TestLoadForBootstrapAllowsMissingJWTSecret(t *testing.T) { viper.Reset() t.Setenv("JWT_SECRET", "") diff --git a/backend/internal/handler/admin/account_codex_import.go b/backend/internal/handler/admin/account_codex_import.go index 01a5fbfa1c..271bd19470 100644 --- a/backend/internal/handler/admin/account_codex_import.go +++ b/backend/internal/handler/admin/account_codex_import.go @@ -115,6 +115,10 @@ func (h *AccountHandler) ImportCodexSession(c *gin.Context) { response.BadRequest(c, "Invalid request: "+err.Error()) return } + if err := service.ValidateOpenAILongContextBillingExtra(service.PlatformOpenAI, req.Extra); err != nil { + response.ErrorFrom(c, err) + return + } if req.Concurrency != nil && *req.Concurrency < 0 { response.BadRequest(c, "concurrency must be >= 0") return diff --git a/backend/internal/handler/admin/account_codex_import_test.go b/backend/internal/handler/admin/account_codex_import_test.go index a52463aa86..96a033d8c3 100644 --- a/backend/internal/handler/admin/account_codex_import_test.go +++ b/backend/internal/handler/admin/account_codex_import_test.go @@ -630,6 +630,7 @@ func TestImportCodexSessionsAccessTokenOnlySameUserUpdatesExisting(t *testing.T) "chatgpt_user_id": "user-1", "access_token": existingToken, }, + Extra: map[string]any{"openai_long_context_billing_enabled": false}, }}) handler := NewAccountHandler(svc, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil) req := CodexSessionImportRequest{SkipDefaultGroupBind: boolPtr(true)} @@ -650,6 +651,9 @@ func TestImportCodexSessionsAccessTokenOnlySameUserUpdatesExisting(t *testing.T) if len(svc.updatedAccounts) != 1 || svc.updatedAccounts[0].id != 10 { t.Fatalf("updated accounts = %+v, want account 10", svc.updatedAccounts) } + if got := svc.updatedAccounts[0].input.Extra["openai_long_context_billing_enabled"]; got != false { + t.Fatalf("openai_long_context_billing_enabled = %v, want false", got) + } } func TestImportCodexSessionsUpgradesAccessTokenOnlyAccountWithRefreshToken(t *testing.T) { diff --git a/backend/internal/handler/admin/account_handler.go b/backend/internal/handler/admin/account_handler.go index a4b0773999..b886728159 100644 --- a/backend/internal/handler/admin/account_handler.go +++ b/backend/internal/handler/admin/account_handler.go @@ -784,6 +784,10 @@ func (h *AccountHandler) Create(c *gin.Context) { response.BadRequest(c, "Invalid request: "+err.Error()) return } + if err := service.ValidateOpenAILongContextBillingExtra(req.Platform, req.Extra); err != nil { + response.ErrorFrom(c, err) + return + } if req.RateMultiplier != nil && *req.RateMultiplier < 0 { response.BadRequest(c, "rate_multiplier must be >= 0") return @@ -1299,6 +1303,10 @@ func (h *AccountHandler) ApplyOAuthCredentials(c *gin.Context) { response.ErrorFrom(c, infraerrors.BadRequest("NOT_OAUTH", "cannot apply oauth credentials to non-OAuth account")) return } + if err := service.ValidateOpenAILongContextBillingExtra(existing.Platform, req.Extra); err != nil { + response.ErrorFrom(c, err) + return + } updatedAccount, err := h.adminService.UpdateAccount(ctx, accountID, &service.UpdateAccountInput{ Type: req.Type, @@ -1592,6 +1600,12 @@ func (h *AccountHandler) BatchCreate(c *gin.Context) { response.BadRequest(c, "Invalid request: "+err.Error()) return } + for _, item := range req.Accounts { + if err := service.ValidateOpenAILongContextBillingExtra(item.Platform, item.Extra); err != nil { + response.ErrorFrom(c, err) + return + } + } executeAdminIdempotentJSON(c, "admin.accounts.batch_create", req, service.DefaultWriteIdempotencyTTL(), func(ctx context.Context) (any, error) { success := 0 diff --git a/backend/internal/handler/admin/account_handler_long_context_billing_test.go b/backend/internal/handler/admin/account_handler_long_context_billing_test.go new file mode 100644 index 0000000000..d50513a3e8 --- /dev/null +++ b/backend/internal/handler/admin/account_handler_long_context_billing_test.go @@ -0,0 +1,165 @@ +package admin + +import ( + "bytes" + "encoding/json" + "net/http" + "net/http/httptest" + "testing" + + infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" + "github.com/Wei-Shaw/sub2api/internal/service" + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" +) + +func TestAccountAdminBoundariesRejectMalformedOpenAILongContextBillingValue(t *testing.T) { + const malformedExtra = `"extra":{"openai_long_context_billing_enabled":"true"}` + + tests := []struct { + name string + method string + path string + body string + mount func(*gin.Engine, *AccountHandler) + setup func(*stubAdminService) + }{ + { + name: "create", + method: http.MethodPost, + path: "/accounts", + body: `{"name":"account","platform":"openai","type":"apikey","credentials":{"api_key":"test"},` + malformedExtra + `}`, + mount: func(router *gin.Engine, handler *AccountHandler) { router.POST("/accounts", handler.Create) }, + }, + { + name: "update", + method: http.MethodPut, + path: "/accounts/1", + body: `{` + malformedExtra + `}`, + mount: func(router *gin.Engine, handler *AccountHandler) { router.PUT("/accounts/:id", handler.Update) }, + setup: func(stub *stubAdminService) { + stub.updateAccountErr = infraerrors.BadRequest("OPENAI_LONG_CONTEXT_BILLING_INVALID", "invalid") + }, + }, + { + name: "bulk update", + method: http.MethodPost, + path: "/accounts/bulk-update", + body: `{"account_ids":[1],` + malformedExtra + `}`, + mount: func(router *gin.Engine, handler *AccountHandler) { + router.POST("/accounts/bulk-update", handler.BulkUpdate) + }, + setup: func(stub *stubAdminService) { + stub.bulkUpdateAccountErr = infraerrors.BadRequest("OPENAI_LONG_CONTEXT_BILLING_INVALID", "invalid") + }, + }, + { + name: "batch create", + method: http.MethodPost, + path: "/accounts/batch", + body: `{"accounts":[{"name":"account","platform":"openai","type":"apikey","credentials":{"api_key":"test"},` + malformedExtra + `}]}`, + mount: func(router *gin.Engine, handler *AccountHandler) { router.POST("/accounts/batch", handler.BatchCreate) }, + }, + { + name: "Codex session import", + method: http.MethodPost, + path: "/accounts/import-codex-session", + body: `{"content":"token",` + malformedExtra + `}`, + mount: func(router *gin.Engine, handler *AccountHandler) { + router.POST("/accounts/import-codex-session", handler.ImportCodexSession) + }, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + gin.SetMode(gin.TestMode) + stub := newStubAdminService() + if tt.setup != nil { + tt.setup(stub) + } + handler := NewAccountHandler(stub, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil) + router := gin.New() + tt.mount(router, handler) + recorder := httptest.NewRecorder() + request := httptest.NewRequest(tt.method, tt.path, bytes.NewBufferString(tt.body)) + request.Header.Set("Content-Type", "application/json") + + router.ServeHTTP(recorder, request) + + require.Equal(t, http.StatusBadRequest, recorder.Code) + var responseBody struct { + Reason string `json:"reason"` + } + require.NoError(t, json.Unmarshal(recorder.Body.Bytes(), &responseBody)) + require.Equal(t, "OPENAI_LONG_CONTEXT_BILLING_INVALID", responseBody.Reason) + }) + } +} + +func TestAccountCreateBoundaryDoesNotApplyOpenAIValidationToOtherPlatforms(t *testing.T) { + gin.SetMode(gin.TestMode) + handler := NewAccountHandler(newStubAdminService(), nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil) + router := gin.New() + router.POST("/accounts", handler.Create) + recorder := httptest.NewRecorder() + request := httptest.NewRequest(http.MethodPost, "/accounts", bytes.NewBufferString( + `{"name":"account","platform":"anthropic","type":"apikey","credentials":{"api_key":"test"},"extra":{"openai_long_context_billing_enabled":"provider-owned"}}`, + )) + request.Header.Set("Content-Type", "application/json") + + router.ServeHTTP(recorder, request) + + require.Equal(t, http.StatusOK, recorder.Code) +} + +func TestApplyOAuthCredentialsRejectsMalformedOpenAILongContextBillingBeforeMutation(t *testing.T) { + gin.SetMode(gin.TestMode) + stub := newStubAdminService() + stub.getAccountResult = &service.Account{ + ID: 1, + Platform: service.PlatformOpenAI, + Type: service.AccountTypeOAuth, + } + handler := NewAccountHandler(stub, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil) + router := gin.New() + router.POST("/accounts/:id/apply-oauth-credentials", handler.ApplyOAuthCredentials) + recorder := httptest.NewRecorder() + request := httptest.NewRequest(http.MethodPost, "/accounts/1/apply-oauth-credentials", bytes.NewBufferString( + `{"type":"oauth","credentials":{"access_token":"new-token"},"extra":{"openai_long_context_billing_enabled":"true"}}`, + )) + request.Header.Set("Content-Type", "application/json") + + router.ServeHTTP(recorder, request) + + require.Equal(t, http.StatusBadRequest, recorder.Code) + var responseBody struct { + Reason string `json:"reason"` + } + require.NoError(t, json.Unmarshal(recorder.Body.Bytes(), &responseBody)) + require.Equal(t, "OPENAI_LONG_CONTEXT_BILLING_INVALID", responseBody.Reason) + require.Zero(t, stub.updateAccountCalls) + require.Zero(t, stub.updateAccountExtraCalls) +} + +func TestOpenAIOAuthCodexPATBoundaryRejectsMalformedOpenAILongContextBillingValueBeforeTokenValidation(t *testing.T) { + gin.SetMode(gin.TestMode) + handler := NewOpenAIOAuthHandler(nil, newStubAdminService(), nil) + router := gin.New() + router.Use(gin.Recovery()) + router.POST("/openai/create-from-codex-pat", handler.CreateAccountFromCodexPAT) + recorder := httptest.NewRecorder() + request := httptest.NewRequest(http.MethodPost, "/openai/create-from-codex-pat", bytes.NewBufferString( + `{"access_token":"token","extra":{"openai_long_context_billing_enabled":1}}`, + )) + request.Header.Set("Content-Type", "application/json") + + router.ServeHTTP(recorder, request) + + require.Equal(t, http.StatusBadRequest, recorder.Code) + var responseBody struct { + Reason string `json:"reason"` + } + require.NoError(t, json.Unmarshal(recorder.Body.Bytes(), &responseBody)) + require.Equal(t, "OPENAI_LONG_CONTEXT_BILLING_INVALID", responseBody.Reason) +} diff --git a/backend/internal/handler/admin/admin_service_stub_test.go b/backend/internal/handler/admin/admin_service_stub_test.go index 7a7cbb473e..5e9c4d517e 100644 --- a/backend/internal/handler/admin/admin_service_stub_test.go +++ b/backend/internal/handler/admin/admin_service_stub_test.go @@ -33,6 +33,9 @@ type stubAdminService struct { createSparkShadowErr error updateAccountErr error bulkUpdateAccountErr error + getAccountResult *service.Account + updateAccountCalls int + updateAccountExtraCalls int checkMixedErr error lastMixedCheck struct { accountID int64 @@ -388,6 +391,9 @@ func (s *stubAdminService) ListOpenAISchedulableAccountsForSchedulerScore(_ cont } func (s *stubAdminService) GetAccount(ctx context.Context, id int64) (*service.Account, error) { + if s.getAccountResult != nil { + return s.getAccountResult, nil + } account := service.Account{ID: id, Name: "account", Status: service.StatusActive} return &account, nil } @@ -413,6 +419,7 @@ func (s *stubAdminService) CreateAccount(ctx context.Context, input *service.Cre } func (s *stubAdminService) UpdateAccount(ctx context.Context, id int64, input *service.UpdateAccountInput) (*service.Account, error) { + s.updateAccountCalls++ if s.updateAccountErr != nil { return nil, s.updateAccountErr } @@ -421,6 +428,7 @@ func (s *stubAdminService) UpdateAccount(ctx context.Context, id int64, input *s } func (s *stubAdminService) UpdateAccountExtra(ctx context.Context, id int64, updates map[string]any) error { + s.updateAccountExtraCalls++ return nil } diff --git a/backend/internal/handler/admin/grok_oauth_handler.go b/backend/internal/handler/admin/grok_oauth_handler.go index a5841b9d5d..dfe5632e30 100644 --- a/backend/internal/handler/admin/grok_oauth_handler.go +++ b/backend/internal/handler/admin/grok_oauth_handler.go @@ -454,7 +454,7 @@ func (h *GrokOAuthHandler) QueryQuota(c *gin.Context) { response.BadRequest(c, "grok quota service is not enabled") return } - result, err := h.quotaService.ProbeUsage(c.Request.Context(), accountID) + result, err := h.quotaService.QueryQuota(c.Request.Context(), accountID) if err != nil { response.ErrorFrom(c, err) return diff --git a/backend/internal/handler/admin/grok_oauth_handler_test.go b/backend/internal/handler/admin/grok_oauth_handler_test.go index 6101a25d35..64ea044aa3 100644 --- a/backend/internal/handler/admin/grok_oauth_handler_test.go +++ b/backend/internal/handler/admin/grok_oauth_handler_test.go @@ -8,6 +8,7 @@ import ( "net/http" "net/http/httptest" "strings" + "sync" "testing" "time" @@ -41,17 +42,35 @@ func (r *grokQuotaHandlerAccountRepo) UpdateExtra(_ context.Context, id int64, u } type grokQuotaHandlerUpstream struct { - resp *http.Response - lastReq *http.Request - lastBody []byte + mu sync.Mutex + requests []*http.Request + bodies [][]byte } func (u *grokQuotaHandlerUpstream) Do(req *http.Request, _ string, _ int64, _ int) (*http.Response, error) { - u.lastReq = req + var body []byte if req.Body != nil { - u.lastBody, _ = io.ReadAll(req.Body) + body, _ = io.ReadAll(req.Body) } - return u.resp, nil + u.mu.Lock() + u.requests = append(u.requests, req) + u.bodies = append(u.bodies, body) + u.mu.Unlock() + if req.URL.Path == "/v1/responses" { + return &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{ + "X-Ratelimit-Limit-Requests": []string{"10"}, + "X-Ratelimit-Remaining-Requests": []string{"8"}, + }, + Body: io.NopCloser(strings.NewReader(`{"id":"resp_probe"}`)), + }, nil + } + payload := `{"config":{"billingPeriodStart":"2026-07-01T00:00:00Z","billingPeriodEnd":"2026-08-01T00:00:00Z"}}` + if req.URL.RawQuery == "format=credits" { + payload = `{"config":{"currentPeriod":{"type":"WEEKLY","start":"2026-07-09T03:25:00Z","end":"2026-07-16T03:25:00Z"}}}` + } + return &http.Response{StatusCode: http.StatusOK, Header: make(http.Header), Body: io.NopCloser(strings.NewReader(payload))}, nil } func (u *grokQuotaHandlerUpstream) DoWithTLS( @@ -77,14 +96,7 @@ func TestGrokOAuthHandlerQueryQuotaProbesUpstream(t *testing.T) { "expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339), }, }} - upstream := &grokQuotaHandlerUpstream{resp: &http.Response{ - StatusCode: http.StatusOK, - Header: http.Header{ - "X-Ratelimit-Limit-Requests": []string{"10"}, - "X-Ratelimit-Remaining-Requests": []string{"8"}, - }, - Body: io.NopCloser(strings.NewReader(`{"id":"resp_probe"}`)), - }} + upstream := &grokQuotaHandlerUpstream{} quotaService := service.NewGrokQuotaService(repo, nil, service.NewGrokTokenProvider(repo, nil), upstream) handler := NewGrokOAuthHandler(nil, nil, quotaService) @@ -95,12 +107,23 @@ func TestGrokOAuthHandlerQueryQuotaProbesUpstream(t *testing.T) { router.ServeHTTP(rec, req) require.Equal(t, http.StatusOK, rec.Code) - require.Contains(t, rec.Body.String(), `"source":"active_probe"`) + require.Contains(t, rec.Body.String(), `"source":"hybrid_probe"`) + require.Contains(t, rec.Body.String(), `"billing":`) + require.Contains(t, rec.Body.String(), `"snapshot":`) require.Contains(t, rec.Body.String(), `"headers_observed":true`) require.NotContains(t, rec.Body.String(), "access-token") - require.Equal(t, xai.DefaultCLIBaseURL+"/responses", upstream.lastReq.URL.String()) - require.Equal(t, "Bearer access-token", upstream.lastReq.Header.Get("Authorization")) - require.Contains(t, string(upstream.lastBody), `"store":false`) + upstream.mu.Lock() + requests := append([]*http.Request(nil), upstream.requests...) + bodies := append([][]byte(nil), upstream.bodies...) + upstream.mu.Unlock() + require.Len(t, requests, 3) + for i, upstreamReq := range requests { + require.Equal(t, "Bearer access-token", upstreamReq.Header.Get("Authorization")) + if upstreamReq.URL.String() == xai.DefaultCLIBaseURL+"/responses" { + require.Contains(t, string(bodies[i]), `"model":"grok-4.5"`) + require.Contains(t, string(bodies[i]), `"store":false`) + } + } require.NotNil(t, repo.updates[42]) } diff --git a/backend/internal/handler/admin/openai_oauth_handler.go b/backend/internal/handler/admin/openai_oauth_handler.go index d7a756bd00..78d57299b6 100644 --- a/backend/internal/handler/admin/openai_oauth_handler.go +++ b/backend/internal/handler/admin/openai_oauth_handler.go @@ -304,6 +304,10 @@ func (h *OpenAIOAuthHandler) CreateAccountFromCodexPAT(c *gin.Context) { response.BadRequest(c, "Invalid request: "+err.Error()) return } + if err := service.ValidateOpenAILongContextBillingExtra(service.PlatformOpenAI, req.Extra); err != nil { + response.ErrorFrom(c, err) + return + } if req.Concurrency != nil && *req.Concurrency < 0 { response.BadRequest(c, "concurrency must be >= 0") return diff --git a/backend/internal/handler/admin/ops_system_log_handler.go b/backend/internal/handler/admin/ops_system_log_handler.go index 9f3c8b893a..1b6af45976 100644 --- a/backend/internal/handler/admin/ops_system_log_handler.go +++ b/backend/internal/handler/admin/ops_system_log_handler.go @@ -15,6 +15,7 @@ import ( type opsSystemLogCleanupRequest struct { StartTime string `json:"start_time"` EndTime string `json:"end_time"` + Host string `json:"host"` Level string `json:"level"` Component string `json:"component"` @@ -56,6 +57,7 @@ func (h *OpsHandler) ListSystemLogs(c *gin.Context) { PageSize: pageSize, StartTime: &start, EndTime: &end, + Host: strings.TrimSpace(c.Query("host")), Level: strings.TrimSpace(c.Query("level")), Component: strings.TrimSpace(c.Query("component")), RequestID: strings.TrimSpace(c.Query("request_id")), @@ -153,6 +155,7 @@ func (h *OpsHandler) CleanupSystemLogs(c *gin.Context) { filter := &service.OpsSystemLogCleanupFilter{ StartTime: start, EndTime: end, + Host: strings.TrimSpace(req.Host), Level: strings.TrimSpace(req.Level), Component: strings.TrimSpace(req.Component), RequestID: strings.TrimSpace(req.RequestID), diff --git a/backend/internal/handler/admin/ops_system_log_handler_test.go b/backend/internal/handler/admin/ops_system_log_handler_test.go index 9557fce442..3390fbe3cb 100644 --- a/backend/internal/handler/admin/ops_system_log_handler_test.go +++ b/backend/internal/handler/admin/ops_system_log_handler_test.go @@ -2,6 +2,7 @@ package admin import ( "bytes" + "context" "encoding/json" "net/http" "net/http/httptest" @@ -19,6 +20,26 @@ type responseEnvelope struct { Data json.RawMessage `json:"data"` } +type opsSystemLogCaptureRepo struct { + service.OpsRepository + listFilter *service.OpsSystemLogFilter + cleanupFilter *service.OpsSystemLogCleanupFilter +} + +func (r *opsSystemLogCaptureRepo) ListSystemLogs(_ context.Context, filter *service.OpsSystemLogFilter) (*service.OpsSystemLogList, error) { + r.listFilter = filter + return &service.OpsSystemLogList{Logs: []*service.OpsSystemLog{}, Page: filter.Page, PageSize: filter.PageSize}, nil +} + +func (r *opsSystemLogCaptureRepo) DeleteSystemLogs(_ context.Context, filter *service.OpsSystemLogCleanupFilter) (int64, error) { + r.cleanupFilter = filter + return 1, nil +} + +func (r *opsSystemLogCaptureRepo) InsertSystemLogCleanupAudit(_ context.Context, _ *service.OpsSystemLogCleanupAudit) error { + return nil +} + func newOpsSystemLogTestRouter(handler *OpsHandler, withUser bool) *gin.Engine { gin.SetMode(gin.TestMode) r := gin.New() @@ -121,6 +142,23 @@ func TestOpsSystemLogHandler_ListSuccess(t *testing.T) { } } +func TestOpsSystemLogHandler_ListAcceptsHost(t *testing.T) { + repo := &opsSystemLogCaptureRepo{} + svc := service.NewOpsService(repo, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil) + h := NewOpsHandler(svc) + r := newOpsSystemLogTestRouter(h, false) + + w := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodGet, "/logs?host=api-node-1", nil) + r.ServeHTTP(w, req) + if w.Code != http.StatusOK { + t.Fatalf("status=%d, want 200", w.Code) + } + if repo.listFilter == nil || repo.listFilter.Host != "api-node-1" { + t.Fatalf("host filter = %+v, want api-node-1", repo.listFilter) + } +} + func TestOpsSystemLogHandler_CleanupUnauthorized(t *testing.T) { svc := service.NewOpsService(nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil) h := NewOpsHandler(svc) @@ -205,6 +243,24 @@ func TestOpsSystemLogHandler_CleanupAcceptsAPIKeyID(t *testing.T) { } } +func TestOpsSystemLogHandler_CleanupAcceptsHost(t *testing.T) { + repo := &opsSystemLogCaptureRepo{} + svc := service.NewOpsService(repo, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil) + h := NewOpsHandler(svc) + r := newOpsSystemLogTestRouter(h, true) + + w := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodPost, "/logs/cleanup", bytes.NewBufferString(`{"host":"api-node-1"}`)) + req.Header.Set("Content-Type", "application/json") + r.ServeHTTP(w, req) + if w.Code != http.StatusOK { + t.Fatalf("status=%d, want 200", w.Code) + } + if repo.cleanupFilter == nil || repo.cleanupFilter.Host != "api-node-1" { + t.Fatalf("host filter = %+v, want api-node-1", repo.cleanupFilter) + } +} + func TestOpsSystemLogHandler_CleanupInvalidAPIKeyID(t *testing.T) { svc := service.NewOpsService(nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil) h := NewOpsHandler(svc) diff --git a/backend/internal/handler/admin/ops_ws_handler.go b/backend/internal/handler/admin/ops_ws_handler.go index 75fd7ea002..e4c42cc9c0 100644 --- a/backend/internal/handler/admin/ops_ws_handler.go +++ b/backend/internal/handler/admin/ops_ws_handler.go @@ -16,6 +16,7 @@ import ( "time" "github.com/Wei-Shaw/sub2api/internal/pkg/logger" + servermiddleware "github.com/Wei-Shaw/sub2api/internal/server/middleware" "github.com/Wei-Shaw/sub2api/internal/service" "github.com/gin-gonic/gin" "github.com/gorilla/websocket" @@ -323,7 +324,7 @@ func (h *OpsHandler) QPSWSHandler(c *gin.Context) { // If realtime monitoring is disabled, prefer a successful WS upgrade followed by a clean close // with a deterministic close code. This prevents clients from spinning on 404/1006 reconnect loops. if !h.opsService.IsRealtimeMonitoringEnabled(c.Request.Context()) { - conn, err := upgrader.Upgrade(c.Writer, c.Request, nil) + conn, err := upgrader.Upgrade(c.Writer, c.Request, servermiddleware.ServerTimingResponseHeader(c)) if err != nil { c.JSON(http.StatusNotFound, gin.H{"error": "ops realtime monitoring is disabled"}) return @@ -358,7 +359,7 @@ func (h *OpsHandler) QPSWSHandler(c *gin.Context) { defer releaseOpsWSIPSlot(clientIP) } - conn, err := upgrader.Upgrade(c.Writer, c.Request, nil) + conn, err := upgrader.Upgrade(c.Writer, c.Request, servermiddleware.ServerTimingResponseHeader(c)) if err != nil { logger.LegacyPrintf("handler.admin.ops_ws", "[OpsWS] upgrade failed: %v", err) return diff --git a/backend/internal/handler/dto/mappers.go b/backend/internal/handler/dto/mappers.go index e770bcf036..3c45c3b95e 100644 --- a/backend/internal/handler/dto/mappers.go +++ b/backend/internal/handler/dto/mappers.go @@ -599,54 +599,55 @@ func usageLogFromServiceUser(l *service.UsageLog) UsageLog { requestedModel = l.Model } return UsageLog{ - ID: l.ID, - UserID: l.UserID, - APIKeyID: l.APIKeyID, - AccountID: l.AccountID, - RequestID: l.RequestID, - Model: requestedModel, - ServiceTier: l.ServiceTier, - ReasoningEffort: l.ReasoningEffort, - InboundEndpoint: l.InboundEndpoint, - GroupID: l.GroupID, - SubscriptionID: l.SubscriptionID, - InputTokens: l.InputTokens, - OutputTokens: l.OutputTokens, - CacheCreationTokens: l.CacheCreationTokens, - CacheReadTokens: l.CacheReadTokens, - CacheCreation5mTokens: l.CacheCreation5mTokens, - CacheCreation1hTokens: l.CacheCreation1hTokens, - InputCost: l.InputCost, - OutputCost: l.OutputCost, - CacheCreationCost: l.CacheCreationCost, - CacheReadCost: l.CacheReadCost, - TotalCost: l.TotalCost, - ActualCost: l.ActualCost, - RateMultiplier: l.RateMultiplier, - BillingType: l.BillingType, - RequestType: requestType.String(), - Stream: stream, - OpenAIWSMode: openAIWSMode, - DurationMs: l.DurationMs, - FirstTokenMs: l.FirstTokenMs, - ImageCount: l.ImageCount, - ImageSize: l.ImageSize, - ImageInputSize: l.ImageInputSize, - ImageOutputSize: l.ImageOutputSize, - ImageOutputTokens: l.ImageOutputTokens, - ImageOutputCost: l.ImageOutputCost, - ImageSizeSource: l.ImageSizeSource, - ImageSizeBreakdown: l.ImageSizeBreakdown, - MediaType: l.MediaType, - UserAgent: l.UserAgent, - IPAddress: l.IPAddress, - CacheTTLOverridden: l.CacheTTLOverridden, - BillingMode: l.BillingMode, - CreatedAt: l.CreatedAt, - User: UserFromServiceShallow(l.User), - APIKey: APIKeyFromService(l.APIKey), - Group: GroupFromServiceShallow(l.Group), - Subscription: UserSubscriptionFromService(l.Subscription), + ID: l.ID, + UserID: l.UserID, + APIKeyID: l.APIKeyID, + AccountID: l.AccountID, + RequestID: l.RequestID, + Model: requestedModel, + ServiceTier: l.ServiceTier, + ReasoningEffort: l.ReasoningEffort, + InboundEndpoint: l.InboundEndpoint, + GroupID: l.GroupID, + SubscriptionID: l.SubscriptionID, + InputTokens: l.InputTokens, + OutputTokens: l.OutputTokens, + CacheCreationTokens: l.CacheCreationTokens, + CacheReadTokens: l.CacheReadTokens, + CacheCreation5mTokens: l.CacheCreation5mTokens, + CacheCreation1hTokens: l.CacheCreation1hTokens, + InputCost: l.InputCost, + OutputCost: l.OutputCost, + CacheCreationCost: l.CacheCreationCost, + CacheReadCost: l.CacheReadCost, + TotalCost: l.TotalCost, + ActualCost: l.ActualCost, + RateMultiplier: l.RateMultiplier, + LongContextBillingApplied: l.LongContextBillingApplied, + BillingType: l.BillingType, + RequestType: requestType.String(), + Stream: stream, + OpenAIWSMode: openAIWSMode, + DurationMs: l.DurationMs, + FirstTokenMs: l.FirstTokenMs, + ImageCount: l.ImageCount, + ImageSize: l.ImageSize, + ImageInputSize: l.ImageInputSize, + ImageOutputSize: l.ImageOutputSize, + ImageOutputTokens: l.ImageOutputTokens, + ImageOutputCost: l.ImageOutputCost, + ImageSizeSource: l.ImageSizeSource, + ImageSizeBreakdown: l.ImageSizeBreakdown, + MediaType: l.MediaType, + UserAgent: l.UserAgent, + IPAddress: l.IPAddress, + CacheTTLOverridden: l.CacheTTLOverridden, + BillingMode: l.BillingMode, + CreatedAt: l.CreatedAt, + User: UserFromServiceShallow(l.User), + APIKey: APIKeyFromService(l.APIKey), + Group: GroupFromServiceShallow(l.Group), + Subscription: UserSubscriptionFromService(l.Subscription), } } diff --git a/backend/internal/handler/dto/types.go b/backend/internal/handler/dto/types.go index 0418e5bc3e..619926c1e4 100644 --- a/backend/internal/handler/dto/types.go +++ b/backend/internal/handler/dto/types.go @@ -482,13 +482,14 @@ type UsageLog struct { CacheCreation5mTokens int `json:"cache_creation_5m_tokens"` CacheCreation1hTokens int `json:"cache_creation_1h_tokens"` - InputCost float64 `json:"input_cost"` - OutputCost float64 `json:"output_cost"` - CacheCreationCost float64 `json:"cache_creation_cost"` - CacheReadCost float64 `json:"cache_read_cost"` - TotalCost float64 `json:"total_cost"` - ActualCost float64 `json:"actual_cost"` - RateMultiplier float64 `json:"rate_multiplier"` + InputCost float64 `json:"input_cost"` + OutputCost float64 `json:"output_cost"` + CacheCreationCost float64 `json:"cache_creation_cost"` + CacheReadCost float64 `json:"cache_read_cost"` + TotalCost float64 `json:"total_cost"` + ActualCost float64 `json:"actual_cost"` + RateMultiplier float64 `json:"rate_multiplier"` + LongContextBillingApplied bool `json:"long_context_billing_applied"` BillingType int8 `json:"billing_type"` RequestType string `json:"request_type"` diff --git a/backend/internal/handler/openai_codex_models_handler.go b/backend/internal/handler/openai_codex_models_handler.go index e64c555d14..1c1357cbfa 100644 --- a/backend/internal/handler/openai_codex_models_handler.go +++ b/backend/internal/handler/openai_codex_models_handler.go @@ -15,11 +15,13 @@ import ( // Codex CLI and the Codex desktop app refresh their model picker from // GET {base_url}/models?client_version=... (custom provider mode) or // GET /backend-api/codex/models (chatgpt_base_url mode). Both routes land -// here. The manifest is proxied verbatim from the ChatGPT backend with a -// schedulable OAuth account's credentials, so clients pointed at the gateway -// see the account's real, always-current model entitlements instead of a -// frozen local cache. +// here. The manifest is proxied verbatim from the selected account's ChatGPT +// backend or custom API key upstream. API key manifests use a short-lived, +// asynchronously revalidated cache to tolerate canceled client requests. func (h *OpenAIGatewayHandler) CodexModels(c *gin.Context) { + if c.Request.Context().Err() != nil { + return + } apiKey, ok := middleware2.GetAPIKeyFromContext(c) if !ok || apiKey.Group == nil { h.errorResponse(c, http.StatusUnauthorized, "invalid_request_error", "API key group is required") @@ -30,24 +32,54 @@ func (h *OpenAIGatewayHandler) CodexModels(c *gin.Context) { return } - account, err := h.gatewayService.SelectAccountForModel(c.Request.Context(), apiKey.GroupID, "", "") - if err != nil { - h.errorResponse(c, http.StatusServiceUnavailable, "upstream_error", "No available OpenAI accounts") - return + maxAccountSwitches := h.maxAccountSwitches + if maxAccountSwitches <= 0 { + maxAccountSwitches = 3 } + failedAccountIDs := make(map[int64]struct{}) + switchCount := 0 + var lastUpstreamErr error - manifest, err := h.gatewayService.FetchCodexModelsManifest(c.Request.Context(), account, c.Query("client_version"), c.GetHeader("If-None-Match")) - if err != nil { - h.errorResponse(c, infraerrors.Code(err), "upstream_error", infraerrors.Message(err)) - return - } + for { + account, err := h.gatewayService.SelectAccountForModelWithExclusions(c.Request.Context(), apiKey.GroupID, "", "", failedAccountIDs) + if err != nil { + if c.Request.Context().Err() != nil { + return + } + if lastUpstreamErr != nil { + h.errorResponse(c, infraerrors.Code(lastUpstreamErr), "upstream_error", infraerrors.Message(lastUpstreamErr)) + return + } + h.errorResponse(c, http.StatusServiceUnavailable, "upstream_error", "No available OpenAI accounts") + return + } - if manifest.ETag != "" { - c.Header("ETag", manifest.ETag) - } - if manifest.NotModified { - c.Status(http.StatusNotModified) + manifest, err := h.gatewayService.FetchCodexModelsManifest(c.Request.Context(), account, c.Query("client_version"), c.GetHeader("If-None-Match")) + if err != nil { + if c.Request.Context().Err() != nil { + return + } + if service.IsRetryableCodexModelsManifestError(err) && switchCount < maxAccountSwitches { + failedAccountIDs[account.ID] = struct{}{} + switchCount++ + lastUpstreamErr = err + continue + } + h.errorResponse(c, infraerrors.Code(err), "upstream_error", infraerrors.Message(err)) + return + } + if c.Request.Context().Err() != nil { + return + } + + if manifest.ETag != "" { + c.Header("ETag", manifest.ETag) + } + if manifest.NotModified { + c.Status(http.StatusNotModified) + return + } + c.Data(http.StatusOK, "application/json", manifest.Body) return } - c.Data(http.StatusOK, "application/json", manifest.Body) } diff --git a/backend/internal/handler/openai_codex_models_handler_test.go b/backend/internal/handler/openai_codex_models_handler_test.go new file mode 100644 index 0000000000..ba74a5869f --- /dev/null +++ b/backend/internal/handler/openai_codex_models_handler_test.go @@ -0,0 +1,288 @@ +package handler + +import ( + "context" + "errors" + "fmt" + "io" + "net" + "net/http" + "net/http/httptest" + "strings" + "sync" + "testing" + + "github.com/Wei-Shaw/sub2api/internal/config" + middleware2 "github.com/Wei-Shaw/sub2api/internal/server/middleware" + "github.com/Wei-Shaw/sub2api/internal/service" + "github.com/gin-gonic/gin" +) + +type codexModelsFailoverAccountRepo struct { + service.AccountRepository + accounts []service.Account +} + +func (r codexModelsFailoverAccountRepo) GetByID(_ context.Context, id int64) (*service.Account, error) { + for i := range r.accounts { + if r.accounts[i].ID == id { + account := r.accounts[i] + return &account, nil + } + } + return nil, service.ErrNoAvailableAccounts +} + +func (r codexModelsFailoverAccountRepo) ListSchedulableByPlatform(_ context.Context, platform string) ([]service.Account, error) { + accounts := make([]service.Account, 0, len(r.accounts)) + for _, account := range r.accounts { + if account.Platform == platform { + accounts = append(accounts, account) + } + } + return accounts, nil +} + +type codexModelsFailoverHTTPUpstream struct { + service.HTTPUpstream + mu sync.Mutex + accountIDs []int64 + firstErr error + firstStatus int + statuses map[int64]int +} + +func (u *codexModelsFailoverHTTPUpstream) Do(_ *http.Request, _ string, accountID int64, _ int) (*http.Response, error) { + u.mu.Lock() + u.accountIDs = append(u.accountIDs, accountID) + u.mu.Unlock() + + status, hasStatus := u.statuses[accountID] + if accountID == 1 || hasStatus { + if u.firstErr != nil { + return nil, u.firstErr + } + if !hasStatus { + status = u.firstStatus + } + if status == 0 { + status = http.StatusServiceUnavailable + } + return &http.Response{ + StatusCode: status, + Status: http.StatusText(status), + Header: make(http.Header), + Body: io.NopCloser(strings.NewReader( + `{"error":{"message":"No available OpenAI accounts","type":"upstream_error"}}`, + )), + }, nil + } + return &http.Response{ + StatusCode: http.StatusOK, + Status: "200 OK", + Header: make(http.Header), + Body: io.NopCloser(strings.NewReader(`{"models":[{"slug":"gpt-5.6-sol"}]}`)), + }, nil +} + +func (u *codexModelsFailoverHTTPUpstream) calls() []int64 { + u.mu.Lock() + defer u.mu.Unlock() + return append([]int64(nil), u.accountIDs...) +} + +func TestCodexModelsCanceledRequestDoesNotWriteResponse(t *testing.T) { + gin.SetMode(gin.TestMode) + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + ctx, cancel := context.WithCancel(context.Background()) + cancel() + c.Request = httptest.NewRequest(http.MethodGet, "/v1/models", nil).WithContext(ctx) + + h := &OpenAIGatewayHandler{} + h.CodexModels(c) + + if c.Writer.Written() { + t.Fatalf("canceled request wrote an HTTP response: status=%d body=%q", recorder.Code, recorder.Body.String()) + } +} + +func TestCodexModelsFailsOverFromRetryableUpstreamStatus(t *testing.T) { + retryableStatuses := []int{ + http.StatusTooManyRequests, + http.StatusInternalServerError, + http.StatusBadGateway, + http.StatusServiceUnavailable, + http.StatusGatewayTimeout, + } + for _, status := range retryableStatuses { + t.Run(http.StatusText(status), func(t *testing.T) { + handler, upstream, groupID := newCodexModelsFailoverTestHandler(status) + recorder := performCodexModelsRequest(t, handler, groupID) + + if got, want := upstream.calls(), []int64{1, 2}; !equalInt64Slices(got, want) { + t.Fatalf("upstream account calls: got %v, want %v", got, want) + } + if recorder.Code != http.StatusOK { + t.Fatalf("status: got %d, want %d; body=%s", recorder.Code, http.StatusOK, recorder.Body.String()) + } + if got, want := recorder.Body.String(), `{"models":[{"slug":"gpt-5.6-sol"}]}`; got != want { + t.Fatalf("body: got %q, want %q", got, want) + } + }) + } +} + +func TestCodexModelsFailsOverFromUpstreamTransportError(t *testing.T) { + handler, upstream, groupID := newCodexModelsFailoverTestHandler(http.StatusServiceUnavailable) + upstream.firstErr = &net.OpError{ + Op: "read", + Net: "tcp", + Err: errors.New("connection reset"), + } + recorder := performCodexModelsRequest(t, handler, groupID) + + if got, want := upstream.calls(), []int64{1, 2}; !equalInt64Slices(got, want) { + t.Fatalf("upstream account calls: got %v, want %v", got, want) + } + if recorder.Code != http.StatusOK { + t.Fatalf("status: got %d, want %d; body=%s", recorder.Code, http.StatusOK, recorder.Body.String()) + } +} + +func TestCodexModelsDoesNotFailOverFromPermanentUpstreamStatus(t *testing.T) { + statuses := []int{ + http.StatusBadRequest, + http.StatusUnauthorized, + http.StatusForbidden, + http.StatusNotFound, + 600, + } + for _, status := range statuses { + t.Run(fmt.Sprintf("status_%d", status), func(t *testing.T) { + handler, upstream, groupID := newCodexModelsFailoverTestHandler(status) + recorder := performCodexModelsRequest(t, handler, groupID) + + if got, want := upstream.calls(), []int64{1}; !equalInt64Slices(got, want) { + t.Fatalf("upstream account calls: got %v, want %v", got, want) + } + if recorder.Code != http.StatusBadGateway { + t.Fatalf("status: got %d, want %d; body=%s", recorder.Code, http.StatusBadGateway, recorder.Body.String()) + } + }) + } +} + +func TestCodexModelsDoesNotFailOverFromUpstreamConfigurationError(t *testing.T) { + handler, upstream, groupID := newCodexModelsFailoverTestHandler(http.StatusServiceUnavailable) + upstream.firstErr = errors.New("invalid proxy URL") + recorder := performCodexModelsRequest(t, handler, groupID) + + if got, want := upstream.calls(), []int64{1}; !equalInt64Slices(got, want) { + t.Fatalf("upstream account calls: got %v, want %v", got, want) + } + if recorder.Code != http.StatusBadGateway { + t.Fatalf("status: got %d, want %d; body=%s", recorder.Code, http.StatusBadGateway, recorder.Body.String()) + } +} + +func TestCodexModelsReturnsLastUpstreamErrorWhenAccountsAreExhausted(t *testing.T) { + handler, upstream, groupID := newCodexModelsFailoverTestHandler(http.StatusServiceUnavailable) + upstream.statuses = map[int64]int{ + 1: http.StatusServiceUnavailable, + 2: http.StatusGatewayTimeout, + } + recorder := performCodexModelsRequest(t, handler, groupID) + + if got, want := upstream.calls(), []int64{1, 2}; !equalInt64Slices(got, want) { + t.Fatalf("upstream account calls: got %v, want %v", got, want) + } + if recorder.Code != http.StatusBadGateway { + t.Fatalf("status: got %d, want %d; body=%s", recorder.Code, http.StatusBadGateway, recorder.Body.String()) + } + if body := recorder.Body.String(); !strings.Contains(body, "upstream error 504") { + t.Fatalf("body does not preserve the last upstream error: %s", body) + } +} + +func TestCodexModelsHonorsAccountSwitchLimit(t *testing.T) { + handler, upstream, groupID := newCodexModelsFailoverTestHandlerWithAccountCount(http.StatusServiceUnavailable, 4, 2) + upstream.statuses = map[int64]int{ + 1: http.StatusServiceUnavailable, + 2: http.StatusBadGateway, + 3: http.StatusGatewayTimeout, + 4: http.StatusInternalServerError, + } + recorder := performCodexModelsRequest(t, handler, groupID) + + if got, want := upstream.calls(), []int64{1, 2, 3}; !equalInt64Slices(got, want) { + t.Fatalf("upstream account calls: got %v, want %v", got, want) + } + if recorder.Code != http.StatusBadGateway { + t.Fatalf("status: got %d, want %d; body=%s", recorder.Code, http.StatusBadGateway, recorder.Body.String()) + } + if body := recorder.Body.String(); !strings.Contains(body, "upstream error 504") { + t.Fatalf("body does not preserve the limit-ending upstream error: %s", body) + } +} + +func newCodexModelsFailoverTestHandler(firstStatus int) (*OpenAIGatewayHandler, *codexModelsFailoverHTTPUpstream, int64) { + return newCodexModelsFailoverTestHandlerWithAccountCount(firstStatus, 2, 3) +} + +func newCodexModelsFailoverTestHandlerWithAccountCount(firstStatus, accountCount, maxSwitches int) (*OpenAIGatewayHandler, *codexModelsFailoverHTTPUpstream, int64) { + gin.SetMode(gin.TestMode) + groupID := int64(42) + accounts := make([]service.Account, 0, accountCount) + for i := 1; i <= accountCount; i++ { + accounts = append(accounts, service.Account{ + ID: int64(i), + Name: fmt.Sprintf("upstream-%d", i), + Platform: service.PlatformOpenAI, + Type: service.AccountTypeAPIKey, + Status: service.StatusActive, + Schedulable: true, + Priority: i - 1, + Concurrency: 1, + Credentials: map[string]any{ + "api_key": fmt.Sprintf("sk-%d", i), + "base_url": fmt.Sprintf("https://upstream-%d.example/v1", i), + }, + }) + } + upstream := &codexModelsFailoverHTTPUpstream{firstStatus: firstStatus} + cfg := &config.Config{RunMode: config.RunModeSimple} + gatewayService := service.NewOpenAIGatewayService( + codexModelsFailoverAccountRepo{accounts: accounts}, + nil, nil, nil, nil, nil, nil, cfg, nil, nil, nil, nil, nil, + upstream, + nil, nil, nil, nil, nil, nil, nil, nil, + ) + return &OpenAIGatewayHandler{gatewayService: gatewayService, maxAccountSwitches: maxSwitches}, upstream, groupID +} + +func performCodexModelsRequest(t *testing.T, handler *OpenAIGatewayHandler, groupID int64) *httptest.ResponseRecorder { + t.Helper() + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + c.Request = httptest.NewRequest(http.MethodGet, "/v1/models?client_version=0.144.0", nil) + c.Set(string(middleware2.ContextKeyAPIKey), &service.APIKey{ + GroupID: &groupID, + Group: &service.Group{ID: groupID, Platform: service.PlatformOpenAI}, + }) + + handler.CodexModels(c) + return recorder +} + +func equalInt64Slices(got, want []int64) bool { + if len(got) != len(want) { + return false + } + for i := range got { + if got[i] != want[i] { + return false + } + } + return true +} diff --git a/backend/internal/handler/openai_gateway_handler_test.go b/backend/internal/handler/openai_gateway_handler_test.go index 781c16b392..e4b594c0b8 100644 --- a/backend/internal/handler/openai_gateway_handler_test.go +++ b/backend/internal/handler/openai_gateway_handler_test.go @@ -415,11 +415,13 @@ func TestResolveOpenAIMessagesDispatchMappedModel(t *testing.T) { SonnetMappedModel: "gpt-5.2", ExactModelMappings: map[string]string{ "claude-sonnet-4-5-20250929": "gpt-5.4-mini-high", + "claude-fable-5": "gpt-5.6-sol", }, }, }, } require.Equal(t, "gpt-5.4-mini", resolveOpenAIMessagesDispatchMappedModel(apiKey, "claude-sonnet-4-5-20250929")) + require.Equal(t, "gpt-5.6-sol", resolveOpenAIMessagesDispatchMappedModel(apiKey, "claude-fable-5")) }) t.Run("uses_family_default_when_no_override", func(t *testing.T) { diff --git a/backend/internal/pkg/antigravity/client.go b/backend/internal/pkg/antigravity/client.go index e318d1cdaf..39b6d2c90c 100644 --- a/backend/internal/pkg/antigravity/client.go +++ b/backend/internal/pkg/antigravity/client.go @@ -17,6 +17,7 @@ import ( "github.com/Wei-Shaw/sub2api/internal/pkg/proxyurl" "github.com/Wei-Shaw/sub2api/internal/pkg/proxyutil" + "github.com/Wei-Shaw/sub2api/internal/pkg/servertiming" ) // ForbiddenError 表示上游返回 403 Forbidden @@ -279,7 +280,6 @@ func NewClient(proxyURL string) (*Client, error) { } client.Transport = transport } - return &Client{ httpClient: client, }, nil @@ -341,7 +341,7 @@ func (c *Client) ExchangeCode(ctx context.Context, code, codeVerifier string) (* } req.Header.Set("Content-Type", "application/x-www-form-urlencoded") - resp, err := c.httpClient.Do(req) + resp, err := servertiming.Do(c.httpClient, req) if err != nil { return nil, fmt.Errorf("token 交换请求失败: %w", err) } @@ -383,7 +383,7 @@ func (c *Client) RefreshToken(ctx context.Context, refreshToken string) (*TokenR } req.Header.Set("Content-Type", "application/x-www-form-urlencoded") - resp, err := c.httpClient.Do(req) + resp, err := servertiming.Do(c.httpClient, req) if err != nil { return nil, fmt.Errorf("token 刷新请求失败: %w", err) } @@ -414,7 +414,7 @@ func (c *Client) GetUserInfo(ctx context.Context, accessToken string) (*UserInfo } req.Header.Set("Authorization", "Bearer "+accessToken) - resp, err := c.httpClient.Do(req) + resp, err := servertiming.Do(c.httpClient, req) if err != nil { return nil, fmt.Errorf("用户信息请求失败: %w", err) } @@ -465,7 +465,7 @@ func (c *Client) LoadCodeAssist(ctx context.Context, accessToken string) (*LoadC req.Header.Set("Content-Type", "application/json") req.Header.Set("User-Agent", GetUserAgentForContext(ctx)) - resp, err := c.httpClient.Do(req) + resp, err := servertiming.Do(c.httpClient, req) if err != nil { lastErr = fmt.Errorf("loadCodeAssist 请求失败: %w", err) if shouldFallbackToNextURL(err, 0) && urlIdx < len(availableURLs)-1 { @@ -544,7 +544,7 @@ func (c *Client) OnboardUser(ctx context.Context, accessToken, tierID string) (s req.Header.Set("Content-Type", "application/json") req.Header.Set("User-Agent", GetUserAgentForContext(ctx)) - resp, err := c.httpClient.Do(req) + resp, err := servertiming.Do(c.httpClient, req) if err != nil { lastErr = fmt.Errorf("onboardUser 请求失败: %w", err) if shouldFallbackToNextURL(err, 0) && urlIdx < len(availableURLs)-1 { @@ -683,7 +683,7 @@ func (c *Client) FetchAvailableModels(ctx context.Context, accessToken, projectI req.Header.Set("Content-Type", "application/json") req.Header.Set("User-Agent", GetUserAgentForContext(ctx)) - resp, err := fetchClient.Do(req) + resp, err := servertiming.Do(fetchClient, req) if err != nil { lastErr = fmt.Errorf("fetchAvailableModels 请求失败: %w", err) if shouldFallbackToNextURL(err, 0) && urlIdx < len(availableURLs)-1 { @@ -842,7 +842,7 @@ func (c *Client) SetUserSettings(ctx context.Context, accessToken string) (*SetU req.Header.Set("X-Goog-Api-Client", "gl-node/22.21.1") req.Host = "daily-cloudcode-pa.googleapis.com" - resp, err := c.httpClient.Do(req) + resp, err := servertiming.Do(c.httpClient, req) if err != nil { return nil, fmt.Errorf("setUserSettings 请求失败: %w", err) } @@ -885,7 +885,7 @@ func (c *Client) FetchUserInfo(ctx context.Context, accessToken, projectID strin req.Header.Set("X-Goog-Api-Client", "gl-node/22.21.1") req.Host = "daily-cloudcode-pa.googleapis.com" - resp, err := c.httpClient.Do(req) + resp, err := servertiming.Do(c.httpClient, req) if err != nil { return nil, fmt.Errorf("fetchUserInfo 请求失败: %w", err) } diff --git a/backend/internal/pkg/httpclient/pool.go b/backend/internal/pkg/httpclient/pool.go index 12804cc67d..22d3c65feb 100644 --- a/backend/internal/pkg/httpclient/pool.go +++ b/backend/internal/pkg/httpclient/pool.go @@ -25,6 +25,7 @@ import ( "github.com/Wei-Shaw/sub2api/internal/pkg/proxyurl" "github.com/Wei-Shaw/sub2api/internal/pkg/proxyutil" + "github.com/Wei-Shaw/sub2api/internal/pkg/servertiming" "github.com/Wei-Shaw/sub2api/internal/util/urlvalidator" ) @@ -92,6 +93,7 @@ func buildClient(opts Options) (*http.Client, error) { if opts.ValidateResolvedIP && !opts.AllowPrivateHosts { rt = newValidatedTransport(transport) } + rt = servertiming.WrapRoundTripper(rt) return &http.Client{ Transport: rt, Timeout: opts.Timeout, diff --git a/backend/internal/pkg/servertiming/collector.go b/backend/internal/pkg/servertiming/collector.go new file mode 100644 index 0000000000..553edede31 --- /dev/null +++ b/backend/internal/pkg/servertiming/collector.go @@ -0,0 +1,348 @@ +package servertiming + +import ( + "context" + "fmt" + "sort" + "strconv" + "strings" + "sync" + "time" +) + +const ( + HeaderName = "Server-Timing" + AdminUIHeader = "X-Admin-UI-Request" + MetricDatabase = "db" + MetricRedis = "redis" + dependencyPrefix = "dep_" + + maxMetricNameLength = 48 + maxIntervals = 2048 + maxHeaderLength = 4096 +) + +type contextKey struct{} + +type interval struct { + start time.Time + end time.Time +} + +type metric struct { + count int64 + intervals []interval +} + +// Collector stores request-scoped timing samples. It is safe for concurrent use. +type Collector struct { + startedAt time.Time + + mu sync.Mutex + metrics map[string]*metric + cacheStatus string +} + +// New creates a collector whose total duration starts at startedAt. +func New(startedAt time.Time) *Collector { + if startedAt.IsZero() { + startedAt = time.Now() + } + return &Collector{ + startedAt: startedAt, + metrics: make(map[string]*metric), + } +} + +// WithCollector attaches a collector to a context. +func WithCollector(ctx context.Context, collector *Collector) context.Context { + if ctx == nil { + ctx = context.Background() + } + if collector == nil { + return ctx + } + return context.WithValue(ctx, contextKey{}, collector) +} + +// FromContext returns the request timing collector, when one is active. +func FromContext(ctx context.Context) (*Collector, bool) { + if ctx == nil { + return nil, false + } + collector, ok := ctx.Value(contextKey{}).(*Collector) + return collector, ok && collector != nil +} + +// Active reports whether timing collection is enabled for this request. +func Active(ctx context.Context) bool { + _, ok := FromContext(ctx) + return ok +} + +// Record adds a completed interval and operation count to a metric. +func Record(ctx context.Context, name string, startedAt, endedAt time.Time, count int) { + collector, ok := FromContext(ctx) + if !ok { + return + } + collector.Record(name, startedAt, endedAt, count) +} + +// RecordInterval adds timing without incrementing the operation count. It is +// useful when one logical operation has multiple blocking driver calls. +func RecordInterval(ctx context.Context, name string, startedAt, endedAt time.Time) { + collector, ok := FromContext(ctx) + if !ok { + return + } + collector.record(name, startedAt, endedAt, 0) +} + +// Record adds a completed interval directly to the collector. +func (c *Collector) Record(name string, startedAt, endedAt time.Time, count int) { + if count <= 0 { + count = 1 + } + c.record(name, startedAt, endedAt, count) +} + +func (c *Collector) record(name string, startedAt, endedAt time.Time, count int) { + name = normalizeMetricName(name) + if c == nil || name == "" || startedAt.IsZero() || endedAt.Before(startedAt) { + return + } + if count < 0 { + count = 0 + } + + c.mu.Lock() + m := c.metrics[name] + if m == nil { + m = &metric{} + c.metrics[name] = m + } + m.count += int64(count) + if len(m.intervals) < maxIntervals { + m.intervals = append(m.intervals, interval{start: startedAt, end: endedAt}) + } + c.mu.Unlock() +} + +// Observe starts a metric span and returns an idempotent completion function. +func Observe(ctx context.Context, name string) func() { + collector, ok := FromContext(ctx) + name = normalizeMetricName(name) + if !ok || name == "" { + return func() {} + } + startedAt := time.Now() + var once sync.Once + return func() { + once.Do(func() { + collector.Record(name, startedAt, time.Now(), 1) + }) + } +} + +// ObserveDependency starts a named external dependency span. +func ObserveDependency(ctx context.Context, module string) func() { + return Observe(ctx, dependencyMetricName(module)) +} + +// RecordDependency records a completed external dependency interval. +func RecordDependency(ctx context.Context, module string, startedAt, endedAt time.Time) { + Record(ctx, dependencyMetricName(module), startedAt, endedAt, 1) +} + +// SetCacheStatus records the response-cache outcome for the request. +func SetCacheStatus(ctx context.Context, status string) { + collector, ok := FromContext(ctx) + if !ok { + return + } + status = normalizeCacheStatus(status) + if status == "" { + return + } + collector.mu.Lock() + collector.cacheStatus = status + collector.mu.Unlock() +} + +// HeaderValue renders a bounded, deterministic Server-Timing header. +func HeaderValue(ctx context.Context, endedAt time.Time, cacheStatus string) string { + collector, ok := FromContext(ctx) + if !ok { + return "" + } + return collector.HeaderValue(endedAt, cacheStatus) +} + +// HeaderValue renders a bounded, deterministic Server-Timing header. +func (c *Collector) HeaderValue(endedAt time.Time, cacheStatus string) string { + if c == nil { + return "" + } + if endedAt.IsZero() { + endedAt = time.Now() + } + if endedAt.Before(c.startedAt) { + endedAt = c.startedAt + } + + c.mu.Lock() + metrics := make(map[string]metric, len(c.metrics)) + allIntervals := make([]interval, 0) + dependencyIntervals := make([]interval, 0) + var dependencyCount int64 + for name, source := range c.metrics { + copied := metric{count: source.count, intervals: append([]interval(nil), source.intervals...)} + metrics[name] = copied + allIntervals = append(allIntervals, copied.intervals...) + if strings.HasPrefix(name, dependencyPrefix) { + dependencyIntervals = append(dependencyIntervals, copied.intervals...) + dependencyCount += copied.count + } + } + storedCacheStatus := c.cacheStatus + c.mu.Unlock() + + total := endedAt.Sub(c.startedAt) + blocked := unionDuration(allIntervals, c.startedAt, endedAt) + app := total - blocked + if app < 0 { + app = 0 + } + + cacheStatus = normalizeCacheStatus(cacheStatus) + if cacheStatus == "" { + cacheStatus = normalizeCacheStatus(storedCacheStatus) + } + if cacheStatus == "" { + cacheStatus = "bypass" + } + + database := metrics[MetricDatabase] + redisMetric := metrics[MetricRedis] + parts := []string{ + "total;dur=" + formatDuration(total), + "app;dur=" + formatDuration(app), + fmt.Sprintf("db;dur=%s;desc=\"queries=%d\"", formatDuration(unionDuration(database.intervals, c.startedAt, endedAt)), database.count), + fmt.Sprintf("redis;dur=%s;desc=\"commands=%d\"", formatDuration(unionDuration(redisMetric.intervals, c.startedAt, endedAt)), redisMetric.count), + "cache;desc=\"" + cacheStatus + "\"", + fmt.Sprintf("deps;dur=%s;desc=\"calls=%d\"", formatDuration(unionDuration(dependencyIntervals, c.startedAt, endedAt)), dependencyCount), + } + + dependencyNames := make([]string, 0) + for name := range metrics { + if strings.HasPrefix(name, dependencyPrefix) { + dependencyNames = append(dependencyNames, name) + } + } + sort.Strings(dependencyNames) + for _, name := range dependencyNames { + m := metrics[name] + part := fmt.Sprintf("%s;dur=%s;desc=\"calls=%d\"", name, formatDuration(unionDuration(m.intervals, c.startedAt, endedAt)), m.count) + candidate := strings.Join(append(parts, part), ", ") + if len(candidate) > maxHeaderLength { + break + } + parts = append(parts, part) + } + + return strings.Join(parts, ", ") +} + +func dependencyMetricName(module string) string { + module = normalizeMetricName(module) + module = strings.TrimPrefix(module, dependencyPrefix) + if module == "" { + module = "http" + } + return dependencyPrefix + module +} + +func normalizeMetricName(name string) string { + name = strings.ToLower(strings.TrimSpace(name)) + if name == "" { + return "" + } + var b strings.Builder + b.Grow(min(len(name), maxMetricNameLength)) + for _, r := range name { + if b.Len() >= maxMetricNameLength { + break + } + switch { + case r >= 'a' && r <= 'z', r >= '0' && r <= '9': + _, _ = b.WriteRune(r) + case r == '_' || r == '-': + _ = b.WriteByte('_') + } + } + return strings.Trim(b.String(), "_") +} + +func normalizeCacheStatus(status string) string { + switch strings.ToLower(strings.TrimSpace(status)) { + case "hit": + return "hit" + case "miss": + return "miss" + case "bypass": + return "bypass" + default: + return "" + } +} + +func unionDuration(intervals []interval, lowerBound, upperBound time.Time) time.Duration { + if len(intervals) == 0 || !upperBound.After(lowerBound) { + return 0 + } + normalized := make([]interval, 0, len(intervals)) + for _, item := range intervals { + start := item.start + end := item.end + if start.Before(lowerBound) { + start = lowerBound + } + if end.After(upperBound) { + end = upperBound + } + if end.After(start) { + normalized = append(normalized, interval{start: start, end: end}) + } + } + if len(normalized) == 0 { + return 0 + } + sort.Slice(normalized, func(i, j int) bool { + return normalized[i].start.Before(normalized[j].start) + }) + + currentStart := normalized[0].start + currentEnd := normalized[0].end + var total time.Duration + for _, item := range normalized[1:] { + if !item.start.After(currentEnd) { + if item.end.After(currentEnd) { + currentEnd = item.end + } + continue + } + total += currentEnd.Sub(currentStart) + currentStart = item.start + currentEnd = item.end + } + total += currentEnd.Sub(currentStart) + return total +} + +func formatDuration(value time.Duration) string { + if value < 0 { + value = 0 + } + return strconv.FormatFloat(float64(value)/float64(time.Millisecond), 'f', 1, 64) +} diff --git a/backend/internal/pkg/servertiming/collector_test.go b/backend/internal/pkg/servertiming/collector_test.go new file mode 100644 index 0000000000..1bb809f9bc --- /dev/null +++ b/backend/internal/pkg/servertiming/collector_test.go @@ -0,0 +1,129 @@ +package servertiming + +import ( + "context" + "fmt" + "strings" + "sync" + "testing" + "time" +) + +func TestCollectorHeaderValueAggregatesIntervals(t *testing.T) { + startedAt := time.Unix(100, 0) + collector := New(startedAt) + collector.Record(MetricDatabase, startedAt.Add(10*time.Millisecond), startedAt.Add(40*time.Millisecond), 2) + collector.Record(MetricRedis, startedAt.Add(30*time.Millisecond), startedAt.Add(50*time.Millisecond), 3) + collector.Record(dependencyMetricName("openai"), startedAt.Add(70*time.Millisecond), startedAt.Add(100*time.Millisecond), 1) + collector.Record(dependencyMetricName("github"), startedAt.Add(60*time.Millisecond), startedAt.Add(90*time.Millisecond), 1) + + got := collector.HeaderValue(startedAt.Add(120*time.Millisecond), "miss") + want := `total;dur=120.0, app;dur=40.0, db;dur=30.0;desc="queries=2", redis;dur=20.0;desc="commands=3", cache;desc="miss", deps;dur=40.0;desc="calls=2", dep_github;dur=30.0;desc="calls=1", dep_openai;dur=30.0;desc="calls=1"` + if got != want { + t.Fatalf("HeaderValue() = %q, want %q", got, want) + } +} + +func TestRecordIntervalDoesNotIncrementCount(t *testing.T) { + startedAt := time.Unix(200, 0) + collector := New(startedAt) + ctx := WithCollector(context.Background(), collector) + + Record(ctx, MetricDatabase, startedAt.Add(10*time.Millisecond), startedAt.Add(20*time.Millisecond), 1) + RecordInterval(ctx, MetricDatabase, startedAt.Add(30*time.Millisecond), startedAt.Add(40*time.Millisecond)) + + header := HeaderValue(ctx, startedAt.Add(100*time.Millisecond), "hit") + if !strings.Contains(header, `db;dur=20.0;desc="queries=1"`) { + t.Fatalf("header %q does not contain one query with both blocking intervals", header) + } + if !strings.Contains(header, "app;dur=80.0") { + t.Fatalf("header %q does not subtract the interval union from app time", header) + } +} + +func TestCollectorCacheStatusFallback(t *testing.T) { + startedAt := time.Unix(300, 0) + collector := New(startedAt) + ctx := WithCollector(context.Background(), collector) + + SetCacheStatus(ctx, " HIT ") + if got := HeaderValue(ctx, startedAt.Add(time.Millisecond), "invalid"); !strings.Contains(got, `cache;desc="hit"`) { + t.Fatalf("HeaderValue() = %q, want stored cache hit", got) + } + + other := New(startedAt) + if got := other.HeaderValue(startedAt.Add(time.Millisecond), "invalid"); !strings.Contains(got, `cache;desc="bypass"`) { + t.Fatalf("HeaderValue() = %q, want cache bypass", got) + } +} + +func TestCollectorSanitizesDependencyMetric(t *testing.T) { + startedAt := time.Unix(400, 0) + collector := New(startedAt) + ctx := WithCollector(context.Background(), collector) + RecordDependency(ctx, "GitHub API\r\nInjected;dur=999", startedAt, startedAt.Add(time.Millisecond)) + + header := HeaderValue(ctx, startedAt.Add(2*time.Millisecond), "bypass") + if strings.ContainsAny(header, "\r\n") || strings.Contains(header, ";dur=999") { + t.Fatalf("unsafe metric content reached header: %q", header) + } + if !strings.Contains(header, "dep_githubapiinjecteddur999;dur=1.0") { + t.Fatalf("sanitized dependency metric missing from header: %q", header) + } +} + +func TestCollectorBoundsHeaderLength(t *testing.T) { + startedAt := time.Unix(500, 0) + collector := New(startedAt) + for i := 0; i < 300; i++ { + collector.Record( + dependencyMetricName(fmt.Sprintf("module_%03d_with_a_deliberately_long_name", i)), + startedAt, + startedAt.Add(time.Millisecond), + 1, + ) + } + + header := collector.HeaderValue(startedAt.Add(2*time.Millisecond), "bypass") + if len(header) > maxHeaderLength { + t.Fatalf("header length = %d, want <= %d", len(header), maxHeaderLength) + } + if !strings.Contains(header, "total;dur=2.0") || !strings.Contains(header, "deps;dur=1.0") { + t.Fatalf("bounded header lost fixed metrics: %q", header) + } +} + +func TestCollectorConcurrentRecording(t *testing.T) { + startedAt := time.Now() + collector := New(startedAt) + ctx := WithCollector(context.Background(), collector) + + const workers = 25 + const recordsPerWorker = 100 + var wg sync.WaitGroup + wg.Add(workers) + for i := 0; i < workers; i++ { + go func() { + defer wg.Done() + for j := 0; j < recordsPerWorker; j++ { + Record(ctx, MetricDatabase, startedAt, startedAt.Add(time.Microsecond), 1) + } + }() + } + wg.Wait() + + header := HeaderValue(ctx, startedAt.Add(time.Millisecond), "bypass") + want := fmt.Sprintf(`queries=%d`, workers*recordsPerWorker) + if !strings.Contains(header, want) { + t.Fatalf("header %q does not contain %q", header, want) + } +} + +func TestContextHelpersHandleMissingCollector(t *testing.T) { + if Active(context.Background()) { + t.Fatal("context without collector reported active") + } + if got := HeaderValue(context.Background(), time.Now(), "hit"); got != "" { + t.Fatalf("HeaderValue() = %q without collector, want empty", got) + } +} diff --git a/backend/internal/pkg/servertiming/http.go b/backend/internal/pkg/servertiming/http.go new file mode 100644 index 0000000000..e326e24302 --- /dev/null +++ b/backend/internal/pkg/servertiming/http.go @@ -0,0 +1,104 @@ +package servertiming + +import ( + "context" + "net/http" + "strings" + "time" +) + +type dependencyModuleKey struct{} + +type timingRoundTripper struct { + base http.RoundTripper +} + +// WithDependencyModule overrides the safe module name used for an outbound call. +func WithDependencyModule(ctx context.Context, module string) context.Context { + if ctx == nil { + ctx = context.Background() + } + module = strings.TrimPrefix(normalizeMetricName(module), dependencyPrefix) + if module == "" { + return ctx + } + return context.WithValue(ctx, dependencyModuleKey{}, module) +} + +// WrapRoundTripper records outbound response-header latency for active requests. +func WrapRoundTripper(base http.RoundTripper) http.RoundTripper { + if base == nil { + base = http.DefaultTransport + } + if _, ok := base.(*timingRoundTripper); ok { + return base + } + return &timingRoundTripper{base: base} +} + +// InstrumentClient returns a shallow client copy with an instrumented transport. +func InstrumentClient(client *http.Client) *http.Client { + if client == nil { + client = &http.Client{} + } + copyClient := *client + copyClient.Transport = WrapRoundTripper(copyClient.Transport) + return ©Client +} + +// Do records response-header latency without changing the client's transport +// type. Use it for clients whose callers inspect or configure *http.Transport. +func Do(client *http.Client, req *http.Request) (*http.Response, error) { + if client == nil { + client = http.DefaultClient + } + if req == nil || !Active(req.Context()) { + return client.Do(req) + } + startedAt := time.Now() + response, err := client.Do(req) + RecordDependency(req.Context(), dependencyModule(req), startedAt, time.Now()) + return response, err +} + +func (t *timingRoundTripper) RoundTrip(req *http.Request) (*http.Response, error) { + if req == nil || !Active(req.Context()) { + return t.base.RoundTrip(req) + } + startedAt := time.Now() + response, err := t.base.RoundTrip(req) + RecordDependency(req.Context(), dependencyModule(req), startedAt, time.Now()) + return response, err +} + +func dependencyModule(req *http.Request) string { + if req != nil { + if module, ok := req.Context().Value(dependencyModuleKey{}).(string); ok && module != "" { + return module + } + } + if req == nil || req.URL == nil { + return "http" + } + host := strings.ToLower(req.URL.Hostname()) + switch { + case strings.Contains(host, "github"): + return "github" + case strings.Contains(host, "openai"): + return "openai" + case strings.Contains(host, "anthropic"): + return "anthropic" + case strings.Contains(host, "generativelanguage") || strings.Contains(host, "gemini"): + return "gemini" + case strings.Contains(host, "cloudcode") || strings.Contains(host, "antigravity"): + return "antigravity" + case strings.Contains(host, "googleapis") || strings.Contains(host, "google"): + return "google" + case strings.Contains(host, "amazonaws") || strings.Contains(host, "cloudflarestorage") || strings.Contains(host, "s3"): + return "s3" + case strings.Contains(host, "stripe") || strings.Contains(host, "airwallex") || strings.Contains(host, "alipay") || strings.Contains(host, "wechatpay") || strings.Contains(host, "paypal"): + return "payment" + default: + return "http" + } +} diff --git a/backend/internal/pkg/servertiming/http_test.go b/backend/internal/pkg/servertiming/http_test.go new file mode 100644 index 0000000000..d37f378414 --- /dev/null +++ b/backend/internal/pkg/servertiming/http_test.go @@ -0,0 +1,168 @@ +package servertiming + +import ( + "context" + "io" + "net/http" + "strings" + "testing" + "time" +) + +type roundTripFunc func(*http.Request) (*http.Response, error) + +func (f roundTripFunc) RoundTrip(req *http.Request) (*http.Response, error) { + return f(req) +} + +type trackingBody struct { + read bool +} + +func (b *trackingBody) Read(_ []byte) (int, error) { + b.read = true + return 0, io.EOF +} + +func (b *trackingBody) Close() error { return nil } + +func TestWrapRoundTripperRecordsResponseHeaderLatency(t *testing.T) { + startedAt := time.Now() + collector := New(startedAt) + body := &trackingBody{} + baseCalled := false + base := roundTripFunc(func(req *http.Request) (*http.Response, error) { + baseCalled = true + return &http.Response{ + StatusCode: http.StatusOK, + Body: body, + Header: make(http.Header), + Request: req, + }, nil + }) + req, err := http.NewRequestWithContext(WithCollector(context.Background(), collector), http.MethodGet, "https://api.github.com/repos/example/project", nil) + if err != nil { + t.Fatal(err) + } + + resp, err := WrapRoundTripper(base).RoundTrip(req) + if err != nil { + t.Fatal(err) + } + defer func() { _ = resp.Body.Close() }() + if !baseCalled { + t.Fatal("base RoundTripper was not called") + } + if body.read { + t.Fatal("RoundTripper instrumentation read the response body; timing must stop at response headers") + } + header := collector.HeaderValue(time.Now(), "bypass") + if !strings.Contains(header, `dep_github;dur=`) || !strings.Contains(header, `deps;dur=`) { + t.Fatalf("dependency metrics missing from header: %q", header) + } +} + +func TestWrapRoundTripperUsesContextModuleOverride(t *testing.T) { + collector := New(time.Now()) + ctx := WithDependencyModule(WithCollector(context.Background(), collector), "data-managementd") + req, err := http.NewRequestWithContext(ctx, http.MethodGet, "https://private.example.test/path", nil) + if err != nil { + t.Fatal(err) + } + base := roundTripFunc(func(req *http.Request) (*http.Response, error) { + return &http.Response{StatusCode: http.StatusOK, Body: http.NoBody, Header: make(http.Header), Request: req}, nil + }) + + if _, err := WrapRoundTripper(base).RoundTrip(req); err != nil { + t.Fatal(err) + } + header := collector.HeaderValue(time.Now(), "bypass") + if !strings.Contains(header, "dep_data_managementd") { + t.Fatalf("module override missing from header: %q", header) + } + if strings.Contains(header, "private.example") { + t.Fatalf("raw host leaked into header: %q", header) + } +} + +func TestWrapRoundTripperSkipsInactiveContext(t *testing.T) { + called := false + base := roundTripFunc(func(req *http.Request) (*http.Response, error) { + called = true + return &http.Response{StatusCode: http.StatusNoContent, Body: http.NoBody, Header: make(http.Header), Request: req}, nil + }) + req, err := http.NewRequest(http.MethodGet, "https://api.openai.com/v1/models", nil) + if err != nil { + t.Fatal(err) + } + if _, err := WrapRoundTripper(base).RoundTrip(req); err != nil { + t.Fatal(err) + } + if !called { + t.Fatal("inactive request did not reach base RoundTripper") + } +} + +func TestDoRecordsWithoutChangingTransportType(t *testing.T) { + base := roundTripFunc(func(req *http.Request) (*http.Response, error) { + return &http.Response{StatusCode: http.StatusOK, Body: http.NoBody, Header: make(http.Header), Request: req}, nil + }) + client := &http.Client{Transport: base} + collector := New(time.Now()) + req, err := http.NewRequestWithContext(WithCollector(context.Background(), collector), http.MethodGet, "https://api.openai.com/v1/models", nil) + if err != nil { + t.Fatal(err) + } + if _, err := Do(client, req); err != nil { + t.Fatal(err) + } + if _, ok := client.Transport.(roundTripFunc); !ok { + t.Fatalf("Do changed client transport type to %T", client.Transport) + } + if header := collector.HeaderValue(time.Now(), "bypass"); !strings.Contains(header, "dep_openai;dur=") { + t.Fatalf("dependency metric missing from header: %q", header) + } +} + +func TestDependencyModuleClassification(t *testing.T) { + tests := map[string]string{ + "https://api.github.com/repos/a/b": "github", + "https://api.openai.com/v1/models": "openai", + "https://api.anthropic.com/v1/messages": "anthropic", + "https://generativelanguage.googleapis.com/v1/models": "gemini", + "https://cloudcode-pa.googleapis.com/v1internal": "antigravity", + "https://storage.googleapis.com/bucket/object": "google", + "https://bucket.s3.amazonaws.com/object": "s3", + "https://api.stripe.com/v1/refunds": "payment", + "https://dependency.example.test/path": "http", + } + for rawURL, want := range tests { + req, err := http.NewRequest(http.MethodGet, rawURL, nil) + if err != nil { + t.Fatalf("NewRequest(%q): %v", rawURL, err) + } + if got := dependencyModule(req); got != want { + t.Errorf("dependencyModule(%q) = %q, want %q", rawURL, got, want) + } + } +} + +func TestClientInstrumentationDoesNotMutateOriginal(t *testing.T) { + base := roundTripFunc(func(req *http.Request) (*http.Response, error) { + return &http.Response{StatusCode: http.StatusOK, Body: http.NoBody, Header: make(http.Header), Request: req}, nil + }) + original := &http.Client{Transport: base, Timeout: time.Second} + instrumented := InstrumentClient(original) + if instrumented == original { + t.Fatal("InstrumentClient returned the original client") + } + if _, ok := original.Transport.(roundTripFunc); !ok { + t.Fatalf("InstrumentClient mutated the original transport to %T", original.Transport) + } + if instrumented.Timeout != original.Timeout { + t.Fatal("InstrumentClient did not preserve client settings") + } + if WrapRoundTripper(instrumented.Transport) != instrumented.Transport { + t.Fatal("WrapRoundTripper wrapped an already instrumented transport twice") + } +} diff --git a/backend/internal/pkg/xai/billing.go b/backend/internal/pkg/xai/billing.go new file mode 100644 index 0000000000..15b9c7e50e --- /dev/null +++ b/backend/internal/pkg/xai/billing.go @@ -0,0 +1,372 @@ +package xai + +import ( + "encoding/json" + "fmt" + "math" + "net/http" + "strconv" + "strings" + "time" +) + +const ( + // CLI client identity required by cli-chat-proxy billing endpoints. + CLITokenAuthHeader = "x-xai-token-auth" + CLITokenAuthValue = "xai-grok-cli" + CLIClientVersionHeader = "x-grok-client-version" + // Keep in sync with https://x.ai/cli/stable. + CLIClientVersion = "0.2.93" + CLIUserAgent = "grok-pager/" + CLIClientVersion + " grok-shell/" + CLIClientVersion + " (macos; aarch64)" + + BillingWeeklyPath = "/billing?format=credits" + BillingMonthlyPath = "/billing" + + SuperGrokLimitCents = 15_000 // $150.00 + SuperGrokHeavyLimitCents = 150_000 // $1,500.00 +) + +// BillingPeriod describes the current weekly/monthly window. +type BillingPeriod struct { + Type string `json:"type,omitempty"` + Start string `json:"start,omitempty"` + End string `json:"end,omitempty"` +} + +// BillingProductUsage is per-product usage inside the weekly credits window. +type BillingProductUsage struct { + Product string `json:"product,omitempty"` + UsagePercent *float64 `json:"usagePercent,omitempty"` +} + +// BillingConfig is the nested config object from /v1/billing responses. +type BillingConfig struct { + CurrentPeriod *BillingPeriod `json:"currentPeriod,omitempty"` + CreditUsagePercent *float64 `json:"creditUsagePercent,omitempty"` + ProductUsage []BillingProductUsage `json:"productUsage,omitempty"` + MonthlyLimit json.RawMessage `json:"monthlyLimit,omitempty"` + Used json.RawMessage `json:"used,omitempty"` + BillingPeriodStart string `json:"billingPeriodStart,omitempty"` + BillingPeriodEnd string `json:"billingPeriodEnd,omitempty"` +} + +// BillingPayload is the top-level body from /v1/billing. +type BillingPayload struct { + Config *BillingConfig `json:"config,omitempty"` +} + +// BillingProductSummary is a normalized product usage row for UI. +type BillingProductSummary struct { + Product string `json:"product"` + UsagePercent *float64 `json:"usage_percent,omitempty"` +} + +// BillingSummary is the merged weekly + monthly billing view. +type BillingSummary struct { + PeriodType string `json:"period_type,omitempty"` // weekly | monthly | unknown + UsagePercent *float64 `json:"usage_percent,omitempty"` + PeriodStart string `json:"period_start,omitempty"` + PeriodEnd string `json:"period_end,omitempty"` + ProductUsage []BillingProductSummary `json:"product_usage,omitempty"` + MonthlyLimitCents *float64 `json:"monthly_limit_cents,omitempty"` + UsedCents *float64 `json:"used_cents,omitempty"` + IncludedUsedCents *float64 `json:"included_used_cents,omitempty"` + BillingPeriodStart string `json:"billing_period_start,omitempty"` + BillingPeriodEnd string `json:"billing_period_end,omitempty"` + UsedPercent *float64 `json:"used_percent,omitempty"` + Plan string `json:"plan,omitempty"` // SuperGrok | SuperGrok Heavy | "" + StatusCode int `json:"status_code,omitempty"` + Source string `json:"source,omitempty"` + FetchedAt string `json:"fetched_at,omitempty"` + UpdatedAt string `json:"updated_at,omitempty"` + WeeklyUpdatedAt string `json:"weekly_updated_at,omitempty"` + MonthlyUpdatedAt string `json:"monthly_updated_at,omitempty"` + Partial bool `json:"partial,omitempty"` + FailedWindows []string `json:"failed_windows,omitempty"` +} + +// BuildBillingURL builds weekly or monthly billing URL against the CLI chat proxy. +func BuildBillingURL(formatCredits bool) string { + base := strings.TrimRight(DefaultCLIBaseURL, "/") + if formatCredits { + return base + BillingWeeklyPath + } + return base + BillingMonthlyPath +} + +// ApplyCLIBillingHeaders sets Authorization + CLI identity headers for billing GETs. +func ApplyCLIBillingHeaders(req *http.Request, accessToken string) { + if req == nil { + return + } + token := strings.TrimSpace(accessToken) + if token != "" { + req.Header.Set("Authorization", "Bearer "+token) + } + req.Header.Set("Accept", "application/json") + req.Header.Set("Content-Type", "application/json") + req.Header.Set(CLITokenAuthHeader, CLITokenAuthValue) + req.Header.Set(CLIClientVersionHeader, CLIClientVersion) + req.Header.Set("User-Agent", CLIUserAgent) +} + +// ParseBillingPayload unmarshals a billing API response body. +func ParseBillingPayload(body []byte) (*BillingPayload, error) { + if len(body) == 0 { + return nil, fmt.Errorf("empty billing body") + } + var payload BillingPayload + if err := json.Unmarshal(body, &payload); err != nil { + return nil, err + } + return &payload, nil +} + +// BuildBillingSummary normalizes a billing config into a UI-friendly summary. +func BuildBillingSummary(config *BillingConfig) *BillingSummary { + if config == nil { + return nil + } + summary := &BillingSummary{} + period := config.CurrentPeriod + periodType := resolvePeriodType(period) + creditUsage := cloneFloat(config.CreditUsagePercent) + + periodStart := "" + periodEnd := "" + if period != nil { + periodStart = strings.TrimSpace(period.Start) + periodEnd = strings.TrimSpace(period.End) + } + if periodStart == "" { + periodStart = strings.TrimSpace(config.BillingPeriodStart) + } + if periodEnd == "" { + periodEnd = strings.TrimSpace(config.BillingPeriodEnd) + } + + products := make([]BillingProductSummary, 0, len(config.ProductUsage)) + for _, item := range config.ProductUsage { + product := strings.TrimSpace(item.Product) + if product == "" { + continue + } + products = append(products, BillingProductSummary{ + Product: product, + UsagePercent: cloneFloat(item.UsagePercent), + }) + } + + monthlyLimit := parseCentValue(config.MonthlyLimit) + used := parseCentValue(config.Used) + billingStart := strings.TrimSpace(config.BillingPeriodStart) + billingEnd := strings.TrimSpace(config.BillingPeriodEnd) + + var includedUsed *float64 + if used != nil { + if monthlyLimit != nil && *monthlyLimit > 0 { + v := math.Min(*used, *monthlyLimit) + includedUsed = &v + } else { + includedUsed = cloneFloat(used) + } + } + + var usedPercent *float64 + if monthlyLimit != nil && *monthlyLimit > 0 && includedUsed != nil { + v := (*includedUsed / *monthlyLimit) * 100 + usedPercent = &v + } + + hasWeekly := creditUsage != nil || periodType == "weekly" || len(products) > 0 + hasMonthly := monthlyLimit != nil || used != nil || (!hasWeekly && billingEnd != "") + if !hasWeekly && !hasMonthly { + return nil + } + + if hasWeekly { + if periodType == "unknown" { + periodType = "weekly" + } + summary.PeriodType = periodType + summary.UsagePercent = creditUsage + summary.PeriodStart = periodStart + summary.PeriodEnd = periodEnd + } else { + // Monthly-only: do not put monthly % into UsagePercent (weekly bar field). + // Frontend weekly bar only renders when PeriodType == weekly. + summary.PeriodType = "monthly" + summary.PeriodStart = billingStart + summary.PeriodEnd = billingEnd + } + summary.ProductUsage = products + summary.MonthlyLimitCents = monthlyLimit + summary.UsedCents = used + summary.IncludedUsedCents = includedUsed + if hasMonthly { + summary.BillingPeriodStart = billingStart + summary.BillingPeriodEnd = billingEnd + } + summary.UsedPercent = usedPercent + summary.Plan = resolvePlan(monthlyLimit) + return summary +} + +// MergeBillingProbeResult updates successful billing domains while retaining +// the previous value for any domain that could not be refreshed. +func MergeBillingProbeResult(previous, weekly, monthly *BillingSummary, weeklyOK, monthlyOK bool) *BillingSummary { + var out BillingSummary + if previous != nil { + out = *previous + previousUpdatedAt := previous.UpdatedAt + if previousUpdatedAt == "" { + previousUpdatedAt = previous.FetchedAt + } + if out.WeeklyUpdatedAt == "" && (out.UsagePercent != nil || len(out.ProductUsage) > 0) { + out.WeeklyUpdatedAt = previousUpdatedAt + } + if out.MonthlyUpdatedAt == "" && (out.MonthlyLimitCents != nil || out.UsedPercent != nil) { + out.MonthlyUpdatedAt = previousUpdatedAt + } + } + now := time.Now().UTC().Format(time.RFC3339) + + if weeklyOK && weekly != nil { + out.PeriodType = weekly.PeriodType + out.UsagePercent = weekly.UsagePercent + out.PeriodStart = weekly.PeriodStart + out.PeriodEnd = weekly.PeriodEnd + out.ProductUsage = weekly.ProductUsage + out.WeeklyUpdatedAt = now + } + if monthlyOK && monthly != nil { + if out.PeriodType == "" { + out.PeriodType = "monthly" + } + out.MonthlyLimitCents = monthly.MonthlyLimitCents + out.UsedCents = monthly.UsedCents + out.IncludedUsedCents = monthly.IncludedUsedCents + out.BillingPeriodStart = monthly.BillingPeriodStart + out.BillingPeriodEnd = monthly.BillingPeriodEnd + out.UsedPercent = monthly.UsedPercent + out.Plan = monthly.Plan + out.MonthlyUpdatedAt = now + } + + out.Partial = !weeklyOK || !monthlyOK + out.FailedWindows = nil + if !weeklyOK { + out.FailedWindows = append(out.FailedWindows, "weekly") + } + if !monthlyOK { + out.FailedWindows = append(out.FailedWindows, "monthly") + } + if !weeklyOK && !monthlyOK && previous == nil { + return nil + } + return &out +} + +// StampBillingSummary sets fetch metadata. +func StampBillingSummary(summary *BillingSummary, statusCode int, source string) *BillingSummary { + if summary == nil { + return nil + } + now := time.Now().UTC().Format(time.RFC3339) + summary.StatusCode = statusCode + summary.Source = source + summary.FetchedAt = now + summary.UpdatedAt = now + return summary +} + +func resolvePeriodType(period *BillingPeriod) string { + if period == nil { + return "unknown" + } + raw := strings.ToLower(strings.TrimSpace(period.Type)) + if strings.Contains(raw, "weekly") { + return "weekly" + } + if strings.Contains(raw, "monthly") { + return "monthly" + } + return "unknown" +} + +func resolvePlan(monthlyLimitCents *float64) string { + if monthlyLimitCents == nil { + return "" + } + // Allow small float noise. + limit := math.Round(*monthlyLimitCents) + switch limit { + case SuperGrokLimitCents: + return "SuperGrok" + case SuperGrokHeavyLimitCents: + return "SuperGrok Heavy" + default: + return "" + } +} + +func parseCentValue(raw json.RawMessage) *float64 { + if len(raw) == 0 || string(raw) == "null" { + return nil + } + // Object form: {"val": 123} + var obj struct { + Val any `json:"val"` + } + if err := json.Unmarshal(raw, &obj); err == nil && obj.Val != nil { + return anyToFloat(obj.Val) + } + // Bare number / string + var n any + if err := json.Unmarshal(raw, &n); err != nil { + return nil + } + return anyToFloat(n) +} + +func anyToFloat(v any) *float64 { + switch n := v.(type) { + case float64: + return &n + case float32: + f := float64(n) + return &f + case int: + f := float64(n) + return &f + case int64: + f := float64(n) + return &f + case json.Number: + f, err := n.Float64() + if err != nil { + return nil + } + return &f + case string: + s := strings.TrimSpace(n) + if s == "" { + return nil + } + f, err := strconv.ParseFloat(s, 64) + if err != nil { + return nil + } + return &f + default: + return nil + } +} + +func cloneFloat(v *float64) *float64 { + if v == nil { + return nil + } + f := *v + return &f +} diff --git a/backend/internal/pkg/xai/billing_test.go b/backend/internal/pkg/xai/billing_test.go new file mode 100644 index 0000000000..1d863f6a39 --- /dev/null +++ b/backend/internal/pkg/xai/billing_test.go @@ -0,0 +1,127 @@ +package xai + +import ( + "encoding/json" + "net/http" + "testing" + + "github.com/stretchr/testify/require" +) + +func TestBuildBillingURL(t *testing.T) { + t.Parallel() + require.Equal(t, "https://cli-chat-proxy.grok.com/v1/billing?format=credits", BuildBillingURL(true)) + require.Equal(t, "https://cli-chat-proxy.grok.com/v1/billing", BuildBillingURL(false)) +} + +func TestApplyCLIBillingHeaders(t *testing.T) { + t.Parallel() + req, err := http.NewRequest(http.MethodGet, BuildBillingURL(true), nil) + require.NoError(t, err) + + ApplyCLIBillingHeaders(req, " token ") + + require.Equal(t, "Bearer token", req.Header.Get("Authorization")) + require.Equal(t, CLITokenAuthValue, req.Header.Get(CLITokenAuthHeader)) + require.Equal(t, CLIClientVersion, req.Header.Get(CLIClientVersionHeader)) + require.Equal(t, "grok-pager/"+CLIClientVersion+" grok-shell/"+CLIClientVersion+" (macos; aarch64)", req.UserAgent()) +} + +func TestBuildBillingSummaryWeeklyAndMonthly(t *testing.T) { + t.Parallel() + + weeklyBody := []byte(`{ + "config": { + "currentPeriod": {"type":"WEEKLY","start":"2026-07-09T03:25:00Z","end":"2026-07-16T03:25:00Z"}, + "creditUsagePercent": 2.0, + "productUsage": [{"product":"Api","usagePercent":2.0}] + } + }`) + monthlyBody := []byte(`{ + "config": { + "monthlyLimit": {"val": 15000}, + "used": {"val": 78}, + "billingPeriodStart": "2026-07-01T00:00:00Z", + "billingPeriodEnd": "2026-08-01T00:00:00Z" + } + }`) + + weeklyPayload, err := ParseBillingPayload(weeklyBody) + require.NoError(t, err) + monthlyPayload, err := ParseBillingPayload(monthlyBody) + require.NoError(t, err) + + weekly := BuildBillingSummary(weeklyPayload.Config) + monthly := BuildBillingSummary(monthlyPayload.Config) + require.NotNil(t, weekly) + require.NotNil(t, monthly) + require.Equal(t, "weekly", weekly.PeriodType) + require.InDelta(t, 2.0, *weekly.UsagePercent, 1e-9) + require.Equal(t, "Api", weekly.ProductUsage[0].Product) + require.Equal(t, "SuperGrok", monthly.Plan) + require.InDelta(t, 15000, *monthly.MonthlyLimitCents, 1e-9) + require.InDelta(t, 78, *monthly.UsedCents, 1e-9) + require.InDelta(t, 0.52, *monthly.UsedPercent, 1e-2) + + merged := MergeBillingProbeResult(nil, weekly, monthly, true, true) + require.Equal(t, "weekly", merged.PeriodType) + require.InDelta(t, 2.0, *merged.UsagePercent, 1e-9) + require.Equal(t, "SuperGrok", merged.Plan) + require.InDelta(t, 15000, *merged.MonthlyLimitCents, 1e-9) + require.Equal(t, "2026-08-01T00:00:00Z", merged.BillingPeriodEnd) +} + +func TestParseCentValueBareNumber(t *testing.T) { + t.Parallel() + raw, _ := json.Marshal(15000) + v := parseCentValue(raw) + require.NotNil(t, v) + require.InDelta(t, 15000, *v, 1e-9) +} + +func TestBuildBillingSummaryMonthlyOnlyKeepsWeeklyUsageEmpty(t *testing.T) { + t.Parallel() + payload, err := ParseBillingPayload([]byte(`{"config":{"monthlyLimit":{"val":15000},"used":{"val":7500},"billingPeriodStart":"2026-07-01T00:00:00Z","billingPeriodEnd":"2026-08-01T00:00:00Z"}}`)) + require.NoError(t, err) + + summary := BuildBillingSummary(payload.Config) + require.NotNil(t, summary) + require.Equal(t, "monthly", summary.PeriodType) + require.Nil(t, summary.UsagePercent) + require.InDelta(t, 50, *summary.UsedPercent, 1e-9) +} + +func TestMergeBillingProbeResultRetainsFailedWindow(t *testing.T) { + t.Parallel() + previous := &BillingSummary{ + PeriodType: "weekly", + UsagePercent: floatPointer(100), + PeriodEnd: "2026-07-16T00:00:00Z", + MonthlyLimitCents: floatPointer(15000), + UsedPercent: floatPointer(20), + BillingPeriodEnd: "2026-08-01T00:00:00Z", + WeeklyUpdatedAt: "2026-07-10T00:00:00Z", + MonthlyUpdatedAt: "2026-07-10T00:00:00Z", + FailedWindows: []string{"monthly"}, + } + monthly := &BillingSummary{ + PeriodType: "monthly", + MonthlyLimitCents: floatPointer(15000), + UsedPercent: floatPointer(30), + BillingPeriodEnd: "2026-08-01T00:00:00Z", + } + + merged := MergeBillingProbeResult(previous, nil, monthly, false, true) + require.Equal(t, "weekly", merged.PeriodType) + require.InDelta(t, 100, *merged.UsagePercent, 1e-9) + require.Equal(t, previous.WeeklyUpdatedAt, merged.WeeklyUpdatedAt) + require.InDelta(t, 30, *merged.UsedPercent, 1e-9) + require.NotEqual(t, previous.MonthlyUpdatedAt, merged.MonthlyUpdatedAt) + require.True(t, merged.Partial) + require.Equal(t, []string{"weekly"}, merged.FailedWindows) + require.Equal(t, []string{"monthly"}, previous.FailedWindows) +} + +func floatPointer(value float64) *float64 { + return &value +} diff --git a/backend/internal/repository/account_repo.go b/backend/internal/repository/account_repo.go index 261be7c14b..6d26d474c6 100644 --- a/backend/internal/repository/account_repo.go +++ b/backend/internal/repository/account_repo.go @@ -61,6 +61,7 @@ var schedulerNeutralExtraKeyPrefixes = []string{ var schedulerNeutralExtraKeys = map[string]struct{}{ "codex_usage_updated_at": {}, + "grok_billing_snapshot": {}, "session_window_utilization": {}, } diff --git a/backend/internal/repository/account_repo_grok_billing_test.go b/backend/internal/repository/account_repo_grok_billing_test.go new file mode 100644 index 0000000000..fb41ae5ffa --- /dev/null +++ b/backend/internal/repository/account_repo_grok_billing_test.go @@ -0,0 +1,16 @@ +package repository + +import ( + "testing" + + "github.com/stretchr/testify/require" +) + +func TestGrokBillingSnapshotIsSchedulerNeutral(t *testing.T) { + t.Parallel() + + require.True(t, isSchedulerNeutralExtraKey("grok_billing_snapshot")) + require.False(t, shouldEnqueueSchedulerOutboxForExtraUpdates(map[string]any{ + "grok_billing_snapshot": map[string]any{"usage_percent": 50}, + })) +} diff --git a/backend/internal/repository/backup_s3_store.go b/backend/internal/repository/backup_s3_store.go index 5d419f574b..2104e1e5d7 100644 --- a/backend/internal/repository/backup_s3_store.go +++ b/backend/internal/repository/backup_s3_store.go @@ -13,6 +13,7 @@ import ( "github.com/aws/aws-sdk-go-v2/credentials" "github.com/aws/aws-sdk-go-v2/service/s3" + "github.com/Wei-Shaw/sub2api/internal/pkg/servertiming" "github.com/Wei-Shaw/sub2api/internal/service" ) @@ -63,12 +64,14 @@ func (s *S3BackupStore) Upload(ctx context.Context, key string, body io.Reader, return 0, fmt.Errorf("read body: %w", err) } + finish := servertiming.ObserveDependency(ctx, "s3") _, err = s.client.PutObject(ctx, &s3.PutObjectInput{ Bucket: &s.bucket, Key: &key, Body: bytes.NewReader(data), ContentType: &contentType, }) + finish() if err != nil { return 0, fmt.Errorf("S3 PutObject: %w", err) } @@ -76,10 +79,12 @@ func (s *S3BackupStore) Upload(ctx context.Context, key string, body io.Reader, } func (s *S3BackupStore) Download(ctx context.Context, key string) (io.ReadCloser, error) { + finish := servertiming.ObserveDependency(ctx, "s3") result, err := s.client.GetObject(ctx, &s3.GetObjectInput{ Bucket: &s.bucket, Key: &key, }) + finish() if err != nil { return nil, fmt.Errorf("S3 GetObject: %w", err) } @@ -87,10 +92,12 @@ func (s *S3BackupStore) Download(ctx context.Context, key string) (io.ReadCloser } func (s *S3BackupStore) Delete(ctx context.Context, key string) error { + finish := servertiming.ObserveDependency(ctx, "s3") _, err := s.client.DeleteObject(ctx, &s3.DeleteObjectInput{ Bucket: &s.bucket, Key: &key, }) + finish() return err } @@ -107,9 +114,11 @@ func (s *S3BackupStore) PresignURL(ctx context.Context, key string, expiry time. } func (s *S3BackupStore) HeadBucket(ctx context.Context) error { + finish := servertiming.ObserveDependency(ctx, "s3") _, err := s.client.HeadBucket(ctx, &s3.HeadBucketInput{ Bucket: &s.bucket, }) + finish() if err != nil { return fmt.Errorf("S3 HeadBucket failed: %w", err) } diff --git a/backend/internal/repository/claude_oauth_service.go b/backend/internal/repository/claude_oauth_service.go index 5c5f27c86a..ec2d426ecb 100644 --- a/backend/internal/repository/claude_oauth_service.go +++ b/backend/internal/repository/claude_oauth_service.go @@ -276,5 +276,5 @@ func createReqClient(proxyURL string) (*req.Client, error) { client.SetProxyURL(trimmed) } - return client, nil + return instrumentReqClient(client), nil } diff --git a/backend/internal/repository/ent.go b/backend/internal/repository/ent.go index 64d321924d..3abb528e98 100644 --- a/backend/internal/repository/ent.go +++ b/backend/internal/repository/ent.go @@ -15,7 +15,7 @@ import ( "entgo.io/ent/dialect" entsql "entgo.io/ent/dialect/sql" - _ "github.com/lib/pq" // PostgreSQL 驱动,通过副作用导入注册驱动 + "github.com/lib/pq" ) // InitEnt 初始化 Ent ORM 客户端并返回客户端实例和底层的 *sql.DB。 @@ -48,9 +48,19 @@ func InitEnt(cfg *config.Config) (*ent.Client, *sql.DB, error) { // 使用 Ent 的 SQL 驱动打开 PostgreSQL 连接。 // dialect.Postgres 指定使用 PostgreSQL 方言进行 SQL 生成。 - drv, err := entsql.Open(dialect.Postgres, dsn) - if err != nil { - return nil, nil, err + var drv *entsql.Driver + if cfg.Server.EnableServerTiming { + connector, err := pq.NewConnector(dsn) + if err != nil { + return nil, nil, err + } + drv = entsql.OpenDB(dialect.Postgres, sql.OpenDB(newServerTimingConnector(connector))) + } else { + var err error + drv, err = entsql.Open(dialect.Postgres, dsn) + if err != nil { + return nil, nil, err + } } applyDBPoolSettings(drv.DB(), cfg) diff --git a/backend/internal/repository/http_upstream.go b/backend/internal/repository/http_upstream.go index bb079b0789..0d5afd500c 100644 --- a/backend/internal/repository/http_upstream.go +++ b/backend/internal/repository/http_upstream.go @@ -25,6 +25,7 @@ import ( "github.com/Wei-Shaw/sub2api/internal/config" "github.com/Wei-Shaw/sub2api/internal/pkg/proxyurl" "github.com/Wei-Shaw/sub2api/internal/pkg/proxyutil" + "github.com/Wei-Shaw/sub2api/internal/pkg/servertiming" "github.com/Wei-Shaw/sub2api/internal/pkg/tlsfingerprint" "github.com/Wei-Shaw/sub2api/internal/service" "github.com/Wei-Shaw/sub2api/internal/util/urlvalidator" @@ -186,7 +187,7 @@ func (s *httpUpstreamService) Do(req *http.Request, proxyURL string, accountID i } // 执行请求 - resp, err := entry.client.Do(req) + resp, err := servertiming.Do(entry.client, req) if err != nil { s.recordOpenAIHTTP2Failure(profile, entry.protocolMode, entry.proxyKey, err) // 请求失败,立即减少计数 @@ -243,7 +244,7 @@ func (s *httpUpstreamService) DoWithTLS(req *http.Request, proxyURL string, acco return nil, err } - resp, err := entry.client.Do(req) + resp, err := servertiming.Do(entry.client, req) if err != nil { atomic.AddInt64(&entry.inFlight, -1) atomic.StoreInt64(&entry.lastUsed, time.Now().UnixNano()) diff --git a/backend/internal/repository/openai_long_context_billing_migration_integration_test.go b/backend/internal/repository/openai_long_context_billing_migration_integration_test.go new file mode 100644 index 0000000000..5f50ed0729 --- /dev/null +++ b/backend/internal/repository/openai_long_context_billing_migration_integration_test.go @@ -0,0 +1,159 @@ +//go:build integration + +package repository + +import ( + "context" + "testing" + + dbmigrations "github.com/Wei-Shaw/sub2api/migrations" + "github.com/stretchr/testify/require" +) + +func TestMigration175EnforcesOpenAILongContextBillingWriteInvariant(t *testing.T) { + tx := testTx(t) + ctx := context.Background() + migrationSQL, err := dbmigrations.FS.ReadFile("175_default_openai_long_context_billing.sql") + require.NoError(t, err) + _, err = tx.ExecContext(ctx, ` +DROP TRIGGER IF EXISTS accounts_propagate_openai_long_context_billing_extra ON accounts; +DROP TRIGGER IF EXISTS accounts_enforce_openai_long_context_billing_extra ON accounts; +`) + require.NoError(t, err) + + var ordinaryID int64 + require.NoError(t, tx.QueryRowContext(ctx, ` +INSERT INTO accounts (name, platform, type, extra) +VALUES ('migration-175-ordinary', 'openai', 'oauth', '{}'::jsonb) +RETURNING id +`).Scan(&ordinaryID)) + + var parentID int64 + require.NoError(t, tx.QueryRowContext(ctx, ` +INSERT INTO accounts (name, platform, type, extra) +VALUES ('migration-175-parent', 'openai', 'oauth', '{"openai_long_context_billing_enabled":false}'::jsonb) +RETURNING id +`).Scan(&parentID)) + + var shadowID int64 + require.NoError(t, tx.QueryRowContext(ctx, ` +INSERT INTO accounts (name, platform, type, extra, parent_account_id, quota_dimension) +VALUES ('migration-175-shadow', 'openai', 'oauth', '{}'::jsonb, $1, 'spark') +RETURNING id +`, parentID).Scan(&shadowID)) + + var malformedLegacyID int64 + require.NoError(t, tx.QueryRowContext(ctx, ` +INSERT INTO accounts (name, platform, type, extra) +VALUES ('migration-175-malformed-legacy', 'openai', 'oauth', '{"openai_long_context_billing_enabled":"false"}'::jsonb) +RETURNING id +`).Scan(&malformedLegacyID)) + + _, err = tx.ExecContext(ctx, string(migrationSQL)) + require.NoError(t, err) + _, err = tx.ExecContext(ctx, string(migrationSQL)) + require.NoError(t, err) + + var ordinaryEnabled bool + require.NoError(t, tx.QueryRowContext(ctx, ` +SELECT (extra->>'openai_long_context_billing_enabled')::boolean +FROM accounts +WHERE id = $1 +`, ordinaryID).Scan(&ordinaryEnabled)) + require.False(t, ordinaryEnabled) + + var shadowEnabled bool + require.NoError(t, tx.QueryRowContext(ctx, ` +SELECT (extra->>'openai_long_context_billing_enabled')::boolean +FROM accounts +WHERE id = $1 +`, shadowID).Scan(&shadowEnabled)) + require.False(t, shadowEnabled) + + var initialShadowOutboxEvents int + require.NoError(t, tx.QueryRowContext(ctx, ` +SELECT COUNT(*) +FROM scheduler_outbox +WHERE event_type = 'account_changed' AND account_id = $1 +`, shadowID).Scan(&initialShadowOutboxEvents)) + require.Equal(t, 1, initialShadowOutboxEvents) + + var malformedLegacyEnabled bool + require.NoError(t, tx.QueryRowContext(ctx, ` +SELECT (extra->>'openai_long_context_billing_enabled')::boolean +FROM accounts +WHERE id = $1 +`, malformedLegacyID).Scan(&malformedLegacyEnabled)) + require.False(t, malformedLegacyEnabled) + _, err = tx.ExecContext(ctx, ` +UPDATE accounts +SET extra = extra || '{"migration_175_unrelated_update":true}'::jsonb +WHERE id = $1 +`, malformedLegacyID) + require.NoError(t, err) + + _, err = tx.ExecContext(ctx, "TRUNCATE scheduler_outbox") + require.NoError(t, err) + _, err = tx.ExecContext(ctx, ` +UPDATE accounts +SET extra = '{"legacy_writer_replaced_extra":true}'::jsonb +WHERE id = $1 +`, parentID) + require.NoError(t, err) + var parentEnabled bool + require.NoError(t, tx.QueryRowContext(ctx, ` +SELECT (extra->>'openai_long_context_billing_enabled')::boolean +FROM accounts +WHERE id = $1 +`, parentID).Scan(&parentEnabled)) + require.False(t, parentEnabled) + require.NoError(t, tx.QueryRowContext(ctx, ` +SELECT (extra->>'openai_long_context_billing_enabled')::boolean +FROM accounts +WHERE id = $1 +`, shadowID).Scan(&shadowEnabled)) + require.False(t, shadowEnabled) + var preservedOptOutEvents int + require.NoError(t, tx.QueryRowContext(ctx, ` +SELECT COUNT(*) +FROM scheduler_outbox +WHERE event_type = 'account_changed' AND account_id = $1 +`, shadowID).Scan(&preservedOptOutEvents)) + require.Zero(t, preservedOptOutEvents) + + require.NoError(t, tx.QueryRowContext(ctx, ` +INSERT INTO accounts (name, platform, type, extra) +VALUES ('migration-175-rolling-writer', 'openai', 'oauth', '{}'::jsonb) +RETURNING (extra->>'openai_long_context_billing_enabled')::boolean +`).Scan(&ordinaryEnabled)) + require.False(t, ordinaryEnabled) + + _, err = tx.ExecContext(ctx, "TRUNCATE scheduler_outbox") + require.NoError(t, err) + _, err = tx.ExecContext(ctx, ` +UPDATE accounts +SET extra = jsonb_set(extra, '{openai_long_context_billing_enabled}', 'true'::jsonb, true) +WHERE id = $1 +`, parentID) + require.NoError(t, err) + require.NoError(t, tx.QueryRowContext(ctx, ` +SELECT (extra->>'openai_long_context_billing_enabled')::boolean +FROM accounts +WHERE id = $1 +`, shadowID).Scan(&shadowEnabled)) + require.True(t, shadowEnabled) + + var shadowOutboxEvents int + require.NoError(t, tx.QueryRowContext(ctx, ` +SELECT COUNT(*) +FROM scheduler_outbox +WHERE event_type = 'account_changed' AND account_id = $1 +`, shadowID).Scan(&shadowOutboxEvents)) + require.Equal(t, 1, shadowOutboxEvents) + + _, err = tx.ExecContext(ctx, ` +INSERT INTO accounts (name, platform, type, extra) +VALUES ('migration-175-malformed', 'openai', 'oauth', '{"openai_long_context_billing_enabled":"false"}'::jsonb) +`) + require.ErrorContains(t, err, "openai_long_context_billing_enabled must be a boolean") +} diff --git a/backend/internal/repository/ops_repo.go b/backend/internal/repository/ops_repo.go index 2129a451c4..900abcf212 100644 --- a/backend/internal/repository/ops_repo.go +++ b/backend/internal/repository/ops_repo.go @@ -718,6 +718,7 @@ func (r *opsRepository) BatchInsertSystemLogs(ctx context.Context, inputs []*ser stmt, err := tx.PrepareContext(ctx, pq.CopyIn( "ops_system_logs", "created_at", + "host", "level", "component", "message", @@ -760,6 +761,7 @@ func (r *opsRepository) BatchInsertSystemLogs(ctx context.Context, inputs []*ser if _, err := stmt.ExecContext( ctx, createdAt.UTC(), + opsNullString(input.Host), level, component, message, @@ -827,6 +829,7 @@ func (r *opsRepository) ListSystemLogs(ctx context.Context, filter *service.OpsS SELECT l.id, l.created_at, + COALESCE(l.host, ''), l.level, COALESCE(l.component, ''), COALESCE(l.message, ''), @@ -859,6 +862,7 @@ LIMIT $` + itoa(len(args)+1) + ` OFFSET $` + itoa(len(args)+2) if err := rows.Scan( &item.ID, &item.CreatedAt, + &item.Host, &item.Level, &item.Component, &item.Message, @@ -1130,6 +1134,11 @@ func buildOpsSystemLogsWhere(filter *service.OpsSystemLogFilter) (string, []any, hasConstraint = true } if filter != nil { + if v := strings.TrimSpace(filter.Host); v != "" { + args = append(args, v) + clauses = append(clauses, "l.host = $"+itoa(len(args))) + hasConstraint = true + } if v := strings.ToLower(strings.TrimSpace(filter.Level)); v != "" { args = append(args, v) clauses = append(clauses, "LOWER(COALESCE(l.level,'')) = $"+itoa(len(args))) @@ -1194,6 +1203,7 @@ func buildOpsSystemLogsCleanupWhere(filter *service.OpsSystemLogCleanupFilter) ( listFilter := &service.OpsSystemLogFilter{ StartTime: filter.StartTime, EndTime: filter.EndTime, + Host: filter.Host, Level: filter.Level, Component: filter.Component, RequestID: filter.RequestID, diff --git a/backend/internal/repository/ops_repo_system_logs_test.go b/backend/internal/repository/ops_repo_system_logs_test.go index 98199f4828..48be3e7256 100644 --- a/backend/internal/repository/ops_repo_system_logs_test.go +++ b/backend/internal/repository/ops_repo_system_logs_test.go @@ -18,6 +18,7 @@ func TestBuildOpsSystemLogsWhere_WithClientRequestIDAndUserID(t *testing.T) { filter := &service.OpsSystemLogFilter{ StartTime: &start, EndTime: &end, + Host: "api-node-1", Level: "warn", Component: "http.access", RequestID: "req-1", @@ -37,8 +38,11 @@ func TestBuildOpsSystemLogsWhere_WithClientRequestIDAndUserID(t *testing.T) { if where == "" { t.Fatalf("where should not be empty") } - if len(args) != 12 { - t.Fatalf("args len = %d, want 12", len(args)) + if len(args) != 13 { + t.Fatalf("args len = %d, want 13", len(args)) + } + if !contains(where, "l.host = $") { + t.Fatalf("where should include host condition: %s", where) } if !contains(where, "COALESCE(l.client_request_id,'') = $") { t.Fatalf("where should include client_request_id condition: %s", where) @@ -68,6 +72,7 @@ func TestBuildOpsSystemLogsCleanupWhere_WithClientRequestIDAndUserID(t *testing. userID := int64(9) apiKeyID := int64(10) filter := &service.OpsSystemLogCleanupFilter{ + Host: "api-node-2", ClientRequestID: "creq-9", UserID: &userID, APIKeyID: &apiKeyID, @@ -77,8 +82,11 @@ func TestBuildOpsSystemLogsCleanupWhere_WithClientRequestIDAndUserID(t *testing. if !hasConstraint { t.Fatalf("expected hasConstraint=true") } - if len(args) != 3 { - t.Fatalf("args len = %d, want 3", len(args)) + if len(args) != 4 { + t.Fatalf("args len = %d, want 4", len(args)) + } + if !contains(where, "l.host = $") { + t.Fatalf("where should include host condition: %s", where) } if !contains(where, "COALESCE(l.client_request_id,'') = $") { t.Fatalf("where should include client_request_id condition: %s", where) diff --git a/backend/internal/repository/redis.go b/backend/internal/repository/redis.go index 2b4ee4e636..0ead4644c1 100644 --- a/backend/internal/repository/redis.go +++ b/backend/internal/repository/redis.go @@ -21,7 +21,11 @@ import ( // 2. MinIdleConns: 保持最小空闲连接,减少冷启动延迟(默认 10) // 3. DialTimeout/ReadTimeout/WriteTimeout: 精确控制各阶段超时 func InitRedis(cfg *config.Config) *redis.Client { - return redis.NewClient(buildRedisOptions(cfg)) + client := redis.NewClient(buildRedisOptions(cfg)) + if cfg.Server.EnableServerTiming { + client.AddHook(serverTimingRedisHook{}) + } + return client } // buildRedisOptions 构建 Redis 连接选项 diff --git a/backend/internal/repository/req_client_pool.go b/backend/internal/repository/req_client_pool.go index 32501f7b19..95ab27ce32 100644 --- a/backend/internal/repository/req_client_pool.go +++ b/backend/internal/repository/req_client_pool.go @@ -2,11 +2,13 @@ package repository import ( "fmt" + "net/http" "strings" "sync" "time" "github.com/Wei-Shaw/sub2api/internal/pkg/proxyurl" + "github.com/Wei-Shaw/sub2api/internal/pkg/servertiming" "github.com/imroc/req/v3" ) @@ -57,6 +59,7 @@ func getSharedReqClient(opts reqClientOptions) (*req.Client, error) { if trimmed != "" { client.SetProxyURL(trimmed) } + client = instrumentReqClient(client) actual, _ := sharedReqClients.LoadOrStore(key, client) if c, ok := actual.(*req.Client); ok { @@ -65,6 +68,17 @@ func getSharedReqClient(opts reqClientOptions) (*req.Client, error) { return client, nil } +func instrumentReqClient(client *req.Client) *req.Client { + if client == nil { + return nil + } + client.GetTransport().WrapRoundTripFunc(func(rt http.RoundTripper) req.HttpRoundTripFunc { + timed := servertiming.WrapRoundTripper(rt) + return timed.RoundTrip + }) + return client +} + func buildReqClientKey(opts reqClientOptions) string { return fmt.Sprintf("%s|%s|%t|%t", strings.TrimSpace(opts.ProxyURL), diff --git a/backend/internal/repository/req_client_pool_test.go b/backend/internal/repository/req_client_pool_test.go index 9067d0129f..3a27841c5a 100644 --- a/backend/internal/repository/req_client_pool_test.go +++ b/backend/internal/repository/req_client_pool_test.go @@ -1,12 +1,17 @@ package repository import ( + "context" + "net/http" + "net/http/httptest" "reflect" + "strings" "sync" "testing" "time" "unsafe" + "github.com/Wei-Shaw/sub2api/internal/pkg/servertiming" "github.com/imroc/req/v3" "github.com/stretchr/testify/require" ) @@ -118,3 +123,20 @@ func TestCreateGeminiReqClient_ForceHTTP2Disabled(t *testing.T) { require.NoError(t, err) require.Equal(t, "", forceHTTPVersion(t, client)) } + +func TestInstrumentReqClientRecordsDependency(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.WriteHeader(http.StatusNoContent) + })) + defer server.Close() + + collector := servertiming.New(time.Now()) + ctx := servertiming.WithCollector(context.Background(), collector) + client := instrumentReqClient(req.C()) + response, err := client.R().SetContext(ctx).Get(server.URL) + require.NoError(t, err) + require.Equal(t, http.StatusNoContent, response.StatusCode) + + header := collector.HeaderValue(time.Now(), "bypass") + require.True(t, strings.Contains(header, "dep_http;dur="), header) +} diff --git a/backend/internal/repository/server_timing_redis.go b/backend/internal/repository/server_timing_redis.go new file mode 100644 index 0000000000..dba35450de --- /dev/null +++ b/backend/internal/repository/server_timing_redis.go @@ -0,0 +1,39 @@ +package repository + +import ( + "context" + "time" + + "github.com/Wei-Shaw/sub2api/internal/pkg/servertiming" + "github.com/redis/go-redis/v9" +) + +type serverTimingRedisHook struct{} + +func (serverTimingRedisHook) DialHook(next redis.DialHook) redis.DialHook { + return next +} + +func (serverTimingRedisHook) ProcessHook(next redis.ProcessHook) redis.ProcessHook { + return func(ctx context.Context, cmd redis.Cmder) error { + if !servertiming.Active(ctx) { + return next(ctx, cmd) + } + startedAt := time.Now() + err := next(ctx, cmd) + servertiming.Record(ctx, servertiming.MetricRedis, startedAt, time.Now(), 1) + return err + } +} + +func (serverTimingRedisHook) ProcessPipelineHook(next redis.ProcessPipelineHook) redis.ProcessPipelineHook { + return func(ctx context.Context, cmds []redis.Cmder) error { + if !servertiming.Active(ctx) { + return next(ctx, cmds) + } + startedAt := time.Now() + err := next(ctx, cmds) + servertiming.Record(ctx, servertiming.MetricRedis, startedAt, time.Now(), len(cmds)) + return err + } +} diff --git a/backend/internal/repository/server_timing_redis_test.go b/backend/internal/repository/server_timing_redis_test.go new file mode 100644 index 0000000000..d1ae47e3b0 --- /dev/null +++ b/backend/internal/repository/server_timing_redis_test.go @@ -0,0 +1,63 @@ +package repository + +import ( + "context" + "errors" + "strings" + "testing" + "time" + + "github.com/Wei-Shaw/sub2api/internal/pkg/servertiming" + "github.com/redis/go-redis/v9" +) + +func TestServerTimingRedisHookRecordsCommands(t *testing.T) { + collector := servertiming.New(time.Now()) + ctx := servertiming.WithCollector(context.Background(), collector) + hook := serverTimingRedisHook{} + + process := hook.ProcessHook(func(context.Context, redis.Cmder) error { + time.Sleep(time.Millisecond) + return errors.New("redis failure") + }) + if err := process(ctx, redis.NewStringCmd(ctx, "get", "sensitive-key")); err == nil { + t.Fatal("ProcessHook did not return the underlying error") + } + + pipeline := hook.ProcessPipelineHook(func(context.Context, []redis.Cmder) error { + time.Sleep(time.Millisecond) + return nil + }) + commands := []redis.Cmder{ + redis.NewStringCmd(ctx, "get", "first-secret"), + redis.NewStringCmd(ctx, "get", "second-secret"), + redis.NewStatusCmd(ctx, "set", "third-secret", "value"), + } + if err := pipeline(ctx, commands); err != nil { + t.Fatal(err) + } + + header := collector.HeaderValue(time.Now(), "bypass") + if !strings.Contains(header, `commands=4`) { + t.Fatalf("header %q does not report one command and a three-command pipeline", header) + } + if strings.Contains(header, "secret") || strings.Contains(header, "get") { + t.Fatalf("Redis command details leaked into header: %q", header) + } +} + +func TestServerTimingRedisHookSkipsInactiveContext(t *testing.T) { + called := false + hook := serverTimingRedisHook{} + process := hook.ProcessHook(func(context.Context, redis.Cmder) error { + called = true + return nil + }) + ctx := context.Background() + if err := process(ctx, redis.NewStringCmd(ctx, "ping")); err != nil { + t.Fatal(err) + } + if !called { + t.Fatal("inactive Redis command did not reach the next hook") + } +} diff --git a/backend/internal/repository/server_timing_sql.go b/backend/internal/repository/server_timing_sql.go new file mode 100644 index 0000000000..062663f08b --- /dev/null +++ b/backend/internal/repository/server_timing_sql.go @@ -0,0 +1,311 @@ +package repository + +import ( + "context" + "database/sql/driver" + "errors" + "io" + "reflect" + "time" + + "github.com/Wei-Shaw/sub2api/internal/pkg/servertiming" +) + +type serverTimingConnector struct { + base driver.Connector +} + +func newServerTimingConnector(base driver.Connector) driver.Connector { + return &serverTimingConnector{base: base} +} + +func (c *serverTimingConnector) Connect(ctx context.Context) (driver.Conn, error) { + startedAt := time.Now() + conn, err := c.base.Connect(ctx) + servertiming.RecordInterval(ctx, servertiming.MetricDatabase, startedAt, time.Now()) + if err != nil { + return nil, err + } + return &serverTimingConn{Conn: conn}, nil +} + +func (c *serverTimingConnector) Driver() driver.Driver { + return c.base.Driver() +} + +type serverTimingConn struct { + driver.Conn +} + +func (c *serverTimingConn) Prepare(query string) (driver.Stmt, error) { + stmt, err := c.Conn.Prepare(query) + if err != nil { + return nil, err + } + return &serverTimingStmt{Stmt: stmt}, nil +} + +func (c *serverTimingConn) PrepareContext(ctx context.Context, query string) (driver.Stmt, error) { + startedAt := time.Now() + var ( + stmt driver.Stmt + err error + ) + if preparer, ok := c.Conn.(driver.ConnPrepareContext); ok { + stmt, err = preparer.PrepareContext(ctx, query) + } else { + stmt, err = c.Conn.Prepare(query) + } + servertiming.Record(ctx, servertiming.MetricDatabase, startedAt, time.Now(), 1) + if err != nil { + return nil, err + } + return &serverTimingStmt{Stmt: stmt}, nil +} + +func (c *serverTimingConn) ExecContext(ctx context.Context, query string, args []driver.NamedValue) (driver.Result, error) { + execer, ok := c.Conn.(driver.ExecerContext) + if !ok { + return nil, driver.ErrSkip + } + startedAt := time.Now() + result, err := execer.ExecContext(ctx, query, args) + servertiming.Record(ctx, servertiming.MetricDatabase, startedAt, time.Now(), 1) + return result, err +} + +func (c *serverTimingConn) QueryContext(ctx context.Context, query string, args []driver.NamedValue) (driver.Rows, error) { + queryer, ok := c.Conn.(driver.QueryerContext) + if !ok { + return nil, driver.ErrSkip + } + startedAt := time.Now() + rows, err := queryer.QueryContext(ctx, query, args) + servertiming.Record(ctx, servertiming.MetricDatabase, startedAt, time.Now(), 1) + if err != nil || rows == nil { + return rows, err + } + return newServerTimingRows(ctx, rows), nil +} + +func (c *serverTimingConn) BeginTx(ctx context.Context, opts driver.TxOptions) (driver.Tx, error) { + startedAt := time.Now() + var ( + tx driver.Tx + err error + ) + if beginner, ok := c.Conn.(driver.ConnBeginTx); ok { + tx, err = beginner.BeginTx(ctx, opts) + } else { + if opts.Isolation != driver.IsolationLevel(0) { + return nil, errors.New("driver does not support non-default isolation") + } + if opts.ReadOnly { + return nil, errors.New("driver does not support read-only transactions") + } + // The wrapper exposes ConnBeginTx, so it must retain database/sql's + // legacy fallback for drivers that only implement Conn.Begin. + tx, err = c.Conn.Begin() //nolint:staticcheck // Required driver compatibility fallback. + } + servertiming.RecordInterval(ctx, servertiming.MetricDatabase, startedAt, time.Now()) + if err != nil || tx == nil { + return tx, err + } + return &serverTimingTx{Tx: tx, ctx: ctx}, nil +} + +func (c *serverTimingConn) Ping(ctx context.Context) error { + if pinger, ok := c.Conn.(driver.Pinger); ok { + startedAt := time.Now() + err := pinger.Ping(ctx) + servertiming.RecordInterval(ctx, servertiming.MetricDatabase, startedAt, time.Now()) + return err + } + return nil +} + +func (c *serverTimingConn) ResetSession(ctx context.Context) error { + if resetter, ok := c.Conn.(driver.SessionResetter); ok { + startedAt := time.Now() + err := resetter.ResetSession(ctx) + servertiming.RecordInterval(ctx, servertiming.MetricDatabase, startedAt, time.Now()) + return err + } + return nil +} + +func (c *serverTimingConn) IsValid() bool { + if validator, ok := c.Conn.(driver.Validator); ok { + return validator.IsValid() + } + return true +} + +func (c *serverTimingConn) CheckNamedValue(value *driver.NamedValue) error { + if checker, ok := c.Conn.(driver.NamedValueChecker); ok { + return checker.CheckNamedValue(value) + } + return driver.ErrSkip +} + +type serverTimingStmt struct { + driver.Stmt +} + +func (s *serverTimingStmt) ExecContext(ctx context.Context, args []driver.NamedValue) (driver.Result, error) { + startedAt := time.Now() + var ( + result driver.Result + err error + ) + if execer, ok := s.Stmt.(driver.StmtExecContext); ok { + result, err = execer.ExecContext(ctx, args) + } else { + var values []driver.Value + values, err = namedValues(args) + if err == nil { + // The wrapper exposes StmtExecContext and must preserve the fallback + // database/sql would use for a legacy driver statement. + result, err = s.Stmt.Exec(values) //nolint:staticcheck // Required driver compatibility fallback. + } + } + servertiming.Record(ctx, servertiming.MetricDatabase, startedAt, time.Now(), 1) + return result, err +} + +func (s *serverTimingStmt) QueryContext(ctx context.Context, args []driver.NamedValue) (driver.Rows, error) { + startedAt := time.Now() + var ( + rows driver.Rows + err error + ) + if queryer, ok := s.Stmt.(driver.StmtQueryContext); ok { + rows, err = queryer.QueryContext(ctx, args) + } else { + var values []driver.Value + values, err = namedValues(args) + if err == nil { + // The wrapper exposes StmtQueryContext and must preserve the fallback + // database/sql would use for a legacy driver statement. + rows, err = s.Stmt.Query(values) //nolint:staticcheck // Required driver compatibility fallback. + } + } + servertiming.Record(ctx, servertiming.MetricDatabase, startedAt, time.Now(), 1) + if err != nil || rows == nil { + return rows, err + } + return newServerTimingRows(ctx, rows), nil +} + +func (s *serverTimingStmt) CheckNamedValue(value *driver.NamedValue) error { + if checker, ok := s.Stmt.(driver.NamedValueChecker); ok { + return checker.CheckNamedValue(value) + } + return driver.ErrSkip +} + +func namedValues(args []driver.NamedValue) ([]driver.Value, error) { + values := make([]driver.Value, len(args)) + for i, arg := range args { + if arg.Name != "" { + return nil, errors.New("named parameters are not supported") + } + values[i] = arg.Value + } + return values, nil +} + +type serverTimingRows struct { + driver.Rows + ctx context.Context +} + +func newServerTimingRows(ctx context.Context, rows driver.Rows) *serverTimingRows { + return &serverTimingRows{Rows: rows, ctx: ctx} +} + +func (r *serverTimingRows) Close() error { + startedAt := time.Now() + err := r.Rows.Close() + servertiming.RecordInterval(r.ctx, servertiming.MetricDatabase, startedAt, time.Now()) + return err +} + +func (r *serverTimingRows) Next(dest []driver.Value) error { + startedAt := time.Now() + err := r.Rows.Next(dest) + servertiming.RecordInterval(r.ctx, servertiming.MetricDatabase, startedAt, time.Now()) + return err +} + +func (r *serverTimingRows) HasNextResultSet() bool { + if rows, ok := r.Rows.(driver.RowsNextResultSet); ok { + return rows.HasNextResultSet() + } + return false +} + +func (r *serverTimingRows) NextResultSet() error { + rows, ok := r.Rows.(driver.RowsNextResultSet) + if !ok { + return io.EOF + } + startedAt := time.Now() + err := rows.NextResultSet() + servertiming.RecordInterval(r.ctx, servertiming.MetricDatabase, startedAt, time.Now()) + return err +} + +func (r *serverTimingRows) ColumnTypeScanType(index int) reflect.Type { + if rows, ok := r.Rows.(driver.RowsColumnTypeScanType); ok { + return rows.ColumnTypeScanType(index) + } + return reflect.TypeOf(new(any)).Elem() +} + +func (r *serverTimingRows) ColumnTypeDatabaseTypeName(index int) string { + if rows, ok := r.Rows.(driver.RowsColumnTypeDatabaseTypeName); ok { + return rows.ColumnTypeDatabaseTypeName(index) + } + return "" +} + +func (r *serverTimingRows) ColumnTypeLength(index int) (int64, bool) { + if rows, ok := r.Rows.(driver.RowsColumnTypeLength); ok { + return rows.ColumnTypeLength(index) + } + return 0, false +} + +func (r *serverTimingRows) ColumnTypeNullable(index int) (bool, bool) { + if rows, ok := r.Rows.(driver.RowsColumnTypeNullable); ok { + return rows.ColumnTypeNullable(index) + } + return false, false +} + +func (r *serverTimingRows) ColumnTypePrecisionScale(index int) (int64, int64, bool) { + if rows, ok := r.Rows.(driver.RowsColumnTypePrecisionScale); ok { + return rows.ColumnTypePrecisionScale(index) + } + return 0, 0, false +} + +type serverTimingTx struct { + driver.Tx + ctx context.Context +} + +func (t *serverTimingTx) Commit() error { + startedAt := time.Now() + err := t.Tx.Commit() + servertiming.RecordInterval(t.ctx, servertiming.MetricDatabase, startedAt, time.Now()) + return err +} + +func (t *serverTimingTx) Rollback() error { + startedAt := time.Now() + err := t.Tx.Rollback() + servertiming.RecordInterval(t.ctx, servertiming.MetricDatabase, startedAt, time.Now()) + return err +} diff --git a/backend/internal/repository/server_timing_sql_test.go b/backend/internal/repository/server_timing_sql_test.go new file mode 100644 index 0000000000..3a8bbbe03e --- /dev/null +++ b/backend/internal/repository/server_timing_sql_test.go @@ -0,0 +1,258 @@ +package repository + +import ( + "context" + "database/sql/driver" + "io" + "regexp" + "strconv" + "strings" + "testing" + "time" + + "github.com/Wei-Shaw/sub2api/internal/pkg/servertiming" +) + +const fakeDriverDelay = 2 * time.Millisecond + +type timingFakeDriver struct{} + +func (timingFakeDriver) Open(string) (driver.Conn, error) { return newTimingFakeConn(), nil } + +type timingFakeConnector struct { + conn driver.Conn +} + +func (c timingFakeConnector) Connect(context.Context) (driver.Conn, error) { + time.Sleep(fakeDriverDelay) + return c.conn, nil +} + +func (timingFakeConnector) Driver() driver.Driver { return timingFakeDriver{} } + +type timingFakeConn struct{} + +func newTimingFakeConn() *timingFakeConn { return &timingFakeConn{} } + +func (c *timingFakeConn) Prepare(string) (driver.Stmt, error) { + time.Sleep(fakeDriverDelay) + return &timingFakeStmt{}, nil +} + +func (c *timingFakeConn) PrepareContext(context.Context, string) (driver.Stmt, error) { + time.Sleep(fakeDriverDelay) + return &timingFakeStmt{}, nil +} + +func (c *timingFakeConn) Close() error { return nil } + +func (c *timingFakeConn) Begin() (driver.Tx, error) { + time.Sleep(fakeDriverDelay) + return &timingFakeTx{}, nil +} + +func (c *timingFakeConn) BeginTx(context.Context, driver.TxOptions) (driver.Tx, error) { + time.Sleep(fakeDriverDelay) + return &timingFakeTx{}, nil +} + +func (c *timingFakeConn) ExecContext(context.Context, string, []driver.NamedValue) (driver.Result, error) { + time.Sleep(fakeDriverDelay) + return driver.RowsAffected(1), nil +} + +func (c *timingFakeConn) QueryContext(context.Context, string, []driver.NamedValue) (driver.Rows, error) { + time.Sleep(fakeDriverDelay) + return &timingFakeRows{values: [][]driver.Value{{"value"}}}, nil +} + +func (c *timingFakeConn) Ping(context.Context) error { + time.Sleep(fakeDriverDelay) + return nil +} + +func (c *timingFakeConn) ResetSession(context.Context) error { + time.Sleep(fakeDriverDelay) + return nil +} + +type timingFakeStmt struct{} + +func (s *timingFakeStmt) Close() error { return nil } +func (s *timingFakeStmt) NumInput() int { return -1 } + +func (s *timingFakeStmt) Exec([]driver.Value) (driver.Result, error) { + time.Sleep(fakeDriverDelay) + return driver.RowsAffected(1), nil +} + +func (s *timingFakeStmt) Query([]driver.Value) (driver.Rows, error) { + time.Sleep(fakeDriverDelay) + return &timingFakeRows{values: [][]driver.Value{{"value"}}}, nil +} + +func (s *timingFakeStmt) ExecContext(context.Context, []driver.NamedValue) (driver.Result, error) { + time.Sleep(fakeDriverDelay) + return driver.RowsAffected(1), nil +} + +func (s *timingFakeStmt) QueryContext(context.Context, []driver.NamedValue) (driver.Rows, error) { + time.Sleep(fakeDriverDelay) + return &timingFakeRows{values: [][]driver.Value{{"value"}}}, nil +} + +type timingFakeRows struct { + values [][]driver.Value + index int +} + +func (r *timingFakeRows) Columns() []string { return []string{"value"} } + +func (r *timingFakeRows) Close() error { + time.Sleep(fakeDriverDelay) + return nil +} + +func (r *timingFakeRows) Next(dest []driver.Value) error { + time.Sleep(fakeDriverDelay) + if r.index >= len(r.values) { + return io.EOF + } + copy(dest, r.values[r.index]) + r.index++ + return nil +} + +type timingFakeTx struct{} + +func (t *timingFakeTx) Commit() error { + time.Sleep(fakeDriverDelay) + return nil +} + +func (t *timingFakeTx) Rollback() error { + time.Sleep(fakeDriverDelay) + return nil +} + +func metricDuration(t *testing.T, header, metric string) float64 { + t.Helper() + re := regexp.MustCompile(`(?:^|, )` + regexp.QuoteMeta(metric) + `;dur=([0-9]+(?:\.[0-9]+)?)`) + match := re.FindStringSubmatch(header) + if len(match) != 2 { + t.Fatalf("metric %q missing from header %q", metric, header) + } + value, err := strconv.ParseFloat(match[1], 64) + if err != nil { + t.Fatalf("parse %s duration: %v", metric, err) + } + return value +} + +func TestServerTimingConnectorRecordsDriverCallsWithoutRowLifetime(t *testing.T) { + startedAt := time.Now() + collector := servertiming.New(startedAt) + ctx := servertiming.WithCollector(context.Background(), collector) + + wrapped := newServerTimingConnector(timingFakeConnector{conn: newTimingFakeConn()}) + rawConn, err := wrapped.Connect(ctx) + if err != nil { + t.Fatal(err) + } + conn, ok := rawConn.(*serverTimingConn) + if !ok { + t.Fatalf("Connect() returned %T, want *serverTimingConn", rawConn) + } + + if _, err := conn.ExecContext(ctx, "sensitive update", nil); err != nil { + t.Fatal(err) + } + rows, err := conn.QueryContext(ctx, "sensitive select", nil) + if err != nil { + t.Fatal(err) + } + values := make([]driver.Value, 1) + if err := rows.Next(values); err != nil { + t.Fatal(err) + } + + // Application work between row reads must remain app time. + time.Sleep(30 * time.Millisecond) + if err := rows.Next(values); err != io.EOF { + t.Fatalf("rows.Next() = %v, want EOF", err) + } + if err := rows.Close(); err != nil { + t.Fatal(err) + } + + header := collector.HeaderValue(time.Now(), "bypass") + if !strings.Contains(header, `queries=2`) { + t.Fatalf("header %q does not report two SQL operations", header) + } + if strings.Contains(header, "sensitive") { + t.Fatalf("SQL text leaked into header: %q", header) + } + if app, db := metricDuration(t, header, "app"), metricDuration(t, header, "db"); app <= db { + t.Fatalf("row processing gap was counted as DB time: app=%.1fms db=%.1fms header=%q", app, db, header) + } +} + +func TestServerTimingPreparedStatementsAndTransactions(t *testing.T) { + collector := servertiming.New(time.Now()) + ctx := servertiming.WithCollector(context.Background(), collector) + conn := &serverTimingConn{Conn: newTimingFakeConn()} + + stmt, err := conn.PrepareContext(ctx, "prepare sensitive statement") + if err != nil { + t.Fatal(err) + } + timedStmt, ok := stmt.(*serverTimingStmt) + if !ok { + t.Fatalf("PrepareContext() returned %T, want *serverTimingStmt", stmt) + } + if _, err := timedStmt.ExecContext(ctx, nil); err != nil { + t.Fatal(err) + } + rows, err := timedStmt.QueryContext(ctx, nil) + if err != nil { + t.Fatal(err) + } + if err := rows.Close(); err != nil { + t.Fatal(err) + } + + tx, err := conn.BeginTx(ctx, driver.TxOptions{}) + if err != nil { + t.Fatal(err) + } + if err := tx.Commit(); err != nil { + t.Fatal(err) + } + if err := conn.Ping(ctx); err != nil { + t.Fatal(err) + } + if err := conn.ResetSession(ctx); err != nil { + t.Fatal(err) + } + + header := collector.HeaderValue(time.Now(), "bypass") + if !strings.Contains(header, `queries=3`) { + t.Fatalf("header %q does not report prepare, exec, and query operations", header) + } + if metricDuration(t, header, "db") <= 0 { + t.Fatalf("DB duration was not recorded: %q", header) + } +} + +func TestNamedValuesRejectNamedParameters(t *testing.T) { + if _, err := namedValues([]driver.NamedValue{{Name: "secret", Value: 1}}); err == nil { + t.Fatal("namedValues accepted a named parameter") + } + values, err := namedValues([]driver.NamedValue{{Ordinal: 1, Value: "value"}}) + if err != nil { + t.Fatal(err) + } + if len(values) != 1 || values[0] != "value" { + t.Fatalf("namedValues() = %#v", values) + } +} diff --git a/backend/internal/repository/usage_log_repo_insert.go b/backend/internal/repository/usage_log_repo_insert.go index dfd8969512..ec09b308a0 100644 --- a/backend/internal/repository/usage_log_repo_insert.go +++ b/backend/internal/repository/usage_log_repo_insert.go @@ -71,6 +71,7 @@ var usageLogInsertArgTypes = [...]string{ "text", // inbound_endpoint "text", // upstream_endpoint "boolean", // cache_ttl_overridden + "boolean", // long_context_billing_applied "bigint", // channel_id "text", // model_mapping_chain "text", // billing_tier @@ -263,6 +264,7 @@ func (r *usageLogRepository) createSingle(ctx context.Context, sqlq sqlExecutor, inbound_endpoint, upstream_endpoint, cache_ttl_overridden, + long_context_billing_applied, channel_id, model_mapping_chain, billing_tier, @@ -275,7 +277,7 @@ func (r *usageLogRepository) createSingle(ctx context.Context, sqlq sqlExecutor, $10, $11, $12, $13, $14, $15, $16, $17, $18, $19, $20, $21, $22, $23, - $24, $25, $26, $27, $28, $29, $30, $31, $32, $33, $34, $35, $36, $37, $38, $39, $40, $41, $42, $43, $44, $45, $46, $47, $48, $49, $50, $51, $52, $53 + $24, $25, $26, $27, $28, $29, $30, $31, $32, $33, $34, $35, $36, $37, $38, $39, $40, $41, $42, $43, $44, $45, $46, $47, $48, $49, $50, $51, $52, $53, $54 ) ON CONFLICT (request_id, api_key_id) DO NOTHING RETURNING id, created_at @@ -714,6 +716,7 @@ func buildUsageLogBatchInsertQuery(keys []string, preparedByKey map[string]usage inbound_endpoint, upstream_endpoint, cache_ttl_overridden, + long_context_billing_applied, channel_id, model_mapping_chain, billing_tier, @@ -722,7 +725,7 @@ func buildUsageLogBatchInsertQuery(keys []string, preparedByKey map[string]usage created_at ) AS (VALUES `) - args := make([]any, 0, len(keys)*53) + args := make([]any, 0, len(keys)*54) argPos := 1 for idx, key := range keys { if idx > 0 { @@ -798,6 +801,7 @@ func buildUsageLogBatchInsertQuery(keys []string, preparedByKey map[string]usage inbound_endpoint, upstream_endpoint, cache_ttl_overridden, + long_context_billing_applied, channel_id, model_mapping_chain, billing_tier, @@ -853,6 +857,7 @@ func buildUsageLogBatchInsertQuery(keys []string, preparedByKey map[string]usage inbound_endpoint, upstream_endpoint, cache_ttl_overridden, + long_context_billing_applied, channel_id, model_mapping_chain, billing_tier, @@ -948,6 +953,7 @@ func buildUsageLogBestEffortInsertQuery(preparedList []usageLogInsertPrepared) ( inbound_endpoint, upstream_endpoint, cache_ttl_overridden, + long_context_billing_applied, channel_id, model_mapping_chain, billing_tier, @@ -956,7 +962,7 @@ func buildUsageLogBestEffortInsertQuery(preparedList []usageLogInsertPrepared) ( created_at ) AS (VALUES `) - args := make([]any, 0, len(preparedList)*53) + args := make([]any, 0, len(preparedList)*54) argPos := 1 for idx, prepared := range preparedList { if idx > 0 { @@ -1029,6 +1035,7 @@ func buildUsageLogBestEffortInsertQuery(preparedList []usageLogInsertPrepared) ( inbound_endpoint, upstream_endpoint, cache_ttl_overridden, + long_context_billing_applied, channel_id, model_mapping_chain, billing_tier, @@ -1084,6 +1091,7 @@ func buildUsageLogBestEffortInsertQuery(preparedList []usageLogInsertPrepared) ( inbound_endpoint, upstream_endpoint, cache_ttl_overridden, + long_context_billing_applied, channel_id, model_mapping_chain, billing_tier, @@ -1147,6 +1155,7 @@ func execUsageLogInsertNoResult(ctx context.Context, sqlq sqlExecutor, prepared inbound_endpoint, upstream_endpoint, cache_ttl_overridden, + long_context_billing_applied, channel_id, model_mapping_chain, billing_tier, @@ -1159,7 +1168,7 @@ func execUsageLogInsertNoResult(ctx context.Context, sqlq sqlExecutor, prepared $10, $11, $12, $13, $14, $15, $16, $17, $18, $19, $20, $21, $22, $23, - $24, $25, $26, $27, $28, $29, $30, $31, $32, $33, $34, $35, $36, $37, $38, $39, $40, $41, $42, $43, $44, $45, $46, $47, $48, $49, $50, $51, $52, $53 + $24, $25, $26, $27, $28, $29, $30, $31, $32, $33, $34, $35, $36, $37, $38, $39, $40, $41, $42, $43, $44, $45, $46, $47, $48, $49, $50, $51, $52, $53, $54 ) ON CONFLICT (request_id, api_key_id) DO NOTHING `, prepared.args...) @@ -1264,6 +1273,7 @@ func prepareUsageLogInsert(log *service.UsageLog) usageLogInsertPrepared { inboundEndpoint, upstreamEndpoint, log.CacheTTLOverridden, + log.LongContextBillingApplied, channelID, modelMappingChain, billingTier, diff --git a/backend/internal/repository/usage_log_repo_query.go b/backend/internal/repository/usage_log_repo_query.go index c178429bab..1fdedd8665 100644 --- a/backend/internal/repository/usage_log_repo_query.go +++ b/backend/internal/repository/usage_log_repo_query.go @@ -19,7 +19,7 @@ import ( "github.com/Wei-Shaw/sub2api/internal/service" ) -const usageLogSelectColumns = "id, user_id, api_key_id, account_id, request_id, model, requested_model, upstream_model, group_id, subscription_id, input_tokens, output_tokens, cache_creation_tokens, cache_read_tokens, cache_creation_5m_tokens, cache_creation_1h_tokens, image_output_tokens, image_output_cost, input_cost, output_cost, cache_creation_cost, cache_read_cost, total_cost, actual_cost, rate_multiplier, account_rate_multiplier, billing_type, request_type, stream, openai_ws_mode, duration_ms, first_token_ms, user_agent, ip_address, image_count, image_size, image_input_size, image_output_size, image_size_source, image_size_breakdown, video_count, video_resolution, video_duration_seconds, service_tier, reasoning_effort, inbound_endpoint, upstream_endpoint, cache_ttl_overridden, channel_id, model_mapping_chain, billing_tier, billing_mode, account_stats_cost, created_at" +const usageLogSelectColumns = "id, user_id, api_key_id, account_id, request_id, model, requested_model, upstream_model, group_id, subscription_id, input_tokens, output_tokens, cache_creation_tokens, cache_read_tokens, cache_creation_5m_tokens, cache_creation_1h_tokens, image_output_tokens, image_output_cost, input_cost, output_cost, cache_creation_cost, cache_read_cost, total_cost, actual_cost, rate_multiplier, account_rate_multiplier, billing_type, request_type, stream, openai_ws_mode, duration_ms, first_token_ms, user_agent, ip_address, image_count, image_size, image_input_size, image_output_size, image_size_source, image_size_breakdown, video_count, video_resolution, video_duration_seconds, service_tier, reasoning_effort, inbound_endpoint, upstream_endpoint, cache_ttl_overridden, long_context_billing_applied, channel_id, model_mapping_chain, billing_tier, billing_mode, account_stats_cost, created_at" func (r *usageLogRepository) GetByID(ctx context.Context, id int64) (log *service.UsageLog, err error) { query := "SELECT " + usageLogSelectColumns + " FROM usage_logs WHERE id = $1" @@ -425,60 +425,61 @@ func (r *usageLogRepository) loadSubscriptions(ctx context.Context, ids []int64) func scanUsageLog(scanner interface{ Scan(...any) error }) (*service.UsageLog, error) { var ( - id int64 - userID int64 - apiKeyID int64 - accountID int64 - requestID sql.NullString - model string - requestedModel sql.NullString - upstreamModel sql.NullString - groupID sql.NullInt64 - subscriptionID sql.NullInt64 - inputTokens int - outputTokens int - cacheCreationTokens int - cacheReadTokens int - cacheCreation5m int - cacheCreation1h int - imageOutputTokens int - imageOutputCost float64 - inputCost float64 - outputCost float64 - cacheCreationCost float64 - cacheReadCost float64 - totalCost float64 - actualCost float64 - rateMultiplier float64 - accountRateMultiplier sql.NullFloat64 - billingType int16 - requestTypeRaw int16 - stream bool - openaiWSMode bool - durationMs sql.NullInt64 - firstTokenMs sql.NullInt64 - userAgent sql.NullString - ipAddress sql.NullString - imageCount int - imageSize sql.NullString - imageInputSize sql.NullString - imageOutputSize sql.NullString - imageSizeSource sql.NullString - imageSizeBreakdown sql.NullString - videoCount int - videoResolution sql.NullString - videoDurationSeconds sql.NullInt64 - serviceTier sql.NullString - reasoningEffort sql.NullString - inboundEndpoint sql.NullString - upstreamEndpoint sql.NullString - cacheTTLOverridden bool - channelID sql.NullInt64 - modelMappingChain sql.NullString - billingTier sql.NullString - billingMode sql.NullString - accountStatsCost sql.NullFloat64 - createdAt time.Time + id int64 + userID int64 + apiKeyID int64 + accountID int64 + requestID sql.NullString + model string + requestedModel sql.NullString + upstreamModel sql.NullString + groupID sql.NullInt64 + subscriptionID sql.NullInt64 + inputTokens int + outputTokens int + cacheCreationTokens int + cacheReadTokens int + cacheCreation5m int + cacheCreation1h int + imageOutputTokens int + imageOutputCost float64 + inputCost float64 + outputCost float64 + cacheCreationCost float64 + cacheReadCost float64 + totalCost float64 + actualCost float64 + rateMultiplier float64 + accountRateMultiplier sql.NullFloat64 + billingType int16 + requestTypeRaw int16 + stream bool + openaiWSMode bool + durationMs sql.NullInt64 + firstTokenMs sql.NullInt64 + userAgent sql.NullString + ipAddress sql.NullString + imageCount int + imageSize sql.NullString + imageInputSize sql.NullString + imageOutputSize sql.NullString + imageSizeSource sql.NullString + imageSizeBreakdown sql.NullString + videoCount int + videoResolution sql.NullString + videoDurationSeconds sql.NullInt64 + serviceTier sql.NullString + reasoningEffort sql.NullString + inboundEndpoint sql.NullString + upstreamEndpoint sql.NullString + cacheTTLOverridden bool + longContextBillingApplied bool + channelID sql.NullInt64 + modelMappingChain sql.NullString + billingTier sql.NullString + billingMode sql.NullString + accountStatsCost sql.NullFloat64 + createdAt time.Time ) if err := scanner.Scan( @@ -530,6 +531,7 @@ func scanUsageLog(scanner interface{ Scan(...any) error }) (*service.UsageLog, e &inboundEndpoint, &upstreamEndpoint, &cacheTTLOverridden, + &longContextBillingApplied, &channelID, &modelMappingChain, &billingTier, @@ -541,34 +543,35 @@ func scanUsageLog(scanner interface{ Scan(...any) error }) (*service.UsageLog, e } log := &service.UsageLog{ - ID: id, - UserID: userID, - APIKeyID: apiKeyID, - AccountID: accountID, - Model: model, - RequestedModel: coalesceTrimmedString(requestedModel, model), - InputTokens: inputTokens, - OutputTokens: outputTokens, - CacheCreationTokens: cacheCreationTokens, - CacheReadTokens: cacheReadTokens, - CacheCreation5mTokens: cacheCreation5m, - CacheCreation1hTokens: cacheCreation1h, - ImageOutputTokens: imageOutputTokens, - ImageOutputCost: imageOutputCost, - InputCost: inputCost, - OutputCost: outputCost, - CacheCreationCost: cacheCreationCost, - CacheReadCost: cacheReadCost, - TotalCost: totalCost, - ActualCost: actualCost, - RateMultiplier: rateMultiplier, - AccountRateMultiplier: nullFloat64Ptr(accountRateMultiplier), - BillingType: int8(billingType), - RequestType: service.RequestTypeFromInt16(requestTypeRaw), - ImageCount: imageCount, - VideoCount: videoCount, - CacheTTLOverridden: cacheTTLOverridden, - CreatedAt: createdAt, + ID: id, + UserID: userID, + APIKeyID: apiKeyID, + AccountID: accountID, + Model: model, + RequestedModel: coalesceTrimmedString(requestedModel, model), + InputTokens: inputTokens, + OutputTokens: outputTokens, + CacheCreationTokens: cacheCreationTokens, + CacheReadTokens: cacheReadTokens, + CacheCreation5mTokens: cacheCreation5m, + CacheCreation1hTokens: cacheCreation1h, + ImageOutputTokens: imageOutputTokens, + ImageOutputCost: imageOutputCost, + InputCost: inputCost, + OutputCost: outputCost, + CacheCreationCost: cacheCreationCost, + CacheReadCost: cacheReadCost, + TotalCost: totalCost, + ActualCost: actualCost, + RateMultiplier: rateMultiplier, + AccountRateMultiplier: nullFloat64Ptr(accountRateMultiplier), + BillingType: int8(billingType), + RequestType: service.RequestTypeFromInt16(requestTypeRaw), + ImageCount: imageCount, + VideoCount: videoCount, + CacheTTLOverridden: cacheTTLOverridden, + LongContextBillingApplied: longContextBillingApplied, + CreatedAt: createdAt, } // 先回填 legacy 字段,再基于 legacy + request_type 计算最终请求类型,保证历史数据兼容。 log.Stream = stream diff --git a/backend/internal/repository/usage_log_repo_request_type_test.go b/backend/internal/repository/usage_log_repo_request_type_test.go index c32ad2b63f..052c319183 100644 --- a/backend/internal/repository/usage_log_repo_request_type_test.go +++ b/backend/internal/repository/usage_log_repo_request_type_test.go @@ -88,6 +88,7 @@ func TestUsageLogRepositoryCreateSyncRequestTypeAndLegacyFields(t *testing.T) { sqlmock.AnyArg(), // inbound_endpoint sqlmock.AnyArg(), // upstream_endpoint log.CacheTTLOverridden, + log.LongContextBillingApplied, sqlmock.AnyArg(), // channel_id sqlmock.AnyArg(), // model_mapping_chain sqlmock.AnyArg(), // billing_tier @@ -174,6 +175,7 @@ func TestUsageLogRepositoryCreate_PersistsServiceTier(t *testing.T) { sqlmock.AnyArg(), sqlmock.AnyArg(), log.CacheTTLOverridden, + log.LongContextBillingApplied, sqlmock.AnyArg(), // channel_id sqlmock.AnyArg(), // model_mapping_chain sqlmock.AnyArg(), // billing_tier @@ -813,6 +815,7 @@ func TestScanUsageLogRequestTypeAndLegacyFallback(t *testing.T) { sql.NullString{}, sql.NullString{}, false, + false, sql.NullInt64{}, sql.NullString{}, sql.NullString{}, @@ -884,6 +887,7 @@ func TestScanUsageLogRequestTypeAndLegacyFallback(t *testing.T) { sql.NullString{}, sql.NullString{}, false, + false, sql.NullInt64{}, // channel_id sql.NullString{}, // model_mapping_chain sql.NullString{}, // billing_tier @@ -939,6 +943,7 @@ func TestScanUsageLogRequestTypeAndLegacyFallback(t *testing.T) { sql.NullString{}, sql.NullString{}, false, + false, sql.NullInt64{}, // channel_id sql.NullString{}, // model_mapping_chain sql.NullString{}, // billing_tier @@ -994,6 +999,7 @@ func TestScanUsageLogRequestTypeAndLegacyFallback(t *testing.T) { sql.NullString{}, sql.NullString{}, false, + false, sql.NullInt64{}, // channel_id sql.NullString{}, // model_mapping_chain sql.NullString{}, // billing_tier diff --git a/backend/internal/server/api_contract_test.go b/backend/internal/server/api_contract_test.go index 372cc46bbf..a5e3fde155 100644 --- a/backend/internal/server/api_contract_test.go +++ b/backend/internal/server/api_contract_test.go @@ -594,6 +594,7 @@ func TestAPIContracts(t *testing.T) { "total_cost": 0.5, "actual_cost": 0.5, "rate_multiplier": 1, + "long_context_billing_applied": false, "billing_type": 0, "stream": true, "duration_ms": 100, diff --git a/backend/internal/server/middleware/cors.go b/backend/internal/server/middleware/cors.go index 03d5d025de..0283d53115 100644 --- a/backend/internal/server/middleware/cors.go +++ b/backend/internal/server/middleware/cors.go @@ -52,7 +52,7 @@ func CORS(cfg config.CORSConfig) gin.HandlerFunc { } allowHeaders := []string{ "Content-Type", "Content-Length", "Accept-Encoding", "X-CSRF-Token", "Authorization", - "accept", "origin", "Cache-Control", "X-Requested-With", "X-API-Key", + "accept", "origin", "Cache-Control", "X-Requested-With", "X-API-Key", "X-Admin-UI-Request", } // OpenAI Node SDK 会发送 x-stainless-* 请求头,需在 CORS 中显式放行。 openAIProperties := []string{ @@ -83,7 +83,7 @@ func CORS(cfg config.CORSConfig) gin.HandlerFunc { } c.Writer.Header().Set("Access-Control-Allow-Headers", allowHeadersValue) c.Writer.Header().Set("Access-Control-Allow-Methods", "POST, OPTIONS, GET, PUT, DELETE, PATCH") - c.Writer.Header().Set("Access-Control-Expose-Headers", "ETag") + c.Writer.Header().Set("Access-Control-Expose-Headers", "ETag, Server-Timing") c.Writer.Header().Set("Access-Control-Max-Age", "86400") } // 处理预检请求 diff --git a/backend/internal/server/middleware/cors_test.go b/backend/internal/server/middleware/cors_test.go index 6d0bea3608..6a61f696df 100644 --- a/backend/internal/server/middleware/cors_test.go +++ b/backend/internal/server/middleware/cors_test.go @@ -103,8 +103,10 @@ func TestCORS_AllowedOrigin_HasAllowHeaders(t *testing.T) { // 应设置 Allow-Headers、Allow-Methods 和 Max-Age assert.NotEmpty(t, w.Header().Get("Access-Control-Allow-Headers"), "允许的 origin 应收到 Allow-Headers") + assert.Contains(t, w.Header().Get("Access-Control-Allow-Headers"), "X-Admin-UI-Request") assert.NotEmpty(t, w.Header().Get("Access-Control-Allow-Methods"), "允许的 origin 应收到 Allow-Methods") + assert.Contains(t, w.Header().Get("Access-Control-Expose-Headers"), "Server-Timing") assert.Equal(t, "86400", w.Header().Get("Access-Control-Max-Age"), "允许的 origin 应收到 Max-Age=86400") assert.Equal(t, "https://allowed.example.com", w.Header().Get("Access-Control-Allow-Origin"), diff --git a/backend/internal/server/middleware/server_timing.go b/backend/internal/server/middleware/server_timing.go new file mode 100644 index 0000000000..2bb21071e0 --- /dev/null +++ b/backend/internal/server/middleware/server_timing.go @@ -0,0 +1,132 @@ +package middleware + +import ( + "net/http" + "strings" + "sync" + "time" + + "github.com/Wei-Shaw/sub2api/internal/pkg/servertiming" + "github.com/gin-gonic/gin" +) + +const ( + snapshotCacheHeader = "X-Snapshot-Cache" + usageCacheHeader = "X-Usage-Stats-Cache" +) + +type serverTimingResponseWriter struct { + gin.ResponseWriter + context *gin.Context + once sync.Once +} + +func (w *serverTimingResponseWriter) Unwrap() http.ResponseWriter { + return w.ResponseWriter +} + +// ServerTiming collects timing only for requests made by the Admin web UI. +func ServerTiming(enabled bool) gin.HandlerFunc { + if !enabled { + return func(c *gin.Context) { + c.Next() + } + } + return func(c *gin.Context) { + if !isAdminUIRequest(c) || c.Request == nil { + c.Next() + return + } + + collector := servertiming.New(time.Now()) + c.Request = c.Request.WithContext(servertiming.WithCollector(c.Request.Context(), collector)) + writer := &serverTimingResponseWriter{ + ResponseWriter: c.Writer, + context: c, + } + c.Writer = writer + c.Next() + writer.finalize() + } +} + +func (w *serverTimingResponseWriter) WriteHeader(statusCode int) { + w.ResponseWriter.WriteHeader(statusCode) +} + +func (w *serverTimingResponseWriter) WriteHeaderNow() { + w.finalize() + w.ResponseWriter.WriteHeaderNow() +} + +func (w *serverTimingResponseWriter) Write(data []byte) (int, error) { + w.finalize() + return w.ResponseWriter.Write(data) +} + +func (w *serverTimingResponseWriter) WriteString(data string) (int, error) { + w.finalize() + return w.ResponseWriter.WriteString(data) +} + +func (w *serverTimingResponseWriter) Flush() { + w.finalize() + w.ResponseWriter.Flush() +} + +func (w *serverTimingResponseWriter) finalize() { + if w == nil { + return + } + w.once.Do(func() { + if value := ServerTimingHeaderValue(w.context); value != "" { + w.ResponseWriter.Header().Set(servertiming.HeaderName, value) + } + }) +} + +// ServerTimingHeaderValue returns a timing value only for an authenticated admin. +func ServerTimingHeaderValue(c *gin.Context) string { + if c == nil || c.Request == nil { + return "" + } + role, ok := GetUserRoleFromContext(c) + if !ok || role != "admin" { + return "" + } + return servertiming.HeaderValue(c.Request.Context(), time.Now(), responseCacheStatus(c.Writer.Header())) +} + +// ServerTimingResponseHeader builds the extra header map required by WebSocket upgrades. +func ServerTimingResponseHeader(c *gin.Context) http.Header { + value := ServerTimingHeaderValue(c) + if value == "" { + return nil + } + return http.Header{servertiming.HeaderName: []string{value}} +} + +func isAdminUIRequest(c *gin.Context) bool { + if c == nil || c.Request == nil || c.Request.URL == nil { + return false + } + if strings.TrimSpace(c.GetHeader(servertiming.AdminUIHeader)) == "1" { + return true + } + path := strings.TrimSpace(c.Request.URL.Path) + return path == "/api/v1/admin" || strings.HasPrefix(path, "/api/v1/admin/") +} + +func responseCacheStatus(header http.Header) string { + for _, name := range []string{snapshotCacheHeader, usageCacheHeader} { + switch strings.ToLower(strings.TrimSpace(header.Get(name))) { + case "hit": + return "hit" + case "miss": + return "miss" + case "bypass": + return "bypass" + } + } + return "bypass" +} diff --git a/backend/internal/server/middleware/server_timing_test.go b/backend/internal/server/middleware/server_timing_test.go new file mode 100644 index 0000000000..c064840ece --- /dev/null +++ b/backend/internal/server/middleware/server_timing_test.go @@ -0,0 +1,188 @@ +package middleware + +import ( + "net/http" + "net/http/httptest" + "strings" + "testing" + "time" + + "github.com/Wei-Shaw/sub2api/internal/pkg/servertiming" + "github.com/gin-gonic/gin" +) + +func runServerTimingRequest( + t *testing.T, + enabled bool, + path string, + marker string, + role string, + handler gin.HandlerFunc, +) *httptest.ResponseRecorder { + t.Helper() + gin.SetMode(gin.TestMode) + engine := gin.New() + engine.Use(ServerTiming(enabled)) + engine.Any("/*path", func(c *gin.Context) { + if role != "" { + c.Set(string(ContextKeyUserRole), role) + } + handler(c) + }) + + recorder := httptest.NewRecorder() + request := httptest.NewRequest(http.MethodGet, path, nil) + if marker != "" { + request.Header.Set(servertiming.AdminUIHeader, marker) + } + engine.ServeHTTP(recorder, request) + return recorder +} + +func TestServerTimingScopesAndRoleGate(t *testing.T) { + tests := []struct { + name string + enabled bool + path string + marker string + role string + wantHeader bool + }{ + {name: "disabled", enabled: false, path: "/api/v1/admin/users", role: "admin"}, + {name: "admin API path", enabled: true, path: "/api/v1/admin/users", role: "admin", wantHeader: true}, + {name: "shared API marked by admin UI", enabled: true, path: "/api/v1/groups/available", marker: "1", role: "admin", wantHeader: true}, + {name: "non admin role", enabled: true, path: "/api/v1/groups/available", marker: "1", role: "user"}, + {name: "unauthenticated public request", enabled: true, path: "/api/v1/settings/public", marker: "1"}, + {name: "unmarked shared API", enabled: true, path: "/api/v1/groups/available", role: "admin"}, + {name: "invalid marker", enabled: true, path: "/api/v1/groups/available", marker: "true", role: "admin"}, + {name: "admin prefix boundary", enabled: true, path: "/api/v1/administrator", role: "admin"}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + recorder := runServerTimingRequest(t, tt.enabled, tt.path, tt.marker, tt.role, func(c *gin.Context) { + c.JSON(http.StatusOK, gin.H{"ok": true}) + }) + header := recorder.Header().Get(servertiming.HeaderName) + if tt.wantHeader && header == "" { + t.Fatalf("%s header missing", servertiming.HeaderName) + } + if !tt.wantHeader && header != "" { + t.Fatalf("unexpected %s header: %q", servertiming.HeaderName, header) + } + if header != "" && (!strings.Contains(header, "total;dur=") || !strings.Contains(header, `cache;desc="bypass"`)) { + t.Fatalf("incomplete timing header: %q", header) + } + }) + } +} + +func TestServerTimingCollectorIsRequestScoped(t *testing.T) { + active := false + recorder := runServerTimingRequest(t, true, "/api/v1/keys", "1", "admin", func(c *gin.Context) { + active = servertiming.Active(c.Request.Context()) + c.Status(http.StatusNoContent) + }) + if !active { + t.Fatal("collector was not attached to marked request context") + } + if recorder.Header().Get(servertiming.HeaderName) == "" { + t.Fatal("timing header missing from status-only response") + } +} + +func TestServerTimingFinalizesBeforeEarlyCommit(t *testing.T) { + recorder := runServerTimingRequest(t, true, "/api/v1/admin/stream", "", "admin", func(c *gin.Context) { + c.Status(http.StatusAccepted) + c.Writer.WriteHeaderNow() + }) + if got := recorder.Header().Get(servertiming.HeaderName); got == "" { + t.Fatal("timing header was not written before response commit") + } +} + +func TestServerTimingFinalizesOnFlush(t *testing.T) { + recorder := runServerTimingRequest(t, true, "/api/v1/admin/export", "", "admin", func(c *gin.Context) { + c.Writer.Flush() + }) + if got := recorder.Header().Get(servertiming.HeaderName); got == "" { + t.Fatal("timing header was not written before stream flush") + } +} + +func TestServerTimingStatusResponses(t *testing.T) { + tests := []struct { + name string + status int + }{ + {name: "not modified", status: http.StatusNotModified}, + {name: "internal error", status: http.StatusInternalServerError}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + recorder := runServerTimingRequest(t, true, "/api/v1/admin/test", "", "admin", func(c *gin.Context) { + c.Status(tt.status) + }) + if recorder.Code != tt.status { + t.Fatalf("status = %d, want %d", recorder.Code, tt.status) + } + if got := recorder.Header().Get(servertiming.HeaderName); got == "" { + t.Fatalf("timing header missing from status %d response", tt.status) + } + }) + } +} + +func TestServerTimingResponseWriterUnwraps(t *testing.T) { + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + baseWriter := c.Writer + writer := &serverTimingResponseWriter{ResponseWriter: baseWriter} + if got := writer.Unwrap(); got != baseWriter { + t.Fatalf("Unwrap() = %T, want original Gin writer", got) + } +} + +func TestServerTimingCacheOutcome(t *testing.T) { + tests := []struct { + name string + headerName string + value string + want string + }{ + {name: "snapshot hit", headerName: snapshotCacheHeader, value: "hit", want: "hit"}, + {name: "usage miss", headerName: usageCacheHeader, value: "MISS", want: "miss"}, + {name: "invalid", headerName: snapshotCacheHeader, value: "stale", want: "bypass"}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + recorder := runServerTimingRequest(t, true, "/api/v1/admin/dashboard", "", "admin", func(c *gin.Context) { + c.Header(tt.headerName, tt.value) + c.JSON(http.StatusOK, gin.H{"ok": true}) + }) + want := `cache;desc="` + tt.want + `"` + if got := recorder.Header().Get(servertiming.HeaderName); !strings.Contains(got, want) { + t.Fatalf("timing header %q does not contain %q", got, want) + } + }) + } +} + +func TestServerTimingResponseHeaderForWebSocket(t *testing.T) { + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + c.Request = httptest.NewRequest(http.MethodGet, "/api/v1/admin/ops/ws/qps", nil) + collector := servertiming.New(time.Now()) + c.Request = c.Request.WithContext(servertiming.WithCollector(c.Request.Context(), collector)) + c.Set(string(ContextKeyUserRole), "admin") + + header := ServerTimingResponseHeader(c) + if header.Get(servertiming.HeaderName) == "" { + t.Fatal("WebSocket response header missing timing value") + } + + c.Set(string(ContextKeyUserRole), "user") + if got := ServerTimingResponseHeader(c); got != nil { + t.Fatalf("non-admin WebSocket received timing header: %#v", got) + } +} diff --git a/backend/internal/server/router.go b/backend/internal/server/router.go index 3d86373779..5fc70149fe 100644 --- a/backend/internal/server/router.go +++ b/backend/internal/server/router.go @@ -60,6 +60,7 @@ func SetupRouter( } return nil })) + r.Use(middleware2.ServerTiming(cfg.Server.EnableServerTiming)) // Serve embedded frontend with settings injection if available if web.HasEmbeddedFrontend() { diff --git a/backend/internal/service/account.go b/backend/internal/service/account.go index 1ab2e21fdf..3e67fa5982 100644 --- a/backend/internal/service/account.go +++ b/backend/internal/service/account.go @@ -83,6 +83,8 @@ type Account struct { type OpenAIEndpointCapability string +const openAILongContextBillingEnabledKey = "openai_long_context_billing_enabled" + const ( OpenAIEndpointCapabilityChatCompletions OpenAIEndpointCapability = "chat_completions" OpenAIEndpointCapabilityEmbeddings OpenAIEndpointCapability = "embeddings" @@ -1192,6 +1194,14 @@ func (a *Account) IsOpenAI() bool { return a.Platform == PlatformOpenAI } +func (a *Account) IsOpenAILongContextBillingEnabled() bool { + if a == nil || !a.IsOpenAI() || a.Extra == nil { + return false + } + enabled, ok := a.Extra[openAILongContextBillingEnabledKey].(bool) + return ok && enabled +} + func (a *Account) IsAnthropic() bool { return a.Platform == PlatformAnthropic } diff --git a/backend/internal/service/account_long_context_billing_test.go b/backend/internal/service/account_long_context_billing_test.go new file mode 100644 index 0000000000..709559d932 --- /dev/null +++ b/backend/internal/service/account_long_context_billing_test.go @@ -0,0 +1,290 @@ +//go:build unit + +package service + +import ( + "context" + "net/http" + "testing" + + infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" + "github.com/stretchr/testify/require" +) + +func TestAccountIsOpenAILongContextBillingEnabled(t *testing.T) { + tests := []struct { + name string + account *Account + want bool + }{ + {name: "nil account is disabled", account: nil, want: false}, + {name: "non OpenAI account is disabled", account: &Account{Platform: PlatformGrok}, want: false}, + {name: "missing extra defaults disabled", account: &Account{Platform: PlatformOpenAI}, want: false}, + {name: "missing key defaults disabled", account: &Account{Platform: PlatformOpenAI, Extra: map[string]any{}}, want: false}, + {name: "explicit true is enabled", account: &Account{Platform: PlatformOpenAI, Extra: map[string]any{"openai_long_context_billing_enabled": true}}, want: true}, + {name: "explicit false is disabled", account: &Account{Platform: PlatformOpenAI, Extra: map[string]any{"openai_long_context_billing_enabled": false}}, want: false}, + {name: "malformed value is disabled", account: &Account{Platform: PlatformOpenAI, Extra: map[string]any{"openai_long_context_billing_enabled": "false"}}, want: false}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + require.Equal(t, tt.want, tt.account.IsOpenAILongContextBillingEnabled()) + }) + } +} + +func TestNormalizeOpenAILongContextBillingExtra(t *testing.T) { + t.Run("OpenAI missing key persists disabled default", func(t *testing.T) { + extra, err := normalizeOpenAILongContextBillingExtra(PlatformOpenAI, nil) + + require.NoError(t, err) + require.Equal(t, false, extra["openai_long_context_billing_enabled"]) + }) + + t.Run("OpenAI explicit false is preserved", func(t *testing.T) { + extra, err := normalizeOpenAILongContextBillingExtra(PlatformOpenAI, map[string]any{"openai_long_context_billing_enabled": false}) + + require.NoError(t, err) + require.Equal(t, false, extra["openai_long_context_billing_enabled"]) + }) + + t.Run("OpenAI malformed value is rejected", func(t *testing.T) { + _, err := normalizeOpenAILongContextBillingExtra(PlatformOpenAI, map[string]any{"openai_long_context_billing_enabled": "false"}) + + require.Error(t, err) + require.Equal(t, http.StatusBadRequest, infraerrors.Code(err)) + }) + + t.Run("non OpenAI extra is unchanged", func(t *testing.T) { + extra, err := normalizeOpenAILongContextBillingExtra(PlatformGrok, nil) + + require.NoError(t, err) + require.Nil(t, extra) + }) + + t.Run("non OpenAI malformed value is ignored", func(t *testing.T) { + extra := map[string]any{openAILongContextBillingEnabledKey: "provider-owned"} + normalized, err := normalizeOpenAILongContextBillingExtra(PlatformAnthropic, extra) + + require.NoError(t, err) + require.Equal(t, extra, normalized) + }) +} + +type longContextBillingRepoStub struct { + accountRepoStub + account *Account + accounts []*Account + createdAccount *Account + updateExtraCalls int + bulkUpdateCalls int +} + +func (r *longContextBillingRepoStub) Create(_ context.Context, account *Account) error { + account.ID = 1 + r.account = account + r.createdAccount = account + return nil +} + +func (r *longContextBillingRepoStub) GetByID(_ context.Context, _ int64) (*Account, error) { + return r.account, nil +} + +func (r *longContextBillingRepoStub) GetByIDs(_ context.Context, _ []int64) ([]*Account, error) { + if r.accounts != nil { + return r.accounts, nil + } + if r.account == nil { + return nil, nil + } + return []*Account{r.account}, nil +} + +func (r *longContextBillingRepoStub) Update(_ context.Context, account *Account) error { + r.account = account + return nil +} + +func (r *longContextBillingRepoStub) UpdateExtra(_ context.Context, _ int64, _ map[string]any) error { + r.updateExtraCalls++ + return nil +} + +func (r *longContextBillingRepoStub) BulkUpdate(_ context.Context, _ []int64, _ AccountBulkUpdate) (int64, error) { + r.bulkUpdateCalls++ + return 1, nil +} + +func TestAdminServiceCreateAccountDefaultsOpenAILongContextBillingDisabled(t *testing.T) { + repo := &longContextBillingRepoStub{} + svc := &adminServiceImpl{accountRepo: repo} + + account, err := svc.CreateAccount(context.Background(), &CreateAccountInput{ + Name: "openai-account", + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Credentials: map[string]any{"api_key": "test"}, + SkipDefaultGroupBind: true, + }) + + require.NoError(t, err) + require.Same(t, account, repo.createdAccount) + require.Equal(t, false, account.Extra[openAILongContextBillingEnabledKey]) +} + +func TestAdminServiceCreateAccountRejectsMalformedOpenAILongContextBillingValue(t *testing.T) { + repo := &longContextBillingRepoStub{} + svc := &adminServiceImpl{accountRepo: repo} + + account, err := svc.CreateAccount(context.Background(), &CreateAccountInput{ + Platform: PlatformOpenAI, + Extra: map[string]any{openAILongContextBillingEnabledKey: "false"}, + }) + + require.Nil(t, account) + require.Equal(t, http.StatusBadRequest, infraerrors.Code(err)) + require.Nil(t, repo.createdAccount) +} + +func TestAdminServiceUpdateAccountPreservesOpenAILongContextBillingOptOutWhenOmitted(t *testing.T) { + repo := &longContextBillingRepoStub{account: &Account{ + ID: 1, + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Extra: map[string]any{openAILongContextBillingEnabledKey: false}, + }} + svc := &adminServiceImpl{accountRepo: repo} + + account, err := svc.UpdateAccount(context.Background(), 1, &UpdateAccountInput{Extra: map[string]any{}}) + + require.NoError(t, err) + require.Equal(t, false, account.Extra[openAILongContextBillingEnabledKey]) +} + +func TestAdminServiceUpdateAccountAllowsExplicitCodexImportOptIn(t *testing.T) { + repo := &longContextBillingRepoStub{account: &Account{ + ID: 1, + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Credentials: map[string]any{"access_token": "old-token"}, + Extra: map[string]any{ + openAILongContextBillingEnabledKey: false, + "import_source": "codex_session", + }, + }} + svc := &adminServiceImpl{accountRepo: repo} + + account, err := svc.UpdateAccount(context.Background(), 1, &UpdateAccountInput{ + Credentials: map[string]any{"access_token": "new-token"}, + Extra: map[string]any{ + openAILongContextBillingEnabledKey: true, + "import_source": "codex_session", + }, + }) + + require.NoError(t, err) + require.Equal(t, true, account.Extra[openAILongContextBillingEnabledKey]) +} + +func TestAdminServiceUpdateAccountAllowsExplicitOptInOutsideCodexImport(t *testing.T) { + repo := &longContextBillingRepoStub{account: &Account{ + ID: 1, + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Extra: map[string]any{ + openAILongContextBillingEnabledKey: false, + "import_source": "codex_session", + }, + }} + svc := &adminServiceImpl{accountRepo: repo} + + account, err := svc.UpdateAccount(context.Background(), 1, &UpdateAccountInput{Extra: map[string]any{ + openAILongContextBillingEnabledKey: true, + "import_source": "codex_session", + }}) + + require.NoError(t, err) + require.Equal(t, true, account.Extra[openAILongContextBillingEnabledKey]) +} + +func TestAdminServiceUpdateAccountRejectsMalformedOpenAILongContextBillingValue(t *testing.T) { + repo := &longContextBillingRepoStub{account: &Account{ID: 1, Platform: PlatformOpenAI}} + svc := &adminServiceImpl{accountRepo: repo} + + account, err := svc.UpdateAccount(context.Background(), 1, &UpdateAccountInput{Extra: map[string]any{ + openAILongContextBillingEnabledKey: 1, + }}) + + require.Nil(t, account) + require.Equal(t, http.StatusBadRequest, infraerrors.Code(err)) +} + +func TestAdminServiceUpdateAccountExtraRejectsMalformedOpenAILongContextBillingValue(t *testing.T) { + repo := &longContextBillingRepoStub{account: &Account{ID: 1, Platform: PlatformOpenAI}} + svc := &adminServiceImpl{accountRepo: repo} + + err := svc.UpdateAccountExtra(context.Background(), 1, map[string]any{ + openAILongContextBillingEnabledKey: "true", + }) + + require.Equal(t, http.StatusBadRequest, infraerrors.Code(err)) + require.Zero(t, repo.updateExtraCalls) +} + +func TestAdminServiceUpdateAccountExtraAllowsProviderOwnedValueForNonOpenAIAccount(t *testing.T) { + repo := &longContextBillingRepoStub{account: &Account{ID: 1, Platform: PlatformAnthropic}} + svc := &adminServiceImpl{accountRepo: repo} + + err := svc.UpdateAccountExtra(context.Background(), 1, map[string]any{ + openAILongContextBillingEnabledKey: "provider-owned", + }) + + require.NoError(t, err) + require.Equal(t, 1, repo.updateExtraCalls) +} + +func TestAdminServiceBulkUpdateAccountsRejectsMalformedOpenAILongContextBillingValue(t *testing.T) { + repo := &longContextBillingRepoStub{account: &Account{ID: 1, Platform: PlatformOpenAI}} + svc := &adminServiceImpl{accountRepo: repo} + + result, err := svc.BulkUpdateAccounts(context.Background(), &BulkUpdateAccountsInput{ + AccountIDs: []int64{1}, + Extra: map[string]any{openAILongContextBillingEnabledKey: []bool{true}}, + }) + + require.Nil(t, result) + require.Equal(t, http.StatusBadRequest, infraerrors.Code(err)) + require.Zero(t, repo.bulkUpdateCalls) +} + +func TestAdminServiceBulkUpdateAccountsAllowsProviderOwnedValueForNonOpenAIAccounts(t *testing.T) { + repo := &longContextBillingRepoStub{account: &Account{ID: 1, Platform: PlatformGrok}} + svc := &adminServiceImpl{accountRepo: repo} + + result, err := svc.BulkUpdateAccounts(context.Background(), &BulkUpdateAccountsInput{ + AccountIDs: []int64{1}, + Extra: map[string]any{openAILongContextBillingEnabledKey: []string{"provider-owned"}}, + }) + + require.NoError(t, err) + require.NotNil(t, result) + require.Equal(t, 1, repo.bulkUpdateCalls) +} + +func TestAdminServiceBulkUpdateAccountsRejectsMalformedValueForMixedTargetsIncludingOpenAI(t *testing.T) { + repo := &longContextBillingRepoStub{accounts: []*Account{ + {ID: 1, Platform: PlatformGrok}, + {ID: 2, Platform: PlatformOpenAI}, + }} + svc := &adminServiceImpl{accountRepo: repo} + + result, err := svc.BulkUpdateAccounts(context.Background(), &BulkUpdateAccountsInput{ + AccountIDs: []int64{1, 2}, + Extra: map[string]any{openAILongContextBillingEnabledKey: "malformed"}, + }) + + require.Nil(t, result) + require.Equal(t, http.StatusBadRequest, infraerrors.Code(err)) + require.Zero(t, repo.bulkUpdateCalls) +} diff --git a/backend/internal/service/account_test_service.go b/backend/internal/service/account_test_service.go index 31549a6b17..222b3f8a4d 100644 --- a/backend/internal/service/account_test_service.go +++ b/backend/internal/service/account_test_service.go @@ -669,7 +669,7 @@ func (s *AccountTestService) testGrokAccountConnection(c *gin.Context, account * testModelID := strings.TrimSpace(modelID) if testModelID == "" { - testModelID = "grok-4.3" + testModelID = grokDefaultResponsesModel } if mapped := strings.TrimSpace(account.GetMappedModel(testModelID)); mapped != "" { testModelID = mapped diff --git a/backend/internal/service/account_test_service_grok_test.go b/backend/internal/service/account_test_service_grok_test.go index 497224b713..4b0890ff44 100644 --- a/backend/internal/service/account_test_service_grok_test.go +++ b/backend/internal/service/account_test_service_grok_test.go @@ -80,6 +80,47 @@ func TestAccountTestService_TestAccountConnection_GrokUsesXAIResponses(t *testin require.Contains(t, rec.Body.String(), `"type":"test_complete"`) } +func TestAccountTestService_TestAccountConnection_GrokDefaultsEmptyModelTo45(t *testing.T) { + gin.SetMode(gin.TestMode) + + account := &Account{ + ID: 16, + Name: "grok-oauth-default-model", + Platform: PlatformGrok, + Type: AccountTypeOAuth, + Status: StatusActive, + Schedulable: true, + Concurrency: 1, + Credentials: map[string]any{ + "access_token": "grok-access-token", + "expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339), + }, + } + repo := &mockAccountRepoForGemini{accountsByID: map[int64]*Account{account.ID: account}} + upstream := &httpUpstreamRecorder{resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"text/event-stream"}}, + Body: io.NopCloser(strings.NewReader( + "data: {\"type\":\"response.output_text.delta\",\"delta\":\"ok\"}\n\n" + + "data: {\"type\":\"response.completed\"}\n\n", + )), + }} + svc := &AccountTestService{ + accountRepo: repo, + grokTokenProvider: NewGrokTokenProvider(repo, nil), + httpUpstream: upstream, + } + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + c.Request = httptest.NewRequest(http.MethodPost, "/api/v1/admin/accounts/16/test", nil) + + err := svc.TestAccountConnection(c, account.ID, "", "", AccountTestModeDefault) + + require.NoError(t, err) + require.Equal(t, grokDefaultResponsesModel, gjson.GetBytes(upstream.lastBody, "model").String()) + require.Contains(t, recorder.Body.String(), `"model":"grok-4.5"`) +} + func TestAccountTestService_Grok429PersistsRateLimitReset(t *testing.T) { gin.SetMode(gin.TestMode) diff --git a/backend/internal/service/account_usage_service.go b/backend/internal/service/account_usage_service.go index 281122d4f0..9966c66515 100644 --- a/backend/internal/service/account_usage_service.go +++ b/backend/internal/service/account_usage_service.go @@ -111,6 +111,7 @@ const ( apiQueryMaxJitter = 800 * time.Millisecond // 用量查询最大随机延迟 windowStatsCacheTTL = 1 * time.Minute openAIProbeCacheTTL = 10 * time.Minute + grokProbeRetryTTL = 1 * time.Minute openAICodexProbeVersion = "0.144.1" ) @@ -122,6 +123,7 @@ type UsageCache struct { apiFlight singleflight.Group // 防止同一账号的并发请求击穿缓存(Anthropic) antigravityFlight singleflight.Group // 防止同一 Antigravity 账号的并发请求击穿缓存 openAIProbeCache sync.Map // accountID -> time.Time + grokProbeCache sync.Map // accountID -> last billing probe attempt } // NewUsageCache 创建 UsageCache 实例 @@ -196,15 +198,18 @@ type UsageInfo struct { AntigravityQuota map[string]*AntigravityModelQuota `json:"antigravity_quota,omitempty"` // Grok / xAI 被动额度快照 - GrokRequestQuota *xai.QuotaWindow `json:"grok_request_quota,omitempty"` - GrokTokenQuota *xai.QuotaWindow `json:"grok_token_quota,omitempty"` - GrokRetryAfterSeconds *int `json:"grok_retry_after_seconds,omitempty"` - GrokEntitlementStatus string `json:"grok_entitlement_status,omitempty"` - GrokQuotaSnapshotState string `json:"grok_quota_snapshot_state,omitempty"` - GrokLastQuotaProbeAt string `json:"grok_last_quota_probe_at,omitempty"` - GrokLastHeadersSeenAt string `json:"grok_last_headers_seen_at,omitempty"` - GrokLastStatusCode int `json:"grok_last_status_code,omitempty"` - GrokLocalUsage *WindowStats `json:"grok_local_usage,omitempty"` + GrokRequestQuota *xai.QuotaWindow `json:"grok_request_quota,omitempty"` + GrokTokenQuota *xai.QuotaWindow `json:"grok_token_quota,omitempty"` + GrokRetryAfterSeconds *int `json:"grok_retry_after_seconds,omitempty"` + GrokEntitlementStatus string `json:"grok_entitlement_status,omitempty"` + GrokQuotaSnapshotState string `json:"grok_quota_snapshot_state,omitempty"` + GrokLastQuotaProbeAt string `json:"grok_last_quota_probe_at,omitempty"` + GrokLastHeadersSeenAt string `json:"grok_last_headers_seen_at,omitempty"` + GrokLastStatusCode int `json:"grok_last_status_code,omitempty"` + GrokLocalUsage *WindowStats `json:"grok_local_usage,omitempty"` + GrokLocalUsage7d *WindowStats `json:"grok_local_usage_7d,omitempty"` + GrokLocalUsageMonthly *WindowStats `json:"grok_local_usage_monthly,omitempty"` + GrokBilling *xai.BillingSummary `json:"grok_billing,omitempty"` // Antigravity 账号级信息 SubscriptionTier string `json:"subscription_tier,omitempty"` // 归一化订阅等级: FREE/PRO/ULTRA/UNKNOWN @@ -287,6 +292,7 @@ type AccountUsageService struct { geminiQuotaService *GeminiQuotaService antigravityQuotaFetcher *AntigravityQuotaFetcher grokQuotaFetcher *GrokQuotaFetcher + grokQuotaService *GrokQuotaService openAIQuotaService *OpenAIQuotaService cache *UsageCache identityCache IdentityCache @@ -301,6 +307,7 @@ func NewAccountUsageService( geminiQuotaService *GeminiQuotaService, antigravityQuotaFetcher *AntigravityQuotaFetcher, grokQuotaFetcher *GrokQuotaFetcher, + grokQuotaService *GrokQuotaService, openAIQuotaService *OpenAIQuotaService, cache *UsageCache, identityCache IdentityCache, @@ -313,6 +320,7 @@ func NewAccountUsageService( geminiQuotaService: geminiQuotaService, antigravityQuotaFetcher: antigravityQuotaFetcher, grokQuotaFetcher: grokQuotaFetcher, + grokQuotaService: grokQuotaService, openAIQuotaService: openAIQuotaService, cache: cache, identityCache: identityCache, @@ -358,8 +366,8 @@ func (s *AccountUsageService) GetUsage(ctx context.Context, accountID int64, for } if account.Platform == PlatformGrok { - usage, err := s.getGrokUsage(ctx, account) - if err == nil { + usage, err := s.getGrokUsage(ctx, account, forceProbe) + if err == nil && usage != nil && usage.Error == "" { s.tryClearRecoverableAccountError(ctx, account) } return usage, err @@ -930,11 +938,19 @@ func (s *AccountUsageService) getAntigravityUsage(ctx context.Context, account * return usage, nil } -func (s *AccountUsageService) getGrokUsage(ctx context.Context, account *Account) (*UsageInfo, error) { +func (s *AccountUsageService) getGrokUsage(ctx context.Context, account *Account, force bool) (*UsageInfo, error) { if s.grokQuotaFetcher == nil { now := time.Now() return &UsageInfo{UpdatedAt: &now}, nil } + if account != nil && account.IsGrokOAuth() && s.grokQuotaService != nil && (force || grokBillingSnapshotNeedsRefresh(account, time.Now())) && s.shouldProbeGrokBilling(account.ID, time.Now(), force) { + result, err := s.grokQuotaService.ProbeBilling(ctx, account.ID) + if err == nil && result != nil && result.Billing != nil { + mergeAccountExtra(account, map[string]any{grokBillingExtraKey: result.Billing}) + } else if err != nil && force { + return nil, err + } + } usage := s.grokQuotaFetcher.BuildUsageInfo(account) if usage.GrokQuotaSnapshotState == "" { if usage.ErrorCode == "quota_unknown" { @@ -948,12 +964,90 @@ func (s *AccountUsageService) getGrokUsage(ctx context.Context, account *Account if stats, err := s.usageLogRepo.GetAccountTodayStats(ctx, account.ID); err == nil && stats != nil { usage.GrokLocalUsage = windowStatsFromAccountStats(stats) } + usage.GrokLocalUsage7d, usage.GrokLocalUsageMonthly = grokLocalUsageForBilling(ctx, s.usageLogRepo, account.ID, usage.GrokBilling, time.Now().UTC()) } enrichUsageWithAccountError(usage, account) return usage, nil } +func grokLocalUsageForBilling( + ctx context.Context, + repo UsageLogRepository, + accountID int64, + billing *xai.BillingSummary, + now time.Time, +) (*WindowStats, *WindowStats) { + var weekly *WindowStats + var monthly *WindowStats + if repo == nil || accountID <= 0 { + return weekly, monthly + } + if start, ok := currentGrokBillingWindow(billing, true, now); ok { + if stats, err := repo.GetAccountWindowStats(ctx, accountID, start); err == nil { + weekly = windowStatsFromAccountStats(stats) + } else { + slog.Warn("grok_window_usage_query_failed", "account_id", accountID, "window_start", start, "error", err) + } + } + if start, ok := currentGrokBillingWindow(billing, false, now); ok { + if stats, err := repo.GetAccountWindowStats(ctx, accountID, start); err == nil { + monthly = windowStatsFromAccountStats(stats) + } else { + slog.Warn("grok_monthly_usage_query_failed", "account_id", accountID, "window_start", start, "error", err) + } + } + return weekly, monthly +} + +func currentGrokBillingWindow(billing *xai.BillingSummary, weekly bool, now time.Time) (time.Time, bool) { + if billing == nil { + return time.Time{}, false + } + startRaw, endRaw := billing.BillingPeriodStart, billing.BillingPeriodEnd + if weekly { + if billing.PeriodType != "weekly" { + return time.Time{}, false + } + startRaw, endRaw = billing.PeriodStart, billing.PeriodEnd + } + start, startErr := parseTime(strings.TrimSpace(startRaw)) + end, endErr := parseTime(strings.TrimSpace(endRaw)) + if startErr != nil || endErr != nil || now.Before(start) || !now.Before(end) { + return time.Time{}, false + } + return start, true +} + +func grokBillingSnapshotNeedsRefresh(account *Account, now time.Time) bool { + if account == nil { + return false + } + billing, err := grokBillingSnapshotFromExtra(account.Extra) + if err != nil || billing == nil || billing.Partial || len(billing.FailedWindows) > 0 { + return true + } + stamp := strings.TrimSpace(billing.UpdatedAt) + if stamp == "" { + stamp = strings.TrimSpace(billing.FetchedAt) + } + updatedAt, err := parseTime(stamp) + return err != nil || now.Sub(updatedAt) >= openAIProbeCacheTTL +} + +func (s *AccountUsageService) shouldProbeGrokBilling(accountID int64, now time.Time, force bool) bool { + if force || s == nil || s.cache == nil || accountID <= 0 { + return true + } + if cached, ok := s.cache.grokProbeCache.Load(accountID); ok { + if ts, ok := cached.(time.Time); ok && now.Sub(ts) < grokProbeRetryTTL { + return false + } + } + s.cache.grokProbeCache.Store(accountID, now) + return true +} + // recalcAntigravityRemainingSeconds 重新计算 Antigravity UsageInfo 中各窗口的 RemainingSeconds // 用于从缓存取出时更新倒计时,避免返回过时的剩余秒数 func recalcAntigravityRemainingSeconds(info *UsageInfo) { diff --git a/backend/internal/service/admin_account.go b/backend/internal/service/admin_account.go index 52e5ce719b..8cb6d8e63b 100644 --- a/backend/internal/service/admin_account.go +++ b/backend/internal/service/admin_account.go @@ -5,6 +5,7 @@ import ( "errors" "fmt" "log/slog" + "maps" "net/http" "strconv" "strings" @@ -68,7 +69,65 @@ func normalizeAccountConcurrency(platform, accountType string, concurrency int) return concurrency } +// ValidateOpenAILongContextBillingExtra validates the OpenAI account billing flag when present. +func ValidateOpenAILongContextBillingExtra(platform string, extra map[string]any) error { + if platform != PlatformOpenAI { + return nil + } + raw, exists := extra[openAILongContextBillingEnabledKey] + if !exists { + return nil + } + if _, ok := raw.(bool); !ok { + return infraerrors.BadRequest( + "OPENAI_LONG_CONTEXT_BILLING_INVALID", + "openai_long_context_billing_enabled must be a boolean", + ) + } + return nil +} + +func normalizeOpenAILongContextBillingExtra(platform string, extra map[string]any) (map[string]any, error) { + if platform != PlatformOpenAI { + return extra, nil + } + if err := ValidateOpenAILongContextBillingExtra(platform, extra); err != nil { + return nil, err + } + + normalized := maps.Clone(extra) + if normalized == nil { + normalized = make(map[string]any, 1) + } + _, exists := normalized[openAILongContextBillingEnabledKey] + if !exists { + normalized[openAILongContextBillingEnabledKey] = false + } + return normalized, nil +} + +func normalizeOpenAILongContextBillingUpdateExtra(account *Account, input *UpdateAccountInput) (map[string]any, error) { + normalized, err := normalizeOpenAILongContextBillingExtra(account.Platform, input.Extra) + if err != nil || account.Platform != PlatformOpenAI { + return normalized, err + } + + _, provided := input.Extra[openAILongContextBillingEnabledKey] + current, hasCurrent := account.Extra[openAILongContextBillingEnabledKey].(bool) + if !provided { + if hasCurrent { + normalized[openAILongContextBillingEnabledKey] = current + } + } + return normalized, nil +} + func (s *adminServiceImpl) CreateAccount(ctx context.Context, input *CreateAccountInput) (*Account, error) { + accountExtra, err := normalizeOpenAILongContextBillingExtra(input.Platform, input.Extra) + if err != nil { + return nil, err + } + // 绑定分组 groupIDs := input.GroupIDs // 如果没有指定分组,自动绑定对应平台的默认分组 @@ -103,7 +162,7 @@ func (s *adminServiceImpl) CreateAccount(ctx context.Context, input *CreateAccou Platform: input.Platform, Type: input.Type, Credentials: input.Credentials, - Extra: input.Extra, + Extra: accountExtra, ProxyID: input.ProxyID, Concurrency: normalizeAccountConcurrency(input.Platform, input.Type, input.Concurrency), Priority: input.Priority, @@ -183,6 +242,13 @@ func (s *adminServiceImpl) UpdateAccount(ctx context.Context, id int64, input *U if err != nil { return nil, err } + var normalizedExtra map[string]any + if input.Extra != nil { + normalizedExtra, err = normalizeOpenAILongContextBillingUpdateExtra(account, input) + if err != nil { + return nil, err + } + } // 安全/身份不变量(影子账号):通用更新路径被 edit/re-auth/refresh/batch 共用, // 必须在此守住,否则仅在创建时的保证可被这些路径绕过。 if account.IsCredentialShadow() { @@ -238,10 +304,10 @@ func (s *adminServiceImpl) UpdateAccount(ctx context.Context, id int64, input *U // 保留配额用量字段,防止编辑账号时意外重置 for _, key := range []string{"quota_used", "quota_daily_used", "quota_daily_start", "quota_weekly_used", "quota_weekly_start"} { if v, ok := account.Extra[key]; ok { - input.Extra[key] = v + normalizedExtra[key] = v } } - account.Extra = input.Extra + account.Extra = normalizedExtra if account.Platform == PlatformAntigravity && wasOveragesEnabled && !account.IsOveragesEnabled() { delete(account.Extra, "antigravity_credits_overages") // 清理旧版 overages 运行态 // 清除 AICredits 限流 key @@ -353,6 +419,15 @@ func (s *adminServiceImpl) UpdateAccount(ctx context.Context, id int64, input *U // UpdateAccountExtra 仅对 Extra JSONB 做 key 级合并,避免覆盖其它运行态键 // (如 model_rate_limits / passive_usage_* 等)。 func (s *adminServiceImpl) UpdateAccountExtra(ctx context.Context, id int64, updates map[string]any) error { + if _, exists := updates[openAILongContextBillingEnabledKey]; exists { + account, err := s.accountRepo.GetByID(ctx, id) + if err != nil { + return err + } + if err := ValidateOpenAILongContextBillingExtra(account.Platform, updates); err != nil { + return err + } + } if len(updates) == 0 { return nil } @@ -386,16 +461,28 @@ func (s *adminServiceImpl) BulkUpdateAccounts(ctx context.Context, input *BulkUp } needMixedChannelCheck := input.GroupIDs != nil && !input.SkipMixedChannelCheck + _, hasLongContextBillingUpdate := input.Extra[openAILongContextBillingEnabledKey] // 预取所有目标账号,供凭据守卫/代理守卫/混合渠道检查共用,避免多次 DB 查询。 var cachedTargets []*Account - if len(input.Credentials) > 0 || input.ProxyID != nil || needMixedChannelCheck { + if len(input.Credentials) > 0 || input.ProxyID != nil || needMixedChannelCheck || hasLongContextBillingUpdate { loaded, err := s.accountRepo.GetByIDs(ctx, input.AccountIDs) if err != nil { return nil, err } cachedTargets = loaded } + if hasLongContextBillingUpdate { + for _, account := range cachedTargets { + if account == nil || account.Platform != PlatformOpenAI { + continue + } + if err := ValidateOpenAILongContextBillingExtra(account.Platform, input.Extra); err != nil { + return nil, err + } + break + } + } // 影子账号绝不持有凭据:批量更新携带凭据时,目标中不得含影子(外审 G5,与单账号 // UpdateAccount 守卫对齐)。覆盖显式 IDs 与 filter 解析出的 IDs(此处 AccountIDs 已解析完成)。 @@ -745,6 +832,9 @@ func (s *adminServiceImpl) CreateShadow(ctx context.Context, parentID int64, opt Priority: priority, Concurrency: concurrency, Schedulable: true, + Extra: map[string]any{ + openAILongContextBillingEnabledKey: parent.IsOpenAILongContextBillingEnabled(), + }, } // 5. 持久化(Create 填充 shadow.ID)。并发竞态:预查(步骤2)放行后另一请求抢先建成,本次会撞 diff --git a/backend/internal/service/admin_service_spark_shadow_test.go b/backend/internal/service/admin_service_spark_shadow_test.go index 6b4017207a..0eda0d93c7 100644 --- a/backend/internal/service/admin_service_spark_shadow_test.go +++ b/backend/internal/service/admin_service_spark_shadow_test.go @@ -157,6 +157,39 @@ func TestCreateShadow(t *testing.T) { require.Error(t, err) } +func TestCreateShadowInheritsParentEffectiveOpenAILongContextBillingValue(t *testing.T) { + tests := []struct { + name string + parentExtra map[string]any + want bool + }{ + {name: "missing parent value defaults disabled", want: false}, + {name: "explicit parent opt-out is inherited", parentExtra: map[string]any{openAILongContextBillingEnabledKey: false}, want: false}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + repo := newSparkShadowRepoStub() + svc := &adminServiceImpl{accountRepo: repo} + parent := &Account{ + Name: "parent", + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Status: StatusActive, + Credentials: map[string]any{"access_token": "token"}, + Extra: tt.parentExtra, + } + require.NoError(t, repo.Create(context.Background(), parent)) + + shadow, err := svc.CreateShadow(context.Background(), parent.ID, ShadowOptions{Name: "shadow"}) + + require.NoError(t, err) + require.Equal(t, tt.want, shadow.Extra[openAILongContextBillingEnabledKey]) + require.Equal(t, tt.want, shadow.IsOpenAILongContextBillingEnabled()) + }) + } +} + // TestCreateShadow_BindGroups は BindGroups の後置呼び出しを検証する。 // 影子账号が指定グループに属し、ListSchedulableByGroupID で取得可能であること。 func TestCreateShadow_BindGroups(t *testing.T) { diff --git a/backend/internal/service/billing_service.go b/backend/internal/service/billing_service.go index 43d7f22d78..7fa69d41ea 100644 --- a/backend/internal/service/billing_service.go +++ b/backend/internal/service/billing_service.go @@ -153,14 +153,15 @@ type UsageTokens struct { // CostBreakdown 费用明细 type CostBreakdown struct { - InputCost float64 - OutputCost float64 - ImageOutputCost float64 - CacheCreationCost float64 - CacheReadCost float64 - TotalCost float64 - ActualCost float64 // 应用倍率后的实际费用 - BillingMode string // 计费模式("token"/"per_request"/"image"),由 CalculateCostUnified 填充 + InputCost float64 + OutputCost float64 + ImageOutputCost float64 + CacheCreationCost float64 + CacheReadCost float64 + TotalCost float64 + ActualCost float64 // 应用倍率后的实际费用 + BillingMode string // 计费模式("token"/"per_request"/"image"),由 CalculateCostUnified 填充 + LongContextBillingApplied bool } // ErrModelPricingUnavailable indicates that none of the configured pricing @@ -865,16 +866,17 @@ func (s *BillingService) GetModelPricingWithChannel(model string, channelPricing // CostInput 统一计费输入 type CostInput struct { - Ctx context.Context - Model string - GroupID *int64 // 用于渠道定价查找 - Tokens UsageTokens - RequestCount int // 按次计费时使用 - SizeTier string // 按次/图片模式的层级标签("1K","2K","4K","HD" 等) - RateMultiplier float64 - ServiceTier string // "priority","flex","" 等 - Resolver *ModelPricingResolver // 定价解析器 - Resolved *ResolvedPricing // 可选:预解析的定价结果(避免重复 Resolve 调用) + Ctx context.Context + Model string + GroupID *int64 // 用于渠道定价查找 + Tokens UsageTokens + RequestCount int // 按次计费时使用 + SizeTier string // 按次/图片模式的层级标签("1K","2K","4K","HD" 等) + RateMultiplier float64 + ServiceTier string // "priority","flex","" 等 + Resolver *ModelPricingResolver // 定价解析器 + Resolved *ResolvedPricing // 可选:预解析的定价结果(避免重复 Resolve 调用) + LongContextBillingEnabled *bool } // CalculateCostUnified 统一计费入口,支持三种计费模式。 @@ -882,7 +884,18 @@ type CostInput struct { func (s *BillingService) CalculateCostUnified(input CostInput) (*CostBreakdown, error) { if input.Resolver == nil { // 无 Resolver,回退到旧路径 - return s.calculateCostInternal(input.Model, input.Tokens, input.RateMultiplier, input.ServiceTier, nil) + applyLongContextBilling := true + if input.LongContextBillingEnabled != nil { + applyLongContextBilling = *input.LongContextBillingEnabled + } + return s.calculateCostInternalWithPolicy( + input.Model, + input.Tokens, + input.RateMultiplier, + input.ServiceTier, + nil, + applyLongContextBilling, + ) } // 优先使用预解析结果,避免重复 Resolve 调用 @@ -929,6 +942,9 @@ func (s *BillingService) calculateTokenCost(resolved *ResolvedPricing, input Cos // 长上下文定价仅在无区间定价时应用(区间定价已包含上下文分层) applyLongCtx := len(resolved.Intervals) == 0 + if input.LongContextBillingEnabled != nil { + applyLongCtx = applyLongCtx && *input.LongContextBillingEnabled + } return s.computeTokenBreakdown(pricing, input.Tokens, input.RateMultiplier, input.ServiceTier, applyLongCtx), nil } @@ -969,7 +985,10 @@ func (s *BillingService) computeTokenBreakdown( tierMultiplier = serviceTierCostMultiplier(serviceTier) } - if applyLongCtx && s.shouldApplySessionLongContextPricing(tokens, pricing) { + longContextPricingEligible := applyLongCtx && s.shouldApplySessionLongContextPricing(tokens, pricing) + var baselineCost *CostBreakdown + if longContextPricingEligible { + baselineCost = s.computeTokenBreakdown(pricing, tokens, rateMultiplier, serviceTier, false) inputPrice *= pricing.LongContextInputMultiplier outputPrice *= pricing.LongContextOutputMultiplier // 缓存读取本质上是输入侧的复用,应与 input 一同应用长上下文倍率; @@ -1033,6 +1052,7 @@ func (s *BillingService) computeTokenBreakdown( bd.TotalCost = bd.InputCost + bd.OutputCost + bd.ImageOutputCost + bd.CacheCreationCost + bd.CacheReadCost bd.ActualCost = bd.TotalCost * rateMultiplier + bd.LongContextBillingApplied = baselineCost != nil && bd.ActualCost > baselineCost.ActualCost return bd } @@ -1092,7 +1112,28 @@ func (s *BillingService) CalculateCostWithServiceTier(model string, tokens Usage return s.calculateCostInternal(model, tokens, rateMultiplier, serviceTier, nil) } +func (s *BillingService) calculateCostWithServiceTierPolicy( + model string, + tokens UsageTokens, + rateMultiplier float64, + serviceTier string, + longContextBillingEnabled bool, +) (*CostBreakdown, error) { + return s.calculateCostInternalWithPolicy(model, tokens, rateMultiplier, serviceTier, nil, longContextBillingEnabled) +} + func (s *BillingService) calculateCostInternal(model string, tokens UsageTokens, rateMultiplier float64, serviceTier string, channelPricing *ChannelModelPricing) (*CostBreakdown, error) { + return s.calculateCostInternalWithPolicy(model, tokens, rateMultiplier, serviceTier, channelPricing, true) +} + +func (s *BillingService) calculateCostInternalWithPolicy( + model string, + tokens UsageTokens, + rateMultiplier float64, + serviceTier string, + channelPricing *ChannelModelPricing, + longContextBillingEnabled bool, +) (*CostBreakdown, error) { var pricing *ModelPricing var err error if channelPricing != nil { @@ -1104,8 +1145,7 @@ func (s *BillingService) calculateCostInternal(model string, tokens UsageTokens, return nil, err } - // 旧路径始终检查长上下文定价(无区间定价概念) - return s.computeTokenBreakdown(pricing, tokens, rateMultiplier, serviceTier, true), nil + return s.computeTokenBreakdown(pricing, tokens, rateMultiplier, serviceTier, longContextBillingEnabled), nil } func (s *BillingService) applyModelSpecificPricingPolicy(model string, pricing *ModelPricing) *ModelPricing { @@ -1236,13 +1276,14 @@ func (s *BillingService) CalculateCostWithLongContext(model string, tokens Usage // 合并成本 return &CostBreakdown{ - InputCost: inRangeCost.InputCost + outRangeCost.InputCost, - OutputCost: inRangeCost.OutputCost, - ImageOutputCost: inRangeCost.ImageOutputCost, - CacheCreationCost: inRangeCost.CacheCreationCost, - CacheReadCost: inRangeCost.CacheReadCost + outRangeCost.CacheReadCost, - TotalCost: inRangeCost.TotalCost + outRangeCost.TotalCost, - ActualCost: inRangeCost.ActualCost + outRangeCost.ActualCost, + InputCost: inRangeCost.InputCost + outRangeCost.InputCost, + OutputCost: inRangeCost.OutputCost, + ImageOutputCost: inRangeCost.ImageOutputCost, + CacheCreationCost: inRangeCost.CacheCreationCost, + CacheReadCost: inRangeCost.CacheReadCost + outRangeCost.CacheReadCost, + TotalCost: inRangeCost.TotalCost + outRangeCost.TotalCost, + ActualCost: inRangeCost.ActualCost + outRangeCost.ActualCost, + LongContextBillingApplied: outRangeCost.ActualCost > 0, }, nil } diff --git a/backend/internal/service/billing_service_test.go b/backend/internal/service/billing_service_test.go index 53014412fd..885da194e3 100644 --- a/backend/internal/service/billing_service_test.go +++ b/backend/internal/service/billing_service_test.go @@ -261,6 +261,23 @@ func TestCalculateCost_OpenAIGPT54LongContextAppliesWholeSessionMultipliers(t *t require.InDelta(t, expectedOutput, cost.OutputCost, 1e-10) require.InDelta(t, expectedInput+expectedOutput, cost.TotalCost, 1e-10) require.InDelta(t, expectedInput+expectedOutput, cost.ActualCost, 1e-10) + require.True(t, cost.LongContextBillingApplied) +} + +func TestCalculateCost_OpenAIGPT54LongContextMarkerRequiresActualCostIncrease(t *testing.T) { + svc := newTestBillingService() + + cost, err := svc.calculateCostWithServiceTierPolicy( + "gpt-5.4-2026-03-05", + UsageTokens{InputTokens: 300000}, + 0, + "", + true, + ) + + require.NoError(t, err) + require.Zero(t, cost.ActualCost) + require.False(t, cost.LongContextBillingApplied) } func TestCalculateCost_OpenAIGPT55ProUsesGPT55PricingPolicy(t *testing.T) { @@ -831,6 +848,17 @@ func TestCalculateCostWithLongContext_AboveThreshold_CacheBelowThreshold(t *test require.True(t, cost.ActualCost > normalCost.ActualCost, "长上下文费用应高于正常费用") } +func TestCalculateCostWithLongContext_MarkerRequiresActualCostIncrease(t *testing.T) { + svc := newTestBillingService() + tokens := UsageTokens{InputTokens: 300000} + + cost, err := svc.CalculateCostWithLongContext("claude-sonnet-4", tokens, 0, 200000, 2.0) + + require.NoError(t, err) + require.Zero(t, cost.ActualCost) + require.False(t, cost.LongContextBillingApplied) +} + func TestCalculateCostWithLongContext_DisabledThreshold(t *testing.T) { svc := newTestBillingService() diff --git a/backend/internal/service/channel_monitor_checker.go b/backend/internal/service/channel_monitor_checker.go index 7fb829a3cb..889b2bbed7 100644 --- a/backend/internal/service/channel_monitor_checker.go +++ b/backend/internal/service/channel_monitor_checker.go @@ -13,6 +13,7 @@ import ( "strings" "time" + "github.com/Wei-Shaw/sub2api/internal/pkg/servertiming" "github.com/tidwall/gjson" ) @@ -34,7 +35,7 @@ func newSSRFSafeHTTPClient(timeout time.Duration) *http.Client { TLSHandshakeTimeout: monitorTLSHandshakeTimeout, ResponseHeaderTimeout: monitorResponseHeaderTimeout, } - return &http.Client{Timeout: timeout, Transport: tr} + return &http.Client{Timeout: timeout, Transport: servertiming.WrapRoundTripper(tr)} } // CheckOptions 承载一次检测的自定义入参。 diff --git a/backend/internal/service/content_moderation.go b/backend/internal/service/content_moderation.go index 6d3b91d205..f633c8ad17 100644 --- a/backend/internal/service/content_moderation.go +++ b/backend/internal/service/content_moderation.go @@ -22,6 +22,7 @@ import ( infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" "github.com/Wei-Shaw/sub2api/internal/pkg/pagination" + "github.com/Wei-Shaw/sub2api/internal/pkg/servertiming" ) const ( @@ -561,7 +562,7 @@ func NewContentModerationService( userRepo: userRepo, authCacheInvalidator: authCacheInvalidator, emailService: emailService, - httpClient: &http.Client{}, + httpClient: servertiming.InstrumentClient(nil), workerCount: maxContentModerationWorkerCount, asyncQueue: make(chan contentModerationTask, maxContentModerationQueueSize), keyHealth: make(map[string]*contentModerationKeyHealth), diff --git a/backend/internal/service/crs_sync_long_context_billing_test.go b/backend/internal/service/crs_sync_long_context_billing_test.go new file mode 100644 index 0000000000..6439f08190 --- /dev/null +++ b/backend/internal/service/crs_sync_long_context_billing_test.go @@ -0,0 +1,169 @@ +//go:build unit + +package service + +import ( + "context" + "encoding/json" + "net/http" + "net/http/httptest" + "testing" + + "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/stretchr/testify/require" +) + +type crsLongContextAccountRepo struct { + AccountRepository + accounts map[string]*Account + nextID int64 +} + +type crsOpenAILongContextSource struct { + collection string + credentials map[string]any + extra map[string]any +} + +func newCRSLongContextAccountRepo(existing ...*Account) *crsLongContextAccountRepo { + repo := &crsLongContextAccountRepo{accounts: make(map[string]*Account)} + for _, account := range existing { + if account == nil { + continue + } + crsID, _ := account.Extra["crs_account_id"].(string) + repo.accounts[crsID] = account + if account.ID > repo.nextID { + repo.nextID = account.ID + } + } + return repo +} + +func (r *crsLongContextAccountRepo) Create(_ context.Context, account *Account) error { + r.nextID++ + account.ID = r.nextID + crsID, _ := account.Extra["crs_account_id"].(string) + r.accounts[crsID] = account + return nil +} + +func (r *crsLongContextAccountRepo) Update(_ context.Context, account *Account) error { + crsID, _ := account.Extra["crs_account_id"].(string) + r.accounts[crsID] = account + return nil +} + +func (r *crsLongContextAccountRepo) GetByCRSAccountID(_ context.Context, crsID string) (*Account, error) { + return r.accounts[crsID], nil +} + +func (r *crsLongContextAccountRepo) ListShadowsByParent(_ context.Context, _ int64) ([]*Account, error) { + return nil, nil +} + +func TestCRSSyncOpenAILongContextBilling(t *testing.T) { + tests := []struct { + name string + collection string + credentials map[string]any + sourceExtra map[string]any + existingExtra map[string]any + wantAction string + wantEnabled bool + }{ + {name: "OAuth create defaults missing value disabled", collection: "openaiOAuthAccounts", credentials: map[string]any{"access_token": "oauth-token"}, wantAction: "created"}, + {name: "OAuth create preserves source true", collection: "openaiOAuthAccounts", credentials: map[string]any{"access_token": "oauth-token"}, sourceExtra: map[string]any{openAILongContextBillingEnabledKey: true}, wantAction: "created", wantEnabled: true}, + {name: "OAuth create preserves source false", collection: "openaiOAuthAccounts", credentials: map[string]any{"access_token": "oauth-token"}, sourceExtra: map[string]any{openAILongContextBillingEnabledKey: false}, wantAction: "created"}, + {name: "OAuth update defaults missing value disabled", collection: "openaiOAuthAccounts", credentials: map[string]any{"access_token": "oauth-token"}, existingExtra: map[string]any{"existing": true}, wantAction: "updated"}, + {name: "OAuth update preserves existing true when source omits value", collection: "openaiOAuthAccounts", credentials: map[string]any{"access_token": "oauth-token"}, existingExtra: map[string]any{openAILongContextBillingEnabledKey: true}, wantAction: "updated", wantEnabled: true}, + {name: "OAuth update preserves existing false when source omits value", collection: "openaiOAuthAccounts", credentials: map[string]any{"access_token": "oauth-token"}, existingExtra: map[string]any{openAILongContextBillingEnabledKey: false}, wantAction: "updated"}, + {name: "OAuth update preserves source true over existing false", collection: "openaiOAuthAccounts", credentials: map[string]any{"access_token": "oauth-token"}, sourceExtra: map[string]any{openAILongContextBillingEnabledKey: true}, existingExtra: map[string]any{openAILongContextBillingEnabledKey: false}, wantAction: "updated", wantEnabled: true}, + {name: "OAuth update preserves source false over existing true", collection: "openaiOAuthAccounts", credentials: map[string]any{"access_token": "oauth-token"}, sourceExtra: map[string]any{openAILongContextBillingEnabledKey: false}, existingExtra: map[string]any{openAILongContextBillingEnabledKey: true}, wantAction: "updated"}, + {name: "OAuth rejects malformed source value", collection: "openaiOAuthAccounts", credentials: map[string]any{"access_token": "oauth-token"}, sourceExtra: map[string]any{openAILongContextBillingEnabledKey: "false"}, wantAction: "failed"}, + {name: "OAuth rejects malformed existing value", collection: "openaiOAuthAccounts", credentials: map[string]any{"access_token": "oauth-token"}, existingExtra: map[string]any{openAILongContextBillingEnabledKey: "false"}, wantAction: "failed"}, + {name: "OAuth update rejects malformed source value", collection: "openaiOAuthAccounts", credentials: map[string]any{"access_token": "oauth-token"}, sourceExtra: map[string]any{openAILongContextBillingEnabledKey: "false"}, existingExtra: map[string]any{openAILongContextBillingEnabledKey: true}, wantAction: "failed"}, + {name: "API key create defaults missing value disabled", collection: "openaiResponsesAccounts", credentials: map[string]any{"api_key": "sk-test"}, wantAction: "created"}, + {name: "API key create preserves source true", collection: "openaiResponsesAccounts", credentials: map[string]any{"api_key": "sk-test"}, sourceExtra: map[string]any{openAILongContextBillingEnabledKey: true}, wantAction: "created", wantEnabled: true}, + {name: "API key create preserves source false", collection: "openaiResponsesAccounts", credentials: map[string]any{"api_key": "sk-test"}, sourceExtra: map[string]any{openAILongContextBillingEnabledKey: false}, wantAction: "created"}, + {name: "API key update defaults missing value disabled", collection: "openaiResponsesAccounts", credentials: map[string]any{"api_key": "sk-test"}, existingExtra: map[string]any{"existing": true}, wantAction: "updated"}, + {name: "API key update preserves existing true when source omits value", collection: "openaiResponsesAccounts", credentials: map[string]any{"api_key": "sk-test"}, existingExtra: map[string]any{openAILongContextBillingEnabledKey: true}, wantAction: "updated", wantEnabled: true}, + {name: "API key update preserves existing false when source omits value", collection: "openaiResponsesAccounts", credentials: map[string]any{"api_key": "sk-test"}, existingExtra: map[string]any{openAILongContextBillingEnabledKey: false}, wantAction: "updated"}, + {name: "API key update preserves source true over existing false", collection: "openaiResponsesAccounts", credentials: map[string]any{"api_key": "sk-test"}, sourceExtra: map[string]any{openAILongContextBillingEnabledKey: true}, existingExtra: map[string]any{openAILongContextBillingEnabledKey: false}, wantAction: "updated", wantEnabled: true}, + {name: "API key update preserves source false over existing true", collection: "openaiResponsesAccounts", credentials: map[string]any{"api_key": "sk-test"}, sourceExtra: map[string]any{openAILongContextBillingEnabledKey: false}, existingExtra: map[string]any{openAILongContextBillingEnabledKey: true}, wantAction: "updated"}, + {name: "API key rejects malformed source value", collection: "openaiResponsesAccounts", credentials: map[string]any{"api_key": "sk-test"}, sourceExtra: map[string]any{openAILongContextBillingEnabledKey: "false"}, wantAction: "failed"}, + {name: "API key rejects malformed existing value", collection: "openaiResponsesAccounts", credentials: map[string]any{"api_key": "sk-test"}, existingExtra: map[string]any{openAILongContextBillingEnabledKey: "false"}, wantAction: "failed"}, + {name: "API key update rejects malformed source value", collection: "openaiResponsesAccounts", credentials: map[string]any{"api_key": "sk-test"}, sourceExtra: map[string]any{openAILongContextBillingEnabledKey: "false"}, existingExtra: map[string]any{openAILongContextBillingEnabledKey: true}, wantAction: "failed"}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + const crsID = "crs-openai-1" + var existing *Account + if tt.existingExtra != nil { + existingExtra := mergeMap(tt.existingExtra, map[string]any{"crs_account_id": crsID}) + accountType := AccountTypeOAuth + if tt.collection == "openaiResponsesAccounts" { + accountType = AccountTypeAPIKey + } + existing = &Account{ID: 41, Platform: PlatformOpenAI, Type: accountType, Extra: existingExtra} + } + repo := newCRSLongContextAccountRepo(existing) + result := runCRSOpenAILongContextSync(t, repo, crsOpenAILongContextSource{ + collection: tt.collection, + credentials: tt.credentials, + extra: tt.sourceExtra, + }) + + require.Len(t, result.Items, 1) + require.Equal(t, tt.wantAction, result.Items[0].Action) + if tt.wantAction == "failed" { + require.Contains(t, result.Items[0].Error, "openai_long_context_billing_enabled must be a boolean") + return + } + stored, ok := repo.accounts[crsID].Extra[openAILongContextBillingEnabledKey] + require.True(t, ok) + require.Equal(t, tt.wantEnabled, stored) + }) + } +} + +func runCRSOpenAILongContextSync(t *testing.T, repo AccountRepository, source crsOpenAILongContextSource) *SyncFromCRSResult { + t.Helper() + account := map[string]any{ + "kind": "openai", + "id": "crs-openai-1", + "name": "OpenAI CRS", + "isActive": true, + "schedulable": true, + "credentials": source.credentials, + } + if source.extra != nil { + account["extra"] = source.extra + } + + server := httptest.NewServer(http.HandlerFunc(func(response http.ResponseWriter, request *http.Request) { + response.Header().Set("Content-Type", "application/json") + if request.URL.Path == "/web/auth/login" { + _, _ = response.Write([]byte(`{"success":true,"token":"admin-token"}`)) + return + } + require.Equal(t, "/admin/sync/export-accounts", request.URL.Path) + require.NoError(t, json.NewEncoder(response).Encode(map[string]any{ + "success": true, + "data": map[string]any{source.collection: []any{account}}, + })) + })) + t.Cleanup(server.Close) + + cfg := &config.Config{} + cfg.Security.URLAllowlist.AllowInsecureHTTP = true + service := NewCRSSyncService(repo, nil, nil, nil, nil, cfg) + result, err := service.SyncFromCRS(context.Background(), SyncFromCRSInput{ + BaseURL: server.URL, + Username: "admin", + Password: "password", + }) + require.NoError(t, err) + return result +} diff --git a/backend/internal/service/crs_sync_service.go b/backend/internal/service/crs_sync_service.go index edf3cd43d2..d0abc74038 100644 --- a/backend/internal/service/crs_sync_service.go +++ b/backend/internal/service/crs_sync_service.go @@ -168,6 +168,7 @@ type crsOpenAIResponsesAccount struct { Status string `json:"status"` Proxy *crsProxy `json:"proxy"` Credentials map[string]any `json:"credentials"` + Extra map[string]any `json:"extra"` } type crsOpenAIOAuthAccount struct { @@ -632,6 +633,18 @@ func (s *CRSSyncService) SyncFromCRS(ctx context.Context, input SyncFromCRSInput result.Items = append(result.Items, item) continue } + var existingExtra map[string]any + if existing != nil { + existingExtra = existing.Extra + } + extra, err = mergeCRSOpenAILongContextBillingExtra(existingExtra, extra) + if err != nil { + item.Action = "failed" + item.Error = err.Error() + result.Failed++ + result.Items = append(result.Items, item) + continue + } if existing == nil { if !shouldCreateAccount(src.ID, selectedSet) { @@ -670,7 +683,7 @@ func (s *CRSSyncService) SyncFromCRS(ctx context.Context, input SyncFromCRSInput continue } - existing.Extra = mergeMap(existing.Extra, extra) + existing.Extra = extra existing.Name = defaultName(src.Name, src.ID) existing.Platform = PlatformOpenAI existing.Type = AccountTypeOAuth @@ -751,11 +764,13 @@ func (s *CRSSyncService) SyncFromCRS(ctx context.Context, input SyncFromCRSInput concurrency := 3 status := mapCRSStatus(src.IsActive, src.Status) - extra := map[string]any{ - "crs_account_id": src.ID, - "crs_kind": src.Kind, - "crs_synced_at": now, + extra := make(map[string]any, len(src.Extra)+3) + for key, value := range src.Extra { + extra[key] = value } + extra["crs_account_id"] = src.ID + extra["crs_kind"] = src.Kind + extra["crs_synced_at"] = now existing, err := s.accountRepo.GetByCRSAccountID(ctx, src.ID) if err != nil { @@ -765,6 +780,18 @@ func (s *CRSSyncService) SyncFromCRS(ctx context.Context, input SyncFromCRSInput result.Items = append(result.Items, item) continue } + var existingExtra map[string]any + if existing != nil { + existingExtra = existing.Extra + } + extra, err = mergeCRSOpenAILongContextBillingExtra(existingExtra, extra) + if err != nil { + item.Action = "failed" + item.Error = err.Error() + result.Failed++ + result.Items = append(result.Items, item) + continue + } if existing == nil { if !shouldCreateAccount(src.ID, selectedSet) { @@ -809,7 +836,7 @@ func (s *CRSSyncService) SyncFromCRS(ctx context.Context, input SyncFromCRSInput continue } - existing.Extra = mergeMap(existing.Extra, extra) + existing.Extra = extra existing.Name = defaultName(src.Name, src.ID) existing.Platform = PlatformOpenAI existing.Type = AccountTypeAPIKey @@ -1098,6 +1125,10 @@ func mergeMap(existing map[string]any, updates map[string]any) map[string]any { return out } +func mergeCRSOpenAILongContextBillingExtra(existing, updates map[string]any) (map[string]any, error) { + return normalizeOpenAILongContextBillingExtra(PlatformOpenAI, mergeMap(existing, updates)) +} + func (s *CRSSyncService) mapOrCreateProxy(ctx context.Context, enabled bool, cached *[]Proxy, src *crsProxy, defaultName string) (*int64, error) { if !enabled || src == nil { return nil, nil diff --git a/backend/internal/service/gateway_usage_billing.go b/backend/internal/service/gateway_usage_billing.go index 8a95915981..61ab3abd2f 100644 --- a/backend/internal/service/gateway_usage_billing.go +++ b/backend/internal/service/gateway_usage_billing.go @@ -947,6 +947,7 @@ func (s *GatewayService) buildRecordUsageLog( usageLog.CacheReadCost = cost.CacheReadCost usageLog.TotalCost = cost.TotalCost usageLog.ActualCost = cost.ActualCost + usageLog.LongContextBillingApplied = cost.LongContextBillingApplied } return usageLog diff --git a/backend/internal/service/grok_quota_fetcher.go b/backend/internal/service/grok_quota_fetcher.go index 0939b78e20..f220fe33b9 100644 --- a/backend/internal/service/grok_quota_fetcher.go +++ b/backend/internal/service/grok_quota_fetcher.go @@ -3,6 +3,8 @@ package service import ( "encoding/json" "fmt" + "net/http" + "strings" "time" "github.com/Wei-Shaw/sub2api/internal/pkg/xai" @@ -24,54 +26,150 @@ func (f *GrokQuotaFetcher) BuildUsageInfo(account *Account) *UsageInfo { } if account == nil { usage.ErrorCode = "quota_unknown" - usage.Error = "Grok quota is unknown until the first upstream response includes xAI rate-limit headers" + usage.Error = "Grok quota is unknown until billing is probed or an upstream response includes xAI rate-limit headers" return usage } + billing, _ := grokBillingSnapshotFromExtra(account.Extra) snapshot, err := grokQuotaSnapshotFromExtra(account.Extra) + if billing != nil { + usage.GrokBilling = billing + if billing.Plan != "" { + usage.SubscriptionTier = billing.Plan + usage.SubscriptionTierRaw = billing.Plan + } + if parsedAt, parseErr := time.Parse(time.RFC3339, billing.UpdatedAt); parseErr == nil { + usage.UpdatedAt = &parsedAt + } + if billing.FetchedAt != "" { + usage.GrokLastQuotaProbeAt = billing.FetchedAt + } + usage.GrokQuotaSnapshotState = "billing_observed" + usage.GrokLastStatusCode = billing.StatusCode + switch billing.StatusCode { + case 401: + usage.NeedsReauth = true + usage.ErrorCode = "unauthenticated" + case 403: + usage.IsForbidden = true + usage.ForbiddenType = "forbidden" + usage.ErrorCode = "forbidden" + case 429: + usage.ErrorCode = "rate_limited" + } + } + if err != nil || snapshot == nil { - usage.ErrorCode = "quota_unknown" - usage.Error = "Grok quota is unknown until the first upstream response includes xAI rate-limit headers" + applyGrokCredentialUsageFallback(usage, account) + if billing == nil { + usage.ErrorCode = "quota_unknown" + usage.Error = "Grok quota is unknown until billing is probed or an upstream response includes xAI rate-limit headers" + } return usage } - if parsedAt, err := time.Parse(time.RFC3339, snapshot.UpdatedAt); err == nil { - usage.UpdatedAt = &parsedAt + if parsedAt, parseErr := time.Parse(time.RFC3339, snapshot.UpdatedAt); parseErr == nil { + if billing == nil || usage.UpdatedAt == nil || parsedAt.After(*usage.UpdatedAt) { + usage.UpdatedAt = &parsedAt + } } usage.GrokRequestQuota = snapshot.Requests usage.GrokTokenQuota = snapshot.Tokens usage.GrokRetryAfterSeconds = snapshot.RetryAfterSeconds - usage.SubscriptionTier = snapshot.SubscriptionTier - usage.SubscriptionTierRaw = snapshot.SubscriptionTier - usage.GrokEntitlementStatus = snapshot.EntitlementStatus - usage.GrokLastQuotaProbeAt = snapshot.LastProbeAt + if usage.SubscriptionTier == "" { + usage.SubscriptionTier = snapshot.SubscriptionTier + usage.SubscriptionTierRaw = snapshot.SubscriptionTier + } + if usage.GrokEntitlementStatus == "" { + usage.GrokEntitlementStatus = snapshot.EntitlementStatus + } + if usage.GrokLastQuotaProbeAt == "" { + usage.GrokLastQuotaProbeAt = snapshot.LastProbeAt + } usage.GrokLastHeadersSeenAt = snapshot.LastHeadersSeenAt - usage.GrokLastStatusCode = snapshot.StatusCode + if snapshot.StatusCode >= http.StatusBadRequest || usage.GrokLastStatusCode == 0 { + usage.GrokLastStatusCode = snapshot.StatusCode + } if snapshot.HasObservedHeaders() { - usage.GrokQuotaSnapshotState = "observed" - } else { + if usage.GrokQuotaSnapshotState == "" { + usage.GrokQuotaSnapshotState = "observed" + } + } else if billing == nil { usage.GrokQuotaSnapshotState = "no_headers" usage.ErrorCode = "quota_unknown" usage.Error = "No xAI quota headers observed on the latest Grok probe" } - switch snapshot.StatusCode { - case 401: - usage.NeedsReauth = true - usage.ErrorCode = "unauthenticated" - case 403: - usage.IsForbidden = true - usage.ForbiddenType = "forbidden" - usage.ErrorCode = "forbidden" - if usage.GrokEntitlementStatus == "" { - usage.GrokEntitlementStatus = "forbidden" + if usage.ErrorCode == "" { + switch snapshot.StatusCode { + case 401: + usage.NeedsReauth = true + usage.ErrorCode = "unauthenticated" + case 403: + usage.IsForbidden = true + usage.ForbiddenType = "forbidden" + usage.ErrorCode = "forbidden" + if usage.GrokEntitlementStatus == "" { + usage.GrokEntitlementStatus = "forbidden" + } + case 429: + usage.ErrorCode = "rate_limited" } - case 429: - usage.ErrorCode = "rate_limited" } + applyGrokCredentialUsageFallback(usage, account) return usage } +func applyGrokCredentialUsageFallback(usage *UsageInfo, account *Account) { + if usage == nil || account == nil { + return + } + if usage.SubscriptionTier == "" { + tier := strings.TrimSpace(account.GetCredential("subscription_tier")) + usage.SubscriptionTier = tier + usage.SubscriptionTierRaw = tier + } + if usage.GrokEntitlementStatus == "" { + usage.GrokEntitlementStatus = strings.TrimSpace(account.GetCredential("entitlement_status")) + } +} + +func grokBillingSnapshotFromExtra(extra map[string]any) (*xai.BillingSummary, error) { + if extra == nil { + return nil, nil + } + raw, ok := extra[grokBillingExtraKey] + if !ok || raw == nil { + return nil, nil + } + switch snapshot := raw.(type) { + case *xai.BillingSummary: + return snapshot, nil + case xai.BillingSummary: + return &snapshot, nil + case map[string]any: + data, err := json.Marshal(snapshot) + if err != nil { + return nil, err + } + var out xai.BillingSummary + if err := json.Unmarshal(data, &out); err != nil { + return nil, err + } + return &out, nil + default: + data, err := json.Marshal(raw) + if err != nil { + return nil, fmt.Errorf("marshal grok billing snapshot: %w", err) + } + var out xai.BillingSummary + if err := json.Unmarshal(data, &out); err != nil { + return nil, err + } + return &out, nil + } +} + func grokQuotaSnapshotFromExtra(extra map[string]any) (*xai.QuotaSnapshot, error) { if extra == nil { return nil, nil diff --git a/backend/internal/service/grok_quota_fetcher_test.go b/backend/internal/service/grok_quota_fetcher_test.go index d2d9c14993..1de9b51c9e 100644 --- a/backend/internal/service/grok_quota_fetcher_test.go +++ b/backend/internal/service/grok_quota_fetcher_test.go @@ -20,7 +20,34 @@ func TestGrokQuotaFetcherBuildUsageInfoUnknownUntilFirstSnapshot(t *testing.T) { usage := NewGrokQuotaFetcher().BuildUsageInfo(&Account{Platform: PlatformGrok, Type: AccountTypeOAuth}) require.Equal(t, "passive", usage.Source) require.Equal(t, "quota_unknown", usage.ErrorCode) - require.Contains(t, usage.Error, "unknown until the first upstream response") + require.Contains(t, usage.Error, "unknown until billing is probed") +} + +func TestGrokQuotaFetcherUsesCredentialTierWhenBillingHasNoPlan(t *testing.T) { + t.Parallel() + + account := &Account{ + Platform: PlatformGrok, + Type: AccountTypeOAuth, + Credentials: map[string]any{ + "subscription_tier": " FREE ", + "entitlement_status": " active ", + }, + Extra: map[string]any{ + grokBillingExtraKey: &xai.BillingSummary{ + PeriodType: "weekly", + StatusCode: http.StatusOK, + UpdatedAt: "2030-01-01T00:00:00Z", + }, + }, + } + + usage := NewGrokQuotaFetcher().BuildUsageInfo(account) + + require.NotNil(t, usage.GrokBilling) + require.Equal(t, "FREE", usage.SubscriptionTier) + require.Equal(t, "FREE", usage.SubscriptionTierRaw) + require.Equal(t, "active", usage.GrokEntitlementStatus) } func TestGrokQuotaFetcherBuildUsageInfoFromSnapshot(t *testing.T) { @@ -68,6 +95,32 @@ func TestGrokQuotaFetcherBuildUsageInfoFromSnapshot(t *testing.T) { require.True(t, usage.UpdatedAt.Equal(time.Date(2030, 1, 1, 0, 0, 0, 0, time.UTC))) } +func TestGrokQuotaFetcherSnapshotErrorOverridesSuccessfulBillingStatus(t *testing.T) { + t.Parallel() + + updatedAt := "2030-01-01T00:00:00Z" + account := &Account{ + Platform: PlatformGrok, + Type: AccountTypeOAuth, + Extra: map[string]any{ + grokBillingExtraKey: &xai.BillingSummary{ + PeriodType: "weekly", + StatusCode: http.StatusOK, + UpdatedAt: updatedAt, + }, + grokQuotaSnapshotExtraKey: &xai.QuotaSnapshot{ + StatusCode: http.StatusTooManyRequests, + UpdatedAt: updatedAt, + }, + }, + } + + usage := NewGrokQuotaFetcher().BuildUsageInfo(account) + + require.Equal(t, "rate_limited", usage.ErrorCode) + require.Equal(t, http.StatusTooManyRequests, usage.GrokLastStatusCode) +} + func TestGrokQuotaFetcherBuildUsageInfoFromNoHeadersProbe(t *testing.T) { t.Parallel() diff --git a/backend/internal/service/grok_quota_service.go b/backend/internal/service/grok_quota_service.go index 17e91dac6f..aa838cb491 100644 --- a/backend/internal/service/grok_quota_service.go +++ b/backend/internal/service/grok_quota_service.go @@ -7,27 +7,36 @@ import ( "io" "log/slog" "net/http" + "strconv" "strings" + "sync" "time" infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" "github.com/Wei-Shaw/sub2api/internal/pkg/xai" + "golang.org/x/sync/singleflight" ) const ( grokQuotaUpstreamTimeout = 20 * time.Second grokQuotaProbeInput = "." - grokQuotaDefaultModel = "grok-4.3" + grokQuotaDefaultModel = grokDefaultResponsesModel + grokBillingExtraKey = "grok_billing_snapshot" ) type GrokQuotaProbeResult struct { - Source string `json:"source"` - Model string `json:"model"` - Snapshot *xai.QuotaSnapshot `json:"snapshot,omitempty"` - StatusCode int `json:"status_code,omitempty"` - HeadersObserved bool `json:"headers_observed"` - ResetSupported bool `json:"reset_supported"` - FetchedAt int64 `json:"fetched_at"` + Source string `json:"source"` + Model string `json:"model,omitempty"` + Billing *xai.BillingSummary `json:"billing,omitempty"` + Snapshot *xai.QuotaSnapshot `json:"snapshot,omitempty"` + LocalUsage7d *WindowStats `json:"local_usage_7d,omitempty"` + LocalUsageMonthly *WindowStats `json:"local_usage_monthly,omitempty"` + StatusCode int `json:"status_code,omitempty"` + HeadersObserved bool `json:"headers_observed"` + ResetSupported bool `json:"reset_supported"` + FetchedAt int64 `json:"fetched_at"` + Persisted bool `json:"persisted"` + ProbeError string `json:"probe_error,omitempty"` } type GrokQuotaResetResult struct { @@ -41,6 +50,8 @@ type GrokQuotaService struct { proxyRepo ProxyRepository tokenProvider *GrokTokenProvider httpUpstream HTTPUpstream + usageLogRepo UsageLogRepository + probeFlight singleflight.Group } func NewGrokQuotaService( @@ -48,16 +59,70 @@ func NewGrokQuotaService( proxyRepo ProxyRepository, tokenProvider *GrokTokenProvider, httpUpstream HTTPUpstream, + usageLogRepos ...UsageLogRepository, ) *GrokQuotaService { + var usageLogRepo UsageLogRepository + if len(usageLogRepos) > 0 { + usageLogRepo = usageLogRepos[0] + } return &GrokQuotaService{ accountRepo: accountRepo, proxyRepo: proxyRepo, tokenProvider: tokenProvider, httpUpstream: httpUpstream, + usageLogRepo: usageLogRepo, } } +// QueryQuota combines xAI billing data with an active quota-header probe for +// Free accounts, whose billing response does not include usage_percent. +func (s *GrokQuotaService) QueryQuota(ctx context.Context, accountID int64) (*GrokQuotaProbeResult, error) { + billingResult, billingErr := s.ProbeBilling(ctx, accountID) + if billingErr == nil && billingResult != nil && grokBillingHasAuthoritativeQuota(billingResult.Billing) { + return billingResult, nil + } + + probeResult, probeErr := s.ProbeUsage(ctx, accountID) + if probeErr != nil { + if billingResult != nil && billingResult.Billing != nil { + billingResult.ProbeError = probeErr.Error() + return billingResult, nil + } + return nil, probeErr + } + if probeResult == nil { + if billingErr != nil { + return nil, billingErr + } + return nil, infraerrors.New(http.StatusBadGateway, "GROK_QUOTA_PROBE_EMPTY", "Grok quota probe returned no result") + } + if billingResult != nil { + probeResult.Source = "hybrid_probe" + probeResult.Billing = billingResult.Billing + probeResult.LocalUsage7d = billingResult.LocalUsage7d + probeResult.LocalUsageMonthly = billingResult.LocalUsageMonthly + probeResult.Persisted = probeResult.Persisted || billingResult.Persisted + } + return probeResult, nil +} + +func grokBillingHasAuthoritativeQuota(billing *xai.BillingSummary) bool { + if billing == nil { + return false + } + return billing.UsagePercent != nil || + billing.UsedPercent != nil || + (billing.MonthlyLimitCents != nil && *billing.MonthlyLimitCents > 0) || + strings.TrimSpace(billing.Plan) != "" +} + func (s *GrokQuotaService) ProbeUsage(ctx context.Context, accountID int64) (*GrokQuotaProbeResult, error) { + return s.runProbeFlight(ctx, "active:"+strconv.FormatInt(accountID, 10), func(sharedCtx context.Context) (*GrokQuotaProbeResult, error) { + return s.probeUsage(sharedCtx, accountID) + }) +} + +func (s *GrokQuotaService) probeUsage(ctx context.Context, accountID int64) (*GrokQuotaProbeResult, error) { account, token, proxyURL, err := s.prepareProbe(ctx, accountID) if err != nil { return nil, err @@ -95,7 +160,7 @@ func (s *GrokQuotaService) ProbeUsage(ctx context.Context, accountID int64) (*Gr if limited { normalizeGrokExhaustedWindowResets(snapshot, resetAt, time.Now()) } - _ = s.accountRepo.UpdateExtra(ctx, account.ID, map[string]any{ + persistErr := s.accountRepo.UpdateExtra(ctx, account.ID, map[string]any{ grokQuotaSnapshotExtraKey: snapshot, }) if limited { @@ -110,6 +175,7 @@ func (s *GrokQuotaService) ProbeUsage(ctx context.Context, accountID int64) (*Gr HeadersObserved: snapshot.HeadersObserved, ResetSupported: false, FetchedAt: time.Now().Unix(), + Persisted: persistErr == nil, } if resp.StatusCode == http.StatusTooManyRequests { return result, nil @@ -123,6 +189,173 @@ func (s *GrokQuotaService) ProbeUsage(ctx context.Context, accountID int64) (*Gr return result, nil } +// ProbeBilling only calls the xAI billing endpoints. Account usage refreshes +// use this method so opening the account list never consumes model quota. +func (s *GrokQuotaService) ProbeBilling(ctx context.Context, accountID int64) (*GrokQuotaProbeResult, error) { + return s.runProbeFlight(ctx, "billing:"+strconv.FormatInt(accountID, 10), func(sharedCtx context.Context) (*GrokQuotaProbeResult, error) { + return s.probeBilling(sharedCtx, accountID) + }) +} + +func (s *GrokQuotaService) probeBilling(ctx context.Context, accountID int64) (*GrokQuotaProbeResult, error) { + account, token, proxyURL, err := s.prepareProbe(ctx, accountID) + if err != nil { + return nil, err + } + + probeCtx, cancel := context.WithTimeout(ctx, grokQuotaUpstreamTimeout) + defer cancel() + type billingResult struct { + summary *xai.BillingSummary + status int + err error + } + var weekly, monthly billingResult + var wg sync.WaitGroup + wg.Add(2) + go func() { + defer wg.Done() + weekly.summary, weekly.status, weekly.err = s.fetchBilling(probeCtx, account, token, proxyURL, true) + }() + go func() { + defer wg.Done() + monthly.summary, monthly.status, monthly.err = s.fetchBilling(probeCtx, account, token, proxyURL, false) + }() + wg.Wait() + + weeklyOK := weekly.summary != nil + monthlyOK := monthly.summary != nil + if !weeklyOK && !monthlyOK { + return nil, mergeGrokBillingProbeErrors(weekly.status, monthly.status, weekly.err, monthly.err) + } + statusCode := preferSuccessfulBillingStatus(weekly.status, monthly.status, weeklyOK, monthlyOK) + previous, _ := grokBillingSnapshotFromExtra(account.Extra) + billing := xai.MergeBillingProbeResult(previous, weekly.summary, monthly.summary, weeklyOK, monthlyOK) + billing = xai.StampBillingSummary(billing, statusCode, "billing_probe") + persistErr := s.accountRepo.UpdateExtra(ctx, account.ID, map[string]any{ + grokBillingExtraKey: billing, + }) + if persistErr != nil { + slog.Warn("grok_billing_persist_failed", "account_id", account.ID, "error", persistErr) + } + localUsage7d, localUsageMonthly := grokLocalUsageForBilling(ctx, s.usageLogRepo, account.ID, billing, time.Now().UTC()) + return &GrokQuotaProbeResult{ + Source: "billing_probe", + Billing: billing, + LocalUsage7d: localUsage7d, + LocalUsageMonthly: localUsageMonthly, + StatusCode: statusCode, + FetchedAt: time.Now().Unix(), + Persisted: persistErr == nil, + }, nil +} + +func (s *GrokQuotaService) runProbeFlight( + ctx context.Context, + key string, + probe func(context.Context) (*GrokQuotaProbeResult, error), +) (*GrokQuotaProbeResult, error) { + if s == nil { + return nil, infraerrors.New(http.StatusInternalServerError, "GROK_QUOTA_NOT_CONFIGURED", "grok quota service is not configured") + } + resultCh := s.probeFlight.DoChan(key, func() (any, error) { + sharedCtx, cancel := context.WithTimeout(context.Background(), grokQuotaUpstreamTimeout+5*time.Second) + defer cancel() + return probe(sharedCtx) + }) + select { + case <-ctx.Done(): + return nil, ctx.Err() + case flightResult := <-resultCh: + if flightResult.Err != nil { + return nil, flightResult.Err + } + result, ok := flightResult.Val.(*GrokQuotaProbeResult) + if !ok || result == nil { + return nil, infraerrors.New(http.StatusInternalServerError, "GROK_QUOTA_PROBE_RESULT_INVALID", "invalid Grok quota probe result") + } + cloned := *result + return &cloned, nil + } +} + +func (s *GrokQuotaService) fetchBilling( + ctx context.Context, + account *Account, + token string, + proxyURL string, + weekly bool, +) (*xai.BillingSummary, int, error) { + req, err := http.NewRequestWithContext(ctx, http.MethodGet, xai.BuildBillingURL(weekly), nil) + if err != nil { + return nil, 0, infraerrors.Newf(http.StatusInternalServerError, "GROK_QUOTA_PROBE_REQUEST_BUILD_FAILED", "failed to build billing request: %v", err) + } + xai.ApplyCLIBillingHeaders(req, token) + resp, err := s.httpUpstream.Do(req, proxyURL, account.ID, maxInt(account.Concurrency, 2)) + if err != nil { + return nil, 0, infraerrors.Newf(http.StatusBadGateway, "GROK_QUOTA_PROBE_REQUEST_FAILED", "billing request failed: %v", err) + } + defer func() { _ = resp.Body.Close() }() + + bodyBytes, _ := io.ReadAll(io.LimitReader(resp.Body, 1<<20)) + if resp.StatusCode == http.StatusTooManyRequests { + return nil, resp.StatusCode, nil + } + if resp.StatusCode >= 400 { + bodyText := truncate(strings.TrimSpace(string(bodyBytes)), 240) + slog.Warn("grok_quota_billing_failed", "account_id", account.ID, "weekly", weekly, "status", resp.StatusCode, "body", bodyText) + return nil, resp.StatusCode, infraerrors.Newf(mapUpstreamStatus(resp.StatusCode), "GROK_QUOTA_PROBE_UPSTREAM_ERROR", "billing returned %d: %s", resp.StatusCode, bodyText) + } + payload, err := xai.ParseBillingPayload(bodyBytes) + if err != nil { + return nil, resp.StatusCode, infraerrors.Newf(http.StatusBadGateway, "GROK_QUOTA_BILLING_PARSE_ERROR", "failed to parse billing body: %v", err) + } + return xai.BuildBillingSummary(payload.Config), resp.StatusCode, nil +} + +func mergeGrokBillingProbeErrors(weeklyStatus, monthlyStatus int, weeklyErr, monthlyErr error) error { + weeklyKey := grokBillingProbeErrorKey(weeklyStatus, weeklyErr) + monthlyKey := grokBillingProbeErrorKey(monthlyStatus, monthlyErr) + if weeklyKey == monthlyKey { + switch { + case weeklyErr != nil: + return weeklyErr + case monthlyErr != nil: + return monthlyErr + case weeklyStatus == http.StatusTooManyRequests: + return infraerrors.New(http.StatusTooManyRequests, "GROK_QUOTA_PROBE_UPSTREAM_ERROR", "billing rate limited") + case weeklyStatus != 0 && weeklyStatus != http.StatusOK: + return infraerrors.New(mapUpstreamStatus(weeklyStatus), "GROK_QUOTA_PROBE_UPSTREAM_ERROR", "xAI billing endpoints returned the same upstream error") + default: + return infraerrors.New(http.StatusBadGateway, "GROK_QUOTA_BILLING_EMPTY", "xAI billing endpoints returned no quota data") + } + } + slog.Warn("grok_quota_probe_parts_failed", "weekly_status", weeklyStatus, "weekly_error", weeklyErr, "monthly_status", monthlyStatus, "monthly_error", monthlyErr) + return infraerrors.New(http.StatusBadGateway, "GROK_QUOTA_PROBE_PARTS_FAILED", "weekly and monthly billing probes failed differently").WithMetadata(map[string]string{ + "weekly_status": strconv.Itoa(weeklyStatus), "monthly_status": strconv.Itoa(monthlyStatus), + }) +} + +func grokBillingProbeErrorKey(status int, err error) string { + if err != nil { + return strconv.Itoa(status) + ":" + strconv.Itoa(infraerrors.Code(err)) + ":" + infraerrors.Reason(err) + } + return strconv.Itoa(status) + ":empty" +} + +func preferSuccessfulBillingStatus(weeklyStatus, monthlyStatus int, weeklyOK, monthlyOK bool) int { + if weeklyOK && weeklyStatus >= 200 && weeklyStatus < 300 { + return weeklyStatus + } + if monthlyOK && monthlyStatus >= 200 && monthlyStatus < 300 { + return monthlyStatus + } + if weeklyStatus != 0 { + return weeklyStatus + } + return monthlyStatus +} + func (s *GrokQuotaService) ResetQuota(ctx context.Context, accountID int64) (*GrokQuotaResetResult, error) { if _, err := s.loadGrokOAuthAccount(ctx, accountID); err != nil { return nil, err diff --git a/backend/internal/service/grok_quota_service_test.go b/backend/internal/service/grok_quota_service_test.go index 2248674899..d49497e69b 100644 --- a/backend/internal/service/grok_quota_service_test.go +++ b/backend/internal/service/grok_quota_service_test.go @@ -6,11 +6,14 @@ import ( "context" "io" "net/http" + "strconv" "strings" + "sync" "testing" "time" infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" + "github.com/Wei-Shaw/sub2api/internal/pkg/usagestats" "github.com/Wei-Shaw/sub2api/internal/pkg/xai" "github.com/stretchr/testify/require" "github.com/tidwall/gjson" @@ -63,6 +66,107 @@ type grokQuotaProxyRepo struct { calls int } +type grokQuotaUsageLogRepo struct { + UsageLogRepository + stats *usagestats.AccountStats + err error + calls int +} + +func (r *grokQuotaUsageLogRepo) GetAccountWindowStats(context.Context, int64, time.Time) (*usagestats.AccountStats, error) { + r.calls++ + return r.stats, r.err +} + +type grokHybridUpstream struct { + httpUpstreamRecorder + mu sync.Mutex + requests []*http.Request + bodies [][]byte + weeklyUsagePercent *float64 + monthlyLimitCents *float64 + activeStatus int + activeHeaders http.Header + billingStarted chan struct{} + billingRelease <-chan struct{} + billingStartOnce sync.Once + billingStatus int + billingHeaders http.Header +} + +func (u *grokHybridUpstream) Do(req *http.Request, _ string, _ int64, _ int) (*http.Response, error) { + var body []byte + if req != nil && req.Body != nil { + body, _ = io.ReadAll(req.Body) + } + u.mu.Lock() + u.requests = append(u.requests, req) + u.bodies = append(u.bodies, body) + u.mu.Unlock() + + if req.URL.Path == "/v1/responses" { + status := u.activeStatus + if status == 0 { + status = http.StatusOK + } + headers := u.activeHeaders + if headers == nil { + headers = http.Header{ + "X-Ratelimit-Limit-Tokens": []string{"2000000"}, + "X-Ratelimit-Remaining-Tokens": []string{"1500000"}, + } + } + return &http.Response{StatusCode: status, Header: headers, Body: io.NopCloser(strings.NewReader(`{"id":"resp_probe"}`))}, nil + } + if u.billingStarted != nil { + u.billingStartOnce.Do(func() { close(u.billingStarted) }) + } + if u.billingRelease != nil { + select { + case <-u.billingRelease: + case <-req.Context().Done(): + return nil, req.Context().Err() + } + } + if u.billingStatus != 0 && u.billingStatus != http.StatusOK { + return &http.Response{ + StatusCode: u.billingStatus, + Header: u.billingHeaders, + Body: io.NopCloser(strings.NewReader(`{"error":{"message":"billing limited"}}`)), + }, nil + } + + if req.URL.RawQuery == "format=credits" { + usage := "" + if u.weeklyUsagePercent != nil { + usage = `,"creditUsagePercent":` + strconv.FormatFloat(*u.weeklyUsagePercent, 'f', -1, 64) + } + payload := `{"config":{"currentPeriod":{"type":"WEEKLY","start":"2026-07-09T03:25:00Z","end":"2026-07-16T03:25:00Z"}` + usage + `}}` + return &http.Response{StatusCode: http.StatusOK, Header: make(http.Header), Body: io.NopCloser(strings.NewReader(payload))}, nil + } + monthlyLimit := "" + if u.monthlyLimitCents != nil { + monthlyLimit = `,"monthlyLimit":{"val":` + strconv.FormatFloat(*u.monthlyLimitCents, 'f', -1, 64) + `}` + } + monthlyPayload := `{"config":{"billingPeriodStart":"2026-07-01T00:00:00Z","billingPeriodEnd":"2026-08-01T00:00:00Z"` + monthlyLimit + `}}` + return &http.Response{ + StatusCode: http.StatusOK, + Header: make(http.Header), + Body: io.NopCloser(strings.NewReader(monthlyPayload)), + }, nil +} + +func (u *grokHybridUpstream) snapshot() ([]*http.Request, [][]byte) { + u.mu.Lock() + defer u.mu.Unlock() + requests := append([]*http.Request(nil), u.requests...) + bodies := make([][]byte, len(u.bodies)) + for i := range u.bodies { + bodies[i] = append([]byte(nil), u.bodies[i]...) + } + return requests, bodies +} + func (r *grokQuotaProxyRepo) GetByID(_ context.Context, id int64) (*Proxy, error) { r.calls++ return r.proxies[id], nil @@ -102,7 +206,7 @@ func TestGrokQuotaServiceProbeUsageStoresHeaders(t *testing.T) { result, err := svc.ProbeUsage(context.Background(), 42) require.NoError(t, err) require.Equal(t, http.StatusOK, result.StatusCode) - require.Equal(t, "grok-4.3", result.Model) + require.Equal(t, "grok-4.5", result.Model) require.True(t, result.HeadersObserved) require.NotNil(t, result.Snapshot) require.True(t, result.Snapshot.HeadersObserved) @@ -115,7 +219,7 @@ func TestGrokQuotaServiceProbeUsageStoresHeaders(t *testing.T) { require.Equal(t, "https://cli-chat-proxy.grok.com/v1/responses", upstream.lastReq.URL.String()) require.Equal(t, "Bearer access-token", upstream.lastReq.Header.Get("Authorization")) require.Equal(t, grokCLIVersion, upstream.lastReq.Header.Get("X-Grok-Client-Version")) - require.Equal(t, "grok-4.3", gjson.GetBytes(upstream.lastBody, "model").String()) + require.Equal(t, "grok-4.5", gjson.GetBytes(upstream.lastBody, "model").String()) require.Contains(t, string(upstream.lastBody), `"max_output_tokens":1`) require.Contains(t, string(upstream.lastBody), `"store":false`) require.NotNil(t, repo.updates[42][grokQuotaSnapshotExtraKey]) @@ -152,8 +256,8 @@ func TestGrokQuotaServiceProbeUsageIgnoresAccountGrokMapping(t *testing.T) { result, err := svc.ProbeUsage(context.Background(), 47) require.NoError(t, err) - require.Equal(t, "grok-4.3", result.Model) - require.Equal(t, "grok-4.3", gjson.GetBytes(upstream.lastBody, "model").String()) + require.Equal(t, "grok-4.5", result.Model) + require.Equal(t, "grok-4.5", gjson.GetBytes(upstream.lastBody, "model").String()) require.NotContains(t, string(upstream.lastBody), "grok-composer") } @@ -185,7 +289,7 @@ func TestGrokQuotaServiceProbeUsageReportsProbeModelOnUpstreamError(t *testing.T _, err := svc.ProbeUsage(context.Background(), 48) require.Error(t, err) require.Equal(t, "GROK_QUOTA_PROBE_UPSTREAM_ERROR", infraerrors.Reason(err)) - require.Contains(t, infraerrors.Message(err), `probe model "grok-4.3"`) + require.Contains(t, infraerrors.Message(err), `probe model "grok-4.5"`) } func TestGrokQuotaServiceProbeUsageLoadsProxyWhenAccountEdgeMissing(t *testing.T) { @@ -308,6 +412,299 @@ func TestGrokQuotaServiceProbeUsageReturnsRateLimitedSnapshot(t *testing.T) { require.Zero(t, repo.tempUnschedCalls) } +func TestGrokQuotaServiceQueryQuotaFreeFallsBackToGrok45(t *testing.T) { + t.Parallel() + + account := &Account{ + ID: 51, Platform: PlatformGrok, Type: AccountTypeOAuth, Concurrency: 1, + Credentials: map[string]any{ + "access_token": "access-token", + "expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339), + }, + } + repo := &grokQuotaAccountRepo{mockAccountRepoForPlatform: &mockAccountRepoForPlatform{ + accountsByID: map[int64]*Account{account.ID: account}, + }} + upstream := &grokHybridUpstream{} + svc := NewGrokQuotaService(repo, nil, NewGrokTokenProvider(repo, nil), upstream) + + result, err := svc.QueryQuota(context.Background(), account.ID) + require.NoError(t, err) + require.Equal(t, "hybrid_probe", result.Source) + require.Equal(t, "grok-4.5", result.Model) + require.NotNil(t, result.Billing) + require.Nil(t, result.Billing.UsagePercent) + require.NotNil(t, result.Snapshot) + require.NotNil(t, result.Snapshot.Tokens) + require.EqualValues(t, 2_000_000, *result.Snapshot.Tokens.Limit) + require.True(t, result.HeadersObserved) + + requests, bodies := upstream.snapshot() + require.Len(t, requests, 3) + responseCalls := 0 + for i, req := range requests { + if req.URL.Path != "/v1/responses" { + continue + } + responseCalls++ + require.Equal(t, http.MethodPost, req.Method) + require.Equal(t, "grok-4.5", gjson.GetBytes(bodies[i], "model").String()) + require.EqualValues(t, 1, gjson.GetBytes(bodies[i], "max_output_tokens").Int()) + } + require.Equal(t, 1, responseCalls) +} + +func TestGrokQuotaServiceQueryQuotaPaidBillingSkipsActiveProbe(t *testing.T) { + t.Parallel() + + account := &Account{ + ID: 52, Platform: PlatformGrok, Type: AccountTypeOAuth, Concurrency: 1, + Credentials: map[string]any{ + "access_token": "access-token", + "expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339), + }, + } + repo := &grokQuotaAccountRepo{mockAccountRepoForPlatform: &mockAccountRepoForPlatform{ + accountsByID: map[int64]*Account{account.ID: account}, + }} + usagePercent := 25.0 + upstream := &grokHybridUpstream{weeklyUsagePercent: &usagePercent} + svc := NewGrokQuotaService(repo, nil, NewGrokTokenProvider(repo, nil), upstream) + + result, err := svc.QueryQuota(context.Background(), account.ID) + require.NoError(t, err) + require.Equal(t, "billing_probe", result.Source) + require.NotNil(t, result.Billing) + require.InDelta(t, usagePercent, *result.Billing.UsagePercent, 1e-9) + require.Nil(t, result.Snapshot) + require.Empty(t, result.Model) + + requests, _ := upstream.snapshot() + require.Len(t, requests, 2) + for _, req := range requests { + require.Equal(t, "/v1/billing", req.URL.Path) + } +} + +func TestGrokQuotaServiceQueryQuotaCustomPaidMonthlyLimitSkipsActiveProbe(t *testing.T) { + t.Parallel() + + account := &Account{ + ID: 57, Platform: PlatformGrok, Type: AccountTypeOAuth, Concurrency: 1, + Credentials: map[string]any{ + "access_token": "access-token", + "expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339), + }, + } + repo := &grokQuotaAccountRepo{mockAccountRepoForPlatform: &mockAccountRepoForPlatform{ + accountsByID: map[int64]*Account{account.ID: account}, + }} + monthlyLimit := 25_000.0 + upstream := &grokHybridUpstream{monthlyLimitCents: &monthlyLimit} + svc := NewGrokQuotaService(repo, nil, NewGrokTokenProvider(repo, nil), upstream) + + result, err := svc.QueryQuota(context.Background(), account.ID) + require.NoError(t, err) + require.Equal(t, "billing_probe", result.Source) + require.NotNil(t, result.Billing) + require.InDelta(t, monthlyLimit, *result.Billing.MonthlyLimitCents, 1e-9) + require.Nil(t, result.Snapshot) + + requests, _ := upstream.snapshot() + require.Len(t, requests, 2) + for _, req := range requests { + require.Equal(t, "/v1/billing", req.URL.Path) + } +} + +func TestGrokLocalUsageForBillingOnlyReturnsAvailableWindows(t *testing.T) { + t.Parallel() + + now := time.Date(2026, 7, 13, 12, 0, 0, 0, time.UTC) + billing := &xai.BillingSummary{ + PeriodType: "weekly", + PeriodStart: now.Add(-4 * 24 * time.Hour).Format(time.RFC3339), + PeriodEnd: now.Add(3 * 24 * time.Hour).Format(time.RFC3339), + } + + t.Run("valid weekly window", func(t *testing.T) { + repo := &grokQuotaUsageLogRepo{stats: &usagestats.AccountStats{Tokens: 1_500_000}} + weekly, monthly := grokLocalUsageForBilling(context.Background(), repo, 57, billing, now) + require.NotNil(t, weekly) + require.EqualValues(t, 1_500_000, weekly.Tokens) + require.Nil(t, monthly) + require.Equal(t, 1, repo.calls) + }) + + t.Run("query failure", func(t *testing.T) { + repo := &grokQuotaUsageLogRepo{err: context.DeadlineExceeded} + weekly, monthly := grokLocalUsageForBilling(context.Background(), repo, 57, billing, now) + require.Nil(t, weekly) + require.Nil(t, monthly) + require.Equal(t, 1, repo.calls) + }) + + t.Run("missing billing window", func(t *testing.T) { + repo := &grokQuotaUsageLogRepo{} + weekly, monthly := grokLocalUsageForBilling(context.Background(), repo, 57, nil, now) + require.Nil(t, weekly) + require.Nil(t, monthly) + require.Zero(t, repo.calls) + }) +} + +func TestAccountUsageServiceGrokRefreshUsesBillingOnly(t *testing.T) { + t.Parallel() + + account := &Account{ + ID: 54, Platform: PlatformGrok, Type: AccountTypeOAuth, Concurrency: 1, + Credentials: map[string]any{ + "access_token": "access-token", + "expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339), + }, + } + repo := &grokQuotaAccountRepo{mockAccountRepoForPlatform: &mockAccountRepoForPlatform{ + accountsByID: map[int64]*Account{account.ID: account}, + }} + upstream := &grokHybridUpstream{} + quotaService := NewGrokQuotaService(repo, nil, NewGrokTokenProvider(repo, nil), upstream) + usageService := &AccountUsageService{ + grokQuotaFetcher: NewGrokQuotaFetcher(), + grokQuotaService: quotaService, + cache: NewUsageCache(), + } + + usage, err := usageService.getGrokUsage(context.Background(), account, false) + require.NoError(t, err) + require.NotNil(t, usage.GrokBilling) + require.Nil(t, usage.GrokBilling.UsagePercent) + + requests, _ := upstream.snapshot() + require.Len(t, requests, 2) + for _, req := range requests { + require.Equal(t, http.MethodGet, req.Method) + require.Equal(t, "/v1/billing", req.URL.Path) + } +} + +func TestGrokQuotaServiceProbeFlightsDeduplicateBillingAndSeparateActive(t *testing.T) { + t.Parallel() + + account := &Account{ + ID: 55, Platform: PlatformGrok, Type: AccountTypeOAuth, Concurrency: 1, + Credentials: map[string]any{ + "access_token": "access-token", + "expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339), + }, + } + repo := &grokQuotaAccountRepo{mockAccountRepoForPlatform: &mockAccountRepoForPlatform{ + accountsByID: map[int64]*Account{account.ID: account}, + }} + billingStarted := make(chan struct{}) + billingRelease := make(chan struct{}) + upstream := &grokHybridUpstream{billingStarted: billingStarted, billingRelease: billingRelease} + svc := NewGrokQuotaService(repo, nil, NewGrokTokenProvider(repo, nil), upstream) + + type probeOutcome struct { + result *GrokQuotaProbeResult + err error + } + billingOutcomes := make(chan probeOutcome, 2) + go func() { + result, err := svc.ProbeBilling(context.Background(), account.ID) + billingOutcomes <- probeOutcome{result: result, err: err} + }() + <-billingStarted + secondStarted := make(chan struct{}) + go func() { + close(secondStarted) + result, err := svc.ProbeBilling(context.Background(), account.ID) + billingOutcomes <- probeOutcome{result: result, err: err} + }() + <-secondStarted + time.Sleep(25 * time.Millisecond) + + activeResult, err := svc.ProbeUsage(context.Background(), account.ID) + require.NoError(t, err) + require.NotNil(t, activeResult.Snapshot) + close(billingRelease) + for range 2 { + outcome := <-billingOutcomes + require.NoError(t, outcome.err) + require.NotNil(t, outcome.result.Billing) + } + + requests, _ := upstream.snapshot() + billingCalls := 0 + activeCalls := 0 + for _, req := range requests { + switch req.URL.Path { + case "/v1/billing": + billingCalls++ + case "/v1/responses": + activeCalls++ + } + } + require.Equal(t, 2, billingCalls) + require.Equal(t, 1, activeCalls) +} + +func TestGrokQuotaServiceBilling429DoesNotPauseModelScheduling(t *testing.T) { + t.Parallel() + + account := &Account{ + ID: 56, Platform: PlatformGrok, Type: AccountTypeOAuth, Concurrency: 1, + Credentials: map[string]any{ + "access_token": "access-token", + "expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339), + }, + } + repo := &grokQuotaAccountRepo{mockAccountRepoForPlatform: &mockAccountRepoForPlatform{ + accountsByID: map[int64]*Account{account.ID: account}, + }} + upstream := &grokHybridUpstream{ + billingStatus: http.StatusTooManyRequests, + billingHeaders: http.Header{"Retry-After": []string{"45"}}, + } + svc := NewGrokQuotaService(repo, nil, NewGrokTokenProvider(repo, nil), upstream) + + result, err := svc.ProbeBilling(context.Background(), account.ID) + + require.Error(t, err) + require.Nil(t, result) + require.Zero(t, repo.rateLimitedCalls) +} + +func TestGrokQuotaServiceQueryQuotaFree429PersistsLimitAndKeepsBilling(t *testing.T) { + t.Parallel() + + account := &Account{ + ID: 53, Platform: PlatformGrok, Type: AccountTypeOAuth, Concurrency: 1, + Credentials: map[string]any{ + "access_token": "access-token", + "expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339), + }, + } + repo := &grokQuotaAccountRepo{mockAccountRepoForPlatform: &mockAccountRepoForPlatform{ + accountsByID: map[int64]*Account{account.ID: account}, + }} + upstream := &grokHybridUpstream{ + activeStatus: http.StatusTooManyRequests, + activeHeaders: http.Header{"Retry-After": []string{"45"}}, + } + svc := NewGrokQuotaService(repo, nil, NewGrokTokenProvider(repo, nil), upstream) + + result, err := svc.QueryQuota(context.Background(), account.ID) + require.NoError(t, err) + require.Equal(t, http.StatusTooManyRequests, result.StatusCode) + require.NotNil(t, result.Billing) + require.NotNil(t, result.Snapshot) + require.Equal(t, 45, *result.Snapshot.RetryAfterSeconds) + require.Equal(t, 1, repo.rateLimitedCalls) + require.Equal(t, account.ID, repo.lastRateLimitedID) + require.WithinDuration(t, time.Now().Add(45*time.Second), repo.lastRateLimitResetAt, time.Second) +} + func TestGrokQuotaServiceResetQuotaUnsupported(t *testing.T) { t.Parallel() diff --git a/backend/internal/service/openai_alpha_search_billing_test.go b/backend/internal/service/openai_alpha_search_billing_test.go index 1251ee43f9..7151725763 100644 --- a/backend/internal/service/openai_alpha_search_billing_test.go +++ b/backend/internal/service/openai_alpha_search_billing_test.go @@ -50,7 +50,7 @@ func TestCalculateOpenAIRecordUsageCostWebSearchPerCall(t *testing.T) { // 即使 token 倍率(含高峰,3.0)更高也不采用。 apiKey := &APIKey{ID: 1, GroupID: &groupID, Group: &Group{ID: groupID, Platform: PlatformOpenAI}} result := &OpenAIForwardResult{Model: "gpt-5.6-sol", UpstreamModel: "gpt-5.6-sol", WebSearchCalls: 1} - cost, err := svc.calculateOpenAIRecordUsageCost(context.Background(), result, apiKey, []string{"gpt-5.6-sol"}, 3.0, 1.0, 1.0, 2.0, UsageTokens{}, "") + cost, err := svc.calculateOpenAIRecordUsageCost(context.Background(), result, apiKey, []string{"gpt-5.6-sol"}, 3.0, 1.0, 1.0, 2.0, UsageTokens{}, "", false) require.NoError(t, err) require.Equal(t, string(BillingModePerRequest), cost.BillingMode) require.InDelta(t, 0.01, cost.TotalCost, 1e-12) @@ -58,7 +58,7 @@ func TestCalculateOpenAIRecordUsageCostWebSearchPerCall(t *testing.T) { // 分组配置单价 0.005 apiKey.Group.WebSearchPricePerCall = float64Ptr(0.005) - cost, err = svc.calculateOpenAIRecordUsageCost(context.Background(), result, apiKey, []string{"gpt-5.6-sol"}, 1.0, 1.0, 1.0, 1.0, UsageTokens{}, "") + cost, err = svc.calculateOpenAIRecordUsageCost(context.Background(), result, apiKey, []string{"gpt-5.6-sol"}, 1.0, 1.0, 1.0, 1.0, UsageTokens{}, "", false) require.NoError(t, err) require.InDelta(t, 0.005, cost.TotalCost, 1e-12) require.InDelta(t, 0.005, cost.ActualCost, 1e-12) @@ -66,7 +66,7 @@ func TestCalculateOpenAIRecordUsageCostWebSearchPerCall(t *testing.T) { // WebSearchCalls = 0 时不得走按次分支(无定价数据会返回 pricing 错误, // 证明回落到了 token 路径而不是被按次分支吞掉)。 result.WebSearchCalls = 0 - _, err = svc.calculateOpenAIRecordUsageCost(context.Background(), result, apiKey, []string{"gpt-5.6-sol"}, 1.0, 1.0, 1.0, 1.0, UsageTokens{InputTokens: 10}, "") + _, err = svc.calculateOpenAIRecordUsageCost(context.Background(), result, apiKey, []string{"gpt-5.6-sol"}, 1.0, 1.0, 1.0, 1.0, UsageTokens{InputTokens: 10}, "", false) require.Error(t, err) } diff --git a/backend/internal/service/openai_codex_models_service.go b/backend/internal/service/openai_codex_models_service.go index 8a919fa2b0..a7331f64c9 100644 --- a/backend/internal/service/openai_codex_models_service.go +++ b/backend/internal/service/openai_codex_models_service.go @@ -2,21 +2,36 @@ package service import ( "context" + "crypto/sha256" + "errors" + "fmt" "io" + "net" "net/http" "net/url" + "sort" "strings" + "sync" "time" infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" "github.com/Wei-Shaw/sub2api/internal/pkg/httpclient" + "golang.org/x/net/http2" + "golang.org/x/sync/singleflight" ) // chatgptCodexModelsURL is the ChatGPT Codex models manifest endpoint. // Package-level variable so tests can point it at a stub server. var chatgptCodexModelsURL = "https://chatgpt.com/backend-api/codex/models" -const codexModelsManifestBodyLimit int64 = 8 << 20 +const ( + codexModelsManifestBodyLimit int64 = 8 << 20 + codexModelsManifestCacheBodyLimit = 1 << 20 + codexModelsManifestCacheMaxEntries = 64 + codexModelsManifestCacheTTL = 30 * time.Second + codexModelsManifestCacheStaleTTL = 5 * time.Minute + codexModelsManifestRequestTimeout = 15 * time.Second +) // CodexModelsManifest carries the raw upstream manifest payload plus caching // metadata so handlers can pass both through to the client untouched. @@ -26,8 +41,180 @@ type CodexModelsManifest struct { NotModified bool } -// FetchCodexModelsManifest fetches the live Codex models manifest from the -// ChatGPT backend using the account's OAuth credentials. +type codexModelsManifestUpstreamError struct { + err error + retryable bool +} + +func (e *codexModelsManifestUpstreamError) Error() string { return e.err.Error() } + +func (e *codexModelsManifestUpstreamError) Unwrap() error { return e.err } + +// IsRetryableCodexModelsManifestError reports whether another selected account +// may succeed without changing the request. Configuration and upstream 4xx +// responses, except 429, are intentionally not retried. +func IsRetryableCodexModelsManifestError(err error) bool { + var upstreamErr *codexModelsManifestUpstreamError + return errors.As(err, &upstreamErr) && upstreamErr.retryable +} + +func isRetryableCodexModelsManifestTransportError(err error) bool { + if err == nil || errors.Is(err, context.Canceled) { + return false + } + if errors.Is(err, context.DeadlineExceeded) || + errors.Is(err, io.EOF) || + errors.Is(err, io.ErrUnexpectedEOF) || + errors.Is(err, net.ErrClosed) { + return true + } + + var opErr *net.OpError + if errors.As(err, &opErr) { + return true + } + var dnsErr *net.DNSError + if errors.As(err, &dnsErr) { + return true + } + var goAwayErr http2.GoAwayError + if errors.As(err, &goAwayErr) { + return true + } + var streamErr http2.StreamError + if errors.As(err, &streamErr) { + return true + } + var connectionErr http2.ConnectionError + if errors.As(err, &connectionErr) { + return true + } + var netErr net.Error + if errors.As(err, &netErr) && netErr.Timeout() { + return true + } + + // net/http uses unexported HTTP/2 error types, so typed matching is not + // possible for errors produced by the standard library transport. + message := strings.ToLower(err.Error()) + if strings.Contains(message, "http2:") && + (strings.Contains(message, "goaway") || + strings.Contains(message, "refused_stream") || + strings.Contains(message, "frame too large")) { + return true + } + if strings.Contains(message, "stream error: stream id ") { + return true + } + for _, code := range []http2.ErrCode{ + http2.ErrCodeNo, + http2.ErrCodeProtocol, + http2.ErrCodeInternal, + http2.ErrCodeFlowControl, + http2.ErrCodeSettingsTimeout, + http2.ErrCodeStreamClosed, + http2.ErrCodeFrameSize, + http2.ErrCodeRefusedStream, + http2.ErrCodeCancel, + http2.ErrCodeCompression, + http2.ErrCodeConnect, + http2.ErrCodeEnhanceYourCalm, + http2.ErrCodeInadequateSecurity, + http2.ErrCodeHTTP11Required, + } { + if strings.Contains(message, "connection error: "+strings.ToLower(code.String())) { + return true + } + } + return false +} + +type codexModelsManifestRequest struct { + url string + headers http.Header + proxyURL string + accountID int64 + credentialAccountID int64 + accountConcurrency int + useAPIKeyUpstream bool +} + +type codexModelsManifestCacheEntry struct { + manifest *CodexModelsManifest + order uint64 + expiresAt time.Time + staleUntil time.Time +} + +type codexModelsManifestCacheState uint8 + +const ( + codexModelsManifestCacheMiss codexModelsManifestCacheState = iota + codexModelsManifestCacheFresh + codexModelsManifestCacheStale +) + +type codexModelsManifestCache struct { + mu sync.Mutex + entries map[string]codexModelsManifestCacheEntry + nextOrder uint64 + refresh singleflight.Group +} + +func (c *codexModelsManifestCache) get(key string, now time.Time) (*CodexModelsManifest, codexModelsManifestCacheState) { + c.mu.Lock() + defer c.mu.Unlock() + entry, ok := c.entries[key] + if !ok { + return nil, codexModelsManifestCacheMiss + } + if !now.Before(entry.staleUntil) { + delete(c.entries, key) + return nil, codexModelsManifestCacheMiss + } + if now.Before(entry.expiresAt) { + return entry.manifest, codexModelsManifestCacheFresh + } + return entry.manifest, codexModelsManifestCacheStale +} + +func (c *codexModelsManifestCache) set(key string, manifest *CodexModelsManifest, now time.Time) { + if manifest == nil || len(manifest.Body) > codexModelsManifestCacheBodyLimit { + return + } + c.mu.Lock() + defer c.mu.Unlock() + if c.entries == nil { + c.entries = make(map[string]codexModelsManifestCacheEntry) + } + if _, exists := c.entries[key]; !exists && len(c.entries) >= codexModelsManifestCacheMaxEntries { + oldestKey := "" + var oldestOrder uint64 + for candidateKey, entry := range c.entries { + if !now.Before(entry.staleUntil) { + delete(c.entries, candidateKey) + continue + } + if oldestKey == "" || entry.order < oldestOrder { + oldestKey = candidateKey + oldestOrder = entry.order + } + } + if len(c.entries) >= codexModelsManifestCacheMaxEntries && oldestKey != "" { + delete(c.entries, oldestKey) + } + } + c.nextOrder++ + c.entries[key] = codexModelsManifestCacheEntry{ + manifest: manifest, + order: c.nextOrder, + expiresAt: now.Add(codexModelsManifestCacheTTL), + staleUntil: now.Add(codexModelsManifestCacheStaleTTL), + } +} + +// FetchCodexModelsManifest fetches the live Codex models manifest from either +// the ChatGPT backend for OAuth accounts or a custom upstream for API key accounts. // // The response body is passed through verbatim: the manifest schema evolves // with Codex client releases, and interpreting it here would force the gateway @@ -41,49 +228,171 @@ func (s *OpenAIGatewayService) FetchCodexModelsManifest(ctx context.Context, acc if err != nil { return nil, infraerrors.Newf(http.StatusInternalServerError, "OPENAI_CODEX_MODELS_CREDENTIALS_FAILED", "resolve credential account: %v", err) } - accessToken := credAccount.GetOpenAIAccessToken() - if accessToken == "" { - return nil, infraerrors.New(http.StatusBadGateway, "OPENAI_CODEX_MODELS_TOKEN_MISSING", "account has no Codex backend access token") - } clientVersion = strings.TrimSpace(clientVersion) if clientVersion == "" { clientVersion = openAICodexProbeVersion } - requestURL := chatgptCodexModelsURL + "?client_version=" + url.QueryEscape(clientVersion) - reqCtx, cancel := context.WithTimeout(ctx, 15*time.Second) - defer cancel() - req, err := http.NewRequestWithContext(reqCtx, http.MethodGet, requestURL, nil) + requestEndpoint := chatgptCodexModelsURL + authToken := "" + useAPIKeyUpstream := false + appendModelsPath := false + switch { + case credAccount.IsOpenAIOAuth(): + authToken = strings.TrimSpace(credAccount.GetOpenAIAccessToken()) + if authToken == "" { + return nil, infraerrors.New(http.StatusBadGateway, "OPENAI_CODEX_MODELS_TOKEN_MISSING", "account has no Codex backend access token") + } + case credAccount.IsOpenAIApiKey(): + baseURL := strings.TrimSpace(credAccount.GetCredential("base_url")) + if baseURL == "" || isOfficialOpenAIModelsBaseURL(baseURL) { + return nil, infraerrors.New( + http.StatusBadGateway, + "OPENAI_CODEX_MODELS_API_KEY_UPSTREAM_UNSUPPORTED", + "Codex models manifest requires a custom API key upstream base URL", + ) + } + authToken = strings.TrimSpace(credAccount.GetOpenAIApiKey()) + if authToken == "" { + return nil, infraerrors.New(http.StatusBadGateway, "OPENAI_CODEX_MODELS_API_KEY_MISSING", "account has no API key for the Codex models upstream") + } + normalizedBaseURL, validateErr := s.validateUpstreamBaseURL(baseURL) + if validateErr != nil { + return nil, infraerrors.Newf(http.StatusBadGateway, "OPENAI_CODEX_MODELS_API_KEY_UPSTREAM_INVALID", "invalid Codex models upstream base URL: %v", validateErr) + } + requestEndpoint = normalizedBaseURL + useAPIKeyUpstream = true + appendModelsPath = true + default: + return nil, infraerrors.Newf(http.StatusBadGateway, "OPENAI_CODEX_MODELS_ACCOUNT_TYPE_UNSUPPORTED", "account type %q cannot fetch the Codex models manifest", credAccount.Type) + } + + requestURL, err := buildCodexModelsManifestURL(requestEndpoint, appendModelsPath, clientVersion) if err != nil { - return nil, infraerrors.Newf(http.StatusInternalServerError, "OPENAI_CODEX_MODELS_REQUEST_FAILED", "create codex models request: %v", err) + if useAPIKeyUpstream { + return nil, infraerrors.Newf(http.StatusBadGateway, "OPENAI_CODEX_MODELS_API_KEY_UPSTREAM_INVALID", "invalid Codex models upstream base URL: %v", err) + } + return nil, infraerrors.Newf(http.StatusInternalServerError, "OPENAI_CODEX_MODELS_REQUEST_FAILED", "parse codex models request URL: %v", err) } - req.Header.Set("Authorization", "Bearer "+accessToken) - req.Header.Set("Accept", "application/json") - req.Header.Set("Originator", "codex_cli_rs") - req.Header.Set("Version", clientVersion) - req.Header.Set("User-Agent", codexCLIUserAgent) - if ifNoneMatch = strings.TrimSpace(ifNoneMatch); ifNoneMatch != "" { - req.Header.Set("If-None-Match", ifNoneMatch) + + headers := make(http.Header) + headers.Set("Authorization", "Bearer "+authToken) + headers.Set("Accept", "application/json") + headers.Set("Originator", "codex_cli_rs") + headers.Set("Version", clientVersion) + headers.Set("User-Agent", codexCLIUserAgent) + if useAPIKeyUpstream { + credAccount.ApplyHeaderOverrides(headers) + } else { + setOpenAIChatGPTAccountHeaders(headers, credAccount) } - setOpenAIChatGPTAccountHeaders(req.Header, credAccount) proxyURL := "" if account.ProxyID != nil && account.Proxy != nil { proxyURL = account.Proxy.URL() } - client, err := httpclient.GetClient(httpclient.Options{ - ProxyURL: proxyURL, - Timeout: 15 * time.Second, - ResponseHeaderTimeout: 10 * time.Second, + + request := codexModelsManifestRequest{ + url: requestURL.String(), + headers: headers, + proxyURL: proxyURL, + accountID: account.ID, + credentialAccountID: credAccount.ID, + accountConcurrency: account.Concurrency, + useAPIKeyUpstream: useAPIKeyUpstream, + } + if useAPIKeyUpstream { + return s.fetchCachedAPIKeyCodexModelsManifest(ctx, request, ifNoneMatch) + } + return s.fetchCodexModelsManifestUpstream(ctx, request, ifNoneMatch) +} + +func (s *OpenAIGatewayService) fetchCachedAPIKeyCodexModelsManifest(ctx context.Context, request codexModelsManifestRequest, ifNoneMatch string) (*CodexModelsManifest, error) { + if err := ctx.Err(); err != nil { + return nil, err + } + cacheKey := buildCodexModelsManifestCacheKey(request) + manifest, state := s.codexModelsManifestCache.get(cacheKey, time.Now()) + if state == codexModelsManifestCacheFresh { + return codexModelsManifestForClient(manifest, ifNoneMatch), nil + } + resultCh := s.refreshCachedAPIKeyCodexModelsManifest(cacheKey, request) + if state == codexModelsManifestCacheStale { + return codexModelsManifestForClient(manifest, ifNoneMatch), nil + } + select { + case <-ctx.Done(): + return nil, ctx.Err() + case result := <-resultCh: + if result.Err != nil { + return nil, result.Err + } + manifest, ok := result.Val.(*CodexModelsManifest) + if !ok || manifest == nil { + return nil, infraerrors.New(http.StatusInternalServerError, "OPENAI_CODEX_MODELS_REQUEST_FAILED", "invalid shared Codex models manifest result") + } + return codexModelsManifestForClient(manifest, ifNoneMatch), nil + } +} + +func (s *OpenAIGatewayService) refreshCachedAPIKeyCodexModelsManifest(cacheKey string, request codexModelsManifestRequest) <-chan singleflight.Result { + return s.codexModelsManifestCache.refresh.DoChan(cacheKey, func() (any, error) { + cached, _ := s.codexModelsManifestCache.get(cacheKey, time.Now()) + ifNoneMatch := "" + if cached != nil { + ifNoneMatch = cached.ETag + } + manifest, err := s.fetchCodexModelsManifestUpstream(context.Background(), request, ifNoneMatch) + if err != nil { + return nil, err + } + if manifest.NotModified && cached != nil { + s.codexModelsManifestCache.set(cacheKey, cached, time.Now()) + return cached, nil + } + if !manifest.NotModified { + s.codexModelsManifestCache.set(cacheKey, manifest, time.Now()) + } + return manifest, nil }) +} + +func (s *OpenAIGatewayService) fetchCodexModelsManifestUpstream(ctx context.Context, request codexModelsManifestRequest, ifNoneMatch string) (*CodexModelsManifest, error) { + reqCtx, cancel := context.WithTimeout(ctx, codexModelsManifestRequestTimeout) + defer cancel() + req, err := http.NewRequestWithContext(reqCtx, http.MethodGet, request.url, nil) if err != nil { - return nil, infraerrors.Newf(http.StatusInternalServerError, "OPENAI_CODEX_MODELS_PROXY_INVALID", "invalid proxy configuration: %v", err) + return nil, infraerrors.Newf(http.StatusInternalServerError, "OPENAI_CODEX_MODELS_REQUEST_FAILED", "create codex models request: %v", err) + } + req.Header = request.headers.Clone() + if ifNoneMatch = strings.TrimSpace(ifNoneMatch); ifNoneMatch != "" { + req.Header.Set("If-None-Match", ifNoneMatch) } - resp, err := client.Do(req) + var resp *http.Response + if request.useAPIKeyUpstream { + if s.httpUpstream == nil { + return nil, infraerrors.New(http.StatusInternalServerError, "OPENAI_CODEX_MODELS_UPSTREAM_NOT_CONFIGURED", "Codex models upstream HTTP client is not configured") + } + req = req.WithContext(WithHTTPUpstreamProfile(req.Context(), HTTPUpstreamProfileOpenAI)) + resp, err = s.httpUpstream.Do(req, request.proxyURL, request.accountID, request.accountConcurrency) + } else { + client, clientErr := httpclient.GetClient(httpclient.Options{ + ProxyURL: request.proxyURL, + Timeout: codexModelsManifestRequestTimeout, + ResponseHeaderTimeout: 10 * time.Second, + }) + if clientErr != nil { + return nil, infraerrors.Newf(http.StatusInternalServerError, "OPENAI_CODEX_MODELS_PROXY_INVALID", "invalid proxy configuration: %v", clientErr) + } + resp, err = client.Do(req) + } if err != nil { - return nil, infraerrors.Newf(http.StatusBadGateway, "OPENAI_CODEX_MODELS_UPSTREAM_FAILED", "codex models manifest request failed: %v", err) + return nil, &codexModelsManifestUpstreamError{ + err: infraerrors.Newf(http.StatusBadGateway, "OPENAI_CODEX_MODELS_UPSTREAM_FAILED", "codex models manifest request failed: %v", err), + retryable: isRetryableCodexModelsManifestTransportError(err), + } } defer func() { _ = resp.Body.Close() }() @@ -96,12 +405,100 @@ func (s *OpenAIGatewayService) FetchCodexModelsManifest(ctx context.Context, acc if message == "" { message = resp.Status } - return nil, infraerrors.Newf(http.StatusBadGateway, "OPENAI_CODEX_MODELS_UPSTREAM_FAILED", "codex models manifest upstream error %d: %s", resp.StatusCode, message) + return nil, &codexModelsManifestUpstreamError{ + err: infraerrors.Newf(http.StatusBadGateway, "OPENAI_CODEX_MODELS_UPSTREAM_FAILED", "codex models manifest upstream error %d: %s", resp.StatusCode, message), + retryable: resp.StatusCode == http.StatusTooManyRequests || + (resp.StatusCode >= http.StatusInternalServerError && resp.StatusCode < 600), + } } body, err := io.ReadAll(io.LimitReader(resp.Body, codexModelsManifestBodyLimit)) if err != nil { - return nil, infraerrors.Newf(http.StatusBadGateway, "OPENAI_CODEX_MODELS_UPSTREAM_FAILED", "read codex models manifest response: %v", err) + return nil, &codexModelsManifestUpstreamError{ + err: infraerrors.Newf(http.StatusBadGateway, "OPENAI_CODEX_MODELS_UPSTREAM_FAILED", "read codex models manifest response: %v", err), + retryable: isRetryableCodexModelsManifestTransportError(err), + } } return &CodexModelsManifest{Body: body, ETag: resp.Header.Get("ETag")}, nil } + +func buildCodexModelsManifestCacheKey(request codexModelsManifestRequest) string { + hasher := sha256.New() + _, _ = fmt.Fprintf(hasher, "%d\n%d\n%s\n%s\n", request.accountID, request.credentialAccountID, request.proxyURL, request.url) + headerNames := make([]string, 0, len(request.headers)) + for name := range request.headers { + headerNames = append(headerNames, name) + } + sort.Strings(headerNames) + for _, name := range headerNames { + _, _ = fmt.Fprintf(hasher, "%s\n", strings.ToLower(name)) + for _, value := range request.headers[name] { + _, _ = fmt.Fprintf(hasher, "%s\n", value) + } + } + return fmt.Sprintf("%x", hasher.Sum(nil)) +} + +func codexModelsManifestForClient(manifest *CodexModelsManifest, ifNoneMatch string) *CodexModelsManifest { + if manifest == nil { + return nil + } + if codexModelsManifestETagMatches(ifNoneMatch, manifest.ETag) { + return &CodexModelsManifest{ETag: manifest.ETag, NotModified: true} + } + return manifest +} + +func codexModelsManifestETagMatches(ifNoneMatch, etag string) bool { + etag = strings.TrimSpace(etag) + if etag == "" { + return false + } + normalize := func(value string) string { + value = strings.TrimSpace(value) + if len(value) >= 2 && strings.EqualFold(value[:2], "W/") { + value = strings.TrimSpace(value[2:]) + } + return value + } + want := normalize(etag) + for _, candidate := range strings.Split(ifNoneMatch, ",") { + candidate = strings.TrimSpace(candidate) + if candidate == "*" || normalize(candidate) == want { + return true + } + } + return false +} + +func isOfficialOpenAIModelsBaseURL(raw string) bool { + parsed, err := url.Parse(strings.TrimSpace(raw)) + if err != nil { + return false + } + hostname := strings.TrimSuffix(parsed.Hostname(), ".") + return strings.EqualFold(hostname, "api.openai.com") +} + +func buildCodexModelsManifestURL(endpoint string, appendModelsPath bool, clientVersion string) (*url.URL, error) { + requestURL, err := url.Parse(endpoint) + if err != nil { + return nil, err + } + if requestURL.Fragment != "" { + return nil, fmt.Errorf("URL fragments are not supported") + } + + query := requestURL.Query() + requestURL.RawQuery = "" + requestURL.ForceQuery = false + if appendModelsPath { + requestURL, err = url.Parse(buildOpenAIModelsURL(requestURL.String())) + if err != nil { + return nil, err + } + } + query.Set("client_version", clientVersion) + requestURL.RawQuery = query.Encode() + return requestURL, nil +} diff --git a/backend/internal/service/openai_codex_models_service_test.go b/backend/internal/service/openai_codex_models_service_test.go index c9eae35629..2c87861f60 100644 --- a/backend/internal/service/openai_codex_models_service_test.go +++ b/backend/internal/service/openai_codex_models_service_test.go @@ -2,11 +2,146 @@ package service import ( "context" + "errors" + "io" + "net" "net/http" "net/http/httptest" + "net/url" + "strings" + "sync" + "sync/atomic" "testing" + "time" + + "github.com/Wei-Shaw/sub2api/internal/config" + infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" + "github.com/Wei-Shaw/sub2api/internal/pkg/tlsfingerprint" + "golang.org/x/net/http2" ) +type codexModelsHTTPUpstreamStub struct { + do func(req *http.Request, proxyURL string, accountID int64, accountConcurrency int) (*http.Response, error) +} + +type codexModelsBlockingBody struct { + ctx context.Context + readStarted chan struct{} + startedOnce *sync.Once + release <-chan struct{} + body *strings.Reader +} + +func (b *codexModelsBlockingBody) Read(p []byte) (int, error) { + b.startedOnce.Do(func() { close(b.readStarted) }) + select { + case <-b.release: + return b.body.Read(p) + case <-b.ctx.Done(): + return 0, b.ctx.Err() + } +} + +func (b *codexModelsBlockingBody) Close() error { return nil } + +func (s *codexModelsHTTPUpstreamStub) Do(req *http.Request, proxyURL string, accountID int64, accountConcurrency int) (*http.Response, error) { + return s.do(req, proxyURL, accountID, accountConcurrency) +} + +func (s *codexModelsHTTPUpstreamStub) DoWithTLS(req *http.Request, proxyURL string, accountID int64, accountConcurrency int, _ *tlsfingerprint.Profile) (*http.Response, error) { + return s.Do(req, proxyURL, accountID, accountConcurrency) +} + +func TestIsRetryableCodexModelsManifestTransportError(t *testing.T) { + tests := []struct { + name string + err error + retryable bool + }{ + {name: "nil", err: nil}, + {name: "configuration error", err: errors.New("invalid proxy URL")}, + {name: "upstream configuration error", err: errors.New("upstream error: invalid proxy")}, + {name: "proxy connection configuration error", err: errors.New("proxy connection error: invalid configuration")}, + {name: "canceled request", err: context.Canceled}, + { + name: "redirect policy error", + err: &url.Error{ + Op: "Get", + URL: "https://upstream.example/v1/models", + Err: errors.New("stopped after 10 redirects"), + }, + }, + {name: "deadline exceeded", err: context.DeadlineExceeded, retryable: true}, + {name: "unexpected EOF", err: io.ErrUnexpectedEOF, retryable: true}, + {name: "closed connection", err: net.ErrClosed, retryable: true}, + { + name: "network operation", + err: &net.OpError{ + Op: "read", + Net: "tcp", + Err: errors.New("connection reset"), + }, + retryable: true, + }, + { + name: "DNS error", + err: &net.DNSError{Err: "temporary failure", Name: "upstream.example"}, + retryable: true, + }, + { + name: "typed HTTP2 GOAWAY", + err: http2.GoAwayError{ErrCode: http2.ErrCodeNo}, + retryable: true, + }, + { + name: "stdlib HTTP2 GOAWAY", + err: errors.New("http2: server sent GOAWAY and closed the connection; LastStreamID=1, ErrCode=NO_ERROR"), + retryable: true, + }, + { + name: "stdlib HTTP2 refused stream", + err: errors.New("stream error: stream ID 3; REFUSED_STREAM"), + retryable: true, + }, + { + name: "stdlib HTTP2 connection error", + err: errors.New(`Get "https://upstream.example/v1/models": connection error: PROTOCOL_ERROR`), + retryable: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if got := isRetryableCodexModelsManifestTransportError(tt.err); got != tt.retryable { + t.Fatalf("retryable = %v, want %v", got, tt.retryable) + } + }) + } +} + +func newCodexModelsAPIKeyTestService(upstream HTTPUpstream) *OpenAIGatewayService { + return &OpenAIGatewayService{ + cfg: &config.Config{Security: config.SecurityConfig{URLAllowlist: config.URLAllowlistConfig{ + Enabled: false, + }}}, + httpUpstream: upstream, + } +} + +func newCodexModelsAPIKeyTestAccount(baseURL string) *Account { + credentials := map[string]any{"api_key": "sk-upstream"} + if baseURL != "" { + credentials["base_url"] = baseURL + } + return &Account{ + ID: 2, + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Credentials: credentials, + Concurrency: 3, + } +} + func newCodexModelsTestAccount() *Account { return &Account{ ID: 1, @@ -136,3 +271,679 @@ func TestFetchCodexModelsManifestMissingToken(t *testing.T) { t.Fatal("expected error for missing access token, got nil") } } + +func TestFetchCodexModelsManifestAPIKeyCustomUpstream(t *testing.T) { + manifestBody := `{"models":[{"slug":"gpt-5.6"}]}` + var gotRequest *http.Request + var gotProxyURL string + var gotAccountID int64 + var gotConcurrency int + upstream := &codexModelsHTTPUpstreamStub{do: func(req *http.Request, proxyURL string, accountID int64, accountConcurrency int) (*http.Response, error) { + gotRequest = req + gotProxyURL = proxyURL + gotAccountID = accountID + gotConcurrency = accountConcurrency + header := make(http.Header) + header.Set("ETag", `W/"api-key-manifest"`) + return &http.Response{ + StatusCode: http.StatusOK, + Header: header, + Body: io.NopCloser(strings.NewReader(manifestBody)), + }, nil + }} + + s := newCodexModelsAPIKeyTestService(upstream) + manifest, err := s.FetchCodexModelsManifest( + context.Background(), + newCodexModelsAPIKeyTestAccount("https://upstream.example/v1"), + "0.144.0", + "", + ) + if err != nil { + t.Fatalf("FetchCodexModelsManifest returned error: %v", err) + } + + if gotRequest == nil { + t.Fatal("expected request to custom API key upstream") + } + if gotRequest.Method != http.MethodGet { + t.Errorf("method: got %q", gotRequest.Method) + } + if gotRequest.URL.String() != "https://upstream.example/v1/models?client_version=0.144.0" { + t.Errorf("request URL: got %q", gotRequest.URL.String()) + } + if gotRequest.Header.Get("Authorization") != "Bearer sk-upstream" { + t.Errorf("authorization header: got %q", gotRequest.Header.Get("Authorization")) + } + if gotRequest.Header.Get("Originator") != "codex_cli_rs" { + t.Errorf("originator header: got %q", gotRequest.Header.Get("Originator")) + } + if gotRequest.Header.Get("Version") != "0.144.0" { + t.Errorf("version header: got %q", gotRequest.Header.Get("Version")) + } + if gotRequest.Header.Get("User-Agent") != codexCLIUserAgent { + t.Errorf("user-agent header: got %q", gotRequest.Header.Get("User-Agent")) + } + if gotRequest.Header.Get("chatgpt-account-id") != "" { + t.Errorf("chatgpt-account-id must not be sent to API key upstream: got %q", gotRequest.Header.Get("chatgpt-account-id")) + } + if gotProxyURL != "" || gotAccountID != 2 || gotConcurrency != 3 { + t.Errorf("upstream routing metadata: proxy=%q account_id=%d concurrency=%d", gotProxyURL, gotAccountID, gotConcurrency) + } + if string(manifest.Body) != manifestBody { + t.Errorf("body not passed through verbatim: got %q", manifest.Body) + } + if manifest.ETag != `W/"api-key-manifest"` { + t.Errorf("etag not passed through: got %q", manifest.ETag) + } +} + +func TestFetchCodexModelsManifestAPIKeySharedRefreshSurvivesCallerCancellation(t *testing.T) { + const manifestBody = `{"models":[{"slug":"gpt-5.6"}]}` + var calls atomic.Int32 + var readStartedOnce sync.Once + readStarted := make(chan struct{}) + deadlineRemaining := make(chan time.Duration, 1) + release := make(chan struct{}) + upstream := &codexModelsHTTPUpstreamStub{do: func(req *http.Request, _ string, _ int64, _ int) (*http.Response, error) { + calls.Add(1) + deadline, ok := req.Context().Deadline() + if !ok { + deadlineRemaining <- 0 + } else { + deadlineRemaining <- time.Until(deadline) + } + return &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Etag": []string{`W/"shared"`}}, + Body: &codexModelsBlockingBody{ + ctx: req.Context(), + readStarted: readStarted, + startedOnce: &readStartedOnce, + release: release, + body: strings.NewReader(manifestBody), + }, + }, nil + }} + + s := newCodexModelsAPIKeyTestService(upstream) + account := newCodexModelsAPIKeyTestAccount("https://upstream.example") + firstCtx, cancelFirst := context.WithCancel(context.Background()) + firstErr := make(chan error, 1) + go func() { + _, err := s.FetchCodexModelsManifest(firstCtx, account, "0.144.0", "") + firstErr <- err + }() + + select { + case <-readStarted: + case <-time.After(time.Second): + t.Fatal("upstream body read did not start") + } + remaining := <-deadlineRemaining + if remaining < 14*time.Second || remaining > codexModelsManifestRequestTimeout { + t.Errorf("detached refresh deadline: got %s, want approximately %s", remaining, codexModelsManifestRequestTimeout) + } + cancelFirst() + select { + case err := <-firstErr: + if !errors.Is(err, context.Canceled) { + t.Fatalf("first caller error: got %v, want context.Canceled", err) + } + case <-time.After(time.Second): + t.Fatal("canceled caller did not return promptly") + } + + secondResult := make(chan struct { + manifest *CodexModelsManifest + err error + }, 1) + go func() { + manifest, err := s.FetchCodexModelsManifest(context.Background(), account, "0.144.0", "") + secondResult <- struct { + manifest *CodexModelsManifest + err error + }{manifest: manifest, err: err} + }() + + time.Sleep(50 * time.Millisecond) + if got := calls.Load(); got != 1 { + t.Errorf("upstream calls before shared refresh completed: got %d, want 1", got) + } + close(release) + select { + case result := <-secondResult: + if result.err != nil { + t.Fatalf("second caller returned error: %v", result.err) + } + if string(result.manifest.Body) != manifestBody { + t.Errorf("second caller body: got %q", result.manifest.Body) + } + case <-time.After(time.Second): + t.Fatal("second caller did not receive shared refresh result") + } + if got := calls.Load(); got != 1 { + t.Errorf("total upstream calls: got %d, want 1", got) + } +} + +func TestFetchCodexModelsManifestAPIKeyConcurrentRequestsShareRefresh(t *testing.T) { + const callers = 8 + var calls atomic.Int32 + started := make(chan struct{}) + var startedOnce sync.Once + release := make(chan struct{}) + upstream := &codexModelsHTTPUpstreamStub{do: func(_ *http.Request, _ string, _ int64, _ int) (*http.Response, error) { + calls.Add(1) + startedOnce.Do(func() { close(started) }) + <-release + return &http.Response{ + StatusCode: http.StatusOK, + Header: make(http.Header), + Body: io.NopCloser(strings.NewReader(`{"models":[]}`)), + }, nil + }} + + s := newCodexModelsAPIKeyTestService(upstream) + account := newCodexModelsAPIKeyTestAccount("https://upstream.example") + begin := make(chan struct{}) + errs := make(chan error, callers) + for i := 0; i < callers; i++ { + go func() { + <-begin + _, err := s.FetchCodexModelsManifest(context.Background(), account, "0.144.0", "") + errs <- err + }() + } + close(begin) + select { + case <-started: + case <-time.After(time.Second): + t.Fatal("upstream request did not start") + } + time.Sleep(50 * time.Millisecond) + if got := calls.Load(); got != 1 { + t.Errorf("concurrent upstream calls: got %d, want 1", got) + } + close(release) + for i := 0; i < callers; i++ { + if err := <-errs; err != nil { + t.Errorf("caller %d returned error: %v", i, err) + } + } +} + +func TestFetchCodexModelsManifestAPIKeyFreshCacheHandlesETagLocally(t *testing.T) { + var calls atomic.Int32 + upstream := &codexModelsHTTPUpstreamStub{do: func(req *http.Request, _ string, _ int64, _ int) (*http.Response, error) { + calls.Add(1) + if got := req.Header.Get("If-None-Match"); got != "" { + t.Errorf("cache refresh must not inherit a caller's If-None-Match: got %q", got) + } + return &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Etag": []string{`W/"cached"`}}, + Body: io.NopCloser(strings.NewReader(`{"models":[]}`)), + }, nil + }} + + s := newCodexModelsAPIKeyTestService(upstream) + account := newCodexModelsAPIKeyTestAccount("https://upstream.example") + if _, err := s.FetchCodexModelsManifest(context.Background(), account, "0.144.0", ""); err != nil { + t.Fatalf("initial fetch returned error: %v", err) + } + manifest, err := s.FetchCodexModelsManifest(context.Background(), account, "0.144.0", `W/"cached"`) + if err != nil { + t.Fatalf("cached fetch returned error: %v", err) + } + if !manifest.NotModified { + t.Fatal("matching cached ETag must return NotModified") + } + if got := calls.Load(); got != 1 { + t.Errorf("upstream calls: got %d, want 1", got) + } +} + +func TestFetchCodexModelsManifestAPIKeyCacheKeyIsolatesRequestIdentity(t *testing.T) { + var calls atomic.Int32 + upstream := &codexModelsHTTPUpstreamStub{do: func(_ *http.Request, _ string, _ int64, _ int) (*http.Response, error) { + calls.Add(1) + return &http.Response{ + StatusCode: http.StatusOK, + Header: make(http.Header), + Body: io.NopCloser(strings.NewReader(`{"models":[]}`)), + }, nil + }} + s := newCodexModelsAPIKeyTestService(upstream) + + base := newCodexModelsAPIKeyTestAccount("https://upstream.example") + fetch := func(account *Account, version string) { + t.Helper() + if _, err := s.FetchCodexModelsManifest(context.Background(), account, version, ""); err != nil { + t.Fatalf("fetch returned error: %v", err) + } + } + fetch(base, "0.144.0") + fetch(base, "0.144.0") + + differentAccount := newCodexModelsAPIKeyTestAccount("https://upstream.example") + differentAccount.ID = 3 + fetch(differentAccount, "0.144.0") + + differentToken := newCodexModelsAPIKeyTestAccount("https://upstream.example") + differentToken.Credentials["api_key"] = "sk-other" + fetch(differentToken, "0.144.0") + + differentUpstream := newCodexModelsAPIKeyTestAccount("https://other-upstream.example") + fetch(differentUpstream, "0.144.0") + fetch(base, "0.145.0") + + differentHeaders := newCodexModelsAPIKeyTestAccount("https://upstream.example") + differentHeaders.Credentials[credKeyHeaderOverrideEnabled] = true + differentHeaders.Credentials[credKeyHeaderOverrides] = map[string]any{"x-tenant": "other"} + fetch(differentHeaders, "0.144.0") + + proxyID := int64(9) + differentProxy := newCodexModelsAPIKeyTestAccount("https://upstream.example") + differentProxy.ProxyID = &proxyID + differentProxy.Proxy = &Proxy{Protocol: "http", Host: "127.0.0.1", Port: 8080} + fetch(differentProxy, "0.144.0") + fetch(differentProxy, "0.144.0") + + if got := calls.Load(); got != 7 { + t.Errorf("isolated upstream calls: got %d, want 7", got) + } +} + +func TestFetchCodexModelsManifestAPIKeyCacheBoundsEntriesAndBodySize(t *testing.T) { + var calls atomic.Int32 + upstream := &codexModelsHTTPUpstreamStub{do: func(req *http.Request, _ string, _ int64, _ int) (*http.Response, error) { + calls.Add(1) + body := `{"models":[]}` + if strings.Contains(req.URL.Host, "large") { + body = strings.Repeat("x", (1<<20)+1) + } + return &http.Response{ + StatusCode: http.StatusOK, + Header: make(http.Header), + Body: io.NopCloser(strings.NewReader(body)), + }, nil + }} + s := newCodexModelsAPIKeyTestService(upstream) + fetch := func(account *Account) { + t.Helper() + if _, err := s.FetchCodexModelsManifest(context.Background(), account, "0.144.0", ""); err != nil { + t.Fatalf("fetch returned error: %v", err) + } + } + + small := newCodexModelsAPIKeyTestAccount("https://small.example") + fetch(small) + fetch(small) + large := newCodexModelsAPIKeyTestAccount("https://large.example") + large.ID = 3 + fetch(large) + fetch(large) + if got := calls.Load(); got != 3 { + t.Fatalf("body-size bounded cache calls: got %d, want 3", got) + } + + for i := int64(10); i < 75; i++ { + account := newCodexModelsAPIKeyTestAccount("https://bounded.example") + account.ID = i + fetch(account) + } + last := newCodexModelsAPIKeyTestAccount("https://bounded.example") + last.ID = 74 + fetch(last) + if got := calls.Load(); got != 68 { + t.Fatalf("most recent cache entry was not retained: calls=%d, want 68", got) + } + first := newCodexModelsAPIKeyTestAccount("https://bounded.example") + first.ID = 10 + fetch(first) + if got := calls.Load(); got != 69 { + t.Errorf("oldest cache entry was not evicted: calls=%d, want 69", got) + } +} + +func TestFetchCodexModelsManifestAPIKeyServesStaleWhileRefreshing(t *testing.T) { + var calls atomic.Int32 + refreshStarted := make(chan struct{}) + releaseRefresh := make(chan struct{}) + upstream := &codexModelsHTTPUpstreamStub{do: func(_ *http.Request, _ string, _ int64, _ int) (*http.Response, error) { + call := calls.Add(1) + body := `{"models":[{"slug":"old"}]}` + if call > 1 { + if call == 2 { + close(refreshStarted) + } + <-releaseRefresh + body = `{"models":[{"slug":"new"}]}` + } + return &http.Response{ + StatusCode: http.StatusOK, + Header: make(http.Header), + Body: io.NopCloser(strings.NewReader(body)), + }, nil + }} + s := newCodexModelsAPIKeyTestService(upstream) + account := newCodexModelsAPIKeyTestAccount("https://upstream.example") + if _, err := s.FetchCodexModelsManifest(context.Background(), account, "0.144.0", ""); err != nil { + t.Fatalf("initial fetch returned error: %v", err) + } + + s.codexModelsManifestCache.mu.Lock() + for key, entry := range s.codexModelsManifestCache.entries { + entry.expiresAt = time.Now().Add(-time.Second) + s.codexModelsManifestCache.entries[key] = entry + } + s.codexModelsManifestCache.mu.Unlock() + + resultCh := make(chan struct { + manifest *CodexModelsManifest + err error + }, 1) + go func() { + manifest, err := s.FetchCodexModelsManifest(context.Background(), account, "0.144.0", "") + resultCh <- struct { + manifest *CodexModelsManifest + err error + }{manifest: manifest, err: err} + }() + select { + case <-refreshStarted: + case <-time.After(time.Second): + t.Fatal("background refresh did not start") + } + + var staleResult struct { + manifest *CodexModelsManifest + err error + } + select { + case staleResult = <-resultCh: + case <-time.After(100 * time.Millisecond): + t.Error("stale manifest was not returned while refresh was blocked") + close(releaseRefresh) + staleResult = <-resultCh + } + if staleResult.err != nil { + t.Fatalf("stale fetch returned error: %v", staleResult.err) + } + if got := string(staleResult.manifest.Body); got != `{"models":[{"slug":"old"}]}` { + t.Errorf("stale body: got %q", got) + } + if got := calls.Load(); got != 2 { + t.Errorf("upstream calls during stale refresh: got %d, want 2", got) + } + + select { + case <-releaseRefresh: + default: + close(releaseRefresh) + } + deadline := time.Now().Add(time.Second) + for { + manifest, err := s.FetchCodexModelsManifest(context.Background(), account, "0.144.0", "") + if err == nil && string(manifest.Body) == `{"models":[{"slug":"new"}]}` { + break + } + if time.Now().After(deadline) { + t.Fatalf("refreshed manifest was not cached: manifest=%v err=%v", manifest, err) + } + time.Sleep(10 * time.Millisecond) + } + if got := calls.Load(); got != 2 { + t.Errorf("stale refresh was not deduplicated: calls=%d, want 2", got) + } +} + +func TestFetchCodexModelsManifestAPIKeyRevalidatesStaleETag(t *testing.T) { + var calls atomic.Int32 + refreshDone := make(chan struct{}) + upstream := &codexModelsHTTPUpstreamStub{do: func(req *http.Request, _ string, _ int64, _ int) (*http.Response, error) { + call := calls.Add(1) + if call == 1 { + header := make(http.Header) + header.Set("ETag", `W/"cached"`) + return &http.Response{ + StatusCode: http.StatusOK, + Header: header, + Body: io.NopCloser(strings.NewReader(`{"models":[{"slug":"cached"}]}`)), + }, nil + } + if got := req.Header.Get("If-None-Match"); got != `W/"cached"` { + t.Errorf("background revalidation If-None-Match: got %q", got) + } + close(refreshDone) + header := make(http.Header) + header.Set("ETag", `W/"cached"`) + return &http.Response{StatusCode: http.StatusNotModified, Header: header, Body: http.NoBody}, nil + }} + s := newCodexModelsAPIKeyTestService(upstream) + account := newCodexModelsAPIKeyTestAccount("https://upstream.example") + if _, err := s.FetchCodexModelsManifest(context.Background(), account, "0.144.0", ""); err != nil { + t.Fatalf("initial fetch returned error: %v", err) + } + s.codexModelsManifestCache.mu.Lock() + for key, entry := range s.codexModelsManifestCache.entries { + entry.expiresAt = time.Now().Add(-time.Second) + s.codexModelsManifestCache.entries[key] = entry + } + s.codexModelsManifestCache.mu.Unlock() + + manifest, err := s.FetchCodexModelsManifest(context.Background(), account, "0.144.0", "") + if err != nil { + t.Fatalf("stale fetch returned error: %v", err) + } + if got := string(manifest.Body); got != `{"models":[{"slug":"cached"}]}` { + t.Fatalf("stale body: got %q", got) + } + select { + case <-refreshDone: + case <-time.After(time.Second): + t.Fatal("ETag revalidation did not complete") + } + + deadline := time.Now().Add(time.Second) + for { + s.codexModelsManifestCache.mu.Lock() + fresh := false + for _, entry := range s.codexModelsManifestCache.entries { + fresh = time.Now().Before(entry.expiresAt) + } + s.codexModelsManifestCache.mu.Unlock() + if fresh { + break + } + if time.Now().After(deadline) { + t.Fatal("304 revalidation did not renew the cached manifest") + } + time.Sleep(10 * time.Millisecond) + } + manifest, err = s.FetchCodexModelsManifest(context.Background(), account, "0.144.0", "") + if err != nil || string(manifest.Body) != `{"models":[{"slug":"cached"}]}` { + t.Fatalf("renewed cached manifest: body=%q err=%v", manifest.Body, err) + } + if got := calls.Load(); got != 2 { + t.Errorf("upstream calls: got %d, want 2", got) + } +} + +func TestFetchCodexModelsManifestAPIKeyColdCacheHandlesNotModifiedLocally(t *testing.T) { + var gotIfNoneMatch string + upstream := &codexModelsHTTPUpstreamStub{do: func(req *http.Request, _ string, _ int64, _ int) (*http.Response, error) { + gotIfNoneMatch = req.Header.Get("If-None-Match") + header := make(http.Header) + header.Set("ETag", `W/"api-key-manifest"`) + return &http.Response{ + StatusCode: http.StatusOK, + Header: header, + Body: io.NopCloser(strings.NewReader(`{"models":[]}`)), + }, nil + }} + + s := newCodexModelsAPIKeyTestService(upstream) + manifest, err := s.FetchCodexModelsManifest( + context.Background(), + newCodexModelsAPIKeyTestAccount("https://upstream.example"), + "0.144.0", + `W/"api-key-manifest"`, + ) + if err != nil { + t.Fatalf("FetchCodexModelsManifest returned error: %v", err) + } + if !manifest.NotModified { + t.Error("expected NotModified to be true") + } + if manifest.ETag != `W/"api-key-manifest"` { + t.Errorf("etag not passed through: got %q", manifest.ETag) + } + if gotIfNoneMatch != "" { + t.Errorf("cold shared refresh must not inherit caller if-none-match: got %q", gotIfNoneMatch) + } +} + +func TestFetchCodexModelsManifestAPIKeyDoesNotCacheUnexpectedColdNotModified(t *testing.T) { + var calls atomic.Int32 + upstream := &codexModelsHTTPUpstreamStub{do: func(req *http.Request, _ string, _ int64, _ int) (*http.Response, error) { + calls.Add(1) + if got := req.Header.Get("If-None-Match"); got != "" { + t.Errorf("cold shared refresh If-None-Match: got %q", got) + } + header := make(http.Header) + header.Set("ETag", `W/"unexpected"`) + return &http.Response{StatusCode: http.StatusNotModified, Header: header, Body: http.NoBody}, nil + }} + s := newCodexModelsAPIKeyTestService(upstream) + account := newCodexModelsAPIKeyTestAccount("https://upstream.example") + for i := 0; i < 2; i++ { + manifest, err := s.FetchCodexModelsManifest(context.Background(), account, "0.144.0", "") + if err != nil { + t.Fatalf("fetch %d returned error: %v", i, err) + } + if !manifest.NotModified { + t.Fatalf("fetch %d: expected upstream NotModified response", i) + } + } + if got := calls.Load(); got != 2 { + t.Errorf("unexpected cold 304 was cached: upstream calls=%d, want 2", got) + } +} + +func TestFetchCodexModelsManifestAPIKeyPreservesBaseURLQuery(t *testing.T) { + var gotURL string + upstream := &codexModelsHTTPUpstreamStub{do: func(req *http.Request, _ string, _ int64, _ int) (*http.Response, error) { + gotURL = req.URL.String() + return &http.Response{ + StatusCode: http.StatusOK, + Header: make(http.Header), + Body: io.NopCloser(strings.NewReader(`{"models":[]}`)), + }, nil + }} + + s := newCodexModelsAPIKeyTestService(upstream) + _, err := s.FetchCodexModelsManifest( + context.Background(), + newCodexModelsAPIKeyTestAccount("https://upstream.example/v1?tenant=acme"), + "0.144.0", + "", + ) + if err != nil { + t.Fatalf("FetchCodexModelsManifest returned error: %v", err) + } + if gotURL != "https://upstream.example/v1/models?client_version=0.144.0&tenant=acme" { + t.Errorf("request URL: got %q", gotURL) + } +} + +func TestFetchCodexModelsManifestAPIKeyRejectsBaseURLFragment(t *testing.T) { + called := false + upstream := &codexModelsHTTPUpstreamStub{do: func(_ *http.Request, _ string, _ int64, _ int) (*http.Response, error) { + called = true + return &http.Response{ + StatusCode: http.StatusOK, + Header: make(http.Header), + Body: io.NopCloser(strings.NewReader(`{"models":[]}`)), + }, nil + }} + + s := newCodexModelsAPIKeyTestService(upstream) + _, err := s.FetchCodexModelsManifest( + context.Background(), + newCodexModelsAPIKeyTestAccount("https://upstream.example/v1#models"), + "0.144.0", + "", + ) + if err == nil { + t.Fatal("expected invalid upstream base URL error, got nil") + } + if infraerrors.Reason(err) != "OPENAI_CODEX_MODELS_API_KEY_UPSTREAM_INVALID" { + t.Errorf("error reason: got %q", infraerrors.Reason(err)) + } + if called { + t.Fatal("fragment-bearing base URL must be rejected before the upstream request") + } +} + +func TestFetchCodexModelsManifestAPIKeyUpstreamError(t *testing.T) { + upstream := &codexModelsHTTPUpstreamStub{do: func(_ *http.Request, _ string, _ int64, _ int) (*http.Response, error) { + return &http.Response{ + StatusCode: http.StatusTooManyRequests, + Status: "429 Too Many Requests", + Header: make(http.Header), + Body: io.NopCloser(strings.NewReader(`{"error":"rate limited"}`)), + }, nil + }} + + s := newCodexModelsAPIKeyTestService(upstream) + _, err := s.FetchCodexModelsManifest( + context.Background(), + newCodexModelsAPIKeyTestAccount("https://upstream.example"), + "0.144.0", + "", + ) + if err == nil { + t.Fatal("expected error for upstream 429, got nil") + } + if infraerrors.Code(err) != http.StatusBadGateway { + t.Errorf("error status: got %d, want %d", infraerrors.Code(err), http.StatusBadGateway) + } + if infraerrors.Reason(err) != "OPENAI_CODEX_MODELS_UPSTREAM_FAILED" { + t.Errorf("error reason: got %q", infraerrors.Reason(err)) + } +} + +func TestFetchCodexModelsManifestAPIKeyRejectsOfficialOpenAIBaseURL(t *testing.T) { + tests := []struct { + name string + baseURL string + }{ + {name: "missing base URL"}, + {name: "official host", baseURL: "https://api.openai.com"}, + {name: "official versioned URL", baseURL: "https://API.OPENAI.COM:443/v1/"}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + s := newCodexModelsAPIKeyTestService(&codexModelsHTTPUpstreamStub{do: func(_ *http.Request, _ string, _ int64, _ int) (*http.Response, error) { + t.Fatal("official OpenAI API key must not be used as a Codex manifest upstream") + return nil, nil + }}) + + _, err := s.FetchCodexModelsManifest( + context.Background(), + newCodexModelsAPIKeyTestAccount(tt.baseURL), + "0.144.0", + "", + ) + if err == nil { + t.Fatal("expected unsupported API key upstream error, got nil") + } + if infraerrors.Reason(err) != "OPENAI_CODEX_MODELS_API_KEY_UPSTREAM_UNSUPPORTED" { + t.Errorf("error reason: got %q", infraerrors.Reason(err)) + } + }) + } +} diff --git a/backend/internal/service/openai_codex_transform.go b/backend/internal/service/openai_codex_transform.go index 3869e97c99..3bde4abcfb 100644 --- a/backend/internal/service/openai_codex_transform.go +++ b/backend/internal/service/openai_codex_transform.go @@ -838,6 +838,9 @@ func ensureOpenAIResponsesImageGenerationTool(reqBody map[string]any) bool { if isCodexSparkModel(firstNonEmptyString(reqBody["model"])) { return false } + if hasOpenAIImageGenerationTool(reqBody) { + return false + } tool := map[string]any{ "type": "image_generation", @@ -855,16 +858,6 @@ func ensureOpenAIResponsesImageGenerationTool(reqBody map[string]any) bool { reqBody["tools"] = []any{tool} return true } - for _, rawTool := range tools { - toolMap, ok := rawTool.(map[string]any) - if !ok { - continue - } - if strings.TrimSpace(firstNonEmptyString(toolMap["type"])) == "image_generation" { - return false - } - } - reqBody["tools"] = append(tools, tool) return true } diff --git a/backend/internal/service/openai_codex_transform_test.go b/backend/internal/service/openai_codex_transform_test.go index b226655eeb..456740136b 100644 --- a/backend/internal/service/openai_codex_transform_test.go +++ b/backend/internal/service/openai_codex_transform_test.go @@ -617,6 +617,65 @@ func TestEnsureOpenAIResponsesImageGenerationTool_PreservesExistingImageTool(t * require.Equal(t, "webp", tool["output_format"]) } +func TestEnsureOpenAIResponsesImageGenerationTool_PreservesImageGenNamespace(t *testing.T) { + tests := []struct { + name string + reqBody map[string]any + }{ + { + name: "top-level tools", + reqBody: map[string]any{ + "model": "gpt-5.5", + "tools": []any{ + map[string]any{ + "type": "namespace", + "name": "image_gen", + "tools": []any{ + map[string]any{"type": "function", "name": "imagegen"}, + }, + }, + }, + }, + }, + { + name: "responses lite additional_tools", + reqBody: map[string]any{ + "model": "gpt-5.5", + "input": []any{ + map[string]any{ + "type": "additional_tools", + "tools": []any{ + map[string]any{ + "type": "namespace", + "name": "image_gen", + "tools": []any{ + map[string]any{"type": "function", "name": "imagegen"}, + }, + }, + }, + }, + }, + }, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + require.True(t, hasOpenAIImageGenerationTool(tt.reqBody)) + + modified := ensureOpenAIResponsesImageGenerationTool(tt.reqBody) + + require.False(t, modified) + tools, _ := tt.reqBody["tools"].([]any) + for _, rawTool := range tools { + tool, ok := rawTool.(map[string]any) + require.True(t, ok) + require.NotEqual(t, "image_generation", firstNonEmptyString(tool["type"])) + } + }) + } +} + func TestApplyCodexImageGenerationBridgeInstructions_AppendsBridgeOnce(t *testing.T) { reqBody := map[string]any{ "model": "gpt-5.4", diff --git a/backend/internal/service/openai_compat_model_test.go b/backend/internal/service/openai_compat_model_test.go index 69b6ddbca2..e1007c507a 100644 --- a/backend/internal/service/openai_compat_model_test.go +++ b/backend/internal/service/openai_compat_model_test.go @@ -124,6 +124,55 @@ func TestApplyOpenAICompatModelNormalization(t *testing.T) { }) } +func TestForwardAsAnthropic_UsesExactFableMessagesDispatchModel(t *testing.T) { + t.Parallel() + gin.SetMode(gin.TestMode) + + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + body := []byte(`{"model":"claude-fable-5","max_tokens":16,"messages":[{"role":"user","content":"hello"}],"stream":false}`) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages", bytes.NewReader(body)) + c.Request.Header.Set("Content-Type", "application/json") + + upstreamBody := strings.Join([]string{ + `data: {"type":"response.completed","response":{"id":"resp_fable","object":"response","model":"gpt-5.6-sol","status":"completed","output":[{"type":"message","id":"msg_fable","role":"assistant","status":"completed","content":[{"type":"output_text","text":"ok"}]}],"usage":{"input_tokens":5,"output_tokens":2,"total_tokens":7}}}`, + "", + "data: [DONE]", + "", + }, "\n") + upstream := &httpUpstreamRecorder{resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"text/event-stream"}, "x-request-id": []string{"rid_fable"}}, + Body: io.NopCloser(strings.NewReader(upstreamBody)), + }} + + svc := &OpenAIGatewayService{ + httpUpstream: upstream, + cfg: &config.Config{Security: config.SecurityConfig{URLAllowlist: config.URLAllowlistConfig{Enabled: false}}}, + } + account := &Account{ + ID: 1, + Name: "openai-oauth", + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Concurrency: 1, + Credentials: map[string]any{ + "access_token": "oauth-token", + "chatgpt_account_id": "chatgpt-acc", + }, + } + + result, err := svc.ForwardAsAnthropic(context.Background(), c, account, body, "", "gpt-5.6-sol") + require.NoError(t, err) + require.NotNil(t, result) + require.Equal(t, "claude-fable-5", result.Model) + require.Equal(t, "gpt-5.6-sol", result.BillingModel) + require.Equal(t, "gpt-5.6-sol", result.UpstreamModel) + require.Equal(t, "gpt-5.6-sol", gjson.GetBytes(upstream.lastBody, "model").String()) + require.NotContains(t, string(upstream.lastBody), "claude-fable-5") + require.Equal(t, "claude-fable-5", gjson.GetBytes(rec.Body.Bytes(), "model").String()) +} + func TestForwardAsAnthropic_NormalizesRoutingAndEffortForGpt54XHigh(t *testing.T) { t.Parallel() gin.SetMode(gin.TestMode) diff --git a/backend/internal/service/openai_gateway_chat_completions_test.go b/backend/internal/service/openai_gateway_chat_completions_test.go index b85ee33947..5186598a70 100644 --- a/backend/internal/service/openai_gateway_chat_completions_test.go +++ b/backend/internal/service/openai_gateway_chat_completions_test.go @@ -98,7 +98,7 @@ func TestNormalizeResponsesBodyServiceTier(t *testing.T) { require.False(t, gjson.GetBytes(body, "service_tier").Exists()) } -func TestForwardAsChatCompletions_UnknownModelDoesNotUseDefaultMappedModel(t *testing.T) { +func TestForwardAsChatCompletions_UnknownModelWithoutMessagesDispatchKeepsRequestedModel(t *testing.T) { gin.SetMode(gin.TestMode) rec := httptest.NewRecorder() @@ -129,7 +129,7 @@ func TestForwardAsChatCompletions_UnknownModelDoesNotUseDefaultMappedModel(t *te }, } - result, err := svc.ForwardAsChatCompletions(context.Background(), c, account, body, "", "gpt-5.4") + result, err := svc.ForwardAsChatCompletions(context.Background(), c, account, body, "", "") require.Error(t, err) require.Nil(t, result) require.Equal(t, "gpt6", gjson.GetBytes(upstream.lastBody, "model").String()) diff --git a/backend/internal/service/openai_gateway_grok.go b/backend/internal/service/openai_gateway_grok.go index 379b586136..1906b83780 100644 --- a/backend/internal/service/openai_gateway_grok.go +++ b/backend/internal/service/openai_gateway_grok.go @@ -23,6 +23,7 @@ const ( grokComposerImageBridgeMaxOutputTokens = 512 grokUpstreamUserAgent = "sub2api-grok/1.0" grokCLIVersion = "0.2.93" + grokDefaultResponsesModel = "grok-4.5" grokRateLimitFallbackCooldown = 2 * time.Minute ) @@ -41,7 +42,7 @@ func (s *OpenAIGatewayService) forwardGrokResponses( upstreamModel := account.GetMappedModel(originalModel) if strings.TrimSpace(upstreamModel) == "" { - upstreamModel = "grok-4.3" + upstreamModel = grokDefaultResponsesModel } cacheIdentity := resolveGrokCacheIdentity(c, body, "", upstreamModel) patchedBody, err := patchGrokResponsesBody(body, upstreamModel) diff --git a/backend/internal/service/openai_gateway_grok_test.go b/backend/internal/service/openai_gateway_grok_test.go index a3edfacf9d..b45e0bccc0 100644 --- a/backend/internal/service/openai_gateway_grok_test.go +++ b/backend/internal/service/openai_gateway_grok_test.go @@ -853,12 +853,12 @@ func TestForwardAsChatCompletionsForGrokStopFallsBackToXAIChatCompletions(t *tes require.Equal(t, http.StatusOK, recorder.Code) } -func TestForwardGrokResponsesStreamingUsesXAIResponsesAndSnapshots(t *testing.T) { +func TestForwardGrokResponsesStreamingDefaultsEmptyModelTo45AndSnapshots(t *testing.T) { gin.SetMode(gin.TestMode) recorder := httptest.NewRecorder() c, _ := gin.CreateTestContext(recorder) - body := []byte(`{"model":"grok","input":"hi","stream":true,"reasoning_effort":"high"}`) + body := []byte(`{"input":"hi","stream":true,"reasoning_effort":"high"}`) c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", bytes.NewReader(body)) c.Request.Header.Set("Content-Type", "application/json") c.Request.Header.Set("OpenAI-Beta", "responses=experimental") @@ -905,7 +905,7 @@ func TestForwardGrokResponsesStreamingUsesXAIResponsesAndSnapshots(t *testing.T) accountRepo: repo, } - result, err := svc.forwardGrokResponses(context.Background(), c, account, body, "grok", true, time.Now()) + result, err := svc.forwardGrokResponses(context.Background(), c, account, body, "", true, time.Now()) require.NoError(t, err) require.Equal(t, xai.DefaultCLIBaseURL+"/responses", upstream.lastReq.URL.String()) require.Equal(t, "Bearer access-token", upstream.lastReq.Header.Get("Authorization")) diff --git a/backend/internal/service/openai_gateway_record_usage_test.go b/backend/internal/service/openai_gateway_record_usage_test.go index 81578f630f..d2eca05d74 100644 --- a/backend/internal/service/openai_gateway_record_usage_test.go +++ b/backend/internal/service/openai_gateway_record_usage_test.go @@ -39,6 +39,17 @@ type openAIRecordUsageBillingRepoStub struct { lastCtxErr error } +type openAIRecordUsageAccountRepoStub struct { + AccountRepository + account *Account + calls int +} + +func (s *openAIRecordUsageAccountRepoStub) GetByID(_ context.Context, _ int64) (*Account, error) { + s.calls++ + return s.account, nil +} + func (s *openAIRecordUsageBillingRepoStub) Apply(ctx context.Context, cmd *UsageBillingCommand) (*UsageBillingApplyResult, error) { s.calls++ s.lastCmd = cmd @@ -1045,7 +1056,7 @@ func TestOpenAIGatewayServiceRecordUsage_GPT56SeparatesCacheWriteForBillingAndSt require.InDelta(t, usageRepo.lastLog.TotalCost*1.1, usageRepo.lastLog.ActualCost, 1e-12) } -func TestOpenAIGatewayServiceRecordUsage_Gpt54LongContextBillsWholeSession(t *testing.T) { +func TestOpenAIGatewayServiceRecordUsage_Gpt54LongContextBillingDisabledByDefault(t *testing.T) { usageRepo := &openAIRecordUsageLogRepoStub{inserted: true} userRepo := &openAIRecordUsageUserRepoStub{} subRepo := &openAIRecordUsageSubRepoStub{} @@ -1063,7 +1074,45 @@ func TestOpenAIGatewayServiceRecordUsage_Gpt54LongContextBillsWholeSession(t *te }, APIKey: &APIKey{ID: 1014}, User: &User{ID: 2014}, - Account: &Account{ID: 3014}, + Account: &Account{ID: 3014, Platform: PlatformOpenAI}, + }) + + require.NoError(t, err) + require.NotNil(t, usageRepo.lastLog) + + expectedInput := 300000 * 2.5e-6 + expectedOutput := 2000 * 15e-6 + require.InDelta(t, expectedInput, usageRepo.lastLog.InputCost, 1e-10) + require.InDelta(t, expectedOutput, usageRepo.lastLog.OutputCost, 1e-10) + require.InDelta(t, expectedInput+expectedOutput, usageRepo.lastLog.TotalCost, 1e-10) + require.InDelta(t, (expectedInput+expectedOutput)*1.1, usageRepo.lastLog.ActualCost, 1e-10) + require.False(t, usageRepo.lastLog.LongContextBillingApplied) + require.Equal(t, 1, userRepo.deductCalls) +} + +func TestOpenAIGatewayServiceRecordUsage_Gpt54LongContextBillingEnabledPerAccount(t *testing.T) { + usageRepo := &openAIRecordUsageLogRepoStub{inserted: true} + userRepo := &openAIRecordUsageUserRepoStub{} + subRepo := &openAIRecordUsageSubRepoStub{} + svc := newOpenAIRecordUsageServiceForTest(usageRepo, userRepo, subRepo, nil) + + err := svc.RecordUsage(context.Background(), &OpenAIRecordUsageInput{ + Result: &OpenAIForwardResult{ + RequestID: "resp_gpt54_long_context_disabled", + Usage: OpenAIUsage{ + InputTokens: 300000, + OutputTokens: 2000, + }, + Model: "gpt-5.4-2026-03-05", + Duration: time.Second, + }, + APIKey: &APIKey{ID: 1015}, + User: &User{ID: 2015}, + Account: &Account{ + ID: 3015, + Platform: PlatformOpenAI, + Extra: map[string]any{"openai_long_context_billing_enabled": true}, + }, }) require.NoError(t, err) @@ -1075,7 +1124,62 @@ func TestOpenAIGatewayServiceRecordUsage_Gpt54LongContextBillsWholeSession(t *te require.InDelta(t, expectedOutput, usageRepo.lastLog.OutputCost, 1e-10) require.InDelta(t, expectedInput+expectedOutput, usageRepo.lastLog.TotalCost, 1e-10) require.InDelta(t, (expectedInput+expectedOutput)*1.1, usageRepo.lastLog.ActualCost, 1e-10) - require.Equal(t, 1, userRepo.deductCalls) + require.True(t, usageRepo.lastLog.LongContextBillingApplied) +} + +func TestOpenAIGatewayServiceRecordUsage_SparkShadowUsesCurrentParentBillingSetting(t *testing.T) { + tests := []struct { + name string + parentEnabled bool + }{ + {name: "parent opt out overrides stale enabled shadow", parentEnabled: false}, + {name: "parent opt in overrides stale disabled shadow", parentEnabled: true}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + usageRepo := &openAIRecordUsageLogRepoStub{inserted: true} + accountRepo := &openAIRecordUsageAccountRepoStub{account: &Account{ + ID: 4016, + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Extra: map[string]any{openAILongContextBillingEnabledKey: tt.parentEnabled}, + }} + svc := newOpenAIRecordUsageServiceForTest( + usageRepo, + &openAIRecordUsageUserRepoStub{}, + &openAIRecordUsageSubRepoStub{}, + nil, + ) + svc.accountRepo = accountRepo + parentID := int64(4016) + + err := svc.RecordUsage(context.Background(), &OpenAIRecordUsageInput{ + Result: &OpenAIForwardResult{ + RequestID: "resp_gpt54_shadow_parent_setting", + Usage: OpenAIUsage{InputTokens: 300000, OutputTokens: 2000}, + Model: "gpt-5.4-2026-03-05", + Duration: time.Second, + }, + APIKey: &APIKey{ID: 1016}, + User: &User{ID: 2016}, + Account: &Account{ + ID: 3016, + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + ParentAccountID: &parentID, + QuotaDimension: QuotaDimensionSpark, + Extra: map[string]any{ + openAILongContextBillingEnabledKey: !tt.parentEnabled, + }, + }, + }) + + require.NoError(t, err) + require.Equal(t, 1, accountRepo.calls) + require.Equal(t, tt.parentEnabled, usageRepo.lastLog.LongContextBillingApplied) + }) + } } func TestOpenAIGatewayServiceRecordUsage_ServiceTierPriorityUsesFastPricing(t *testing.T) { diff --git a/backend/internal/service/openai_gateway_request_body.go b/backend/internal/service/openai_gateway_request_body.go index 935a32f58b..d56b11fdcc 100644 --- a/backend/internal/service/openai_gateway_request_body.go +++ b/backend/internal/service/openai_gateway_request_body.go @@ -365,15 +365,57 @@ func newOpenAIRequestView(body []byte) openAIRequestView { if len(body) == 0 { return openAIRequestView{} } - return openAIRequestView{ - body: body, - Model: strings.TrimSpace(gjson.GetBytes(body, "model").String()), - Stream: gjson.GetBytes(body, "stream").Bool(), - PromptCacheKey: strings.TrimSpace(gjson.GetBytes(body, "prompt_cache_key").String()), - PreviousResponseID: strings.TrimSpace(gjson.GetBytes(body, "previous_response_id").String()), - ServiceTier: strings.TrimSpace(gjson.GetBytes(body, "service_tier").String()), - ReasoningEffort: strings.TrimSpace(gjson.GetBytes(body, "reasoning.effort").String()), - } + + const ( + modelField uint8 = 1 << iota + streamField + promptCacheKeyField + previousResponseIDField + serviceTierField + reasoningField + allRequestViewFields = modelField | streamField | promptCacheKeyField | + previousResponseIDField | serviceTierField | reasoningField + ) + + view := openAIRequestView{body: body} + var seen uint8 + // parseRawJSONView reads body without copying; view keeps body alive for extracted strings. + parseRawJSONView(body).ForEach(func(key, value gjson.Result) bool { + switch key.Str { + case "model": + if seen&modelField == 0 { + view.Model = strings.TrimSpace(value.String()) + seen |= modelField + } + case "stream": + if seen&streamField == 0 { + view.Stream = value.Bool() + seen |= streamField + } + case "prompt_cache_key": + if seen&promptCacheKeyField == 0 { + view.PromptCacheKey = strings.TrimSpace(value.String()) + seen |= promptCacheKeyField + } + case "previous_response_id": + if seen&previousResponseIDField == 0 { + view.PreviousResponseID = strings.TrimSpace(value.String()) + seen |= previousResponseIDField + } + case "service_tier": + if seen&serviceTierField == 0 { + view.ServiceTier = strings.TrimSpace(value.String()) + seen |= serviceTierField + } + case "reasoning": + if seen&reasoningField == 0 { + view.ReasoningEffort = strings.TrimSpace(value.Get("effort").String()) + seen |= reasoningField + } + } + return seen != allRequestViewFields + }) + return view } // Decode 保留阶段一既有 full-map 行为;后续阶段会把调用点下沉到复杂分支。 diff --git a/backend/internal/service/openai_gateway_service.go b/backend/internal/service/openai_gateway_service.go index 29c7d968a2..6817aa7e98 100644 --- a/backend/internal/service/openai_gateway_service.go +++ b/backend/internal/service/openai_gateway_service.go @@ -404,6 +404,7 @@ type OpenAIGatewayService struct { openaiWSRetryMetrics openAIWSRetryMetrics responseHeaderFilter *responseheaders.CompiledHeaderFilter codexSnapshotThrottle *accountWriteThrottle + codexModelsManifestCache codexModelsManifestCache openaiCompatSessionResponses sync.Map openaiCompatAnthropicDigestSessions sync.Map } diff --git a/backend/internal/service/openai_gateway_service_hotpath_test.go b/backend/internal/service/openai_gateway_service_hotpath_test.go index 1dde60c9f0..326fde534d 100644 --- a/backend/internal/service/openai_gateway_service_hotpath_test.go +++ b/backend/internal/service/openai_gateway_service_hotpath_test.go @@ -27,6 +27,33 @@ func TestOpenAIRequestView_ExtractsRawScalars(t *testing.T) { require.Equal(t, "medium", view.ReasoningEffort) } +func TestOpenAIRequestView_ExtractsFieldsAfterLargeInput(t *testing.T) { + body := []byte(`{"model":"gpt-5","input":[{"content":"` + strings.Repeat("payload", 1024) + `"}],"stream":true,"prompt_cache_key":"session-1","previous_response_id":"resp-1","service_tier":"flex","reasoning":{"effort":"high"}}`) + + view := newOpenAIRequestView(body) + + require.Equal(t, "gpt-5", view.Model) + require.True(t, view.Stream) + require.Equal(t, "session-1", view.PromptCacheKey) + require.Equal(t, "resp-1", view.PreviousResponseID) + require.Equal(t, "flex", view.ServiceTier) + require.Equal(t, "high", view.ReasoningEffort) +} + +func TestOpenAIRequestView_KeepsFirstDuplicateField(t *testing.T) { + view := newOpenAIRequestView([]byte(`{"model":"gpt-5","model":"gpt-5.1","reasoning":{"effort":"low"},"reasoning":{"effort":"high"}}`)) + + require.Equal(t, "gpt-5", view.Model) + require.Equal(t, "low", view.ReasoningEffort) +} + +func TestOpenAIRequestView_KeepsLenientPrefixExtraction(t *testing.T) { + view := newOpenAIRequestView([]byte(`{"model":"gpt-5","stream":true,"input":[`)) + + require.Equal(t, "gpt-5", view.Model) + require.True(t, view.Stream) +} + func TestOpenAIRequestView_DecodeKeepsFullMapBehavior(t *testing.T) { view := newOpenAIRequestView([]byte(`{"model":"gpt-5","stream":true,"input":[{"type":"message","content":"hi"}]}`)) diff --git a/backend/internal/service/openai_gateway_usage.go b/backend/internal/service/openai_gateway_usage.go index 9431b73ecb..410b4d944d 100644 --- a/backend/internal/service/openai_gateway_usage.go +++ b/backend/internal/service/openai_gateway_usage.go @@ -178,7 +178,27 @@ func (s *OpenAIGatewayService) RecordUsage(ctx context.Context, input *OpenAIRec if result.ServiceTier != nil { serviceTier = strings.TrimSpace(*result.ServiceTier) } - cost, err = s.calculateOpenAIRecordUsageCost(ctx, result, apiKey, billingModels, multiplier, imageMultiplier, videoMultiplier, baseMultiplier, tokens, serviceTier) + billingAccount := account + if account.IsShadow() { + billingAccount, err = resolveCredentialAccount(ctx, s.accountRepo, account) + if err != nil { + return err + } + } + longContextBillingEnabled := billingAccount.IsOpenAILongContextBillingEnabled() + cost, err = s.calculateOpenAIRecordUsageCost( + ctx, + result, + apiKey, + billingModels, + multiplier, + imageMultiplier, + videoMultiplier, + baseMultiplier, + tokens, + serviceTier, + longContextBillingEnabled, + ) if err != nil { if !isUsagePricingUnavailableError(err) { return err @@ -257,6 +277,7 @@ func (s *OpenAIGatewayService) RecordUsage(ctx context.Context, input *OpenAIRec usageLog.CacheReadCost = cost.CacheReadCost usageLog.TotalCost = cost.TotalCost usageLog.ActualCost = cost.ActualCost + usageLog.LongContextBillingApplied = cost.LongContextBillingApplied } if isVideoUsage && (cost == nil || cost.BillingMode != string(BillingModeToken)) { usageLog.RateMultiplier = videoMultiplier @@ -366,6 +387,7 @@ func (s *OpenAIGatewayService) calculateOpenAIRecordUsageCost( webSearchMultiplier float64, tokens UsageTokens, serviceTier string, + longContextBillingEnabled bool, ) (*CostBreakdown, error) { billingModel := firstUsageBillingModel(billingModels) if result != nil && result.WebSearchCalls > 0 { @@ -395,7 +417,15 @@ func (s *OpenAIGatewayService) calculateOpenAIRecordUsageCost( if candidate == "" { continue } - cost, err := s.calculateOpenAIRecordUsageTokenCost(ctx, apiKey, candidate, multiplier, tokens, serviceTier) + cost, err := s.calculateOpenAIRecordUsageTokenCost( + ctx, + apiKey, + candidate, + multiplier, + tokens, + serviceTier, + longContextBillingEnabled, + ) if err == nil { return cost, nil } @@ -443,21 +473,29 @@ func (s *OpenAIGatewayService) calculateOpenAIRecordUsageTokenCost( multiplier float64, tokens UsageTokens, serviceTier string, + longContextBillingEnabled bool, ) (*CostBreakdown, error) { if s.resolver != nil && apiKey.Group != nil { gid := apiKey.Group.ID return s.billingService.CalculateCostUnified(CostInput{ - Ctx: ctx, - Model: billingModel, - GroupID: &gid, - Tokens: tokens, - RequestCount: 1, - RateMultiplier: multiplier, - ServiceTier: serviceTier, - Resolver: s.resolver, + Ctx: ctx, + Model: billingModel, + GroupID: &gid, + Tokens: tokens, + RequestCount: 1, + RateMultiplier: multiplier, + ServiceTier: serviceTier, + Resolver: s.resolver, + LongContextBillingEnabled: &longContextBillingEnabled, }) } - return s.billingService.CalculateCostWithServiceTier(billingModel, tokens, multiplier, serviceTier) + return s.billingService.calculateCostWithServiceTierPolicy( + billingModel, + tokens, + multiplier, + serviceTier, + longContextBillingEnabled, + ) } func (s *OpenAIGatewayService) calculateOpenAIImageCost( diff --git a/backend/internal/service/openai_image_generation_controls_test.go b/backend/internal/service/openai_image_generation_controls_test.go index af0cdf669c..49ba51fcb0 100644 --- a/backend/internal/service/openai_image_generation_controls_test.go +++ b/backend/internal/service/openai_image_generation_controls_test.go @@ -283,6 +283,44 @@ func TestOpenAIGatewayServiceForward_ChannelBridgeOverrideEnablesCodexInjection( require.Contains(t, instructions, "image_generation") } +func TestOpenAIGatewayServiceForward_CodexBridgeDoesNotInjectHostedToolAlongsideImageGenNamespace(t *testing.T) { + gin.SetMode(gin.TestMode) + + upstream := &httpUpstreamRecorder{ + resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader(`{"id":"resp_namespace_image","model":"gpt-5.5","usage":{"input_tokens":1,"output_tokens":1}}`)), + }, + } + svc := newOpenAIImageGenerationControlTestService(upstream) + svc.cfg.Gateway.CodexImageGenerationBridgeEnabled = true + c, _ := newOpenAIImageGenerationControlTestContext(true, "codex_cli_rs/0.144.1") + account := newOpenAIImageGenerationControlTestAccount() + body := []byte(`{ + "model":"gpt-5.5", + "stream":false, + "tools":[ + {"type":"function","name":"shell","parameters":{"type":"object"}}, + {"type":"namespace","name":"image_gen","tools":[{"type":"function","name":"imagegen"}]} + ], + "input":[ + {"type":"message","role":"user","content":[{"type":"input_text","text":"draw a cat"}]}, + {"type":"additional_tools","tools":[{"type":"namespace","name":"image_gen","tools":[{"type":"function","name":"imagegen"}]}]} + ], + "tool_choice":"auto" + }`) + + result, err := svc.Forward(context.Background(), c, account, body) + + require.NoError(t, err) + require.NotNil(t, result) + require.NotNil(t, upstream.lastReq) + require.False(t, gjson.GetBytes(upstream.lastBody, `tools.#(type=="image_generation")`).Exists()) + require.Equal(t, "namespace", gjson.GetBytes(upstream.lastBody, `tools.#(name=="image_gen").type`).String()) + require.Equal(t, "namespace", gjson.GetBytes(upstream.lastBody, `input.#(type=="additional_tools").tools.#(name=="image_gen").type`).String()) +} + func TestOpenAIGatewayServiceForward_CodexBridgePreservesExistingToolChoice(t *testing.T) { gin.SetMode(gin.TestMode) diff --git a/backend/internal/service/openai_messages_dispatch_test.go b/backend/internal/service/openai_messages_dispatch_test.go index bafd36449b..db7804a4f3 100644 --- a/backend/internal/service/openai_messages_dispatch_test.go +++ b/backend/internal/service/openai_messages_dispatch_test.go @@ -37,3 +37,25 @@ func TestGroupResolveMessagesDispatchModel_GrokMapsClaudeFamilyToGrok(t *testing require.Empty(t, group.ResolveMessagesDispatchModel("grok")) require.Empty(t, group.ResolveMessagesDispatchModel("gpt-5.3-codex")) } + +func TestSanitizeGroupMessagesDispatchFields_ClearsNonOpenAIPlatform(t *testing.T) { + t.Parallel() + + group := &Group{ + Platform: PlatformAnthropic, + AllowMessagesDispatch: true, + DefaultMappedModel: "gpt-5.6-sol", + MessagesDispatchModelConfig: OpenAIMessagesDispatchModelConfig{ + SonnetMappedModel: "gpt-5.3-codex", + ExactModelMappings: map[string]string{ + "claude-fable-5": "gpt-5.6-sol", + }, + }, + } + + sanitizeGroupMessagesDispatchFields(group) + + require.False(t, group.AllowMessagesDispatch) + require.Empty(t, group.DefaultMappedModel) + require.Equal(t, OpenAIMessagesDispatchModelConfig{}, group.MessagesDispatchModelConfig) +} diff --git a/backend/internal/service/openai_model_mapping.go b/backend/internal/service/openai_model_mapping.go index cb7a8ca84b..8ba1d6fe1b 100644 --- a/backend/internal/service/openai_model_mapping.go +++ b/backend/internal/service/openai_model_mapping.go @@ -3,19 +3,20 @@ package service import "strings" // resolveOpenAIForwardModel 解析 OpenAI 兼容转发使用的模型。 -// defaultMappedModel 只服务于 /v1/messages 的 Claude 系列显式调度映射, -// 不作为普通 OpenAI 请求的未知模型兜底。 -func resolveOpenAIForwardModel(account *Account, requestedModel, defaultMappedModel string) string { +// messagesDispatchMappedModel 是调用方已为 /v1/messages 解析的显式调度结果; +// 普通 OpenAI 请求必须传空,避免将分组配置作为通用模型兜底。 +func resolveOpenAIForwardModel(account *Account, requestedModel, messagesDispatchMappedModel string) string { + messagesDispatchMappedModel = strings.TrimSpace(messagesDispatchMappedModel) if account == nil { - if defaultMappedModel != "" && claudeMessagesDispatchFamily(requestedModel) != "" { - return defaultMappedModel + if messagesDispatchMappedModel != "" { + return messagesDispatchMappedModel } return requestedModel } mappedModel, matched := account.ResolveMappedModel(requestedModel) - if !matched && defaultMappedModel != "" && claudeMessagesDispatchFamily(requestedModel) != "" { - return defaultMappedModel + if !matched && messagesDispatchMappedModel != "" { + return messagesDispatchMappedModel } return mappedModel } diff --git a/backend/internal/service/openai_model_mapping_test.go b/backend/internal/service/openai_model_mapping_test.go index f2ceb3551c..7107a706ad 100644 --- a/backend/internal/service/openai_model_mapping_test.go +++ b/backend/internal/service/openai_model_mapping_test.go @@ -4,159 +4,156 @@ import "testing" func TestResolveOpenAIForwardModel(t *testing.T) { tests := []struct { - name string - account *Account - requestedModel string - defaultMappedModel string - expectedModel string + name string + account *Account + requestedModel string + messagesDispatchMappedModel string + expectedModel string }{ { - name: "uses messages dispatch default for claude model", + name: "uses messages dispatch model for known claude family", account: &Account{ Credentials: map[string]any{}, }, - requestedModel: "claude-opus-4-6", - defaultMappedModel: "gpt-4o-mini", - expectedModel: "gpt-4o-mini", + requestedModel: "claude-opus-4-6", + messagesDispatchMappedModel: "gpt-4o-mini", + expectedModel: "gpt-4o-mini", }, { - name: "does not fall back to group default for invalid gpt model", + name: "uses exact messages dispatch model for unknown claude family", account: &Account{ Credentials: map[string]any{}, }, - requestedModel: "gpt6", - defaultMappedModel: "gpt-5.4", - expectedModel: "gpt6", + requestedModel: "claude-fable-5", + messagesDispatchMappedModel: " gpt-5.6-sol ", + expectedModel: "gpt-5.6-sol", }, { - name: "preserves explicit gpt-5.4 instead of group default", + name: "nil account uses messages dispatch model", + requestedModel: "claude-fable-5", + messagesDispatchMappedModel: "gpt-5.6-sol", + expectedModel: "gpt-5.6-sol", + }, + { + name: "nil account without messages dispatch keeps requested model", + requestedModel: "claude-fable-5", + expectedModel: "claude-fable-5", + }, + { + name: "ordinary unknown gpt model has no messages dispatch fallback", account: &Account{ Credentials: map[string]any{}, }, - requestedModel: "gpt-5.4", - defaultMappedModel: "gpt-4o-mini", - expectedModel: "gpt-5.4", + requestedModel: "gpt6", + expectedModel: "gpt6", }, { - name: "preserves exact passthrough mapping instead of group default", + name: "account exact mapping overrides messages dispatch model", account: &Account{ Credentials: map[string]any{ "model_mapping": map[string]any{ - "gpt-5.4": "gpt-5.4", + "claude-fable-5": "gpt-5.5", }, }, }, - requestedModel: "gpt-5.4", - defaultMappedModel: "gpt-4o-mini", - expectedModel: "gpt-5.4", + requestedModel: "claude-fable-5", + messagesDispatchMappedModel: "gpt-5.6-sol", + expectedModel: "gpt-5.5", }, { - name: "preserves wildcard passthrough mapping instead of group default", + name: "account wildcard mapping overrides messages dispatch model", account: &Account{ Credentials: map[string]any{ "model_mapping": map[string]any{ - "gpt-*": "gpt-5.4", + "claude-*": "gpt-5.4", }, }, }, - requestedModel: "gpt-5.4", - defaultMappedModel: "gpt-4o-mini", - expectedModel: "gpt-5.4", + requestedModel: "claude-fable-5", + messagesDispatchMappedModel: "gpt-5.6-sol", + expectedModel: "gpt-5.4", }, { - name: "uses account remap when explicit target differs", + name: "account passthrough mapping overrides messages dispatch model", account: &Account{ Credentials: map[string]any{ "model_mapping": map[string]any{ - "gpt-5": "gpt-5.4", + "claude-fable-5": "claude-fable-5", }, }, }, - requestedModel: "gpt-5", - defaultMappedModel: "gpt-4o-mini", - expectedModel: "gpt-5.4", + requestedModel: "claude-fable-5", + messagesDispatchMappedModel: "gpt-5.6-sol", + expectedModel: "claude-fable-5", }, { - name: "preserves codex spark instead of group default", + name: "ordinary codex spark request keeps requested model", account: &Account{ Credentials: map[string]any{}, }, - requestedModel: "gpt-5.3-codex-spark", - defaultMappedModel: "gpt-5.4", - expectedModel: "gpt-5.3-codex-spark", + requestedModel: "gpt-5.3-codex-spark", + expectedModel: "gpt-5.3-codex-spark", }, { - name: "preserves gpt-5.5 instead of group default", + name: "ordinary gpt-5.5 request keeps requested model", account: &Account{ Credentials: map[string]any{}, }, - requestedModel: "gpt-5.5", - defaultMappedModel: "gpt-5.4", - expectedModel: "gpt-5.5", + requestedModel: "gpt-5.5", + expectedModel: "gpt-5.5", }, { - name: "preserves gpt-5.5-pro instead of group default", + name: "ordinary gpt-5.5-pro request keeps requested model", account: &Account{ Credentials: map[string]any{}, }, - requestedModel: "gpt-5.5-pro", - defaultMappedModel: "gpt-5.5", - expectedModel: "gpt-5.5-pro", + requestedModel: "gpt-5.5-pro", + expectedModel: "gpt-5.5-pro", }, { - name: "preserves compact-spelled gpt5.5 instead of group default", + name: "ordinary compact-spelled gpt5.5 request keeps requested model", account: &Account{ Credentials: map[string]any{}, }, - requestedModel: "gpt5.5", - defaultMappedModel: "gpt-5.4", - expectedModel: "gpt5.5", + requestedModel: "gpt5.5", + expectedModel: "gpt5.5", }, { - name: "preserves openai namespaced gpt-5.5 instead of group default", + name: "ordinary namespaced gpt-5.5 request keeps requested model", account: &Account{ Credentials: map[string]any{}, }, - requestedModel: "openai/gpt-5.5", - defaultMappedModel: "gpt-5.4", - expectedModel: "openai/gpt-5.5", + requestedModel: "openai/gpt-5.5", + expectedModel: "openai/gpt-5.5", }, { - name: "preserves compact gpt-5.5 instead of group default", + name: "ordinary compact gpt-5.5 request keeps requested model", account: &Account{ Credentials: map[string]any{}, }, - requestedModel: "gpt-5.5-openai-compact", - defaultMappedModel: "gpt-5.4", - expectedModel: "gpt-5.5-openai-compact", + requestedModel: "gpt-5.5-openai-compact", + expectedModel: "gpt-5.5-openai-compact", + }, + { + name: "whitespace-only messages dispatch model is ignored", + account: &Account{ + Credentials: map[string]any{}, + }, + requestedModel: "gpt-5.5", + messagesDispatchMappedModel: " ", + expectedModel: "gpt-5.5", }, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - if got := resolveOpenAIForwardModel(tt.account, tt.requestedModel, tt.defaultMappedModel); got != tt.expectedModel { + if got := resolveOpenAIForwardModel(tt.account, tt.requestedModel, tt.messagesDispatchMappedModel); got != tt.expectedModel { t.Fatalf("resolveOpenAIForwardModel(...) = %q, want %q", got, tt.expectedModel) } }) } } -func TestResolveOpenAIForwardModel_PreventsClaudeModelFromFallingBackToGpt54(t *testing.T) { - account := &Account{ - Credentials: map[string]any{}, - } - - withoutDefault := resolveOpenAIForwardModel(account, "claude-opus-4-6", "") - if withoutDefault != "claude-opus-4-6" { - t.Fatalf("resolveOpenAIForwardModel(...) = %q, want %q", withoutDefault, "claude-opus-4-6") - } - - withDefault := resolveOpenAIForwardModel(account, "claude-opus-4-6", "gpt-5.4") - if withDefault != "gpt-5.4" { - t.Fatalf("resolveOpenAIForwardModel(...) = %q, want %q", withDefault, "gpt-5.4") - } -} - func TestResolveOpenAICompactForwardModel(t *testing.T) { tests := []struct { name string diff --git a/backend/internal/service/openai_ws_http_bridge.go b/backend/internal/service/openai_ws_http_bridge.go index 0afb9181a1..a5aba67c6d 100644 --- a/backend/internal/service/openai_ws_http_bridge.go +++ b/backend/internal/service/openai_ws_http_bridge.go @@ -431,7 +431,7 @@ func resolveGrokWSUpstreamModel(account *Account, body []byte, originalModel str } } if upstreamModel == "" { - upstreamModel = "grok-4.3" + upstreamModel = grokDefaultResponsesModel } return upstreamModel } diff --git a/backend/internal/service/openai_ws_http_bridge_test.go b/backend/internal/service/openai_ws_http_bridge_test.go index 0105c7a331..d2046b9006 100644 --- a/backend/internal/service/openai_ws_http_bridge_test.go +++ b/backend/internal/service/openai_ws_http_bridge_test.go @@ -178,6 +178,51 @@ func TestOpenAIWSHTTPBridgeRelaysSSEFramesAsWebSocketMessages(t *testing.T) { require.True(t, gjson.GetBytes(upstream.lastBody, "stream").Bool()) } +func TestProxyOpenAIWSHTTPBridgeTurnForGrokDefaultsEmptyModelTo45(t *testing.T) { + gin.SetMode(gin.TestMode) + + upstream := &httpUpstreamRecorder{resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"text/event-stream"}}, + Body: io.NopCloser(strings.NewReader(strings.Join([]string{ + `data: {"type":"response.created","response":{"id":"resp_grok_default","model":"grok-4.5"}}`, + "", + `data: {"type":"response.completed","response":{"id":"resp_grok_default","model":"grok-4.5","usage":{"input_tokens":1,"output_tokens":1}}}`, + "", + }, "\n"))), + }} + svc := &OpenAIGatewayService{ + cfg: &config.Config{Gateway: config.GatewayConfig{MaxLineSize: defaultMaxLineSize}}, + httpUpstream: upstream, + } + account := &Account{ + ID: 72, + Platform: PlatformGrok, + Type: AccountTypeOAuth, + Concurrency: 1, + Credentials: map[string]any{"base_url": xai.DefaultCLIBaseURL}, + } + payload := []byte(`{"type":"response.create","generate":true,"stream":true,"input":"hi"}`) + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + c.Request = httptest.NewRequest(http.MethodGet, "/v1/responses", nil) + var events [][]byte + + result, err := svc.proxyOpenAIWSHTTPBridgeTurn( + context.Background(), c, account, "access-token", payload, len(payload), + "", "", "", "", "", 1, + func(message []byte) error { + events = append(events, append([]byte(nil), message...)) + return nil + }, + ) + + require.NoError(t, err) + require.NotNil(t, result) + require.Equal(t, grokDefaultResponsesModel, gjson.GetBytes(upstream.lastBody, "model").String()) + require.Len(t, events, 2) +} + func TestProxyResponsesWebSocketFromClientForGrokUsesXAIHTTPBridge(t *testing.T) { gin.SetMode(gin.TestMode) diff --git a/backend/internal/service/ops_models.go b/backend/internal/service/ops_models.go index e33dcf82a8..d95bbadbb8 100644 --- a/backend/internal/service/ops_models.go +++ b/backend/internal/service/ops_models.go @@ -8,6 +8,7 @@ import ( type OpsSystemLog struct { ID int64 `json:"id"` CreatedAt time.Time `json:"created_at"` + Host string `json:"host"` Level string `json:"level"` Component string `json:"component"` Message string `json:"message"` diff --git a/backend/internal/service/ops_port.go b/backend/internal/service/ops_port.go index 46d171c7c3..2b73d4a694 100644 --- a/backend/internal/service/ops_port.go +++ b/backend/internal/service/ops_port.go @@ -194,6 +194,7 @@ type OpsInsertSystemMetricsInput struct { type OpsInsertSystemLogInput struct { CreatedAt time.Time + Host string Level string Component string Message string @@ -210,6 +211,7 @@ type OpsInsertSystemLogInput struct { type OpsSystemLogFilter struct { StartTime *time.Time EndTime *time.Time + Host string Level string Component string @@ -230,6 +232,7 @@ type OpsSystemLogFilter struct { type OpsSystemLogCleanupFilter struct { StartTime *time.Time EndTime *time.Time + Host string Level string Component string diff --git a/backend/internal/service/ops_system_log_service.go b/backend/internal/service/ops_system_log_service.go index b3be37e8ae..b96ae89d92 100644 --- a/backend/internal/service/ops_system_log_service.go +++ b/backend/internal/service/ops_system_log_service.go @@ -89,6 +89,7 @@ func marshalSystemLogCleanupConditions(filter *OpsSystemLogCleanupFilter) string return "{}" } payload := map[string]any{ + "host": strings.TrimSpace(filter.Host), "level": strings.TrimSpace(filter.Level), "component": strings.TrimSpace(filter.Component), "request_id": strings.TrimSpace(filter.RequestID), diff --git a/backend/internal/service/ops_system_log_service_test.go b/backend/internal/service/ops_system_log_service_test.go index 8b5a84c1f0..e8c6199f17 100644 --- a/backend/internal/service/ops_system_log_service_test.go +++ b/backend/internal/service/ops_system_log_service_test.go @@ -101,6 +101,7 @@ func TestOpsServiceCleanupSystemLogs_SuccessAndAudit(t *testing.T) { now := time.Now().UTC() filter := &OpsSystemLogCleanupFilter{ StartTime: &now, + Host: "api-node-1", Level: "warn", RequestID: "req-1", ClientRequestID: "creq-1", @@ -119,6 +120,9 @@ func TestOpsServiceCleanupSystemLogs_SuccessAndAudit(t *testing.T) { if audit == nil { t.Fatalf("expected cleanup audit") } + if !strings.Contains(audit.Conditions, `"host":"api-node-1"`) { + t.Fatalf("audit conditions should include host: %s", audit.Conditions) + } if !strings.Contains(audit.Conditions, `"client_request_id":"creq-1"`) { t.Fatalf("audit conditions should include client_request_id: %s", audit.Conditions) } diff --git a/backend/internal/service/ops_system_log_sink.go b/backend/internal/service/ops_system_log_sink.go index 2ff273be53..2e6f5515c8 100644 --- a/backend/internal/service/ops_system_log_sink.go +++ b/backend/internal/service/ops_system_log_sink.go @@ -27,6 +27,7 @@ type OpsSystemLogSinkHealth struct { type OpsSystemLogSink struct { opsRepo OpsRepository + host string queue chan *logger.LogEvent @@ -45,10 +46,14 @@ type OpsSystemLogSink struct { lastError atomic.Value } +const maxSystemLogHostLength = 255 + func NewOpsSystemLogSink(opsRepo OpsRepository) *OpsSystemLogSink { ctx, cancel := context.WithCancel(context.Background()) + rawHost, err := os.Hostname() s := &OpsSystemLogSink{ opsRepo: opsRepo, + host: normalizeSystemLogHost(rawHost, err), queue: make(chan *logger.LogEvent, 5000), batchSize: 200, flushInterval: time.Second, @@ -59,6 +64,18 @@ func NewOpsSystemLogSink(opsRepo OpsRepository) *OpsSystemLogSink { return s } +func normalizeSystemLogHost(host string, err error) string { + host = strings.TrimSpace(host) + if err != nil || host == "" { + return "unknown" + } + runes := []rune(host) + if len(runes) > maxSystemLogHostLength { + return string(runes[:maxSystemLogHostLength]) + } + return host +} + func (s *OpsSystemLogSink) Start() { if s == nil || s.opsRepo == nil { return @@ -220,6 +237,7 @@ func (s *OpsSystemLogSink) flushBatch(baseCtx context.Context, batch []*logger.L inputs = append(inputs, &OpsInsertSystemLogInput{ CreatedAt: createdAt, + Host: s.host, Level: strings.ToLower(strings.TrimSpace(event.Level)), Component: component, Message: message, diff --git a/backend/internal/service/ops_system_log_sink_test.go b/backend/internal/service/ops_system_log_sink_test.go index b43d44c32e..0d15f1a662 100644 --- a/backend/internal/service/ops_system_log_sink_test.go +++ b/backend/internal/service/ops_system_log_sink_test.go @@ -140,6 +140,7 @@ func TestOpsSystemLogSink_StartStopAndFlushSuccess(t *testing.T) { } sink := NewOpsSystemLogSink(repo) + sink.host = "api-node-1" sink.batchSize = 1 sink.flushInterval = 10 * time.Millisecond sink.Start() @@ -172,6 +173,9 @@ func TestOpsSystemLogSink_StartStopAndFlushSuccess(t *testing.T) { t.Fatalf("captured len = %d, want 1", len(captured)) } item := captured[0] + if item.Host != "api-node-1" { + t.Fatalf("host = %q, want api-node-1", item.Host) + } if item.RequestID != "req-1" || item.ClientRequestID != "creq-1" { t.Fatalf("unexpected request ids: %+v", item) } @@ -324,3 +328,20 @@ func TestOpsSystemLogSink_HelperFunctions(t *testing.T) { } } } + +func TestNormalizeSystemLogHost(t *testing.T) { + if got := normalizeSystemLogHost(" api-node-1 ", nil); got != "api-node-1" { + t.Fatalf("trimmed host = %q, want api-node-1", got) + } + if got := normalizeSystemLogHost("", nil); got != "unknown" { + t.Fatalf("empty host = %q, want unknown", got) + } + if got := normalizeSystemLogHost("api-node-1", errors.New("hostname unavailable")); got != "unknown" { + t.Fatalf("errored host = %q, want unknown", got) + } + longHost := strings.Repeat("节", maxSystemLogHostLength+1) + got := normalizeSystemLogHost(longHost, nil) + if runeCount := len([]rune(got)); runeCount != maxSystemLogHostLength { + t.Fatalf("truncated host rune count = %d, want %d", runeCount, maxSystemLogHostLength) + } +} diff --git a/backend/internal/service/payment_order.go b/backend/internal/service/payment_order.go index 04feb8002a..9b7ee08990 100644 --- a/backend/internal/service/payment_order.go +++ b/backend/internal/service/payment_order.go @@ -16,6 +16,7 @@ import ( "github.com/Wei-Shaw/sub2api/internal/payment" "github.com/Wei-Shaw/sub2api/internal/payment/provider" infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" + "github.com/Wei-Shaw/sub2api/internal/pkg/servertiming" "github.com/shopspring/decimal" ) @@ -445,7 +446,9 @@ func (s *PaymentService) invokeProvider(ctx context.Context, order *dbent.Paymen IsMobile: req.IsMobile, ReturnURL: providerReturnURL, }, sel, outTradeNo, payAmountStr, subject) + finishProviderCall := servertiming.ObserveDependency(ctx, "payment") pr, err := prov.CreatePayment(ctx, providerReq) + finishProviderCall() if err != nil { slog.Error("[PaymentService] CreatePayment failed", "provider", sel.ProviderKey, "instance", sel.InstanceID, "error", err) if appErr := new(infraerrors.ApplicationError); errors.As(err, &appErr) { diff --git a/backend/internal/service/payment_order_lifecycle.go b/backend/internal/service/payment_order_lifecycle.go index 8ed18797dd..46a2e00605 100644 --- a/backend/internal/service/payment_order_lifecycle.go +++ b/backend/internal/service/payment_order_lifecycle.go @@ -13,6 +13,7 @@ import ( "github.com/Wei-Shaw/sub2api/ent/paymentorder" "github.com/Wei-Shaw/sub2api/internal/payment" infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" + "github.com/Wei-Shaw/sub2api/internal/pkg/servertiming" ) // --- Cancel & Expire --- @@ -157,7 +158,9 @@ func (s *PaymentService) checkPaidWithOptions(ctx context.Context, o *dbent.Paym if queryRef == "" { return "" } + finishProviderCall := servertiming.ObserveDependency(ctx, "payment") resp, err := prov.QueryOrder(ctx, queryRef) + finishProviderCall() if err != nil { slog.Warn("query upstream failed", "orderID", o.ID, "error", err) return "" @@ -199,7 +202,9 @@ func (s *PaymentService) checkPaidWithOptions(ctx context.Context, o *dbent.Paym return "" } if cp, ok := prov.(payment.CancelableProvider); ok { + finishProviderCall := servertiming.ObserveDependency(ctx, "payment") _ = cp.CancelPayment(ctx, queryRef) + finishProviderCall() } return "" } @@ -208,7 +213,9 @@ func requeryPaidOrderOnce(ctx context.Context, prov payment.Provider, queryRef s if prov == nil || strings.TrimSpace(queryRef) == "" { return nil, false } + finishProviderCall := servertiming.ObserveDependency(ctx, "payment") resp, err := prov.QueryOrder(ctx, queryRef) + finishProviderCall() if err != nil { slog.Warn("query upstream retry failed", "queryRef", queryRef, "error", err) return nil, false diff --git a/backend/internal/service/payment_refund.go b/backend/internal/service/payment_refund.go index 91822680ed..bc073a2c34 100644 --- a/backend/internal/service/payment_refund.go +++ b/backend/internal/service/payment_refund.go @@ -19,6 +19,7 @@ import ( "github.com/Wei-Shaw/sub2api/internal/payment" "github.com/Wei-Shaw/sub2api/internal/payment/provider" infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" + "github.com/Wei-Shaw/sub2api/internal/pkg/servertiming" ) // --- Refund Flow --- @@ -347,12 +348,14 @@ func (s *PaymentService) gwRefund(ctx context.Context, p *RefundPlan) (*payment. }) return nil, err } + finishProviderCall := servertiming.ObserveDependency(ctx, "payment") resp, err := prov.Refund(ctx, payment.RefundRequest{ TradeNo: p.Order.PaymentTradeNo, OrderID: p.Order.OutTradeNo, Amount: formatGatewayRefundAmount(p.GatewayAmount, p.Order), Reason: p.Reason, }) + finishProviderCall() if err != nil { if resp != nil && strings.TrimSpace(resp.Status) == payment.ProviderStatusPending { return resp, nil @@ -417,12 +420,14 @@ func (s *PaymentService) QueryAndFinalizeRefund(ctx context.Context, oid int64) } pendingDetail := s.latestRefundPendingDetail(ctx, oid) + finishProviderCall := servertiming.ObserveDependency(ctx, "payment") resp, err := queryProvider.QueryRefund(ctx, payment.RefundQueryRequest{ TradeNo: o.PaymentTradeNo, OrderID: o.OutTradeNo, RefundID: pendingDetail.RefundID, Amount: formatGatewayRefundAmount(o.RefundAmount, o), }) + finishProviderCall() if err != nil { return nil, fmt.Errorf("query refund: %w", err) } diff --git a/backend/internal/service/usage_log.go b/backend/internal/service/usage_log.go index 62e48fc8f9..0adcc04a94 100644 --- a/backend/internal/service/usage_log.go +++ b/backend/internal/service/usage_log.go @@ -142,13 +142,14 @@ type UsageLog struct { ImageOutputTokens int ImageOutputCost float64 - InputCost float64 - OutputCost float64 - CacheCreationCost float64 - CacheReadCost float64 - TotalCost float64 - ActualCost float64 - RateMultiplier float64 + InputCost float64 + OutputCost float64 + CacheCreationCost float64 + CacheReadCost float64 + TotalCost float64 + ActualCost float64 + RateMultiplier float64 + LongContextBillingApplied bool // AccountRateMultiplier 账号计费倍率快照(nil 表示历史数据,按 1.0 处理) AccountRateMultiplier *float64 // AccountStatsCost 账号统计定价预计算费用(nil = 使用默认公式 total_cost × account_rate_multiplier) diff --git a/backend/internal/service/vertex_service_account.go b/backend/internal/service/vertex_service_account.go index 256695ded5..7ccbeee43c 100644 --- a/backend/internal/service/vertex_service_account.go +++ b/backend/internal/service/vertex_service_account.go @@ -18,6 +18,7 @@ import ( "github.com/Wei-Shaw/sub2api/internal/pkg/proxyurl" "github.com/Wei-Shaw/sub2api/internal/pkg/proxyutil" + "github.com/Wei-Shaw/sub2api/internal/pkg/servertiming" "github.com/golang-jwt/jwt/v5" ) @@ -195,7 +196,7 @@ func vertexServiceAccountProxyURL(account *Account) string { func newVertexServiceAccountHTTPClient(proxyURL string) (*http.Client, error) { proxyURL = strings.TrimSpace(proxyURL) if proxyURL == "" { - return &http.Client{Timeout: 15 * time.Second}, nil + return servertiming.InstrumentClient(&http.Client{Timeout: 15 * time.Second}), nil } _, parsedProxy, err := proxyurl.Parse(proxyURL) @@ -211,7 +212,7 @@ func newVertexServiceAccountHTTPClient(proxyURL string) (*http.Client, error) { if err := proxyutil.ConfigureTransportProxy(transport, parsedProxy); err != nil { return nil, err } - return &http.Client{Timeout: 15 * time.Second, Transport: transport}, nil + return servertiming.InstrumentClient(&http.Client{Timeout: 15 * time.Second, Transport: transport}), nil } func exchangeVertexServiceAccountToken(ctx context.Context, key *vertexServiceAccountKey, proxyURL string) (string, time.Duration, error) { diff --git a/backend/internal/service/vertex_service_account_test.go b/backend/internal/service/vertex_service_account_test.go index d77a1988e9..68c756eaa2 100644 --- a/backend/internal/service/vertex_service_account_test.go +++ b/backend/internal/service/vertex_service_account_test.go @@ -13,6 +13,7 @@ import ( "testing" "time" + "github.com/Wei-Shaw/sub2api/internal/pkg/servertiming" "github.com/stretchr/testify/require" "github.com/tidwall/gjson" ) @@ -101,6 +102,24 @@ func TestVertexServiceAccountProxyURL(t *testing.T) { require.Empty(t, vertexServiceAccountProxyURL(&Account{ProxyID: &proxyID})) } +func TestVertexServiceAccountHTTPClientRecordsDependency(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.WriteHeader(http.StatusNoContent) + })) + defer server.Close() + + client, err := newVertexServiceAccountHTTPClient("") + require.NoError(t, err) + collector := servertiming.New(time.Now()) + ctx := servertiming.WithCollector(context.Background(), collector) + request, err := http.NewRequestWithContext(ctx, http.MethodGet, server.URL, nil) + require.NoError(t, err) + response, err := client.Do(request) + require.NoError(t, err) + require.NoError(t, response.Body.Close()) + require.Contains(t, collector.HeaderValue(time.Now(), "bypass"), "dep_http;dur=") +} + func TestExchangeVertexServiceAccountTokenUsesProxy(t *testing.T) { privateKey, err := rsa.GenerateKey(rand.Reader, 2048) require.NoError(t, err) diff --git a/backend/internal/service/wire.go b/backend/internal/service/wire.go index 7258ff05a3..d5d9124ec2 100644 --- a/backend/internal/service/wire.go +++ b/backend/internal/service/wire.go @@ -140,8 +140,9 @@ func ProvideGrokQuotaService( proxyRepo ProxyRepository, tokenProvider *GrokTokenProvider, httpUpstream HTTPUpstream, + usageLogRepo UsageLogRepository, ) *GrokQuotaService { - return NewGrokQuotaService(accountRepo, proxyRepo, tokenProvider, httpUpstream) + return NewGrokQuotaService(accountRepo, proxyRepo, tokenProvider, httpUpstream, usageLogRepo) } // ProvideGeminiTokenProvider creates GeminiTokenProvider with OAuthRefreshAPI injection diff --git a/backend/migrations/174_add_usage_log_long_context_billing.sql b/backend/migrations/174_add_usage_log_long_context_billing.sql new file mode 100644 index 0000000000..090403c310 --- /dev/null +++ b/backend/migrations/174_add_usage_log_long_context_billing.sql @@ -0,0 +1,4 @@ +-- Snapshot whether long-context pricing changed token prices for a request so +-- usage history can explain the applied charge without inferring from totals. +ALTER TABLE usage_logs + ADD COLUMN IF NOT EXISTS long_context_billing_applied BOOLEAN NOT NULL DEFAULT FALSE; diff --git a/backend/migrations/175_add_ops_system_logs_host.sql b/backend/migrations/175_add_ops_system_logs_host.sql new file mode 100644 index 0000000000..e5f9f7299c --- /dev/null +++ b/backend/migrations/175_add_ops_system_logs_host.sql @@ -0,0 +1,3 @@ +-- Track the application host that emitted each indexed system log. +ALTER TABLE ops_system_logs + ADD COLUMN IF NOT EXISTS host VARCHAR(255); diff --git a/backend/migrations/175_default_openai_long_context_billing.sql b/backend/migrations/175_default_openai_long_context_billing.sql new file mode 100644 index 0000000000..cccbea4108 --- /dev/null +++ b/backend/migrations/175_default_openai_long_context_billing.sql @@ -0,0 +1,162 @@ +-- Keep mixed-version writers consistent before backfilling rows that already exist. +CREATE OR REPLACE FUNCTION public.enforce_openai_long_context_billing_extra() +RETURNS TRIGGER +LANGUAGE plpgsql +AS $$ +DECLARE + parent_effective_value JSONB; +BEGIN + IF NEW.platform IS DISTINCT FROM 'openai' THEN + RETURN NEW; + END IF; + + NEW.extra := COALESCE(NEW.extra, '{}'::jsonb); + IF NEW.parent_account_id IS NOT NULL AND NEW.quota_dimension = 'spark' THEN + SELECT CASE + WHEN parent.platform IS DISTINCT FROM 'openai' THEN 'false'::jsonb + WHEN NOT (COALESCE(parent.extra, '{}'::jsonb) ? 'openai_long_context_billing_enabled') THEN 'false'::jsonb + WHEN jsonb_typeof(parent.extra->'openai_long_context_billing_enabled') = 'boolean' + THEN parent.extra->'openai_long_context_billing_enabled' + ELSE 'false'::jsonb + END + INTO parent_effective_value + FROM accounts AS parent + WHERE parent.id = NEW.parent_account_id; + + NEW.extra := jsonb_set( + NEW.extra, + '{openai_long_context_billing_enabled}', + COALESCE(parent_effective_value, 'false'::jsonb), + true + ); + ELSIF NOT (NEW.extra ? 'openai_long_context_billing_enabled') + AND TG_OP = 'UPDATE' + AND OLD.platform = 'openai' + AND jsonb_typeof(OLD.extra->'openai_long_context_billing_enabled') = 'boolean' THEN + NEW.extra := jsonb_set( + NEW.extra, + '{openai_long_context_billing_enabled}', + OLD.extra->'openai_long_context_billing_enabled', + true + ); + ELSIF NOT (NEW.extra ? 'openai_long_context_billing_enabled') THEN + NEW.extra := jsonb_set( + NEW.extra, + '{openai_long_context_billing_enabled}', + 'false'::jsonb, + true + ); + END IF; + + IF jsonb_typeof(NEW.extra->'openai_long_context_billing_enabled') IS DISTINCT FROM 'boolean' THEN + RAISE EXCEPTION 'openai_long_context_billing_enabled must be a boolean' + USING ERRCODE = '22023'; + END IF; + RETURN NEW; +END; +$$; + +CREATE OR REPLACE FUNCTION public.propagate_openai_long_context_billing_extra_to_shadows() +RETURNS TRIGGER +LANGUAGE plpgsql +AS $$ +BEGIN + WITH updated_shadows AS ( + UPDATE accounts AS shadow + SET extra = jsonb_set( + COALESCE(shadow.extra, '{}'::jsonb), + '{openai_long_context_billing_enabled}', + NEW.extra->'openai_long_context_billing_enabled', + true + ) + WHERE shadow.parent_account_id = NEW.id + AND shadow.platform = 'openai' + AND shadow.quota_dimension = 'spark' + AND shadow.extra->'openai_long_context_billing_enabled' + IS DISTINCT FROM NEW.extra->'openai_long_context_billing_enabled' + RETURNING shadow.id + ) + INSERT INTO scheduler_outbox (event_type, account_id) + SELECT 'account_changed', id + FROM updated_shadows; + RETURN NULL; +END; +$$; + +DROP TRIGGER IF EXISTS accounts_enforce_openai_long_context_billing_extra ON accounts; +CREATE TRIGGER accounts_enforce_openai_long_context_billing_extra +BEFORE INSERT OR UPDATE OF platform, extra, parent_account_id, quota_dimension +ON accounts +FOR EACH ROW +EXECUTE FUNCTION public.enforce_openai_long_context_billing_extra(); + +DROP TRIGGER IF EXISTS accounts_propagate_openai_long_context_billing_extra ON accounts; +CREATE TRIGGER accounts_propagate_openai_long_context_billing_extra +AFTER UPDATE OF platform, extra +ON accounts +FOR EACH ROW +WHEN ( + NEW.platform = 'openai' + AND NEW.parent_account_id IS NULL + AND ( + OLD.platform IS DISTINCT FROM NEW.platform + OR OLD.extra->'openai_long_context_billing_enabled' + IS DISTINCT FROM NEW.extra->'openai_long_context_billing_enabled' + ) +) +EXECUTE FUNCTION public.propagate_openai_long_context_billing_extra_to_shadows(); + +UPDATE accounts +SET extra = jsonb_set( + COALESCE(extra, '{}'::jsonb), + '{openai_long_context_billing_enabled}', + 'false'::jsonb, + true +) +WHERE platform = 'openai' + AND COALESCE(extra, '{}'::jsonb) ? 'openai_long_context_billing_enabled' + AND jsonb_typeof(extra->'openai_long_context_billing_enabled') IS DISTINCT FROM 'boolean'; + +UPDATE accounts +SET extra = jsonb_set( + COALESCE(extra, '{}'::jsonb), + '{openai_long_context_billing_enabled}', + 'false'::jsonb, + true +) +WHERE platform = 'openai' + AND parent_account_id IS NULL + AND NOT (COALESCE(extra, '{}'::jsonb) ? 'openai_long_context_billing_enabled'); + +WITH shadow_values AS ( + SELECT + shadow.id, + CASE + WHEN parent.platform IS DISTINCT FROM 'openai' THEN 'false'::jsonb + WHEN NOT (COALESCE(parent.extra, '{}'::jsonb) ? 'openai_long_context_billing_enabled') THEN 'false'::jsonb + WHEN jsonb_typeof(parent.extra->'openai_long_context_billing_enabled') = 'boolean' + THEN parent.extra->'openai_long_context_billing_enabled' + ELSE 'false'::jsonb + END AS effective_value + FROM accounts AS shadow + JOIN accounts AS parent ON parent.id = shadow.parent_account_id + WHERE shadow.platform = 'openai' + AND shadow.quota_dimension = 'spark' +), +updated_shadows AS ( + UPDATE accounts AS shadow + SET extra = jsonb_set( + COALESCE(shadow.extra, '{}'::jsonb), + '{openai_long_context_billing_enabled}', + shadow_values.effective_value, + true + ) + FROM shadow_values + WHERE shadow.id = shadow_values.id + AND shadow.extra->'openai_long_context_billing_enabled' + IS DISTINCT FROM shadow_values.effective_value + RETURNING shadow.id +) +INSERT INTO scheduler_outbox (event_type, account_id) +SELECT 'account_changed', id +FROM updated_shadows; diff --git a/backend/migrations/175a_add_ops_system_logs_host_index_notx.sql b/backend/migrations/175a_add_ops_system_logs_host_index_notx.sql new file mode 100644 index 0000000000..ec2705e49b --- /dev/null +++ b/backend/migrations/175a_add_ops_system_logs_host_index_notx.sql @@ -0,0 +1,2 @@ +CREATE INDEX CONCURRENTLY IF NOT EXISTS idx_ops_system_logs_host_created_at + ON ops_system_logs (host, created_at DESC); diff --git a/backend/migrations/openai_long_context_billing_migration_test.go b/backend/migrations/openai_long_context_billing_migration_test.go new file mode 100644 index 0000000000..212ac15d9d --- /dev/null +++ b/backend/migrations/openai_long_context_billing_migration_test.go @@ -0,0 +1,36 @@ +package migrations + +import ( + "testing" + + "github.com/stretchr/testify/require" +) + +func TestMigration175DefaultsOrdinaryOpenAIAndInheritsForSparkShadows(t *testing.T) { + content, err := FS.ReadFile("175_default_openai_long_context_billing.sql") + require.NoError(t, err) + + sql := string(content) + require.Contains(t, sql, "parent_account_id IS NULL") + require.Contains(t, sql, "quota_dimension = 'spark'") + require.Contains(t, sql, "parent.extra") + require.Contains(t, sql, "jsonb_typeof") + require.Contains(t, sql, "openai_long_context_billing_enabled") +} + +func TestMigration175GuardsMixedVersionAccountWrites(t *testing.T) { + content, err := FS.ReadFile("175_default_openai_long_context_billing.sql") + require.NoError(t, err) + + sql := string(content) + require.Contains(t, sql, "RETURNS TRIGGER") + require.Contains(t, sql, "BEFORE INSERT OR UPDATE") + require.Contains(t, sql, "CREATE TRIGGER") + require.Contains(t, sql, "must be a boolean") + require.Contains(t, sql, "INSERT INTO scheduler_outbox") + require.Contains(t, sql, "'account_changed'") + require.Contains(t, sql, "jsonb_typeof(extra->'openai_long_context_billing_enabled') IS DISTINCT FROM 'boolean'") + require.Contains(t, sql, "WITH shadow_values AS") + require.Contains(t, sql, "TG_OP = 'UPDATE'") + require.Contains(t, sql, "OLD.extra->'openai_long_context_billing_enabled'") +} diff --git a/deploy/.env.example b/deploy/.env.example index f68257df9f..1d1a6ec881 100644 --- a/deploy/.env.example +++ b/deploy/.env.example @@ -23,6 +23,9 @@ SERVER_PORT=8080 # Server mode: release or debug SERVER_MODE=release +# Return Server-Timing for authenticated requests made by the Admin web UI +ENABLE_SERVER_TIMING=false + # Apple container image overrides (ignored by Docker Compose). Pin release tags # or digests for repeatable operator-managed deployments. APPLE_CONTAINER_SUB2API_IMAGE=weishaw/sub2api:latest diff --git a/deploy/config.example.yaml b/deploy/config.example.yaml index 954263d9c4..dfc584a59a 100644 --- a/deploy/config.example.yaml +++ b/deploy/config.example.yaml @@ -20,6 +20,9 @@ server: # Mode: "debug" for development, "release" for production # 运行模式:"debug" 用于开发,"release" 用于生产环境 mode: "release" + # Return Server-Timing for authenticated requests made by the Admin web UI + # 为管理端 Web 页面发出的已认证请求返回 Server-Timing + enable_server_timing: false # Frontend base URL used to generate external links in emails (e.g. password reset) # 用于生成邮件中的外部链接(例如:重置密码链接)的前端基础地址 # Example: "https://example.com" diff --git a/deploy/docker-compose.dev.yml b/deploy/docker-compose.dev.yml index 6f5b3f56f3..43f5dd3f60 100644 --- a/deploy/docker-compose.dev.yml +++ b/deploy/docker-compose.dev.yml @@ -26,6 +26,7 @@ services: - SERVER_HOST=0.0.0.0 - SERVER_PORT=8080 - SERVER_MODE=debug + - ENABLE_SERVER_TIMING=${ENABLE_SERVER_TIMING:-false} - RUN_MODE=${RUN_MODE:-standard} - DATABASE_HOST=postgres - DATABASE_PORT=5432 diff --git a/deploy/docker-compose.local.yml b/deploy/docker-compose.local.yml index 042752e857..5fb161603b 100644 --- a/deploy/docker-compose.local.yml +++ b/deploy/docker-compose.local.yml @@ -51,6 +51,7 @@ services: - SERVER_HOST=0.0.0.0 - SERVER_PORT=8080 - SERVER_MODE=${SERVER_MODE:-release} + - ENABLE_SERVER_TIMING=${ENABLE_SERVER_TIMING:-false} - RUN_MODE=${RUN_MODE:-standard} # ======================================================================= diff --git a/deploy/docker-compose.standalone.yml b/deploy/docker-compose.standalone.yml index 2e1d335624..40ed4751d6 100644 --- a/deploy/docker-compose.standalone.yml +++ b/deploy/docker-compose.standalone.yml @@ -37,6 +37,7 @@ services: - SERVER_HOST=0.0.0.0 - SERVER_PORT=8080 - SERVER_MODE=${SERVER_MODE:-release} + - ENABLE_SERVER_TIMING=${ENABLE_SERVER_TIMING:-false} - RUN_MODE=${RUN_MODE:-standard} # ======================================================================= diff --git a/deploy/docker-compose.yml b/deploy/docker-compose.yml index 22713c59aa..6aecdcfa5a 100644 --- a/deploy/docker-compose.yml +++ b/deploy/docker-compose.yml @@ -47,6 +47,7 @@ services: - SERVER_HOST=0.0.0.0 - SERVER_PORT=8080 - SERVER_MODE=${SERVER_MODE:-release} + - ENABLE_SERVER_TIMING=${ENABLE_SERVER_TIMING:-false} - RUN_MODE=${RUN_MODE:-standard} # ======================================================================= diff --git a/frontend/src/api/__tests__/adminUIRequest.spec.ts b/frontend/src/api/__tests__/adminUIRequest.spec.ts new file mode 100644 index 0000000000..9064a52f1a --- /dev/null +++ b/frontend/src/api/__tests__/adminUIRequest.spec.ts @@ -0,0 +1,38 @@ +import { describe, expect, it } from 'vitest' + +import { + ADMIN_UI_REQUEST_HEADER, + shouldMarkAdminUIRequest, +} from '@/api/adminUIRequest' + +describe('Admin UI request marker', () => { + it('uses the stable request header name', () => { + expect(ADMIN_UI_REQUEST_HEADER).toBe('X-Admin-UI-Request') + }) + + it.each([ + '/admin', + '/admin/users', + '/api/v1/admin', + '/api/v1/admin/accounts?status=active', + 'https://api.example.test/api/v1/admin/dashboard', + ])('marks Admin API request %s before page navigation', (requestURL) => { + expect(shouldMarkAdminUIRequest(requestURL, '/login')).toBe(true) + }) + + it.each(['/keys', '/groups/available', '/auth/me', '/announcements'])( + 'marks shared request %s while an Admin page is active', + (requestURL) => { + expect(shouldMarkAdminUIRequest(requestURL, '/admin/dashboard')).toBe(true) + } + ) + + it.each([ + ['/keys', '/dashboard'], + ['/api/v1/administer', '/dashboard'], + ['/keys', '/administrator'], + ['', '/'], + ])('does not mark request %s on page %s', (requestURL, pagePath) => { + expect(shouldMarkAdminUIRequest(requestURL, pagePath)).toBe(false) + }) +}) diff --git a/frontend/src/api/__tests__/client.spec.ts b/frontend/src/api/__tests__/client.spec.ts index a0a05410d4..b275cca34b 100644 --- a/frontend/src/api/__tests__/client.spec.ts +++ b/frontend/src/api/__tests__/client.spec.ts @@ -12,6 +12,7 @@ describe('API Client', () => { beforeEach(async () => { localStorage.clear() + window.history.replaceState({}, '', '/') // 每次测试重新导入以获取干净的模块状态 vi.resetModules() const mod = await import('@/api/client') @@ -120,6 +121,55 @@ describe('API Client', () => { const config = adapter.mock.calls[0][0] expect(config.withCredentials).toBe(true) }) + + it('Admin API 在进入管理页面前也带 Admin UI 标记', async () => { + const adapter = vi.fn().mockResolvedValue({ + status: 200, + data: { code: 0, data: {} }, + headers: {}, + config: {}, + statusText: 'OK', + }) + apiClient.defaults.adapter = adapter + + await apiClient.get('/admin/users') + + const config = adapter.mock.calls[0][0] + expect(config.headers.get('X-Admin-UI-Request')).toBe('1') + }) + + it('管理页面调用共享 API 时带 Admin UI 标记', async () => { + window.history.replaceState({}, '', '/admin/dashboard') + const adapter = vi.fn().mockResolvedValue({ + status: 200, + data: { code: 0, data: {} }, + headers: {}, + config: {}, + statusText: 'OK', + }) + apiClient.defaults.adapter = adapter + + await apiClient.get('/groups/available') + + const config = adapter.mock.calls[0][0] + expect(config.headers.get('X-Admin-UI-Request')).toBe('1') + }) + + it('普通用户页面调用共享 API 时不带 Admin UI 标记', async () => { + const adapter = vi.fn().mockResolvedValue({ + status: 200, + data: { code: 0, data: {} }, + headers: {}, + config: {}, + statusText: 'OK', + }) + apiClient.defaults.adapter = adapter + + await apiClient.get('/groups/available') + + const config = adapter.mock.calls[0][0] + expect(config.headers.get('X-Admin-UI-Request')).toBeFalsy() + }) }) // --- 响应拦截器 --- diff --git a/frontend/src/api/admin/grok.ts b/frontend/src/api/admin/grok.ts index 6f8f945f2c..a61169551c 100644 --- a/frontend/src/api/admin/grok.ts +++ b/frontend/src/api/admin/grok.ts @@ -4,6 +4,9 @@ */ import { apiClient } from '../client' +import type { GrokBillingSummary, GrokQuotaWindow, WindowStats } from '@/types' + +export type { GrokBillingSummary, GrokQuotaWindow } from '@/types' export interface GrokAuthUrlResponse { auth_url: string @@ -79,13 +82,6 @@ export function getGrokSSOImportTimeout(keyCount: number): number { return batches * GROK_SSO_IMPORT_TIMEOUT_PER_BATCH_MS + GROK_SSO_IMPORT_TIMEOUT_BUFFER_MS } -export interface GrokQuotaWindow { - limit?: number | null - remaining?: number | null - reset_unix?: number | null - reset_at?: string | null -} - export interface GrokQuotaSnapshot { requests?: GrokQuotaWindow | null tokens?: GrokQuotaWindow | null @@ -102,13 +98,18 @@ export interface GrokQuotaSnapshot { } export interface GrokQuotaProbeResult { - source: 'active_probe' - model: string + source: 'active_probe' | 'billing_probe' | 'hybrid_probe' + model?: string + billing?: GrokBillingSummary | null snapshot?: GrokQuotaSnapshot | null + local_usage_7d?: WindowStats | null + local_usage_monthly?: WindowStats | null status_code?: number headers_observed: boolean reset_supported: boolean fetched_at: number + persisted?: boolean + probe_error?: string } export interface GrokQuotaResetResult { diff --git a/frontend/src/api/admin/ops.ts b/frontend/src/api/admin/ops.ts index c7cbc64a4b..3284ef7e15 100644 --- a/frontend/src/api/admin/ops.ts +++ b/frontend/src/api/admin/ops.ts @@ -828,6 +828,7 @@ export interface OpsRuntimeLogConfig { export interface OpsSystemLog { id: number created_at: string + host: string level: string component: string message: string @@ -849,6 +850,7 @@ export interface OpsSystemLogQuery { time_range?: '5m' | '30m' | '1h' | '6h' | '24h' | '7d' | '30d' start_time?: string end_time?: string + host?: string level?: string component?: string request_id?: string @@ -864,6 +866,7 @@ export interface OpsSystemLogQuery { export interface OpsSystemLogCleanupRequest { start_time?: string end_time?: string + host?: string level?: string component?: string request_id?: string diff --git a/frontend/src/api/adminUIRequest.ts b/frontend/src/api/adminUIRequest.ts new file mode 100644 index 0000000000..2d60e2987d --- /dev/null +++ b/frontend/src/api/adminUIRequest.ts @@ -0,0 +1,27 @@ +export const ADMIN_UI_REQUEST_HEADER = 'X-Admin-UI-Request' + +function isAdminPath(path: string): boolean { + return ( + path === '/admin' || + path.startsWith('/admin/') || + path === '/api/v1/admin' || + path.startsWith('/api/v1/admin/') + ) +} + +function requestPath(rawURL: string): string { + const value = rawURL.trim() + if (!value) return '' + try { + const origin = typeof window !== 'undefined' ? window.location.origin : 'http://localhost' + return new URL(value, origin).pathname + } catch { + return value.split(/[?#]/, 1)[0] + } +} + +export function shouldMarkAdminUIRequest(requestURL: string, pagePath?: string): boolean { + const currentPath = + pagePath ?? (typeof window !== 'undefined' ? window.location.pathname : '') + return isAdminPath(requestPath(requestURL)) || isAdminPath(currentPath) +} diff --git a/frontend/src/api/client.ts b/frontend/src/api/client.ts index 5df969f188..a2b4d2f650 100644 --- a/frontend/src/api/client.ts +++ b/frontend/src/api/client.ts @@ -6,6 +6,7 @@ import axios, { AxiosInstance, AxiosError, InternalAxiosRequestConfig, AxiosResponse } from 'axios' import type { ApiResponse } from '@/types' import { getLocale } from '@/i18n' +import { ADMIN_UI_REQUEST_HEADER, shouldMarkAdminUIRequest } from './adminUIRequest' import { getAPIBaseURL } from './url' export { buildApiUrl, buildGatewayUrl } from './url' @@ -74,6 +75,10 @@ apiClient.interceptors.request.use( config.params.timezone = getUserTimezone() } + if (config.headers && shouldMarkAdminUIRequest(String(config.url || ''))) { + config.headers[ADMIN_UI_REQUEST_HEADER] = '1' + } + return config }, (error) => { diff --git a/frontend/src/components/account/AccountUsageCell.vue b/frontend/src/components/account/AccountUsageCell.vue index 5d7af74fc6..d751bf0429 100644 --- a/frontend/src/components/account/AccountUsageCell.vue +++ b/frontend/src/components/account/AccountUsageCell.vue @@ -382,7 +382,15 @@ + +
{{ t('admin.accounts.usageWindow.grokRetryAfter', { time: grokRetryAfterLabel }) }}
@@ -409,7 +424,7 @@
{{ grokQuotaStatusLine }}
- +
-
@@ -602,6 +617,7 @@ import { ref, computed, onMounted, onBeforeUnmount, onUnmounted, watch } from 'vue' import { useI18n } from 'vue-i18n' import { adminAPI } from '@/api/admin' +import type { GrokQuotaProbeResult } from '@/api/admin/grok' import type { Account, AccountUsageInfo, GeminiCredentials, WindowStats } from '@/types' import { buildOpenAIUsageRefreshKey } from '@/utils/accountUsageRefresh' import { enqueueUsageRequest } from '@/utils/usageLoadQueue' @@ -614,6 +630,8 @@ import GrokQuotaProbeCell from './GrokQuotaProbeCell.vue' // Module-level cache shared across all AccountUsageCell instances const _usageCache = new Map() const USAGE_CACHE_TTL = 5 * 60 * 1000 // 5 minutes +// xAI Free billing exposes a window without usage_percent, so estimate it from local tokens. +const GROK_FREE_TOKEN_LIMIT = 2_000_000 const props = withDefaults( defineProps<{ @@ -1047,9 +1065,62 @@ const makeGrokQuotaBar = (quota?: { limit?: number | null; remaining?: number | const grokRequestQuotaBar = computed(() => makeGrokQuotaBar(usageInfo.value?.grok_request_quota)) const grokTokenQuotaBar = computed(() => makeGrokQuotaBar(usageInfo.value?.grok_token_quota)) +const grokLocalUsage = computed(() => + props.todayStats || + usageInfo.value?.grok_local_usage || + usageInfo.value?.grok_local_usage_7d || + usageInfo.value?.grok_local_usage_monthly || + null +) +const grokFreeQuotaUsage = computed(() => + usageInfo.value?.grok_local_usage_7d || + props.todayStats || + usageInfo.value?.grok_local_usage || + null +) +const grokBilling = computed(() => usageInfo.value?.grok_billing || null) +const grokWeeklyBillingBar = computed((): GrokQuotaBarInfo | null => { + const billing = grokBilling.value + if (billing?.period_type?.toLowerCase() !== 'weekly' || billing.usage_percent == null) { + return null + } + return { + utilization: Math.min(100, Math.max(0, billing.usage_percent)), + resetsAt: billing.period_end || null + } +}) +const grokPlanLabelIsFree = (value: string) => value.includes('free') || value.includes('basic') +const grokPlanLabelIsPaid = (value: string) => { + return value !== '' && !grokPlanLabelIsFree(value) && !value.includes('unknown') +} +const grokIsFree = computed(() => { + if (props.account.platform !== 'grok' || props.account.type !== 'oauth') return false + const billing = grokBilling.value + if ( + billing?.usage_percent != null || + billing?.used_percent != null || + (billing?.monthly_limit_cents != null && billing.monthly_limit_cents > 0) + ) return false + + const plan = (billing?.plan || '').trim().toLowerCase() + const tier = (usageInfo.value?.subscription_tier || '').trim().toLowerCase() + const entitlement = (usageInfo.value?.grok_entitlement_status || '').toLowerCase() + if (grokPlanLabelIsPaid(plan) || grokPlanLabelIsPaid(tier)) return false + if ( + grokPlanLabelIsFree(plan) || + grokPlanLabelIsFree(tier) || + grokPlanLabelIsFree(entitlement) + ) return true + return billing != null +}) +const grokFreeTokenBar = computed(() => { + if (!grokIsFree.value || !grokFreeQuotaUsage.value) return null + const used = Math.max(0, grokFreeQuotaUsage.value.tokens || 0) + return { utilization: Math.min(100, (used / GROK_FREE_TOKEN_LIMIT) * 100) } +}) const grokQuotaUnknown = computed(() => { if (props.account.platform !== 'grok') return false - if (grokRequestQuotaBar.value || grokTokenQuotaBar.value) return false + if (grokBilling.value || grokFreeTokenBar.value || grokRequestQuotaBar.value || grokTokenQuotaBar.value) return false return usageInfo.value?.grok_quota_snapshot_state !== 'observed' }) const grokQuotaUnknownLabel = computed(() => { @@ -1080,7 +1151,6 @@ const grokQuotaStatusLine = computed(() => { } return parts.length > 0 ? parts.join(' | ') : null }) -const grokLocalUsage = computed(() => usageInfo.value?.grok_local_usage || props.todayStats || null) const grokEntitlementLabel = computed(() => { const status = (usageInfo.value?.grok_entitlement_status || '').trim() return status || null @@ -1283,6 +1353,34 @@ const loadActiveUsage = async () => { } } +const handleGrokProbed = (result: GrokQuotaProbeResult) => { + const current = usageInfo.value + if (!current) return + const snapshot = result.snapshot + const merged: AccountUsageInfo = { + ...current, + grok_billing: result.billing ?? current.grok_billing, + grok_local_usage_7d: result.local_usage_7d ?? current.grok_local_usage_7d, + grok_local_usage_monthly: result.local_usage_monthly ?? current.grok_local_usage_monthly, + grok_request_quota: snapshot?.requests ?? current.grok_request_quota, + grok_token_quota: snapshot?.tokens ?? current.grok_token_quota, + grok_retry_after_seconds: snapshot?.retry_after_seconds ?? current.grok_retry_after_seconds, + grok_entitlement_status: snapshot?.entitlement_status || current.grok_entitlement_status, + grok_quota_snapshot_state: result.billing + ? 'billing_observed' + : snapshot?.headers_observed + ? 'observed' + : current.grok_quota_snapshot_state, + grok_last_quota_probe_at: result.billing?.fetched_at ?? snapshot?.last_probe_at ?? current.grok_last_quota_probe_at, + grok_last_headers_seen_at: snapshot?.last_headers_seen_at ?? current.grok_last_headers_seen_at, + grok_last_status_code: result.status_code ?? snapshot?.status_code ?? current.grok_last_status_code, + error: result.billing || snapshot ? undefined : current.error, + error_code: result.billing || snapshot ? undefined : current.error_code + } + usageInfo.value = merged + _usageCache.set(props.account.id, { data: merged, ts: Date.now() }) +} + // ===== API Key quota progress bars ===== interface QuotaBarInfo { diff --git a/frontend/src/components/account/CreateAccountModal.vue b/frontend/src/components/account/CreateAccountModal.vue index a7eca7b351..ea9deb7ba6 100644 --- a/frontend/src/components/account/CreateAccountModal.vue +++ b/frontend/src/components/account/CreateAccountModal.vue @@ -2824,6 +2824,38 @@ +
+
+
+ +

+ {{ t('admin.accounts.openai.longContextBillingDesc') }} +

+
+ +
+
+
{ const interceptWarmupRequests = ref(false) const autoPauseOnExpired = ref(true) const openaiPassthroughEnabled = ref(false) +const openAILongContextBillingEnabled = ref(false) +const openAILongContextBillingTouched = ref(false) const openAICompactMode = ref('auto') const openAIResponsesMode = ref('auto') const openAIEndpointCapabilities = ref(['chat_completions', 'embeddings']) @@ -3693,6 +3727,11 @@ const anthropicPassthroughEnabled = ref(false) const anthropicAPIKeyAuthScheme = ref('x_api_key') const webSearchEmulationMode = ref('default') const webSearchGlobalEnabled = ref(false) + +const toggleOpenAILongContextBilling = () => { + openAILongContextBillingEnabled.value = !openAILongContextBillingEnabled.value + openAILongContextBillingTouched.value = true +} const { globalEnabled: quotaNotifyGlobalEnabled, state: quotaNotifyState, @@ -4537,6 +4576,8 @@ const resetForm = () => { interceptWarmupRequests.value = false autoPauseOnExpired.value = true openaiPassthroughEnabled.value = false + openAILongContextBillingEnabled.value = false + openAILongContextBillingTouched.value = false openAICompactMode.value = 'auto' openAIResponsesMode.value = 'auto' openAIEndpointCapabilities.value = ['chat_completions', 'embeddings'] @@ -4619,6 +4660,7 @@ const buildOpenAIExtra = (base?: Record): Record): Record 0 ? extra : undefined } +const buildOpenAICodexImportExtra = (): Record | undefined => { + const extra = buildOpenAIExtra() + if (!extra) { + return undefined + } + if (!openAILongContextBillingTouched.value) { + delete extra.openai_long_context_billing_enabled + } + return Object.keys(extra).length > 0 ? extra : undefined +} + const buildAnthropicExtra = (base?: Record): Record | undefined => { if (form.platform !== 'anthropic' || accountCategory.value !== 'apikey') { return base @@ -5409,7 +5462,7 @@ const handleOpenAIImportCodexSession = async (content: string) => { oauthClient.error.value = '' try { - const extra = buildOpenAIExtra() + const extra = buildOpenAICodexImportExtra() const result = await adminAPI.accounts.importCodexSession({ content: trimmed, name: form.name, @@ -5487,7 +5540,7 @@ const handleOpenAIImportCodexPAT = async (accessToken: string) => { oauthClient.error.value = '' try { - const extra = buildOpenAIExtra() + const extra = buildOpenAICodexImportExtra() await adminAPI.accounts.createOpenAICodexPAT({ access_token: trimmed, name: form.name, diff --git a/frontend/src/components/account/EditAccountModal.vue b/frontend/src/components/account/EditAccountModal.vue index eedd76343a..cfc2fed151 100644 --- a/frontend/src/components/account/EditAccountModal.vue +++ b/frontend/src/components/account/EditAccountModal.vue @@ -1786,7 +1786,39 @@ />
- + +
+
+
+ +

+ {{ t('admin.accounts.openai.longContextBillingDesc') }} +

+
+ +
+
+
('') const openAICompactMode = ref('auto') @@ -3216,6 +3249,7 @@ const syncFormFromAccount = (newAccount: Account | null) => { // Load OpenAI passthrough toggle (OpenAI OAuth/SetupToken/API Key) openaiPassthroughEnabled.value = false + openAILongContextBillingEnabled.value = false editPlanType.value = '' openAICompactMode.value = 'auto' openAIResponsesMode.value = 'auto' @@ -3231,6 +3265,8 @@ const syncFormFromAccount = (newAccount: Account | null) => { webSearchEmulationMode.value = 'default' if (newAccount.platform === 'openai' && (newAccount.type === 'oauth' || newAccount.type === 'setup-token' || newAccount.type === 'apikey')) { openaiPassthroughEnabled.value = extra?.openai_passthrough === true || extra?.openai_oauth_passthrough === true + const longContextBillingValue = extra?.openai_long_context_billing_enabled + openAILongContextBillingEnabled.value = longContextBillingValue === true // plan_type 手动覆盖仅 OAuth 有实际调度语义(IsOpenAIChatGPTSubscription 要求 oauth),故只对 oauth 回填 editPlanType.value = newAccount.type === 'oauth' ? readPlanType(newAccount.credentials as Record | undefined) @@ -4401,6 +4437,11 @@ const handleSubmit = async () => { delete newExtra.openai_passthrough delete newExtra.openai_oauth_passthrough } + if (isSparkShadow.value) { + delete newExtra.openai_long_context_billing_enabled + } else { + newExtra.openai_long_context_billing_enabled = openAILongContextBillingEnabled.value + } if (openAICompactMode.value === 'auto') { delete newExtra.openai_compact_mode } else { diff --git a/frontend/src/components/account/GrokQuotaProbeCell.vue b/frontend/src/components/account/GrokQuotaProbeCell.vue index 183ab3e7ba..fa6fe29bec 100644 --- a/frontend/src/components/account/GrokQuotaProbeCell.vue +++ b/frontend/src/components/account/GrokQuotaProbeCell.vue @@ -55,6 +55,8 @@ const props = defineProps<{ account: Account }>() +const emit = defineEmits<{ probed: [result: GrokQuotaProbeResult] }>() + const { t } = useI18n() const visible = computed(() => props.account.platform === 'grok' && props.account.type === 'oauth') @@ -92,18 +94,27 @@ const retryAfterLabel = computed(() => { const summary = computed(() => { const snapshot = data.value?.snapshot if (!data.value) return '' - if (!snapshot) return t('admin.accounts.usageWindow.grokNoHeaders') - const parts = [ - formatWindow(t('admin.accounts.usageWindow.grokRequests'), snapshot.requests), - formatWindow(t('admin.accounts.usageWindow.grokTokens'), snapshot.tokens) - ].filter(Boolean) + const billing = data.value.billing + const parts: Array = [] + if (billing?.period_type?.toLowerCase() === 'weekly' && billing.usage_percent != null) { + parts.push(t('admin.accounts.usageWindow.grokWeeklyUsage', { + percent: Math.round(Math.min(100, Math.max(0, billing.usage_percent))) + })) + } + if (snapshot) { + parts.push( + formatWindow(t('admin.accounts.usageWindow.grokRequests'), snapshot.requests), + formatWindow(t('admin.accounts.usageWindow.grokTokens'), snapshot.tokens) + ) + } if (retryAfterLabel.value) { parts.push(t('admin.accounts.usageWindow.grokRetryAfter', { time: retryAfterLabel.value })) } - if (snapshot.entitlement_status) { + if (snapshot?.entitlement_status) { parts.push(snapshot.entitlement_status) } - return parts.length > 0 ? parts.join(' | ') : t('admin.accounts.usageWindow.grokNoHeaders') + const visibleParts = parts.filter((part): part is string => Boolean(part)) + return visibleParts.length > 0 ? visibleParts.join(' | ') : t('admin.accounts.usageWindow.grokNoHeaders') }) const truncatedError = computed(() => { @@ -117,6 +128,8 @@ const handleProbe = async () => { error.value = null try { data.value = await adminAPI.grok.queryQuota(props.account.id) + error.value = data.value.probe_error || null + emit('probed', data.value) } catch (e) { error.value = extractErrorMessage(e) } finally { diff --git a/frontend/src/components/account/__tests__/AccountUsageCell.spec.ts b/frontend/src/components/account/__tests__/AccountUsageCell.spec.ts index 2abf6513ca..988c4fc7db 100644 --- a/frontend/src/components/account/__tests__/AccountUsageCell.spec.ts +++ b/frontend/src/components/account/__tests__/AccountUsageCell.spec.ts @@ -660,6 +660,339 @@ describe('AccountUsageCell', () => { expect(wrapper.text()).toContain('admin.accounts.usageWindow.grokTokens|25|true') }) + it('Grok OAuth uses the official weekly billing percentage when available', async () => { + getUsage.mockResolvedValue({ + grok_billing: { + period_type: 'weekly', + usage_percent: 37, + period_end: '2026-07-16T03:25:00Z', + plan: 'SuperGrok' + }, + grok_local_usage: { + requests: 5, + tokens: 2_200_000, + cost: 4.42, + standard_cost: 4.42, + user_cost: 0.44 + }, + grok_request_quota: { limit: 100, remaining: 100 }, + grok_token_quota: { limit: 2_000_000, remaining: 2_000_000 } + }) + + const wrapper = mount(AccountUsageCell, { + props: { + account: makeAccount({ id: 4201, platform: 'grok', type: 'oauth', extra: {} }) + }, + global: { + stubs: { + UsageProgressBar: { + props: ['label', 'utilization', 'resetsAt', 'remainingCapacity'], + template: '
{{ label }}|{{ utilization }}|{{ resetsAt }}|{{ remainingCapacity }}
' + }, + AccountQuotaInfo: true, + GrokQuotaProbeCell: true + } + } + }) + + await flushPromises() + + expect(wrapper.text()).toContain('7d|37|2026-07-16T03:25:00Z') + expect(wrapper.text()).not.toContain('admin.accounts.usageWindow.grokRequests|') + expect(wrapper.text()).not.toContain('admin.accounts.usageWindow.grokTokens|') + expect(wrapper.text()).not.toContain('2M|') + }) + + it.each([ + { tokens: 0, expected: 0, compact: '0' }, + { tokens: 1_000_000, expected: 50, compact: '1.0M' }, + { tokens: 2_000_000, expected: 100, compact: '2.0M' }, + { tokens: 2_200_000, expected: 100, compact: '2.2M' } + ])('Grok Free derives its 2M quota from local tokens: $tokens -> $expected%', async ({ tokens, expected, compact }) => { + getUsage.mockResolvedValue({ + grok_billing: { + period_type: 'weekly', + usage_percent: null, + plan: '' + }, + grok_local_usage: { + requests: 5, + tokens, + cost: 0, + standard_cost: 0, + user_cost: 0 + }, + grok_request_quota: { limit: 100, remaining: 100 }, + grok_token_quota: { limit: 2_000_000, remaining: 2_000_000 } + }) + + const wrapper = mount(AccountUsageCell, { + props: { + account: makeAccount({ id: 4300 + expected, platform: 'grok', type: 'oauth', extra: {} }) + }, + global: { + stubs: { + UsageProgressBar: { + props: ['label', 'utilization'], + template: '
{{ label }}|{{ utilization }}
' + }, + AccountQuotaInfo: true, + GrokQuotaProbeCell: true + } + } + }) + + await flushPromises() + + expect(wrapper.text()).toContain(`2M|${expected}`) + expect(wrapper.findAll('span').filter((node) => node.text() === compact)).toHaveLength(1) + expect(wrapper.findAll('.usage-bar')).toHaveLength(1) + expect(wrapper.text()).not.toContain('admin.accounts.usageWindow.grokRequests|') + expect(wrapper.text()).not.toContain('admin.accounts.usageWindow.grokTokens|') + }) + + it('Grok Free uses the weekly billing window instead of today-only usage', async () => { + getUsage.mockResolvedValue({ + grok_billing: { period_type: 'weekly', usage_percent: null, plan: '' }, + grok_local_usage: { + requests: 2, + tokens: 200_000, + cost: 0, + standard_cost: 0 + }, + grok_local_usage_7d: { + requests: 12, + tokens: 1_500_000, + cost: 0, + standard_cost: 0 + } + }) + + const wrapper = mount(AccountUsageCell, { + props: { + account: makeAccount({ id: 4398, platform: 'grok', type: 'oauth', extra: {} }) + }, + global: { + stubs: { + UsageProgressBar: { + props: ['label', 'utilization'], + template: '
{{ label }}|{{ utilization }}
' + }, + AccountQuotaInfo: true, + GrokQuotaProbeCell: true + } + } + }) + + await flushPromises() + + expect(wrapper.text()).toContain('2M|75') + expect(wrapper.text()).toContain('200.0K') + }) + + it('Grok Free falls back to refreshed today stats when weekly usage is unavailable', async () => { + getUsage.mockResolvedValue({ + grok_billing: { period_type: 'weekly', usage_percent: null, plan: '' }, + grok_local_usage: { + requests: 1, + tokens: 250_000, + cost: 0, + standard_cost: 0, + user_cost: 0 + } + }) + + const wrapper = mount(AccountUsageCell, { + props: { + account: makeAccount({ id: 4399, platform: 'grok', type: 'oauth', extra: {} }), + todayStats: { + requests: 4, + tokens: 1_000_000, + cost: 0, + standard_cost: 0, + user_cost: 0 + } + }, + global: { + stubs: { + UsageProgressBar: { + props: ['label', 'utilization'], + template: '
{{ label }}|{{ utilization }}
' + }, + AccountQuotaInfo: true, + GrokQuotaProbeCell: true + } + } + }) + + await flushPromises() + + expect(wrapper.text()).toContain('2M|50') + expect(wrapper.text()).toContain('1.0M') + expect(wrapper.text()).not.toContain('250K') + }) + + it('Grok paid plans are not mistaken for Free when weekly usage is temporarily missing', async () => { + getUsage.mockResolvedValue({ + grok_billing: { + period_type: 'weekly', + usage_percent: null, + plan: 'SuperGrok Heavy' + }, + grok_entitlement_status: 'free', + grok_local_usage: { + requests: 2, + tokens: 2_000_000, + cost: 1, + standard_cost: 1 + }, + grok_token_quota: { limit: 1_000, remaining: 250 } + }) + + const wrapper = mount(AccountUsageCell, { + props: { + account: makeAccount({ id: 4401, platform: 'grok', type: 'oauth', extra: {} }) + }, + global: { + stubs: { + UsageProgressBar: { + props: ['label', 'utilization'], + template: '
{{ label }}|{{ utilization }}
' + }, + AccountQuotaInfo: true, + GrokQuotaProbeCell: true + } + } + }) + + await flushPromises() + + expect(wrapper.text()).toContain('admin.accounts.usageWindow.grokTokens|25') + expect(wrapper.text()).not.toContain('2M|') + }) + + it('Grok custom paid monthly limits override stale Free entitlement', async () => { + getUsage.mockResolvedValue({ + grok_billing: { + period_type: 'weekly', + usage_percent: null, + monthly_limit_cents: 25_000, + plan: '' + }, + grok_entitlement_status: 'free', + grok_local_usage: { + requests: 2, + tokens: 2_000_000, + cost: 1, + standard_cost: 1 + }, + grok_token_quota: { limit: 1_000, remaining: 250 } + }) + + const wrapper = mount(AccountUsageCell, { + props: { + account: makeAccount({ id: 4402, platform: 'grok', type: 'oauth', extra: {} }) + }, + global: { + stubs: { + UsageProgressBar: { + props: ['label', 'utilization'], + template: '
{{ label }}|{{ utilization }}
' + }, + AccountQuotaInfo: true, + GrokQuotaProbeCell: true + } + } + }) + + await flushPromises() + + expect(wrapper.text()).toContain('admin.accounts.usageWindow.grokTokens|25') + expect(wrapper.text()).not.toContain('2M|') + }) + + it('Grok credential Free tier keeps the 2M fallback when billing is unavailable', async () => { + getUsage.mockResolvedValue({ + subscription_tier: 'FREE', + grok_local_usage: { + requests: 3, + tokens: 1_000_000, + cost: 0, + standard_cost: 0 + } + }) + + const wrapper = mount(AccountUsageCell, { + props: { + account: makeAccount({ id: 4403, platform: 'grok', type: 'oauth', extra: {} }) + }, + global: { + stubs: { + UsageProgressBar: { + props: ['label', 'utilization'], + template: '
{{ label }}|{{ utilization }}
' + }, + AccountQuotaInfo: true, + GrokQuotaProbeCell: true + } + } + }) + + await flushPromises() + + expect(wrapper.text()).toContain('2M|50') + }) + + it('Grok manual probes merge billing, quota headers, and local usage', async () => { + getUsage.mockResolvedValue({ + grok_quota_snapshot_state: 'no_headers', + error: 'stale error', + error_code: 'quota_unknown' + }) + + const wrapper = mount(AccountUsageCell, { + props: { + account: makeAccount({ id: 4501, platform: 'grok', type: 'oauth', extra: {} }) + }, + global: { + stubs: { + UsageProgressBar: { + props: ['label', 'utilization', 'resetsAt'], + template: '
{{ label }}|{{ utilization }}|{{ resetsAt }}
' + }, + AccountQuotaInfo: true, + GrokQuotaProbeCell: { + emits: ['probed'], + template: `` + } + } + } + }) + + await flushPromises() + await wrapper.get('.probe').trigger('click') + + expect(wrapper.text()).toContain('7d|42|2026-07-17T00:00:00Z') + expect(wrapper.text()).toContain('1.0M') + expect(wrapper.text()).toContain('ACTIVE') + expect(wrapper.text()).not.toContain('stale error') + }) + it('Key 账号在 today stats loading 时显示骨架屏', async () => { const wrapper = mount(AccountUsageCell, { props: { diff --git a/frontend/src/components/account/__tests__/CreateAccountModal.spec.ts b/frontend/src/components/account/__tests__/CreateAccountModal.spec.ts new file mode 100644 index 0000000000..62c97d35a4 --- /dev/null +++ b/frontend/src/components/account/__tests__/CreateAccountModal.spec.ts @@ -0,0 +1,213 @@ +import { defineComponent } from 'vue' +import { flushPromises, mount } from '@vue/test-utils' +import { beforeEach, describe, expect, it, vi } from 'vitest' + +const { + createAccountMock, + importCodexSessionMock, + createOpenAICodexPATMock, +} = vi.hoisted(() => ({ + createAccountMock: vi.fn(), + importCodexSessionMock: vi.fn(), + createOpenAICodexPATMock: vi.fn(), +})) + +vi.mock('@/stores/app', () => ({ + useAppStore: () => ({ + showError: vi.fn(), + showSuccess: vi.fn(), + showWarning: vi.fn(), + }), +})) + +vi.mock('@/stores/auth', () => ({ + useAuthStore: () => ({ isSimpleMode: true }), +})) + +vi.mock('@/api/admin', () => ({ + adminAPI: { + accounts: { + create: createAccountMock, + checkMixedChannelRisk: vi.fn().mockResolvedValue({ has_risk: false }), + importCodexSession: importCodexSessionMock, + createOpenAICodexPAT: createOpenAICodexPATMock, + }, + settings: { + getWebSearchEmulationConfig: vi.fn().mockResolvedValue({ enabled: false, providers: [] }), + getSettings: vi.fn().mockResolvedValue({}), + }, + tlsFingerprintProfiles: { + list: vi.fn().mockResolvedValue([]), + }, + }, +})) + +vi.mock('@/api/admin/accounts', () => ({ + getAntigravityDefaultModelMapping: vi.fn().mockResolvedValue([]), +})) + +vi.mock('vue-i18n', async () => { + const actual = await vi.importActual('vue-i18n') + return { + ...actual, + useI18n: () => ({ t: (key: string) => key }), + } +}) + +import CreateAccountModal from '../CreateAccountModal.vue' + +const BaseDialogStub = defineComponent({ + name: 'BaseDialog', + props: { show: { type: Boolean, default: false } }, + template: '
', +}) + +const OAuthAuthorizationFlowStub = defineComponent({ + name: 'OAuthAuthorizationFlow', + emits: ['import-codex-session', 'import-codex-pat'], + template: ` +
+ + +
+ `, +}) + +function mountModal() { + return mount(CreateAccountModal, { + props: { show: true, proxies: [], groups: [] }, + global: { + stubs: { + BaseDialog: BaseDialogStub, + OAuthAuthorizationFlow: OAuthAuthorizationFlowStub, + ConfirmDialog: true, + Select: true, + Icon: true, + PlatformIcon: true, + ProxySelector: true, + ProxyAdBanner: true, + GroupSelector: true, + ModelWhitelistSelector: true, + QuotaLimitCard: true, + }, + }, + }) +} + +async function selectButtonByText(wrapper: ReturnType, text: string) { + const button = wrapper.findAll('button').find((candidate) => candidate.text().includes(text)) + expect(button).toBeDefined() + await button?.trigger('click') +} + +async function submitApiKeyAccount(platform: 'openai' | 'anthropic', enableLongContextBilling = false) { + const wrapper = mountModal() + await selectButtonByText(wrapper, platform === 'openai' ? 'OpenAI' : 'admin.accounts.claudeConsole') + if (platform === 'openai') { + await selectButtonByText(wrapper, 'API Key') + } + await wrapper.get('form#create-account-form input[type="text"]').setValue(`${platform} account`) + await wrapper.get('form#create-account-form input[type="password"]').setValue('test-api-key') + if (enableLongContextBilling) { + await wrapper.get('[data-testid="openai-long-context-billing-toggle"]').trigger('click') + } + await wrapper.get('form#create-account-form').trigger('submit.prevent') + await flushPromises() +} + +async function openCodexImportStep(toggleClicks = 0) { + const wrapper = mountModal() + await selectButtonByText(wrapper, 'OpenAI') + for (let click = 0; click < toggleClicks; click += 1) { + await wrapper.get('[data-testid="openai-long-context-billing-toggle"]').trigger('click') + } + await wrapper.get('form#create-account-form input[type="text"]').setValue('Codex import') + await wrapper.get('form#create-account-form').trigger('submit.prevent') + return wrapper +} + +describe('CreateAccountModal OpenAI long-context billing', () => { + beforeEach(() => { + createAccountMock.mockReset().mockResolvedValue({}) + importCodexSessionMock.mockReset().mockResolvedValue({ + created: 1, + updated: 0, + skipped: 0, + failed: 0, + errors: [], + warnings: [], + }) + createOpenAICodexPATMock.mockReset().mockResolvedValue({}) + }) + + it('sends false explicitly for normal OpenAI account creation by default', async () => { + await submitApiKeyAccount('openai') + + expect(createAccountMock).toHaveBeenCalledTimes(1) + expect(createAccountMock.mock.calls[0]?.[0]?.extra?.openai_long_context_billing_enabled).toBe(false) + }) + + it('sends true explicitly when OpenAI long-context billing is enabled', async () => { + await submitApiKeyAccount('openai', true) + + expect(createAccountMock).toHaveBeenCalledTimes(1) + expect(createAccountMock.mock.calls[0]?.[0]?.extra?.openai_long_context_billing_enabled).toBe(true) + }) + + it('omits the OpenAI setting for non-OpenAI account creation', async () => { + await submitApiKeyAccount('anthropic') + + expect(createAccountMock).toHaveBeenCalledTimes(1) + expect(createAccountMock.mock.calls[0]?.[0]?.extra?.openai_long_context_billing_enabled).toBeUndefined() + }) + + it('leaves Codex session import billing ownership to the backend', async () => { + const wrapper = await openCodexImportStep() + await wrapper.get('[data-testid="import-codex-session"]').trigger('click') + await flushPromises() + + expect(importCodexSessionMock).toHaveBeenCalledTimes(1) + expect(importCodexSessionMock.mock.calls[0]?.[0]?.extra?.openai_long_context_billing_enabled).toBeUndefined() + }) + + it('leaves Codex PAT import billing ownership to the backend', async () => { + const wrapper = await openCodexImportStep() + await wrapper.get('[data-testid="import-codex-pat"]').trigger('click') + await flushPromises() + + expect(createOpenAICodexPATMock).toHaveBeenCalledTimes(1) + expect(createOpenAICodexPATMock.mock.calls[0]?.[0]?.extra?.openai_long_context_billing_enabled).toBeUndefined() + }) + + it('sends explicit true for Codex session import after the toggle is enabled', async () => { + const wrapper = await openCodexImportStep(1) + await wrapper.get('[data-testid="import-codex-session"]').trigger('click') + await flushPromises() + + expect(importCodexSessionMock.mock.calls[0]?.[0]?.extra?.openai_long_context_billing_enabled).toBe(true) + }) + + it('sends explicit false for Codex session import after the toggle is changed back', async () => { + const wrapper = await openCodexImportStep(2) + await wrapper.get('[data-testid="import-codex-session"]').trigger('click') + await flushPromises() + + expect(importCodexSessionMock.mock.calls[0]?.[0]?.extra?.openai_long_context_billing_enabled).toBe(false) + }) + + it('sends explicit true for Codex PAT import after the toggle is enabled', async () => { + const wrapper = await openCodexImportStep(1) + await wrapper.get('[data-testid="import-codex-pat"]').trigger('click') + await flushPromises() + + expect(createOpenAICodexPATMock.mock.calls[0]?.[0]?.extra?.openai_long_context_billing_enabled).toBe(true) + }) + + it('sends explicit false for Codex PAT import after the toggle is changed back', async () => { + const wrapper = await openCodexImportStep(2) + await wrapper.get('[data-testid="import-codex-pat"]').trigger('click') + await flushPromises() + + expect(createOpenAICodexPATMock.mock.calls[0]?.[0]?.extra?.openai_long_context_billing_enabled).toBe(false) + }) +}) diff --git a/frontend/src/components/account/__tests__/EditAccountModal.spec.ts b/frontend/src/components/account/__tests__/EditAccountModal.spec.ts index 44691200f9..b3a583d102 100644 --- a/frontend/src/components/account/__tests__/EditAccountModal.spec.ts +++ b/frontend/src/components/account/__tests__/EditAccountModal.spec.ts @@ -395,6 +395,105 @@ describe('EditAccountModal', () => { }) }) + it('loads and submits the per-account OpenAI long-context billing toggle', async () => { + const account = buildAccount() + account.extra = { + openai_long_context_billing_enabled: true + } + updateAccountMock.mockReset() + checkMixedChannelRiskMock.mockReset() + checkMixedChannelRiskMock.mockResolvedValue({ has_risk: false }) + updateAccountMock.mockResolvedValue(account) + + const wrapper = mountModal(account) + const toggle = wrapper.get('[data-testid="openai-long-context-billing-toggle"]') + expect(toggle.attributes('aria-checked')).toBe('true') + + await toggle.trigger('click') + await wrapper.get('form#edit-account-form').trigger('submit.prevent') + + expect(updateAccountMock).toHaveBeenCalledTimes(1) + expect(updateAccountMock.mock.calls[0]?.[1]?.extra?.openai_long_context_billing_enabled).toBe(false) + }) + + it('defaults legacy OpenAI accounts to long-context billing disabled', async () => { + const account = buildAccount() + updateAccountMock.mockReset() + checkMixedChannelRiskMock.mockReset() + checkMixedChannelRiskMock.mockResolvedValue({ has_risk: false }) + updateAccountMock.mockResolvedValue(account) + + const wrapper = mountModal(account) + const toggle = wrapper.get('[data-testid="openai-long-context-billing-toggle"]') + expect(toggle.attributes('aria-checked')).toBe('false') + + await wrapper.get('form#edit-account-form').trigger('submit.prevent') + + expect(updateAccountMock).toHaveBeenCalledTimes(1) + expect(updateAccountMock.mock.calls[0]?.[1]?.extra?.openai_long_context_billing_enabled).toBe(false) + }) + + it('does not render or submit the long-context billing toggle for Spark shadow accounts', async () => { + const account = buildOpenAISparkShadowAccount() + account.extra = { + openai_long_context_billing_enabled: false + } + updateAccountMock.mockReset() + checkMixedChannelRiskMock.mockReset() + checkMixedChannelRiskMock.mockResolvedValue({ has_risk: false }) + updateAccountMock.mockResolvedValue(account) + const wrapper = mountModal(account) + + expect(wrapper.find('[data-testid="openai-long-context-billing-toggle"]').exists()).toBe(false) + + await wrapper.get('form#edit-account-form').trigger('submit.prevent') + + expect(updateAccountMock).toHaveBeenCalledTimes(1) + expect(updateAccountMock.mock.calls[0]?.[1]?.extra).not.toHaveProperty( + 'openai_long_context_billing_enabled' + ) + }) + + it('preserves an explicit OpenAI long-context billing opt-out', async () => { + const account = buildAccount() + account.extra = { + openai_long_context_billing_enabled: false + } + updateAccountMock.mockReset() + checkMixedChannelRiskMock.mockReset() + checkMixedChannelRiskMock.mockResolvedValue({ has_risk: false }) + updateAccountMock.mockResolvedValue(account) + + const wrapper = mountModal(account) + const toggle = wrapper.get('[data-testid="openai-long-context-billing-toggle"]') + expect(toggle.attributes('aria-checked')).toBe('false') + + await wrapper.get('form#edit-account-form').trigger('submit.prevent') + + expect(updateAccountMock).toHaveBeenCalledTimes(1) + expect(updateAccountMock.mock.calls[0]?.[1]?.extra?.openai_long_context_billing_enabled).toBe(false) + }) + + it('fails closed for malformed OpenAI long-context billing values', async () => { + const account = buildAccount() + account.extra = { + openai_long_context_billing_enabled: 'false' + } + updateAccountMock.mockReset() + checkMixedChannelRiskMock.mockReset() + checkMixedChannelRiskMock.mockResolvedValue({ has_risk: false }) + updateAccountMock.mockResolvedValue(account) + + const wrapper = mountModal(account) + + expect(wrapper.get('[data-testid="openai-long-context-billing-toggle"]').attributes('aria-checked')).toBe('false') + + await wrapper.get('form#edit-account-form').trigger('submit.prevent') + + expect(updateAccountMock).toHaveBeenCalledTimes(1) + expect(updateAccountMock.mock.calls[0]?.[1]?.extra?.openai_long_context_billing_enabled).toBe(false) + }) + it('loads and submits Grok OAuth model mapping edits', async () => { const account = buildGrokOAuthAccount() updateAccountMock.mockReset() diff --git a/frontend/src/components/account/__tests__/GrokQuotaProbeCell.spec.ts b/frontend/src/components/account/__tests__/GrokQuotaProbeCell.spec.ts new file mode 100644 index 0000000000..9b431b339e --- /dev/null +++ b/frontend/src/components/account/__tests__/GrokQuotaProbeCell.spec.ts @@ -0,0 +1,54 @@ +import { flushPromises, mount } from '@vue/test-utils' +import { beforeEach, describe, expect, it, vi } from 'vitest' +import GrokQuotaProbeCell from '../GrokQuotaProbeCell.vue' +import type { Account } from '@/types' + +const { queryQuota } = vi.hoisted(() => ({ + queryQuota: vi.fn() +})) + +vi.mock('@/api/admin', () => ({ + adminAPI: { + grok: { queryQuota } + } +})) + +vi.mock('vue-i18n', () => ({ + useI18n: () => ({ + t: (key: string, params?: Record) => + params?.percent == null ? key : `${key}:${params.percent}` + }) +})) + +const account = { + id: 99, + platform: 'grok', + type: 'oauth' +} as Account + +describe('GrokQuotaProbeCell', () => { + beforeEach(() => { + queryQuota.mockReset() + }) + + it('keeps billing data while exposing a failed Free quota fallback', async () => { + queryQuota.mockResolvedValue({ + source: 'hybrid_probe', + billing: { period_type: 'weekly', usage_percent: null }, + headers_observed: false, + reset_supported: false, + fetched_at: 1, + probe_error: 'upstream returned 402 for probe model "grok-4.5"' + }) + const wrapper = mount(GrokQuotaProbeCell, { props: { account } }) + + await wrapper.get('button').trigger('click') + await flushPromises() + + expect(wrapper.text()).toContain('upstream returned 402 for probe model "grok-4.5"') + expect(wrapper.emitted('probed')?.[0]?.[0]).toMatchObject({ + billing: { period_type: 'weekly', usage_percent: null }, + probe_error: 'upstream returned 402 for probe model "grok-4.5"' + }) + }) +}) diff --git a/frontend/src/components/admin/account/AccountTestModal.vue b/frontend/src/components/admin/account/AccountTestModal.vue index 0a0e3dd9ae..0a8f853ebb 100644 --- a/frontend/src/components/admin/account/AccountTestModal.vue +++ b/frontend/src/components/admin/account/AccountTestModal.vue @@ -250,6 +250,7 @@ import TextArea from '@/components/common/TextArea.vue' import { Icon } from '@/components/icons' import { useClipboard } from '@/composables/useClipboard' import { buildApiUrl } from '@/api/client' +import { ADMIN_UI_REQUEST_HEADER } from '@/api/adminUIRequest' import { adminAPI } from '@/api/admin' import type { Account, ClaudeModel } from '@/types' @@ -438,7 +439,8 @@ const startTest = async () => { method: 'POST', headers: { Authorization: `Bearer ${localStorage.getItem('auth_token')}`, - 'Content-Type': 'application/json' + 'Content-Type': 'application/json', + [ADMIN_UI_REQUEST_HEADER]: '1' }, body: JSON.stringify(requestBody), signal: abortController.signal diff --git a/frontend/src/components/admin/usage/UsageTable.vue b/frontend/src/components/admin/usage/UsageTable.vue index ff4464454c..623ac705d3 100644 --- a/frontend/src/components/admin/usage/UsageTable.vue +++ b/frontend/src/components/admin/usage/UsageTable.vue @@ -168,6 +168,11 @@
${{ row.actual_cost?.toFixed(6) || '0.000000' }} + x2
{ } as DOMRect) }) + it('marks only usage rows that actually applied long-context billing', () => { + const wrapper = mount(UsageTable, { + props: { + data: [ + { + ...baseImageRow, + request_id: 'req-long-context-enabled', + long_context_billing_applied: true, + }, + { + ...baseImageRow, + request_id: 'req-long-context-disabled', + long_context_billing_applied: false, + }, + ], + loading: false, + columns: [], + }, + global: { + stubs: { + DataTable: DataTableStub, + EmptyState: true, + Icon: true, + Teleport: true, + }, + }, + }) + + expect(wrapper.findAll('[data-testid="long-context-billing-marker"]')).toHaveLength(1) + expect(wrapper.get('[data-testid="long-context-billing-marker"]').text()).toBe('x2') + }) + it('shows service tier and billing breakdown in cost tooltip', async () => { const row = { request_id: 'req-admin-1', diff --git a/frontend/src/i18n/locales/en/admin/accounts.ts b/frontend/src/i18n/locales/en/admin/accounts.ts index e326aace99..42b3f19a2c 100644 --- a/frontend/src/i18n/locales/en/admin/accounts.ts +++ b/frontend/src/i18n/locales/en/admin/accounts.ts @@ -402,6 +402,9 @@ export default { oauthPassthrough: 'Auto passthrough (auth only)', oauthPassthroughDesc: 'When enabled, this OpenAI account uses automatic passthrough: the gateway forwards request/response as-is and only swaps auth, while keeping billing/concurrency/audit and necessary safety filtering.', + longContextBilling: 'API long-context pricing', + longContextBillingDesc: + 'Disabled by default. Enable only when this account\'s upstream charges OpenAI API long-context rates above the model threshold.', responsesWebsocketsV2: 'Responses WebSocket v2', responsesWebsocketsV2Desc: 'Disabled by default. Enable to allow responses_websockets_v2 capability (still gated by global and account-type switches).', @@ -1210,6 +1213,7 @@ export default { claude: 'Claude', grokRequests: 'Req', grokTokens: 'Tok', + grokWeeklyUsage: 'Weekly {percent}%', grokUnknown: 'Grok quota is unknown until the first upstream response includes xAI rate-limit headers.', grokRetryAfter: 'Retry after {time}', grokProbe: 'Probe', diff --git a/frontend/src/i18n/locales/en/admin/ops.ts b/frontend/src/i18n/locales/en/admin/ops.ts index 88e997ed82..588f673518 100644 --- a/frontend/src/i18n/locales/en/admin/ops.ts +++ b/frontend/src/i18n/locales/en/admin/ops.ts @@ -50,6 +50,7 @@ export default { timeRange: 'Time range', startTime: 'Start time (optional)', endTime: 'End time (optional)', + host: 'Host', component: 'Component', componentPlaceholder: 'e.g. http.access', keyId: 'KEY ID', diff --git a/frontend/src/i18n/locales/zh/admin/accounts.ts b/frontend/src/i18n/locales/zh/admin/accounts.ts index 77ab494542..0ce8fef57f 100644 --- a/frontend/src/i18n/locales/zh/admin/accounts.ts +++ b/frontend/src/i18n/locales/zh/admin/accounts.ts @@ -317,6 +317,7 @@ export default { claude: 'Claude', grokRequests: '请求', grokTokens: 'Token', + grokWeeklyUsage: '周额度已用 {percent}%', grokUnknown: 'Grok 配额需等待首次上游响应返回 xAI rate-limit 头后显示。', grokRetryAfter: '{time} 后重试', grokProbe: '探测', @@ -505,6 +506,8 @@ export default { oauthPassthrough: '自动透传(仅替换认证)', oauthPassthroughDesc: '开启后,该 OpenAI 账号将自动透传请求与响应,仅替换认证并保留计费/并发/审计及必要安全过滤;如遇兼容性问题可随时关闭回滚。', + longContextBilling: 'API 长上下文计费', + longContextBillingDesc: '默认关闭。仅当该账号的上游会按模型阈值收取 OpenAI API 长上下文费率时开启。', responsesWebsocketsV2: 'Responses WebSocket v2', responsesWebsocketsV2Desc: '默认关闭。开启后可启用 responses_websockets_v2 协议能力(受网关全局开关与账号类型开关约束)。', diff --git a/frontend/src/i18n/locales/zh/admin/ops.ts b/frontend/src/i18n/locales/zh/admin/ops.ts index 97b974fd8f..830d797a31 100644 --- a/frontend/src/i18n/locales/zh/admin/ops.ts +++ b/frontend/src/i18n/locales/zh/admin/ops.ts @@ -50,6 +50,7 @@ export default { timeRange: '时间范围', startTime: '开始时间(可选)', endTime: '结束时间(可选)', + host: 'Host', component: '组件', componentPlaceholder: '例如 http.access', keyId: 'KEY ID', diff --git a/frontend/src/types/index.ts b/frontend/src/types/index.ts index b7095ebe03..7cddbace60 100644 --- a/frontend/src/types/index.ts +++ b/frontend/src/types/index.ts @@ -1020,10 +1020,38 @@ export interface AntigravityModelQuota { } export interface GrokQuotaWindow { - limit?: number - remaining?: number - reset_unix?: number - reset_at?: string + limit?: number | null + remaining?: number | null + reset_unix?: number | null + reset_at?: string | null +} + +export interface GrokBillingProductUsage { + product: string + usage_percent?: number | null +} + +export interface GrokBillingSummary { + period_type?: string + usage_percent?: number | null + period_start?: string + period_end?: string + product_usage?: GrokBillingProductUsage[] + monthly_limit_cents?: number | null + used_cents?: number | null + included_used_cents?: number | null + billing_period_start?: string + billing_period_end?: string + used_percent?: number | null + plan?: string + status_code?: number + source?: string + fetched_at?: string + updated_at?: string + weekly_updated_at?: string + monthly_updated_at?: string + partial?: boolean + failed_windows?: string[] } export interface AccountUsageInfo { @@ -1049,6 +1077,11 @@ export interface AccountUsageInfo { grok_last_headers_seen_at?: string grok_last_status_code?: number grok_local_usage?: WindowStats | null + grok_local_usage_7d?: WindowStats | null + grok_local_usage_monthly?: WindowStats | null + grok_billing?: GrokBillingSummary | null + subscription_tier?: string + subscription_tier_raw?: string ai_credits?: Array<{ credit_type?: string amount?: number @@ -1348,6 +1381,7 @@ export interface UsageLog { total_cost: number actual_cost: number rate_multiplier: number + long_context_billing_applied: boolean billing_type: number request_type?: UsageRequestType diff --git a/frontend/src/views/admin/ops/components/OpsSystemLogTable.vue b/frontend/src/views/admin/ops/components/OpsSystemLogTable.vue index 34aedb463a..5aac985809 100644 --- a/frontend/src/views/admin/ops/components/OpsSystemLogTable.vue +++ b/frontend/src/views/admin/ops/components/OpsSystemLogTable.vue @@ -48,6 +48,7 @@ const filters = reactive({ time_range: '1h' as '5m' | '30m' | '1h' | '6h' | '24h' | '7d' | '30d', start_time: '', end_time: '', + host: '', level: '', component: '', request_id: '', @@ -175,6 +176,7 @@ const buildQuery = () => { } if (filters.start_time) query.start_time = toRFC3339(filters.start_time) if (filters.end_time) query.end_time = toRFC3339(filters.end_time) + if (filters.host.trim()) query.host = filters.host.trim() if (filters.level.trim()) query.level = filters.level.trim() if (filters.component.trim()) query.component = filters.component.trim() if (filters.request_id.trim()) query.request_id = filters.request_id.trim() @@ -288,6 +290,7 @@ const cleanupCurrentFilter = async () => { const payload = { start_time: toRFC3339(filters.start_time), end_time: toRFC3339(filters.end_time), + host: filters.host.trim() || undefined, level: filters.level.trim() || undefined, component: filters.component.trim() || undefined, request_id: filters.request_id.trim() || undefined, @@ -313,6 +316,7 @@ const resetFilters = () => { filters.time_range = '1h' filters.start_time = '' filters.end_time = '' + filters.host = '' filters.level = '' filters.component = '' filters.request_id = '' @@ -454,6 +458,10 @@ onMounted(async () => { {{ t('admin.ops.systemLogs.component') }} +