diff --git a/backend/ent/group.go b/backend/ent/group.go index 3b1d2a4a51..110f06742a 100644 --- a/backend/ent/group.go +++ b/backend/ent/group.go @@ -97,6 +97,10 @@ type Group struct { AudioTtsPricePerMillionChars *float64 `json:"audio_tts_price_per_million_chars,omitempty"` // STT 每小时价格(USD) AudioSttPricePerHour *float64 `json:"audio_stt_price_per_hour,omitempty"` + // 是否按上下文长度应用模型阶梯价格 + LongContextPricingEnabled bool `json:"long_context_pricing_enabled,omitempty"` + // 分组逐模型定价;优先级高于渠道和内置定价 + ModelPricing json.RawMessage `json:"model_pricing,omitempty"` // 是否仅允许 Claude Code 客户端 ClaudeCodeOnly bool `json:"claude_code_only,omitempty"` // 非 Claude Code 请求降级使用的分组 ID @@ -245,9 +249,9 @@ func (*Group) scanValues(columns []string) ([]any, error) { values := make([]any, len(columns)) for i := range columns { switch columns[i] { - case group.FieldVideoModelPrices, group.FieldModelRouting, group.FieldSupportedModelScopes, group.FieldMessagesDispatchModelConfig, group.FieldModelsListConfig, group.FieldReasoningEffortMappings: + case group.FieldVideoModelPrices, group.FieldModelPricing, group.FieldModelRouting, group.FieldSupportedModelScopes, group.FieldMessagesDispatchModelConfig, group.FieldModelsListConfig, group.FieldReasoningEffortMappings: values[i] = new([]byte) - case group.FieldPeakRateEnabled, group.FieldIsExclusive, group.FieldAllowImageGeneration, group.FieldAllowBatchImageGeneration, group.FieldImageRateIndependent, group.FieldVideoRateIndependent, group.FieldClaudeCodeOnly, group.FieldModelRoutingEnabled, group.FieldMcpXMLInject, group.FieldAllowMessagesDispatch, group.FieldAllowLive, group.FieldRequireOauthOnly, group.FieldRequirePrivacySet, group.FieldProfitControlEnabled: + case group.FieldPeakRateEnabled, group.FieldIsExclusive, group.FieldAllowImageGeneration, group.FieldAllowBatchImageGeneration, group.FieldImageRateIndependent, group.FieldVideoRateIndependent, group.FieldLongContextPricingEnabled, group.FieldClaudeCodeOnly, group.FieldModelRoutingEnabled, group.FieldMcpXMLInject, group.FieldAllowMessagesDispatch, group.FieldAllowLive, group.FieldRequireOauthOnly, group.FieldRequirePrivacySet, group.FieldProfitControlEnabled: values[i] = new(sql.NullBool) case group.FieldRateMultiplier, group.FieldPeakRateMultiplier, group.FieldDailyLimitUsd, group.FieldWeeklyLimitUsd, group.FieldMonthlyLimitUsd, group.FieldImageRateMultiplier, group.FieldImagePrice1k, group.FieldImagePrice2k, group.FieldImagePrice4k, group.FieldBatchImageDiscountMultiplier, group.FieldBatchImageHoldMultiplier, group.FieldVideoRateMultiplier, group.FieldVideoPrice480p, group.FieldVideoPrice720p, group.FieldVideoPrice1080p, group.FieldWebSearchPricePerCall, group.FieldSearchPricePer1k, group.FieldAudioRealtimePricePerMin, group.FieldAudioTtsPricePerMillionChars, group.FieldAudioSttPricePerHour, group.FieldProfitMinMargin, group.FieldProfitSafetyBuffer: values[i] = new(sql.NullFloat64) @@ -531,6 +535,20 @@ func (_m *Group) assignValues(columns []string, values []any) error { _m.AudioSttPricePerHour = new(float64) *_m.AudioSttPricePerHour = value.Float64 } + case group.FieldLongContextPricingEnabled: + if value, ok := values[i].(*sql.NullBool); !ok { + return fmt.Errorf("unexpected type %T for field long_context_pricing_enabled", values[i]) + } else if value.Valid { + _m.LongContextPricingEnabled = value.Bool + } + case group.FieldModelPricing: + if value, ok := values[i].(*[]byte); !ok { + return fmt.Errorf("unexpected type %T for field model_pricing", values[i]) + } else if value != nil && len(*value) > 0 { + if err := json.Unmarshal(*value, &_m.ModelPricing); err != nil { + return fmt.Errorf("unmarshal field model_pricing: %w", err) + } + } case group.FieldClaudeCodeOnly: if value, ok := values[i].(*sql.NullBool); !ok { return fmt.Errorf("unexpected type %T for field claude_code_only", values[i]) @@ -896,6 +914,12 @@ func (_m *Group) String() string { builder.WriteString(fmt.Sprintf("%v", *v)) } builder.WriteString(", ") + builder.WriteString("long_context_pricing_enabled=") + builder.WriteString(fmt.Sprintf("%v", _m.LongContextPricingEnabled)) + builder.WriteString(", ") + builder.WriteString("model_pricing=") + builder.WriteString(fmt.Sprintf("%v", _m.ModelPricing)) + builder.WriteString(", ") builder.WriteString("claude_code_only=") builder.WriteString(fmt.Sprintf("%v", _m.ClaudeCodeOnly)) builder.WriteString(", ") diff --git a/backend/ent/group/group.go b/backend/ent/group/group.go index 76a07b96a0..5ec04d32ae 100644 --- a/backend/ent/group/group.go +++ b/backend/ent/group/group.go @@ -94,6 +94,10 @@ const ( FieldAudioTtsPricePerMillionChars = "audio_tts_price_per_million_chars" // FieldAudioSttPricePerHour holds the string denoting the audio_stt_price_per_hour field in the database. FieldAudioSttPricePerHour = "audio_stt_price_per_hour" + // FieldLongContextPricingEnabled holds the string denoting the long_context_pricing_enabled field in the database. + FieldLongContextPricingEnabled = "long_context_pricing_enabled" + // FieldModelPricing holds the string denoting the model_pricing field in the database. + FieldModelPricing = "model_pricing" // FieldClaudeCodeOnly holds the string denoting the claude_code_only field in the database. FieldClaudeCodeOnly = "claude_code_only" // FieldFallbackGroupID holds the string denoting the fallback_group_id field in the database. @@ -250,6 +254,8 @@ var Columns = []string{ FieldAudioRealtimePricePerMin, FieldAudioTtsPricePerMillionChars, FieldAudioSttPricePerHour, + FieldLongContextPricingEnabled, + FieldModelPricing, FieldClaudeCodeOnly, FieldFallbackGroupID, FieldFallbackGroupIDOnInvalidRequest, @@ -364,6 +370,8 @@ var ( AudioTtsPricePerMillionCharsValidator func(float64) error // AudioSttPricePerHourValidator is a validator for the "audio_stt_price_per_hour" field. It is called by the builders before save. AudioSttPricePerHourValidator func(float64) error + // DefaultLongContextPricingEnabled holds the default value on creation for the "long_context_pricing_enabled" field. + DefaultLongContextPricingEnabled bool // DefaultClaudeCodeOnly holds the default value on creation for the "claude_code_only" field. DefaultClaudeCodeOnly bool // DefaultModelRoutingEnabled holds the default value on creation for the "model_routing_enabled" field. @@ -604,6 +612,11 @@ func ByAudioSttPricePerHour(opts ...sql.OrderTermOption) OrderOption { return sql.OrderByField(FieldAudioSttPricePerHour, opts...).ToFunc() } +// ByLongContextPricingEnabled orders the results by the long_context_pricing_enabled field. +func ByLongContextPricingEnabled(opts ...sql.OrderTermOption) OrderOption { + return sql.OrderByField(FieldLongContextPricingEnabled, opts...).ToFunc() +} + // ByClaudeCodeOnly orders the results by the claude_code_only field. func ByClaudeCodeOnly(opts ...sql.OrderTermOption) OrderOption { return sql.OrderByField(FieldClaudeCodeOnly, opts...).ToFunc() diff --git a/backend/ent/group/where.go b/backend/ent/group/where.go index 630ec582e8..685a48a922 100644 --- a/backend/ent/group/where.go +++ b/backend/ent/group/where.go @@ -245,6 +245,11 @@ func AudioSttPricePerHour(v float64) predicate.Group { return predicate.Group(sql.FieldEQ(FieldAudioSttPricePerHour, v)) } +// LongContextPricingEnabled applies equality check predicate on the "long_context_pricing_enabled" field. It's identical to LongContextPricingEnabledEQ. +func LongContextPricingEnabled(v bool) predicate.Group { + return predicate.Group(sql.FieldEQ(FieldLongContextPricingEnabled, v)) +} + // ClaudeCodeOnly applies equality check predicate on the "claude_code_only" field. It's identical to ClaudeCodeOnlyEQ. func ClaudeCodeOnly(v bool) predicate.Group { return predicate.Group(sql.FieldEQ(FieldClaudeCodeOnly, v)) @@ -2045,6 +2050,26 @@ func AudioSttPricePerHourNotNil() predicate.Group { return predicate.Group(sql.FieldNotNull(FieldAudioSttPricePerHour)) } +// LongContextPricingEnabledEQ applies the EQ predicate on the "long_context_pricing_enabled" field. +func LongContextPricingEnabledEQ(v bool) predicate.Group { + return predicate.Group(sql.FieldEQ(FieldLongContextPricingEnabled, v)) +} + +// LongContextPricingEnabledNEQ applies the NEQ predicate on the "long_context_pricing_enabled" field. +func LongContextPricingEnabledNEQ(v bool) predicate.Group { + return predicate.Group(sql.FieldNEQ(FieldLongContextPricingEnabled, v)) +} + +// ModelPricingIsNil applies the IsNil predicate on the "model_pricing" field. +func ModelPricingIsNil() predicate.Group { + return predicate.Group(sql.FieldIsNull(FieldModelPricing)) +} + +// ModelPricingNotNil applies the NotNil predicate on the "model_pricing" field. +func ModelPricingNotNil() predicate.Group { + return predicate.Group(sql.FieldNotNull(FieldModelPricing)) +} + // ClaudeCodeOnlyEQ applies the EQ predicate on the "claude_code_only" field. func ClaudeCodeOnlyEQ(v bool) predicate.Group { return predicate.Group(sql.FieldEQ(FieldClaudeCodeOnly, v)) diff --git a/backend/ent/group_create.go b/backend/ent/group_create.go index 3246b92a50..158dca1e40 100644 --- a/backend/ent/group_create.go +++ b/backend/ent/group_create.go @@ -4,6 +4,7 @@ package ent import ( "context" + "encoding/json" "errors" "fmt" "time" @@ -559,6 +560,26 @@ func (_c *GroupCreate) SetNillableAudioSttPricePerHour(v *float64) *GroupCreate return _c } +// SetLongContextPricingEnabled sets the "long_context_pricing_enabled" field. +func (_c *GroupCreate) SetLongContextPricingEnabled(v bool) *GroupCreate { + _c.mutation.SetLongContextPricingEnabled(v) + return _c +} + +// SetNillableLongContextPricingEnabled sets the "long_context_pricing_enabled" field if the given value is not nil. +func (_c *GroupCreate) SetNillableLongContextPricingEnabled(v *bool) *GroupCreate { + if v != nil { + _c.SetLongContextPricingEnabled(*v) + } + return _c +} + +// SetModelPricing sets the "model_pricing" field. +func (_c *GroupCreate) SetModelPricing(v json.RawMessage) *GroupCreate { + _c.mutation.SetModelPricing(v) + return _c +} + // SetClaudeCodeOnly sets the "claude_code_only" field. func (_c *GroupCreate) SetClaudeCodeOnly(v bool) *GroupCreate { _c.mutation.SetClaudeCodeOnly(v) @@ -1042,6 +1063,10 @@ func (_c *GroupCreate) defaults() error { v := group.DefaultVideoRateMultiplier _c.mutation.SetVideoRateMultiplier(v) } + if _, ok := _c.mutation.LongContextPricingEnabled(); !ok { + v := group.DefaultLongContextPricingEnabled + _c.mutation.SetLongContextPricingEnabled(v) + } if _, ok := _c.mutation.ClaudeCodeOnly(); !ok { v := group.DefaultClaudeCodeOnly _c.mutation.SetClaudeCodeOnly(v) @@ -1237,6 +1262,9 @@ func (_c *GroupCreate) check() error { return &ValidationError{Name: "audio_stt_price_per_hour", err: fmt.Errorf(`ent: validator failed for field "Group.audio_stt_price_per_hour": %w`, err)} } } + if _, ok := _c.mutation.LongContextPricingEnabled(); !ok { + return &ValidationError{Name: "long_context_pricing_enabled", err: errors.New(`ent: missing required field "Group.long_context_pricing_enabled"`)} + } if _, ok := _c.mutation.ClaudeCodeOnly(); !ok { return &ValidationError{Name: "claude_code_only", err: errors.New(`ent: missing required field "Group.claude_code_only"`)} } @@ -1484,6 +1512,14 @@ func (_c *GroupCreate) createSpec() (*Group, *sqlgraph.CreateSpec) { _spec.SetField(group.FieldAudioSttPricePerHour, field.TypeFloat64, value) _node.AudioSttPricePerHour = &value } + if value, ok := _c.mutation.LongContextPricingEnabled(); ok { + _spec.SetField(group.FieldLongContextPricingEnabled, field.TypeBool, value) + _node.LongContextPricingEnabled = value + } + if value, ok := _c.mutation.ModelPricing(); ok { + _spec.SetField(group.FieldModelPricing, field.TypeJSON, value) + _node.ModelPricing = value + } if value, ok := _c.mutation.ClaudeCodeOnly(); ok { _spec.SetField(group.FieldClaudeCodeOnly, field.TypeBool, value) _node.ClaudeCodeOnly = value @@ -2396,6 +2432,36 @@ func (u *GroupUpsert) ClearAudioSttPricePerHour() *GroupUpsert { return u } +// SetLongContextPricingEnabled sets the "long_context_pricing_enabled" field. +func (u *GroupUpsert) SetLongContextPricingEnabled(v bool) *GroupUpsert { + u.Set(group.FieldLongContextPricingEnabled, v) + return u +} + +// UpdateLongContextPricingEnabled sets the "long_context_pricing_enabled" field to the value that was provided on create. +func (u *GroupUpsert) UpdateLongContextPricingEnabled() *GroupUpsert { + u.SetExcluded(group.FieldLongContextPricingEnabled) + return u +} + +// SetModelPricing sets the "model_pricing" field. +func (u *GroupUpsert) SetModelPricing(v json.RawMessage) *GroupUpsert { + u.Set(group.FieldModelPricing, v) + return u +} + +// UpdateModelPricing sets the "model_pricing" field to the value that was provided on create. +func (u *GroupUpsert) UpdateModelPricing() *GroupUpsert { + u.SetExcluded(group.FieldModelPricing) + return u +} + +// ClearModelPricing clears the value of the "model_pricing" field. +func (u *GroupUpsert) ClearModelPricing() *GroupUpsert { + u.SetNull(group.FieldModelPricing) + return u +} + // SetClaudeCodeOnly sets the "claude_code_only" field. func (u *GroupUpsert) SetClaudeCodeOnly(v bool) *GroupUpsert { u.Set(group.FieldClaudeCodeOnly, v) @@ -3534,6 +3600,41 @@ func (u *GroupUpsertOne) ClearAudioSttPricePerHour() *GroupUpsertOne { }) } +// SetLongContextPricingEnabled sets the "long_context_pricing_enabled" field. +func (u *GroupUpsertOne) SetLongContextPricingEnabled(v bool) *GroupUpsertOne { + return u.Update(func(s *GroupUpsert) { + s.SetLongContextPricingEnabled(v) + }) +} + +// UpdateLongContextPricingEnabled sets the "long_context_pricing_enabled" field to the value that was provided on create. +func (u *GroupUpsertOne) UpdateLongContextPricingEnabled() *GroupUpsertOne { + return u.Update(func(s *GroupUpsert) { + s.UpdateLongContextPricingEnabled() + }) +} + +// SetModelPricing sets the "model_pricing" field. +func (u *GroupUpsertOne) SetModelPricing(v json.RawMessage) *GroupUpsertOne { + return u.Update(func(s *GroupUpsert) { + s.SetModelPricing(v) + }) +} + +// UpdateModelPricing sets the "model_pricing" field to the value that was provided on create. +func (u *GroupUpsertOne) UpdateModelPricing() *GroupUpsertOne { + return u.Update(func(s *GroupUpsert) { + s.UpdateModelPricing() + }) +} + +// ClearModelPricing clears the value of the "model_pricing" field. +func (u *GroupUpsertOne) ClearModelPricing() *GroupUpsertOne { + return u.Update(func(s *GroupUpsert) { + s.ClearModelPricing() + }) +} + // SetClaudeCodeOnly sets the "claude_code_only" field. func (u *GroupUpsertOne) SetClaudeCodeOnly(v bool) *GroupUpsertOne { return u.Update(func(s *GroupUpsert) { @@ -4889,6 +4990,41 @@ func (u *GroupUpsertBulk) ClearAudioSttPricePerHour() *GroupUpsertBulk { }) } +// SetLongContextPricingEnabled sets the "long_context_pricing_enabled" field. +func (u *GroupUpsertBulk) SetLongContextPricingEnabled(v bool) *GroupUpsertBulk { + return u.Update(func(s *GroupUpsert) { + s.SetLongContextPricingEnabled(v) + }) +} + +// UpdateLongContextPricingEnabled sets the "long_context_pricing_enabled" field to the value that was provided on create. +func (u *GroupUpsertBulk) UpdateLongContextPricingEnabled() *GroupUpsertBulk { + return u.Update(func(s *GroupUpsert) { + s.UpdateLongContextPricingEnabled() + }) +} + +// SetModelPricing sets the "model_pricing" field. +func (u *GroupUpsertBulk) SetModelPricing(v json.RawMessage) *GroupUpsertBulk { + return u.Update(func(s *GroupUpsert) { + s.SetModelPricing(v) + }) +} + +// UpdateModelPricing sets the "model_pricing" field to the value that was provided on create. +func (u *GroupUpsertBulk) UpdateModelPricing() *GroupUpsertBulk { + return u.Update(func(s *GroupUpsert) { + s.UpdateModelPricing() + }) +} + +// ClearModelPricing clears the value of the "model_pricing" field. +func (u *GroupUpsertBulk) ClearModelPricing() *GroupUpsertBulk { + return u.Update(func(s *GroupUpsert) { + s.ClearModelPricing() + }) +} + // SetClaudeCodeOnly sets the "claude_code_only" field. func (u *GroupUpsertBulk) SetClaudeCodeOnly(v bool) *GroupUpsertBulk { return u.Update(func(s *GroupUpsert) { diff --git a/backend/ent/group_update.go b/backend/ent/group_update.go index de5474e069..beca0c71f0 100644 --- a/backend/ent/group_update.go +++ b/backend/ent/group_update.go @@ -4,6 +4,7 @@ package ent import ( "context" + "encoding/json" "errors" "fmt" "time" @@ -787,6 +788,38 @@ func (_u *GroupUpdate) ClearAudioSttPricePerHour() *GroupUpdate { return _u } +// SetLongContextPricingEnabled sets the "long_context_pricing_enabled" field. +func (_u *GroupUpdate) SetLongContextPricingEnabled(v bool) *GroupUpdate { + _u.mutation.SetLongContextPricingEnabled(v) + return _u +} + +// SetNillableLongContextPricingEnabled sets the "long_context_pricing_enabled" field if the given value is not nil. +func (_u *GroupUpdate) SetNillableLongContextPricingEnabled(v *bool) *GroupUpdate { + if v != nil { + _u.SetLongContextPricingEnabled(*v) + } + return _u +} + +// SetModelPricing sets the "model_pricing" field. +func (_u *GroupUpdate) SetModelPricing(v json.RawMessage) *GroupUpdate { + _u.mutation.SetModelPricing(v) + return _u +} + +// AppendModelPricing appends value to the "model_pricing" field. +func (_u *GroupUpdate) AppendModelPricing(v json.RawMessage) *GroupUpdate { + _u.mutation.AppendModelPricing(v) + return _u +} + +// ClearModelPricing clears the value of the "model_pricing" field. +func (_u *GroupUpdate) ClearModelPricing() *GroupUpdate { + _u.mutation.ClearModelPricing() + return _u +} + // SetClaudeCodeOnly sets the "claude_code_only" field. func (_u *GroupUpdate) SetClaudeCodeOnly(v bool) *GroupUpdate { _u.mutation.SetClaudeCodeOnly(v) @@ -1697,6 +1730,20 @@ func (_u *GroupUpdate) sqlSave(ctx context.Context) (_node int, err error) { if _u.mutation.AudioSttPricePerHourCleared() { _spec.ClearField(group.FieldAudioSttPricePerHour, field.TypeFloat64) } + if value, ok := _u.mutation.LongContextPricingEnabled(); ok { + _spec.SetField(group.FieldLongContextPricingEnabled, field.TypeBool, value) + } + if value, ok := _u.mutation.ModelPricing(); ok { + _spec.SetField(group.FieldModelPricing, field.TypeJSON, value) + } + if value, ok := _u.mutation.AppendedModelPricing(); ok { + _spec.AddModifier(func(u *sql.UpdateBuilder) { + sqljson.Append(u, group.FieldModelPricing, value) + }) + } + if _u.mutation.ModelPricingCleared() { + _spec.ClearField(group.FieldModelPricing, field.TypeJSON) + } if value, ok := _u.mutation.ClaudeCodeOnly(); ok { _spec.SetField(group.FieldClaudeCodeOnly, field.TypeBool, value) } @@ -2862,6 +2909,38 @@ func (_u *GroupUpdateOne) ClearAudioSttPricePerHour() *GroupUpdateOne { return _u } +// SetLongContextPricingEnabled sets the "long_context_pricing_enabled" field. +func (_u *GroupUpdateOne) SetLongContextPricingEnabled(v bool) *GroupUpdateOne { + _u.mutation.SetLongContextPricingEnabled(v) + return _u +} + +// SetNillableLongContextPricingEnabled sets the "long_context_pricing_enabled" field if the given value is not nil. +func (_u *GroupUpdateOne) SetNillableLongContextPricingEnabled(v *bool) *GroupUpdateOne { + if v != nil { + _u.SetLongContextPricingEnabled(*v) + } + return _u +} + +// SetModelPricing sets the "model_pricing" field. +func (_u *GroupUpdateOne) SetModelPricing(v json.RawMessage) *GroupUpdateOne { + _u.mutation.SetModelPricing(v) + return _u +} + +// AppendModelPricing appends value to the "model_pricing" field. +func (_u *GroupUpdateOne) AppendModelPricing(v json.RawMessage) *GroupUpdateOne { + _u.mutation.AppendModelPricing(v) + return _u +} + +// ClearModelPricing clears the value of the "model_pricing" field. +func (_u *GroupUpdateOne) ClearModelPricing() *GroupUpdateOne { + _u.mutation.ClearModelPricing() + return _u +} + // SetClaudeCodeOnly sets the "claude_code_only" field. func (_u *GroupUpdateOne) SetClaudeCodeOnly(v bool) *GroupUpdateOne { _u.mutation.SetClaudeCodeOnly(v) @@ -3802,6 +3881,20 @@ func (_u *GroupUpdateOne) sqlSave(ctx context.Context) (_node *Group, err error) if _u.mutation.AudioSttPricePerHourCleared() { _spec.ClearField(group.FieldAudioSttPricePerHour, field.TypeFloat64) } + if value, ok := _u.mutation.LongContextPricingEnabled(); ok { + _spec.SetField(group.FieldLongContextPricingEnabled, field.TypeBool, value) + } + if value, ok := _u.mutation.ModelPricing(); ok { + _spec.SetField(group.FieldModelPricing, field.TypeJSON, value) + } + if value, ok := _u.mutation.AppendedModelPricing(); ok { + _spec.AddModifier(func(u *sql.UpdateBuilder) { + sqljson.Append(u, group.FieldModelPricing, value) + }) + } + if _u.mutation.ModelPricingCleared() { + _spec.ClearField(group.FieldModelPricing, field.TypeJSON) + } if value, ok := _u.mutation.ClaudeCodeOnly(); ok { _spec.SetField(group.FieldClaudeCodeOnly, field.TypeBool, value) } diff --git a/backend/ent/migrate/schema.go b/backend/ent/migrate/schema.go index 491f038bc7..316870e5a5 100644 --- a/backend/ent/migrate/schema.go +++ b/backend/ent/migrate/schema.go @@ -934,6 +934,8 @@ var ( {Name: "audio_realtime_price_per_min", Type: field.TypeFloat64, Nullable: true, SchemaType: map[string]string{"postgres": "decimal(20,8)"}}, {Name: "audio_tts_price_per_million_chars", Type: field.TypeFloat64, Nullable: true, SchemaType: map[string]string{"postgres": "decimal(20,8)"}}, {Name: "audio_stt_price_per_hour", Type: field.TypeFloat64, Nullable: true, SchemaType: map[string]string{"postgres": "decimal(20,8)"}}, + {Name: "long_context_pricing_enabled", Type: field.TypeBool, Default: true}, + {Name: "model_pricing", Type: field.TypeJSON, Nullable: true, SchemaType: map[string]string{"postgres": "jsonb"}}, {Name: "claude_code_only", Type: field.TypeBool, Default: false}, {Name: "fallback_group_id", Type: field.TypeInt64, Nullable: true}, {Name: "fallback_group_id_on_invalid_request", Type: field.TypeInt64, Nullable: true}, @@ -990,7 +992,7 @@ var ( { Name: "group_sort_order", Unique: false, - Columns: []*schema.Column{GroupsColumns[47]}, + Columns: []*schema.Column{GroupsColumns[49]}, }, { Name: "idx_groups_duplicate_operation_id_active", diff --git a/backend/ent/mutation.go b/backend/ent/mutation.go index 99e8a2cd95..da0500d9be 100644 --- a/backend/ent/mutation.go +++ b/backend/ent/mutation.go @@ -21907,6 +21907,9 @@ type GroupMutation struct { addaudio_tts_price_per_million_chars *float64 audio_stt_price_per_hour *float64 addaudio_stt_price_per_hour *float64 + long_context_pricing_enabled *bool + model_pricing *json.RawMessage + appendmodel_pricing json.RawMessage claude_code_only *bool fallback_group_id *int64 addfallback_group_id *int64 @@ -24130,6 +24133,107 @@ func (m *GroupMutation) ResetAudioSttPricePerHour() { delete(m.clearedFields, group.FieldAudioSttPricePerHour) } +// SetLongContextPricingEnabled sets the "long_context_pricing_enabled" field. +func (m *GroupMutation) SetLongContextPricingEnabled(b bool) { + m.long_context_pricing_enabled = &b +} + +// LongContextPricingEnabled returns the value of the "long_context_pricing_enabled" field in the mutation. +func (m *GroupMutation) LongContextPricingEnabled() (r bool, exists bool) { + v := m.long_context_pricing_enabled + if v == nil { + return + } + return *v, true +} + +// OldLongContextPricingEnabled returns the old "long_context_pricing_enabled" field's value of the Group entity. +// If the Group 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 *GroupMutation) OldLongContextPricingEnabled(ctx context.Context) (v bool, err error) { + if !m.op.Is(OpUpdateOne) { + return v, errors.New("OldLongContextPricingEnabled is only allowed on UpdateOne operations") + } + if m.id == nil || m.oldValue == nil { + return v, errors.New("OldLongContextPricingEnabled requires an ID field in the mutation") + } + oldValue, err := m.oldValue(ctx) + if err != nil { + return v, fmt.Errorf("querying old value for OldLongContextPricingEnabled: %w", err) + } + return oldValue.LongContextPricingEnabled, nil +} + +// ResetLongContextPricingEnabled resets all changes to the "long_context_pricing_enabled" field. +func (m *GroupMutation) ResetLongContextPricingEnabled() { + m.long_context_pricing_enabled = nil +} + +// SetModelPricing sets the "model_pricing" field. +func (m *GroupMutation) SetModelPricing(jm json.RawMessage) { + m.model_pricing = &jm + m.appendmodel_pricing = nil +} + +// ModelPricing returns the value of the "model_pricing" field in the mutation. +func (m *GroupMutation) ModelPricing() (r json.RawMessage, exists bool) { + v := m.model_pricing + if v == nil { + return + } + return *v, true +} + +// OldModelPricing returns the old "model_pricing" field's value of the Group entity. +// If the Group 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 *GroupMutation) OldModelPricing(ctx context.Context) (v json.RawMessage, err error) { + if !m.op.Is(OpUpdateOne) { + return v, errors.New("OldModelPricing is only allowed on UpdateOne operations") + } + if m.id == nil || m.oldValue == nil { + return v, errors.New("OldModelPricing requires an ID field in the mutation") + } + oldValue, err := m.oldValue(ctx) + if err != nil { + return v, fmt.Errorf("querying old value for OldModelPricing: %w", err) + } + return oldValue.ModelPricing, nil +} + +// AppendModelPricing adds jm to the "model_pricing" field. +func (m *GroupMutation) AppendModelPricing(jm json.RawMessage) { + m.appendmodel_pricing = append(m.appendmodel_pricing, jm...) +} + +// AppendedModelPricing returns the list of values that were appended to the "model_pricing" field in this mutation. +func (m *GroupMutation) AppendedModelPricing() (json.RawMessage, bool) { + if len(m.appendmodel_pricing) == 0 { + return nil, false + } + return m.appendmodel_pricing, true +} + +// ClearModelPricing clears the value of the "model_pricing" field. +func (m *GroupMutation) ClearModelPricing() { + m.model_pricing = nil + m.appendmodel_pricing = nil + m.clearedFields[group.FieldModelPricing] = struct{}{} +} + +// ModelPricingCleared returns if the "model_pricing" field was cleared in this mutation. +func (m *GroupMutation) ModelPricingCleared() bool { + _, ok := m.clearedFields[group.FieldModelPricing] + return ok +} + +// ResetModelPricing resets all changes to the "model_pricing" field. +func (m *GroupMutation) ResetModelPricing() { + m.model_pricing = nil + m.appendmodel_pricing = nil + delete(m.clearedFields, group.FieldModelPricing) +} + // SetClaudeCodeOnly sets the "claude_code_only" field. func (m *GroupMutation) SetClaudeCodeOnly(b bool) { m.claude_code_only = &b @@ -25435,7 +25539,7 @@ func (m *GroupMutation) Type() string { // order to get all numeric fields that were incremented/decremented, call // AddedFields(). func (m *GroupMutation) Fields() []string { - fields := make([]string, 0, 60) + fields := make([]string, 0, 62) if m.created_at != nil { fields = append(fields, group.FieldCreatedAt) } @@ -25553,6 +25657,12 @@ func (m *GroupMutation) Fields() []string { if m.audio_stt_price_per_hour != nil { fields = append(fields, group.FieldAudioSttPricePerHour) } + if m.long_context_pricing_enabled != nil { + fields = append(fields, group.FieldLongContextPricingEnabled) + } + if m.model_pricing != nil { + fields = append(fields, group.FieldModelPricing) + } if m.claude_code_only != nil { fields = append(fields, group.FieldClaudeCodeOnly) } @@ -25702,6 +25812,10 @@ func (m *GroupMutation) Field(name string) (ent.Value, bool) { return m.AudioTtsPricePerMillionChars() case group.FieldAudioSttPricePerHour: return m.AudioSttPricePerHour() + case group.FieldLongContextPricingEnabled: + return m.LongContextPricingEnabled() + case group.FieldModelPricing: + return m.ModelPricing() case group.FieldClaudeCodeOnly: return m.ClaudeCodeOnly() case group.FieldFallbackGroupID: @@ -25831,6 +25945,10 @@ func (m *GroupMutation) OldField(ctx context.Context, name string) (ent.Value, e return m.OldAudioTtsPricePerMillionChars(ctx) case group.FieldAudioSttPricePerHour: return m.OldAudioSttPricePerHour(ctx) + case group.FieldLongContextPricingEnabled: + return m.OldLongContextPricingEnabled(ctx) + case group.FieldModelPricing: + return m.OldModelPricing(ctx) case group.FieldClaudeCodeOnly: return m.OldClaudeCodeOnly(ctx) case group.FieldFallbackGroupID: @@ -26155,6 +26273,20 @@ func (m *GroupMutation) SetField(name string, value ent.Value) error { } m.SetAudioSttPricePerHour(v) return nil + case group.FieldLongContextPricingEnabled: + v, ok := value.(bool) + if !ok { + return fmt.Errorf("unexpected type %T for field %s", value, name) + } + m.SetLongContextPricingEnabled(v) + return nil + case group.FieldModelPricing: + v, ok := value.(json.RawMessage) + if !ok { + return fmt.Errorf("unexpected type %T for field %s", value, name) + } + m.SetModelPricing(v) + return nil case group.FieldClaudeCodeOnly: v, ok := value.(bool) if !ok { @@ -26713,6 +26845,9 @@ func (m *GroupMutation) ClearedFields() []string { if m.FieldCleared(group.FieldAudioSttPricePerHour) { fields = append(fields, group.FieldAudioSttPricePerHour) } + if m.FieldCleared(group.FieldModelPricing) { + fields = append(fields, group.FieldModelPricing) + } if m.FieldCleared(group.FieldFallbackGroupID) { fields = append(fields, group.FieldFallbackGroupID) } @@ -26790,6 +26925,9 @@ func (m *GroupMutation) ClearField(name string) error { case group.FieldAudioSttPricePerHour: m.ClearAudioSttPricePerHour() return nil + case group.FieldModelPricing: + m.ClearModelPricing() + return nil case group.FieldFallbackGroupID: m.ClearFallbackGroupID() return nil @@ -26924,6 +27062,12 @@ func (m *GroupMutation) ResetField(name string) error { case group.FieldAudioSttPricePerHour: m.ResetAudioSttPricePerHour() return nil + case group.FieldLongContextPricingEnabled: + m.ResetLongContextPricingEnabled() + return nil + case group.FieldModelPricing: + m.ResetModelPricing() + return nil case group.FieldClaudeCodeOnly: m.ResetClaudeCodeOnly() return nil diff --git a/backend/ent/runtime/runtime.go b/backend/ent/runtime/runtime.go index 8042a835a9..4be8d374d6 100644 --- a/backend/ent/runtime/runtime.go +++ b/backend/ent/runtime/runtime.go @@ -1133,80 +1133,84 @@ func init() { groupDescAudioSttPricePerHour := groupFields[35].Descriptor() // group.AudioSttPricePerHourValidator is a validator for the "audio_stt_price_per_hour" field. It is called by the builders before save. group.AudioSttPricePerHourValidator = groupDescAudioSttPricePerHour.Validators[0].(func(float64) error) + // groupDescLongContextPricingEnabled is the schema descriptor for long_context_pricing_enabled field. + groupDescLongContextPricingEnabled := groupFields[36].Descriptor() + // group.DefaultLongContextPricingEnabled holds the default value on creation for the long_context_pricing_enabled field. + group.DefaultLongContextPricingEnabled = groupDescLongContextPricingEnabled.Default.(bool) // groupDescClaudeCodeOnly is the schema descriptor for claude_code_only field. - groupDescClaudeCodeOnly := groupFields[36].Descriptor() + groupDescClaudeCodeOnly := groupFields[38].Descriptor() // group.DefaultClaudeCodeOnly holds the default value on creation for the claude_code_only field. group.DefaultClaudeCodeOnly = groupDescClaudeCodeOnly.Default.(bool) // groupDescModelRoutingEnabled is the schema descriptor for model_routing_enabled field. - groupDescModelRoutingEnabled := groupFields[40].Descriptor() + groupDescModelRoutingEnabled := groupFields[42].Descriptor() // group.DefaultModelRoutingEnabled holds the default value on creation for the model_routing_enabled field. group.DefaultModelRoutingEnabled = groupDescModelRoutingEnabled.Default.(bool) // groupDescMcpXMLInject is the schema descriptor for mcp_xml_inject field. - groupDescMcpXMLInject := groupFields[41].Descriptor() + groupDescMcpXMLInject := groupFields[43].Descriptor() // group.DefaultMcpXMLInject holds the default value on creation for the mcp_xml_inject field. group.DefaultMcpXMLInject = groupDescMcpXMLInject.Default.(bool) // groupDescSupportedModelScopes is the schema descriptor for supported_model_scopes field. - groupDescSupportedModelScopes := groupFields[42].Descriptor() + groupDescSupportedModelScopes := groupFields[44].Descriptor() // group.DefaultSupportedModelScopes holds the default value on creation for the supported_model_scopes field. group.DefaultSupportedModelScopes = groupDescSupportedModelScopes.Default.([]string) // groupDescSortOrder is the schema descriptor for sort_order field. - groupDescSortOrder := groupFields[43].Descriptor() + groupDescSortOrder := groupFields[45].Descriptor() // group.DefaultSortOrder holds the default value on creation for the sort_order field. group.DefaultSortOrder = groupDescSortOrder.Default.(int) // groupDescAllowMessagesDispatch is the schema descriptor for allow_messages_dispatch field. - groupDescAllowMessagesDispatch := groupFields[44].Descriptor() + groupDescAllowMessagesDispatch := groupFields[46].Descriptor() // group.DefaultAllowMessagesDispatch holds the default value on creation for the allow_messages_dispatch field. group.DefaultAllowMessagesDispatch = groupDescAllowMessagesDispatch.Default.(bool) // groupDescAllowLive is the schema descriptor for allow_live field. - groupDescAllowLive := groupFields[45].Descriptor() + groupDescAllowLive := groupFields[47].Descriptor() // group.DefaultAllowLive holds the default value on creation for the allow_live field. group.DefaultAllowLive = groupDescAllowLive.Default.(bool) // groupDescRequireOauthOnly is the schema descriptor for require_oauth_only field. - groupDescRequireOauthOnly := groupFields[46].Descriptor() + groupDescRequireOauthOnly := groupFields[48].Descriptor() // group.DefaultRequireOauthOnly holds the default value on creation for the require_oauth_only field. group.DefaultRequireOauthOnly = groupDescRequireOauthOnly.Default.(bool) // groupDescRequirePrivacySet is the schema descriptor for require_privacy_set field. - groupDescRequirePrivacySet := groupFields[47].Descriptor() + groupDescRequirePrivacySet := groupFields[49].Descriptor() // group.DefaultRequirePrivacySet holds the default value on creation for the require_privacy_set field. group.DefaultRequirePrivacySet = groupDescRequirePrivacySet.Default.(bool) // groupDescDefaultMappedModel is the schema descriptor for default_mapped_model field. - groupDescDefaultMappedModel := groupFields[48].Descriptor() + groupDescDefaultMappedModel := groupFields[50].Descriptor() // group.DefaultDefaultMappedModel holds the default value on creation for the default_mapped_model field. group.DefaultDefaultMappedModel = groupDescDefaultMappedModel.Default.(string) // group.DefaultMappedModelValidator is a validator for the "default_mapped_model" field. It is called by the builders before save. group.DefaultMappedModelValidator = groupDescDefaultMappedModel.Validators[0].(func(string) error) // groupDescMessagesDispatchModelConfig is the schema descriptor for messages_dispatch_model_config field. - groupDescMessagesDispatchModelConfig := groupFields[49].Descriptor() + groupDescMessagesDispatchModelConfig := groupFields[51].Descriptor() // group.DefaultMessagesDispatchModelConfig holds the default value on creation for the messages_dispatch_model_config field. group.DefaultMessagesDispatchModelConfig = groupDescMessagesDispatchModelConfig.Default.(domain.OpenAIMessagesDispatchModelConfig) // groupDescModelsListConfig is the schema descriptor for models_list_config field. - groupDescModelsListConfig := groupFields[50].Descriptor() + groupDescModelsListConfig := groupFields[52].Descriptor() // group.DefaultModelsListConfig holds the default value on creation for the models_list_config field. group.DefaultModelsListConfig = groupDescModelsListConfig.Default.(domain.GroupModelsListConfig) // groupDescRpmLimit is the schema descriptor for rpm_limit field. - groupDescRpmLimit := groupFields[51].Descriptor() + groupDescRpmLimit := groupFields[53].Descriptor() // group.DefaultRpmLimit holds the default value on creation for the rpm_limit field. group.DefaultRpmLimit = groupDescRpmLimit.Default.(int) // groupDescMaxReasoningEffort is the schema descriptor for max_reasoning_effort field. - groupDescMaxReasoningEffort := groupFields[52].Descriptor() + groupDescMaxReasoningEffort := groupFields[54].Descriptor() // group.DefaultMaxReasoningEffort holds the default value on creation for the max_reasoning_effort field. group.DefaultMaxReasoningEffort = groupDescMaxReasoningEffort.Default.(string) // group.MaxReasoningEffortValidator is a validator for the "max_reasoning_effort" field. It is called by the builders before save. group.MaxReasoningEffortValidator = groupDescMaxReasoningEffort.Validators[0].(func(string) error) // groupDescReasoningEffortMappings is the schema descriptor for reasoning_effort_mappings field. - groupDescReasoningEffortMappings := groupFields[53].Descriptor() + groupDescReasoningEffortMappings := groupFields[55].Descriptor() // group.DefaultReasoningEffortMappings holds the default value on creation for the reasoning_effort_mappings field. group.DefaultReasoningEffortMappings = groupDescReasoningEffortMappings.Default.([]domain.ReasoningEffortMapping) // groupDescProfitControlEnabled is the schema descriptor for profit_control_enabled field. - groupDescProfitControlEnabled := groupFields[54].Descriptor() + groupDescProfitControlEnabled := groupFields[56].Descriptor() // group.DefaultProfitControlEnabled holds the default value on creation for the profit_control_enabled field. group.DefaultProfitControlEnabled = groupDescProfitControlEnabled.Default.(bool) // groupDescProfitMinMargin is the schema descriptor for profit_min_margin field. - groupDescProfitMinMargin := groupFields[55].Descriptor() + groupDescProfitMinMargin := groupFields[57].Descriptor() // group.DefaultProfitMinMargin holds the default value on creation for the profit_min_margin field. group.DefaultProfitMinMargin = groupDescProfitMinMargin.Default.(float64) // groupDescProfitSafetyBuffer is the schema descriptor for profit_safety_buffer field. - groupDescProfitSafetyBuffer := groupFields[56].Descriptor() + groupDescProfitSafetyBuffer := groupFields[58].Descriptor() // group.DefaultProfitSafetyBuffer holds the default value on creation for the profit_safety_buffer field. group.DefaultProfitSafetyBuffer = groupDescProfitSafetyBuffer.Default.(float64) idempotencyrecordMixin := schema.IdempotencyRecord{}.Mixin() diff --git a/backend/ent/schema/group.go b/backend/ent/schema/group.go index 85f86891b8..f195ac51db 100644 --- a/backend/ent/schema/group.go +++ b/backend/ent/schema/group.go @@ -1,6 +1,8 @@ package schema import ( + "encoding/json" + "github.com/Wei-Shaw/sub2api/ent/schema/mixins" "github.com/Wei-Shaw/sub2api/internal/domain" @@ -185,6 +187,13 @@ func (Group) Fields() []ent.Field { Min(0). SchemaType(map[string]string{dialect.Postgres: "decimal(20,8)"}). Comment("STT 每小时价格(USD)"), + field.Bool("long_context_pricing_enabled"). + Default(true). + Comment("是否按上下文长度应用模型阶梯价格;默认开启以保持官方/渠道长上下文价"), + field.JSON("model_pricing", json.RawMessage{}). + Optional(). + SchemaType(map[string]string{dialect.Postgres: "jsonb"}). + Comment("分组逐模型定价;优先级高于渠道和内置定价"), // Claude Code 客户端限制 (added by migration 029) field.Bool("claude_code_only"). diff --git a/backend/internal/handler/admin/group_handler.go b/backend/internal/handler/admin/group_handler.go index 8ddbfef183..b8b9587f72 100644 --- a/backend/internal/handler/admin/group_handler.go +++ b/backend/internal/handler/admin/group_handler.go @@ -96,15 +96,17 @@ func NewGroupHandler(adminService service.AdminService, dashboardService *servic // CreateGroupRequest represents create group request type CreateGroupRequest struct { - Name string `json:"name" binding:"required"` - Description string `json:"description"` - Platform string `json:"platform" binding:"omitempty,oneof=anthropic openai gemini antigravity grok composite"` - RateMultiplier float64 `json:"rate_multiplier"` - IsExclusive bool `json:"is_exclusive"` - SubscriptionType string `json:"subscription_type" binding:"omitempty,oneof=standard subscription"` - DailyLimitUSD optionalLimitField `json:"daily_limit_usd"` - WeeklyLimitUSD optionalLimitField `json:"weekly_limit_usd"` - MonthlyLimitUSD optionalLimitField `json:"monthly_limit_usd"` + Name string `json:"name" binding:"required"` + Description string `json:"description"` + Platform string `json:"platform" binding:"omitempty,oneof=anthropic openai gemini antigravity grok composite"` + RateMultiplier float64 `json:"rate_multiplier"` + IsExclusive bool `json:"is_exclusive"` + SubscriptionType string `json:"subscription_type" binding:"omitempty,oneof=standard subscription"` + DailyLimitUSD optionalLimitField `json:"daily_limit_usd"` + WeeklyLimitUSD optionalLimitField `json:"weekly_limit_usd"` + MonthlyLimitUSD optionalLimitField `json:"monthly_limit_usd"` + LongContextPricingEnabled bool `json:"long_context_pricing_enabled"` + ModelPricing []service.ChannelModelPricing `json:"model_pricing"` // 图片生成计费配置(antigravity 和 gemini 平台使用,负数表示清除配置) AllowImageGeneration bool `json:"allow_image_generation"` AllowBatchImageGeneration bool `json:"allow_batch_image_generation"` @@ -162,16 +164,18 @@ type CreateGroupRequest struct { // UpdateGroupRequest represents update group request type UpdateGroupRequest struct { - Name string `json:"name"` - Description *string `json:"description"` - Platform string `json:"platform" binding:"omitempty,oneof=anthropic openai gemini antigravity grok composite"` - RateMultiplier *float64 `json:"rate_multiplier"` - IsExclusive *bool `json:"is_exclusive"` - Status string `json:"status" binding:"omitempty,oneof=active inactive"` - SubscriptionType string `json:"subscription_type" binding:"omitempty,oneof=standard subscription"` - DailyLimitUSD optionalLimitField `json:"daily_limit_usd"` - WeeklyLimitUSD optionalLimitField `json:"weekly_limit_usd"` - MonthlyLimitUSD optionalLimitField `json:"monthly_limit_usd"` + Name string `json:"name"` + Description *string `json:"description"` + Platform string `json:"platform" binding:"omitempty,oneof=anthropic openai gemini antigravity grok composite"` + RateMultiplier *float64 `json:"rate_multiplier"` + IsExclusive *bool `json:"is_exclusive"` + Status string `json:"status" binding:"omitempty,oneof=active inactive"` + SubscriptionType string `json:"subscription_type" binding:"omitempty,oneof=standard subscription"` + DailyLimitUSD optionalLimitField `json:"daily_limit_usd"` + WeeklyLimitUSD optionalLimitField `json:"weekly_limit_usd"` + MonthlyLimitUSD optionalLimitField `json:"monthly_limit_usd"` + LongContextPricingEnabled *bool `json:"long_context_pricing_enabled"` + ModelPricing *[]service.ChannelModelPricing `json:"model_pricing"` // 图片生成计费配置(antigravity 和 gemini 平台使用,负数表示清除配置) AllowImageGeneration *bool `json:"allow_image_generation"` AllowBatchImageGeneration *bool `json:"allow_batch_image_generation"` @@ -508,6 +512,8 @@ func (h *GroupHandler) Create(c *gin.Context) { DailyLimitUSD: req.DailyLimitUSD.ToServiceInput(), WeeklyLimitUSD: req.WeeklyLimitUSD.ToServiceInput(), MonthlyLimitUSD: req.MonthlyLimitUSD.ToServiceInput(), + LongContextPricingEnabled: req.LongContextPricingEnabled, + ModelPricing: req.ModelPricing, AllowImageGeneration: req.AllowImageGeneration, AllowBatchImageGeneration: req.AllowBatchImageGeneration, ImageRateIndependent: req.ImageRateIndependent, @@ -635,6 +641,8 @@ func (h *GroupHandler) Update(c *gin.Context) { DailyLimitUSD: req.DailyLimitUSD.ToServiceInput(), WeeklyLimitUSD: req.WeeklyLimitUSD.ToServiceInput(), MonthlyLimitUSD: req.MonthlyLimitUSD.ToServiceInput(), + LongContextPricingEnabled: req.LongContextPricingEnabled, + ModelPricing: req.ModelPricing, AllowImageGeneration: req.AllowImageGeneration, AllowBatchImageGeneration: req.AllowBatchImageGeneration, ImageRateIndependent: req.ImageRateIndependent, diff --git a/backend/internal/handler/dto/mappers.go b/backend/internal/handler/dto/mappers.go index 0c7d892297..d5a9403a23 100644 --- a/backend/internal/handler/dto/mappers.go +++ b/backend/internal/handler/dto/mappers.go @@ -149,6 +149,7 @@ func GroupFromServiceAdmin(g *service.Group) *AdminGroup { ProfitControlEnabled: g.ProfitControlEnabled, ProfitMinMargin: g.ProfitMinMargin, ProfitSafetyBuffer: g.ProfitSafetyBuffer, + ModelPricing: g.ModelPricing, ModelRouting: g.ModelRouting, ModelRoutingEnabled: g.ModelRoutingEnabled, MCPXMLInject: g.MCPXMLInject, @@ -184,6 +185,7 @@ func groupFromServiceBase(g *service.Group) Group { DailyLimitUSD: g.DailyLimitUSD, WeeklyLimitUSD: g.WeeklyLimitUSD, MonthlyLimitUSD: g.MonthlyLimitUSD, + LongContextPricingEnabled: g.LongContextPricingEnabled, AllowImageGeneration: g.AllowImageGeneration, AllowBatchImageGeneration: g.AllowBatchImageGeneration, ImageRateIndependent: g.ImageRateIndependent, diff --git a/backend/internal/handler/dto/types.go b/backend/internal/handler/dto/types.go index bc2c502c1f..2cd623f7d6 100644 --- a/backend/internal/handler/dto/types.go +++ b/backend/internal/handler/dto/types.go @@ -96,10 +96,11 @@ type Group struct { IsExclusive bool `json:"is_exclusive"` Status string `json:"status"` - SubscriptionType string `json:"subscription_type"` - DailyLimitUSD *float64 `json:"daily_limit_usd"` - WeeklyLimitUSD *float64 `json:"weekly_limit_usd"` - MonthlyLimitUSD *float64 `json:"monthly_limit_usd"` + SubscriptionType string `json:"subscription_type"` + DailyLimitUSD *float64 `json:"daily_limit_usd"` + WeeklyLimitUSD *float64 `json:"weekly_limit_usd"` + MonthlyLimitUSD *float64 `json:"monthly_limit_usd"` + LongContextPricingEnabled bool `json:"long_context_pricing_enabled"` // 图片生成计费配置(仅 antigravity 平台使用) AllowImageGeneration bool `json:"allow_image_generation"` @@ -164,9 +165,10 @@ type AdminGroup struct { // 分组利润控制(五个 token 平台分组可启用;margin/buffer 为小数存储)。 // 仅管理员可见:这三个字段与同响应中的 rate_multiplier 相乘即可反推出 // 运营方的上游成本上限,属于内部经营信息,不得下放到 dto.Group。 - ProfitControlEnabled bool `json:"profit_control_enabled"` - ProfitMinMargin float64 `json:"profit_min_margin"` - ProfitSafetyBuffer float64 `json:"profit_safety_buffer"` + ProfitControlEnabled bool `json:"profit_control_enabled"` + ProfitMinMargin float64 `json:"profit_min_margin"` + ProfitSafetyBuffer float64 `json:"profit_safety_buffer"` + ModelPricing []service.ChannelModelPricing `json:"model_pricing"` // 模型路由配置(仅 anthropic 平台使用) ModelRouting map[string][]int64 `json:"model_routing"` diff --git a/backend/internal/handler/gateway_handler.go b/backend/internal/handler/gateway_handler.go index 1e2d0fbe3d..07315d5efc 100644 --- a/backend/internal/handler/gateway_handler.go +++ b/backend/internal/handler/gateway_handler.go @@ -1249,7 +1249,7 @@ func writeGrokModelsList(c *gin.Context, modelIDs []string) { func grokModelSupportsConfigurableReasoning(modelID string) bool { switch strings.ToLower(strings.TrimSpace(modelID)) { - case "grok-4.5", "grok-4.5-latest", "grok", "grok-latest", "grok-build", "grok-build-latest", "grok-build-0.1": + case "grok-4.6", "grok-4.6-latest", "grok-4.5", "grok-4.5-latest", "grok", "grok-latest", "grok-build", "grok-build-latest", "grok-build-0.1": return true default: return false diff --git a/backend/internal/handler/gateway_web_search.go b/backend/internal/handler/gateway_web_search.go index a0a249f252..741cc322f1 100644 --- a/backend/internal/handler/gateway_web_search.go +++ b/backend/internal/handler/gateway_web_search.go @@ -28,12 +28,8 @@ const ( ) func (h *GatewayHandler) WebSearch(c *gin.Context) { - type webSearchReq struct { - Query string `json:"query" binding:"required"` - MaxResults int `json:"max_results"` - } - - var req webSearchReq + isXSearch := c.GetBool("grok_x_search_endpoint") + var req grokStandaloneSearchRequest if err := c.ShouldBindJSON(&req); err != nil { c.JSON(http.StatusBadRequest, gin.H{"error": gin.H{ "type": "invalid_request_error", @@ -41,7 +37,28 @@ func (h *GatewayHandler) WebSearch(c *gin.Context) { }}) return } - req.MaxResults = normalizeGrokWebSearchMaxResults(req.MaxResults) + query := strings.TrimSpace(req.Query) + if query == "" { + query = strings.TrimSpace(req.Input) + } + if query == "" { + c.JSON(http.StatusBadRequest, gin.H{"error": gin.H{ + "type": "invalid_request_error", + "message": "query is required", + }}) + return + } + req.Query = query + maxResults := 0 + if req.MaxResults != nil { + maxResults = *req.MaxResults + } + maxResults = normalizeGrokWebSearchMaxResults(maxResults) + searchModel := resolveGrokStandaloneSearchModel() + searchLabel := "web_search" + if isXSearch { + searchLabel = "x_search" + } apiKey, ok := middleware2.GetAPIKeyFromContext(c) if !ok || apiKey == nil { @@ -55,7 +72,7 @@ func (h *GatewayHandler) WebSearch(c *gin.Context) { if apiKey.Group == nil || apiKey.Group.Platform != "grok" { c.JSON(http.StatusBadRequest, gin.H{"error": gin.H{ "type": "invalid_request_error", - "message": "web search is only supported for grok groups", + "message": searchLabel + " is only supported for grok groups", }}) return } @@ -79,7 +96,7 @@ func (h *GatewayHandler) WebSearch(c *gin.Context) { "role": "user", "content": req.Query, }}, }) - if decision := h.checkSecurityAudit(c, reqLog, apiKey, subject, service.ContentModerationProtocolOpenAIChat, xai.DefaultTextModel, auditBody); decision != nil && !decision.AllowNextStage { + if decision := h.checkSecurityAudit(c, reqLog, apiKey, subject, service.ContentModerationProtocolOpenAIChat, searchModel, auditBody); decision != nil && !decision.AllowNextStage { status := decision.HTTPStatus if status == 0 { status = http.StatusForbidden @@ -123,7 +140,7 @@ func (h *GatewayHandler) WebSearch(c *gin.Context) { // First attempt + up to 3 failover accounts (max 4 total). for attempt := 0; attempt < 4; attempt++ { selected, selectErr := h.gatewayService.SelectAccountWithLoadAwareness( - c.Request.Context(), groupID, "", xai.DefaultTextModel, failedAccounts, "", 0, + c.Request.Context(), groupID, "", searchModel, failedAccounts, "", 0, ) if selectErr != nil { if attempt == 0 { @@ -159,7 +176,11 @@ func (h *GatewayHandler) WebSearch(c *gin.Context) { account = selected.Account accountReleaseFunc = release - nativeResp, providerName, err = h.doGrokNativeWebSearch(c.Request.Context(), c, account, req.Query, req.MaxResults) + if isXSearch { + nativeResp, providerName, err = h.doGrokNativeXSearch(c.Request.Context(), c, account, req, searchModel, maxResults) + } else { + nativeResp, providerName, err = h.doGrokNativeWebSearch(c.Request.Context(), c, account, req.Query, maxResults, searchModel) + } if err == nil { break } @@ -198,7 +219,7 @@ func (h *GatewayHandler) WebSearch(c *gin.Context) { quotaPlatform := service.QuotaPlatform(c.Request.Context(), apiKey) // Request IDs are billing idempotency keys, so they must be unique per invocation. // Query/IP/UA hashes would collapse repeated identical searches into one charge. - searchRequestID := "web_search:" + uuid.NewString() + searchRequestID := searchLabel + ":" + uuid.NewString() if apiKey.Group != nil { if p := apiKey.Group.GetSearchPricePer1k(); p != nil && *p == 0 { logger.L().With( @@ -211,7 +232,7 @@ func (h *GatewayHandler) WebSearch(c *gin.Context) { if err := h.gatewayService.RecordUsage(ctx, &service.RecordUsageInput{ Result: &service.ForwardResult{ RequestID: searchRequestID, - Model: "grok-web-search", + Model: "grok-" + strings.ReplaceAll(searchLabel, "_", "-"), SearchCount: 1, Duration: 0, }, @@ -240,7 +261,7 @@ func (h *GatewayHandler) WebSearch(c *gin.Context) { "query": req.Query, "results": nativeResp.Results, "provider": providerName, - "max_results": req.MaxResults, + "max_results": maxResults, }) } @@ -299,13 +320,13 @@ func (h *GatewayHandler) acquireWebSearchAccountSlot( // doGrokNativeWebSearch executes web search using the Grok account's native capability // by calling the responses endpoint with web_search tool, then normalizes sources to unified format. -func (h *GatewayHandler) doGrokNativeWebSearch(ctx context.Context, c *gin.Context, account *service.Account, query string, maxResults int) (*websearch.SearchResponse, string, error) { +func (h *GatewayHandler) doGrokNativeWebSearch(ctx context.Context, c *gin.Context, account *service.Account, query string, maxResults int, model string) (*websearch.SearchResponse, string, error) { maxResults = normalizeGrokWebSearchMaxResults(maxResults) // Build a minimal responses request that triggers Grok web search tool. // Ask for structured metadata because xAI action.sources commonly contains URLs only. searchBody := map[string]any{ - "model": xai.DefaultTextModel, + "model": xai.ResolveDefaultTextModel(model), "input": buildGrokWebSearchPrompt(query, maxResults), "tools": []map[string]any{{"type": "web_search"}}, "include": []string{"web_search_call.action.sources"}, @@ -329,6 +350,23 @@ func (h *GatewayHandler) doGrokNativeWebSearch(ctx context.Context, c *gin.Conte }, "grok-native", nil } +func (h *GatewayHandler) doGrokNativeXSearch(ctx context.Context, c *gin.Context, account *service.Account, req grokStandaloneSearchRequest, model string, maxResults int) (*websearch.SearchResponse, string, error) { + maxResults = normalizeGrokWebSearchMaxResults(maxResults) + bodyBytes, err := buildGrokXSearchResponsesBody(req, model) + if err != nil { + return nil, "", err + } + respBytes, err := h.gatewayService.DoGrokNativeResponsesJSON(ctx, account, bodyBytes) + if err != nil { + return nil, "", err + } + results := extractGrokWebSearchSources(respBytes, maxResults) + return &websearch.SearchResponse{ + Results: results, + Query: req.Query, + }, "grok-native", nil +} + func normalizeGrokWebSearchMaxResults(maxResults int) int { if maxResults <= 0 { return defaultGrokWebSearchResults @@ -377,7 +415,8 @@ func extractGrokWebSearchSources(body []byte, maxResults int) []websearch.Search output := gjson.GetBytes(body, "output") output.ForEach(func(_, item gjson.Result) bool { - if item.Get("type").String() == "web_search_call" { + callType := item.Get("type").String() + if callType == "web_search_call" || callType == "x_search_call" { sources := item.Get("action.sources") if sources.IsArray() { sources.ForEach(func(_, src gjson.Result) bool { diff --git a/backend/internal/handler/grok_audio.go b/backend/internal/handler/grok_audio.go index 542d70d423..4ac1ae53bf 100644 --- a/backend/internal/handler/grok_audio.go +++ b/backend/internal/handler/grok_audio.go @@ -89,7 +89,7 @@ func (h *OpenAIGatewayHandler) GrokRealtime(c *gin.Context) { model = "grok-voice-latest" } started := time.Now() - proxyErr := h.gatewayService.ProxyGrokRealtime(c.Request.Context(), c, conn, selection.Account, token, model) + audioObserved, proxyErr := h.gatewayService.ProxyGrokRealtime(c.Request.Context(), c, conn, selection.Account, token, model) elapsed := time.Since(started) if proxyErr != nil { reqLog.Info("grok_realtime.proxy_failed", zap.Error(proxyErr)) @@ -98,20 +98,23 @@ func (h *OpenAIGatewayHandler) GrokRealtime(c *gin.Context) { return } } - // A relay normally returns a close error when either side closes normally. - // Those sessions still consumed upstream audio time and must be billed. - if elapsed > 0 { - result := &service.OpenAIForwardResult{ - // One durable id per WS session so retries cannot collapse or double under client ids. - RequestID: service.StableGrokRealtimeBillingRequestID(""), - Model: model, - Duration: elapsed, - AudioUsage: &service.AudioUsage{Mode: "realtime", DurationOrUnits: elapsed.Minutes()}, - } + if result := grokRealtimeBillingResult(model, elapsed, audioObserved); result != nil { h.recordGrokVoiceUsage(c, apiKey, selection.Account, subscription, "realtime", nil, result) } } +func grokRealtimeBillingResult(model string, elapsed time.Duration, audioObserved bool) *service.OpenAIForwardResult { + if !audioObserved || elapsed <= 0 { + return nil + } + return &service.OpenAIForwardResult{ + RequestID: service.StableGrokRealtimeBillingRequestID(""), + Model: model, + Duration: elapsed, + AudioUsage: &service.AudioUsage{Mode: "realtime", DurationOrUnits: elapsed.Minutes()}, + } +} + func isExpectedGrokRealtimeClose(err error) bool { if err == nil { return true diff --git a/backend/internal/handler/grok_audio_billing_test.go b/backend/internal/handler/grok_audio_billing_test.go index 6b8a6c81be..8ec5d08af0 100644 --- a/backend/internal/handler/grok_audio_billing_test.go +++ b/backend/internal/handler/grok_audio_billing_test.go @@ -4,6 +4,7 @@ package handler import ( "testing" + "time" coderws "github.com/coder/websocket" ) @@ -23,3 +24,29 @@ func TestIsExpectedGrokRealtimeClose(t *testing.T) { t.Fatal("policy violations must not be treated as billable normal closes") } } + +func TestGrokRealtimeBillingResultRequiresObservedAudio(t *testing.T) { + if grokRealtimeBillingResult("grok-voice-latest", time.Second, false) != nil { + t.Fatal("a session without observed audio must not be billed") + } + if grokRealtimeBillingResult("grok-voice-latest", 0, true) != nil { + t.Fatal("zero-duration sessions must not be billed") + } +} + +func TestGrokRealtimeBillingResultUsesForcedUniqueID(t *testing.T) { + first := grokRealtimeBillingResult("grok-voice-latest", 90*time.Second, true) + second := grokRealtimeBillingResult("grok-voice-latest", 90*time.Second, true) + if first == nil || second == nil { + t.Fatal("observed audio sessions should be billable") + } + if first.RequestID == "" { + t.Fatalf("unexpected billing request ID %q", first.RequestID) + } + if first.RequestID == second.RequestID { + t.Fatal("independent realtime connections must not share a billing request ID") + } + if first.AudioUsage == nil || first.AudioUsage.Mode != "realtime" || first.AudioUsage.DurationOrUnits != 1.5 { + t.Fatalf("unexpected audio usage: %#v", first.AudioUsage) + } +} diff --git a/backend/internal/handler/openai_x_search.go b/backend/internal/handler/openai_x_search.go new file mode 100644 index 0000000000..a583cec0bf --- /dev/null +++ b/backend/internal/handler/openai_x_search.go @@ -0,0 +1,79 @@ +package handler + +import ( + "encoding/json" + "fmt" + "strings" + + "github.com/Wei-Shaw/sub2api/internal/pkg/xai" + "github.com/gin-gonic/gin" +) + +type grokStandaloneSearchRequest struct { + Query string `json:"query"` + Input string `json:"input"` + MaxResults *int `json:"max_results"` + AllowedXHandles []string `json:"allowed_x_handles"` + ExcludedXHandles []string `json:"excluded_x_handles"` + FromDate string `json:"from_date"` + ToDate string `json:"to_date"` + EnableImageUnderstanding *bool `json:"enable_image_understanding"` + EnableVideoUnderstanding *bool `json:"enable_video_understanding"` +} + +// XSearch marks the standalone endpoint so WebSearch can use native x_search +// while retaining its dedicated per-call billing contract. +func (h *GatewayHandler) XSearch(c *gin.Context) { + c.Set("grok_x_search_endpoint", true) + h.WebSearch(c) +} + +func resolveGrokStandaloneSearchModel() string { + return xai.ResolveDefaultTextModel(xai.RuntimeModelMappingOptions().DefaultText) +} + +func buildGrokXSearchResponsesBody(req grokStandaloneSearchRequest, model string) ([]byte, error) { + input := strings.TrimSpace(req.Query) + if input == "" { + input = strings.TrimSpace(req.Input) + } + tool := map[string]any{"type": "x_search"} + if len(req.AllowedXHandles) > 0 { + tool["allowed_x_handles"] = req.AllowedXHandles + } + if len(req.ExcludedXHandles) > 0 { + tool["excluded_x_handles"] = req.ExcludedXHandles + } + if strings.TrimSpace(req.FromDate) != "" { + tool["from_date"] = strings.TrimSpace(req.FromDate) + } + if strings.TrimSpace(req.ToDate) != "" { + tool["to_date"] = strings.TrimSpace(req.ToDate) + } + if req.EnableImageUnderstanding != nil { + tool["enable_image_understanding"] = *req.EnableImageUnderstanding + } + if req.EnableVideoUnderstanding != nil { + tool["enable_video_understanding"] = *req.EnableVideoUnderstanding + } + maxResults := 0 + if req.MaxResults != nil { + maxResults = *req.MaxResults + } + return json.Marshal(map[string]any{ + "model": xai.ResolveDefaultTextModel(model), + "input": buildGrokXSearchPrompt(input, maxResults), + "tools": []map[string]any{tool}, + "tool_choice": "required", + "include": []string{"x_search_call.action.sources"}, + "store": false, + "stream": false, + }) +} + +func buildGrokXSearchPrompt(query string, maxResults int) string { + return fmt.Sprintf(`Search X for the user query below. Return ONLY valid JSON with this exact shape: {"results":[{"url":"https://...","title":"post or page title","snippet":"concise factual summary"}]}. Return at most %d unique results. Every URL must be an actual x_search source. Populate a non-empty title and snippet for every result. Do not wrap the JSON in markdown. + +User query: +%s`, normalizeGrokWebSearchMaxResults(maxResults), query) +} diff --git a/backend/internal/handler/openai_x_search_test.go b/backend/internal/handler/openai_x_search_test.go new file mode 100644 index 0000000000..6747e6c803 --- /dev/null +++ b/backend/internal/handler/openai_x_search_test.go @@ -0,0 +1,58 @@ +package handler + +import ( + "testing" + + "github.com/Wei-Shaw/sub2api/internal/pkg/xai" + "github.com/stretchr/testify/require" + "github.com/tidwall/gjson" +) + +func TestBuildGrokXSearchResponsesBody(t *testing.T) { + t.Parallel() + understandImages := true + understandVideos := false + body, err := buildGrokXSearchResponsesBody(grokStandaloneSearchRequest{ + Query: "latest posts from xAI", + AllowedXHandles: []string{"xai"}, + ExcludedXHandles: []string{"spam"}, + FromDate: "2026-08-01", + ToDate: "2026-08-10", + EnableImageUnderstanding: &understandImages, + EnableVideoUnderstanding: &understandVideos, + }, xai.DefaultTextModel) + require.NoError(t, err) + require.Equal(t, xai.DefaultTextModel, gjson.GetBytes(body, "model").String()) + require.Contains(t, gjson.GetBytes(body, "input").String(), "latest posts from xAI") + require.Contains(t, gjson.GetBytes(body, "input").String(), "Return ONLY valid JSON") + require.Equal(t, "x_search_call.action.sources", gjson.GetBytes(body, "include.0").String()) + require.Equal(t, "required", gjson.GetBytes(body, "tool_choice").String()) + require.Equal(t, "x_search", gjson.GetBytes(body, "tools.0.type").String()) + require.Equal(t, "xai", gjson.GetBytes(body, "tools.0.allowed_x_handles.0").String()) + require.Equal(t, "spam", gjson.GetBytes(body, "tools.0.excluded_x_handles.0").String()) + require.Equal(t, "2026-08-01", gjson.GetBytes(body, "tools.0.from_date").String()) + require.Equal(t, "2026-08-10", gjson.GetBytes(body, "tools.0.to_date").String()) + require.True(t, gjson.GetBytes(body, "tools.0.enable_image_understanding").Bool()) + require.False(t, gjson.GetBytes(body, "tools.0.enable_video_understanding").Bool()) + require.False(t, gjson.GetBytes(body, "store").Bool()) + require.False(t, gjson.GetBytes(body, "stream").Bool()) +} + +func TestBuildGrokXSearchResponsesBodyAcceptsInputAlias(t *testing.T) { + t.Parallel() + body, err := buildGrokXSearchResponsesBody(grokStandaloneSearchRequest{Input: "latest posts from xAI"}, xai.DefaultTextModel) + require.NoError(t, err) + require.Contains(t, gjson.GetBytes(body, "input").String(), "latest posts from xAI") +} + +func TestResolveGrokStandaloneSearchModelUsesRuntimeDefault(t *testing.T) { + original := xai.RuntimeModelMappingOptions() + t.Cleanup(func() { xai.SetRuntimeModelMappingOptions(original) }) + xai.SetRuntimeModelMappingOptions(xai.ModelMappingOptions{DefaultText: "grok-4.6"}) + + model := resolveGrokStandaloneSearchModel() + body, err := buildGrokXSearchResponsesBody(grokStandaloneSearchRequest{Query: "latest posts from xAI"}, model) + require.NoError(t, err) + require.Equal(t, "grok-4.6", model) + require.Equal(t, model, gjson.GetBytes(body, "model").String()) +} diff --git a/backend/internal/pkg/apicompat/chatcompletions_responses_bridge.go b/backend/internal/pkg/apicompat/chatcompletions_responses_bridge.go index 4e20fe4216..70dee035d8 100644 --- a/backend/internal/pkg/apicompat/chatcompletions_responses_bridge.go +++ b/backend/internal/pkg/apicompat/chatcompletions_responses_bridge.go @@ -63,6 +63,9 @@ func ResponsesToChatCompletionsRequest(req *ResponsesRequest) (*ChatCompletionsR if tool.Function != nil { declared[tool.Function.Name] = true } + if strings.EqualFold(strings.TrimSpace(tool.Type), "x_search") { + declared["x_search"] = true + } } if tc := responsesToolChoiceToChatToolChoice(req.ToolChoice, declared); len(tc) > 0 { out.ToolChoice = tc @@ -847,6 +850,16 @@ func responsesToolsToChatTools(tools []ResponsesTool) ([]ChatTool, error) { return nil, err } out = append(out, flattened...) + case "x_search": + out = append(out, ChatTool{ + Type: "x_search", + AllowedXHandles: tool.AllowedXHandles, + ExcludedXHandles: tool.ExcludedXHandles, + FromDate: tool.FromDate, + ToDate: tool.ToDate, + EnableImageUnderstanding: tool.EnableImageUnderstanding, + EnableVideoUnderstanding: tool.EnableVideoUnderstanding, + }) } // 其余类型(web_search、image_generation 等服务端工具)在 chat 上游没有 // 对应能力,维持丢弃。 @@ -948,6 +961,15 @@ func responsesToolChoiceToChatToolChoice(raw json.RawMessage, declared map[strin } var name string switch rawString(choice["type"]) { + case "x_search": + if !declared["x_search"] { + return nil + } + out, err := json.Marshal(map[string]any{"type": "x_search"}) + if err != nil { + return raw + } + return out case "tool_search": // tool_search 未被丢弃而是降级为同名 function 代理(见 // responsesToolsToChatTools),强制选择它同样降级为 function 选择, diff --git a/backend/internal/pkg/apicompat/chatcompletions_responses_bridge_custom_tools_test.go b/backend/internal/pkg/apicompat/chatcompletions_responses_bridge_custom_tools_test.go index 6265a64d0a..945900d6d6 100644 --- a/backend/internal/pkg/apicompat/chatcompletions_responses_bridge_custom_tools_test.go +++ b/backend/internal/pkg/apicompat/chatcompletions_responses_bridge_custom_tools_test.go @@ -661,6 +661,20 @@ func TestResponsesToChatCompletionsRequest_DropsToolChoiceForDroppedTool(t *test require.Len(t, out.Tools, 1) assert.Empty(t, out.ToolChoice, "指向被丢弃服务端工具的 tool_choice 必须丢弃") + out, err = ResponsesToChatCompletionsRequest(&ResponsesRequest{ + Model: "glm-5.2", + Input: json.RawMessage(`"hi"`), + Tools: []ResponsesTool{ + {Type: "function", Name: "wait", Parameters: json.RawMessage(`{"type":"object","properties":{}}`)}, + {Type: "web_search"}, + {Type: "x_search"}, + }, + ToolChoice: json.RawMessage(`{"type":"function","name":"web_search"}`), + }) + require.NoError(t, err) + require.Len(t, out.Tools, 2) + assert.Empty(t, out.ToolChoice, "surviving x_search must not keep a function tool_choice named web_search") + // 具名选择指向不存在的工具名。 out, err = ResponsesToChatCompletionsRequest(&ResponsesRequest{ Model: "glm-5.2", diff --git a/backend/internal/pkg/apicompat/chatcompletions_to_responses.go b/backend/internal/pkg/apicompat/chatcompletions_to_responses.go index 0f65d217e8..6a3c751f13 100644 --- a/backend/internal/pkg/apicompat/chatcompletions_to_responses.go +++ b/backend/internal/pkg/apicompat/chatcompletions_to_responses.go @@ -419,6 +419,18 @@ func convertChatToolsToResponses(tools []ChatTool, functions []ChatFunction) []R var out []ResponsesTool for _, t := range tools { + if strings.EqualFold(strings.TrimSpace(t.Type), "x_search") { + out = append(out, ResponsesTool{ + Type: "x_search", + AllowedXHandles: t.AllowedXHandles, + ExcludedXHandles: t.ExcludedXHandles, + FromDate: t.FromDate, + ToDate: t.ToDate, + EnableImageUnderstanding: t.EnableImageUnderstanding, + EnableVideoUnderstanding: t.EnableVideoUnderstanding, + }) + continue + } if t.Type != "function" || t.Function == nil { continue } diff --git a/backend/internal/pkg/apicompat/chatcompletions_x_search_test.go b/backend/internal/pkg/apicompat/chatcompletions_x_search_test.go new file mode 100644 index 0000000000..2ffbea2a3b --- /dev/null +++ b/backend/internal/pkg/apicompat/chatcompletions_x_search_test.go @@ -0,0 +1,79 @@ +package apicompat + +import ( + "encoding/json" + "testing" + + "github.com/stretchr/testify/require" +) + +func TestChatCompletionsToResponsesPreservesXSearchTool(t *testing.T) { + enabled := true + req := &ChatCompletionsRequest{ + Model: "grok-4.5", + Messages: []ChatMessage{ + {Role: "user", Content: json.RawMessage(`"latest xAI post"`)}, + }, + Tools: []ChatTool{{ + Type: "x_search", + AllowedXHandles: []string{"xai"}, + ExcludedXHandles: []string{"spam"}, + FromDate: "2026-08-01", + ToDate: "2026-08-10", + EnableImageUnderstanding: &enabled, + EnableVideoUnderstanding: &enabled, + }}, + ToolChoice: json.RawMessage(`{"type":"x_search"}`), + } + + resp, err := ChatCompletionsToResponses(req) + require.NoError(t, err) + require.Len(t, resp.Tools, 1) + require.Equal(t, "x_search", resp.Tools[0].Type) + require.Equal(t, []string{"xai"}, resp.Tools[0].AllowedXHandles) + require.Equal(t, []string{"spam"}, resp.Tools[0].ExcludedXHandles) + require.Equal(t, "2026-08-01", resp.Tools[0].FromDate) + require.Equal(t, "2026-08-10", resp.Tools[0].ToDate) + require.NotNil(t, resp.Tools[0].EnableImageUnderstanding) + require.True(t, *resp.Tools[0].EnableImageUnderstanding) + require.NotNil(t, resp.Tools[0].EnableVideoUnderstanding) + require.True(t, *resp.Tools[0].EnableVideoUnderstanding) + require.JSONEq(t, `{"type":"x_search"}`, string(resp.ToolChoice)) +} + +func TestResponsesToChatCompletionsPreservesXSearchTool(t *testing.T) { + enabled := true + req := &ResponsesRequest{ + Model: "grok-4.5", + Input: json.RawMessage(`"latest xAI post"`), + Tools: []ResponsesTool{{ + Type: "x_search", + AllowedXHandles: []string{"xai"}, + ExcludedXHandles: []string{"spam"}, + FromDate: "2026-08-01", + ToDate: "2026-08-10", + EnableImageUnderstanding: &enabled, + EnableVideoUnderstanding: &enabled, + }}, + ToolChoice: json.RawMessage(`{"type":"x_search"}`), + } + + chat, err := ResponsesToChatCompletionsRequest(req) + require.NoError(t, err) + require.Len(t, chat.Tools, 1) + require.Equal(t, "x_search", chat.Tools[0].Type) + require.Equal(t, []string{"xai"}, chat.Tools[0].AllowedXHandles) + require.Equal(t, []string{"spam"}, chat.Tools[0].ExcludedXHandles) + require.JSONEq(t, `{"type":"x_search"}`, string(chat.ToolChoice)) +} + +func TestResponsesToChatCompletionsXSearchToolChoiceString(t *testing.T) { + chat, err := ResponsesToChatCompletionsRequest(&ResponsesRequest{ + Model: "grok-4.5", + Input: json.RawMessage(`"latest xAI post"`), + Tools: []ResponsesTool{{Type: "x_search"}}, + ToolChoice: json.RawMessage(`"x_search"`), + }) + require.NoError(t, err) + require.JSONEq(t, `"x_search"`, string(chat.ToolChoice)) +} diff --git a/backend/internal/pkg/apicompat/types.go b/backend/internal/pkg/apicompat/types.go index c484a7e589..bcc8cd481e 100644 --- a/backend/internal/pkg/apicompat/types.go +++ b/backend/internal/pkg/apicompat/types.go @@ -302,7 +302,7 @@ type ResponsesContentPart struct { // ResponsesTool describes a tool in the Responses API. type ResponsesTool struct { - Type string `json:"type"` // "function" | "custom" | "web_search" | "local_shell" etc. + Type string `json:"type"` // "function" | "custom" | "web_search" | "x_search" | "local_shell" etc. Name string `json:"name,omitempty"` Description string `json:"description,omitempty"` Parameters json.RawMessage `json:"parameters,omitempty"` @@ -311,6 +311,14 @@ type ResponsesTool struct { // type=namespace 的子工具列表(tools 与 children 二选一,语义相同)。 Tools []ResponsesTool `json:"tools,omitempty"` Children []ResponsesTool `json:"children,omitempty"` + + // type=x_search + AllowedXHandles []string `json:"allowed_x_handles,omitempty"` + ExcludedXHandles []string `json:"excluded_x_handles,omitempty"` + FromDate string `json:"from_date,omitempty"` + ToDate string `json:"to_date,omitempty"` + EnableImageUnderstanding *bool `json:"enable_image_understanding,omitempty"` + EnableVideoUnderstanding *bool `json:"enable_video_understanding,omitempty"` } // UnmarshalJSON 容忍字符串形式的工具声明:codex 会以 "name" 简写声明 custom 工具, @@ -675,8 +683,16 @@ type ChatImageURL struct { // ChatTool describes a tool available to the model. type ChatTool struct { - Type string `json:"type"` // "function" + Type string `json:"type"` // "function" | "x_search" Function *ChatFunction `json:"function,omitempty"` + + // type=x_search + AllowedXHandles []string `json:"allowed_x_handles,omitempty"` + ExcludedXHandles []string `json:"excluded_x_handles,omitempty"` + FromDate string `json:"from_date,omitempty"` + ToDate string `json:"to_date,omitempty"` + EnableImageUnderstanding *bool `json:"enable_image_understanding,omitempty"` + EnableVideoUnderstanding *bool `json:"enable_video_understanding,omitempty"` } // ChatFunction describes a function tool definition. diff --git a/backend/internal/pkg/xai/models.go b/backend/internal/pkg/xai/models.go index c25f489317..990a6d28e0 100644 --- a/backend/internal/pkg/xai/models.go +++ b/backend/internal/pkg/xai/models.go @@ -83,6 +83,7 @@ func (o ModelMappingOptions) defaultText() string { var defaultModels = []Model{ // Text + {ID: "grok-4.6", Object: "model", Type: "model", OwnedBy: "xai", DisplayName: "Grok 4.6"}, {ID: "grok-4.5", Object: "model", Type: "model", OwnedBy: "xai", DisplayName: "Grok 4.5"}, {ID: "grok-4.3", Object: "model", Type: "model", OwnedBy: "xai", DisplayName: "Grok 4.3"}, {ID: "grok-3-mini", Object: "model", Type: "model", OwnedBy: "xai", DisplayName: "Grok 3 Mini"}, @@ -106,6 +107,8 @@ var defaultModels = []Model{ var grokTextResponsesModelAliases = map[string]string{ "grok": DefaultTextModel, "grok-latest": DefaultTextModel, + "grok-4.6": "grok-4.6", + "grok-4.6-latest": "grok-4.6", "grok-4.5": DefaultTextModel, "grok-4.5-latest": DefaultTextModel, "grok-4.3": "grok-4.3", diff --git a/backend/internal/pkg/xai/models_test.go b/backend/internal/pkg/xai/models_test.go index 98aef08cfc..a01d993ccb 100644 --- a/backend/internal/pkg/xai/models_test.go +++ b/backend/internal/pkg/xai/models_test.go @@ -52,11 +52,20 @@ func TestCanonicalImagineVideoModel(t *testing.T) { func TestIsGrokModelID(t *testing.T) { t.Parallel() require.True(t, IsGrokModelID("grok-4.5")) + require.True(t, IsGrokModelID("grok-4.6")) require.True(t, IsGrokModelID("x-ai/grok-4.3")) require.False(t, IsGrokModelID("gpt-5")) require.False(t, IsGrokModelID("claude-sonnet-4")) } +func TestDefaultModelsIncludesGrok46(t *testing.T) { + t.Parallel() + ids := DefaultModelIDs() + require.Contains(t, ids, "grok-4.6") + require.Equal(t, "grok-4.6", ResolveGrokTextResponsesModelID("grok-4.6")) + require.Equal(t, "grok-4.6", ResolveGrokTextResponsesModelID("grok-4.6-latest")) +} + func TestResolveGrokTextResponsesModelID(t *testing.T) { t.Parallel() require.Equal(t, "grok-4.5", ResolveGrokTextResponsesModelID("")) diff --git a/backend/internal/pkg/xai/oauth_test.go b/backend/internal/pkg/xai/oauth_test.go index 3919f4576b..90036ff8f2 100644 --- a/backend/internal/pkg/xai/oauth_test.go +++ b/backend/internal/pkg/xai/oauth_test.go @@ -352,6 +352,8 @@ func TestDefaultModelMappingIncludesGrokAliases(t *testing.T) { mapping := DefaultModelMapping() require.Equal(t, "grok-4.5", mapping["grok"]) require.Equal(t, "grok-4.5", mapping["grok-latest"]) + require.Equal(t, "grok-4.6", mapping["grok-4.6"]) + require.Equal(t, "grok-4.6", mapping["grok-4.6-latest"]) require.Equal(t, "grok-4.5", mapping["grok-4.5"]) require.Equal(t, "grok-4.5", mapping["grok-4.5-latest"]) require.Equal(t, "grok-build-0.1", mapping["grok-build"]) diff --git a/backend/internal/pkg/xai/quota.go b/backend/internal/pkg/xai/quota.go index 11133a05b3..9825f66343 100644 --- a/backend/internal/pkg/xai/quota.go +++ b/backend/internal/pkg/xai/quota.go @@ -43,6 +43,12 @@ type QuotaSnapshot struct { LastProbeAt string `json:"last_probe_at,omitempty"` LastHeadersSeenAt string `json:"last_headers_seen_at,omitempty"` UpdatedAt string `json:"updated_at"` + // Model is the upstream id that produced these rate-limit headers. + Model string `json:"model,omitempty"` + // PlanFrom45Responses is inferred from a grok-4.5 Responses window + // (8300/53M = Heavy). Carried across later non-4.5 overwrites. + PlanFrom45Responses string `json:"plan_from_45_responses,omitempty"` + PlanFrom45ResponsesAt string `json:"plan_from_45_responses_at,omitempty"` } func (s *QuotaSnapshot) HasObservedHeaders() bool { diff --git a/backend/internal/pkg/xai/subscription_tier.go b/backend/internal/pkg/xai/subscription_tier.go new file mode 100644 index 0000000000..9444e6a3f4 --- /dev/null +++ b/backend/internal/pkg/xai/subscription_tier.go @@ -0,0 +1,279 @@ +package xai + +import ( + "encoding/json" + "strconv" + "strings" + "time" +) + +// GrokQuotaSignalMaxAge bounds how long a grok-4.5 Responses window can +// influence SuperGrok vs Heavy inference. +const GrokQuotaSignalMaxAge = 24 * time.Hour + +const ( + grok45ResponsesModel = "grok-4.5" + grokHeavyQuotaRequestLimit int64 = 8_300 + grokHeavyQuotaTokenLimit int64 = 53_000_000 +) + +// MapJWTSubscriptionTier maps prod_auth.SubscriptionTier numeric JWT claims +// to stable snake_case keys used by Grok Build / Mixpanel. +func MapJWTSubscriptionTier(tier uint64) string { + switch tier { + case 0: + return "free" + case 1: + return "supergrok" + case 2: + return "x_basic" + case 3: + return "x_premium" + case 4: + return "x_premium_plus" + case 5: + return "supergrok_heavy" + case 6: + return "supergrok_lite" + case 7: + return "supergrok_plus" + default: + return strconv.FormatUint(tier, 10) + } +} + +// NormalizeSubscriptionTier canonicalizes display names, /user strings, and +// JWT-derived keys onto the same snake_case identifiers. +func NormalizeSubscriptionTier(raw string) string { + t := strings.ToLower(strings.TrimSpace(raw)) + t = strings.ReplaceAll(t, "-", "_") + t = strings.Join(strings.Fields(t), "_") + switch t { + case "free", "grok_free", "grokfree", "free_tier", "freetier", "grok_basic", "grokbasic": + return "free" + case "supergrok", "grokpro": + return "supergrok" + case "supergrok_lite", "supergroklite": + return "supergrok_lite" + case "supergrok_heavy", "supergrokheavy": + return "supergrok_heavy" + case "supergrok_pro", "supergrokpro": + return "supergrok_pro" + case "supergrok_plus", "supergrokplus": + return "supergrok_plus" + case "x_basic", "xbasic", "basic": + return "x_basic" + case "x_premium", "xpremium": + return "x_premium" + case "x_premium_plus", "xpremiumplus", "x_premium+": + return "x_premium_plus" + default: + return t + } +} + +// SubscriptionTierFromJWT decodes an access token payload (no signature check) +// and maps the numeric or string `tier` claim. +func SubscriptionTierFromJWT(jwt string) string { + claims := DecodeJWTClaims(jwt) + if claims == nil { + return "" + } + raw, ok := claims["tier"] + if !ok || raw == nil { + return "" + } + switch v := raw.(type) { + case float64: + if v < 0 { + return "" + } + return MapJWTSubscriptionTier(uint64(v)) + case json.Number: + n, err := v.Int64() + if err != nil || n < 0 { + return NormalizeSubscriptionTier(v.String()) + } + return MapJWTSubscriptionTier(uint64(n)) + case string: + trimmed := strings.TrimSpace(v) + if trimmed == "" { + return "" + } + if n, err := strconv.ParseUint(trimmed, 10, 64); err == nil { + return MapJWTSubscriptionTier(n) + } + return NormalizeSubscriptionTier(trimmed) + default: + return "" + } +} + +// CanonicalGrokPlan resolves SuperGrok vs Heavy when the provider label is +// ambiguous (SuperGrokPro). JWT numeric claims are applied by the caller first. +// Monthly $150/$1500 limits still win when present. +// Rate-limit windows are only used when they came from grok-4.5 Responses. +func CanonicalGrokPlan(monthlyLimitCents *float64, subscriptionTier string, quota *QuotaSnapshot) string { + if plan := resolvePlan(monthlyLimitCents); plan != "" { + return NormalizeSubscriptionTier(plan) + } + + normalized := NormalizeSubscriptionTier(subscriptionTier) + switch normalized { + case "free", "x_basic": + return "free" + case "supergrok_heavy": + return "supergrok_heavy" + case "supergrok_lite": + return "supergrok_lite" + case "supergrok_plus": + return "supergrok_plus" + } + + if isAmbiguousGrokPaidPlan(normalized) { + if hint := Grok45ResponsesPlanHint(quota, time.Time{}); hint != "" { + return hint + } + return "supergrok" + } + return "" +} + +func isAmbiguousGrokPaidPlan(normalized string) bool { + switch normalized { + case "supergrok", "supergrok_pro", "paid", "pro": + return true + default: + return false + } +} + +// IsGrok45ResponsesQuotaModel reports whether model is the grok-4.5 Responses +// id (or a dated grok-4.5-* variant). Empty and other families are false. +func IsGrok45ResponsesQuotaModel(model string) bool { + m := strings.ToLower(strings.TrimSpace(StripGrokProviderPrefix(model))) + return m == grok45ResponsesModel || strings.HasPrefix(m, grok45ResponsesModel+"-") +} + +// Grok45ResponsesPlanHint returns SuperGrok / Heavy inferred from a grok-4.5 +// Responses window. Other models' limits are ignored. +func Grok45ResponsesPlanHint(quota *QuotaSnapshot, now time.Time) string { + if quota == nil { + return "" + } + if plan := NormalizeSubscriptionTier(quota.PlanFrom45Responses); plan == "supergrok" || plan == "supergrok_heavy" { + if isQuotaTimestampFresh(quota.PlanFrom45ResponsesAt, now) { + return plan + } + } + if !IsGrok45ResponsesQuotaModel(quota.Model) || !IsQuotaSnapshotFresh(quota, now) { + return "" + } + if quotaLooksLikeGrokHeavy(quota) { + return "supergrok_heavy" + } + return "" +} + +// ApplyGrok45ResponsesPlanSignal records a grok-4.5 Heavy/SuperGrok hint, or +// copies the previous 4.5 hint when this observation is a different model. +func (s *QuotaSnapshot) ApplyGrok45ResponsesPlanSignal(prev *QuotaSnapshot) { + if s == nil { + return + } + observedAt := firstNonEmptyQuotaTime(s.LastHeadersSeenAt, s.UpdatedAt) + if IsGrok45ResponsesQuotaModel(s.Model) && quotaHasLimitWindow(s) { + if quotaLooksLikeGrokHeavy(s) { + s.PlanFrom45Responses = "supergrok_heavy" + s.PlanFrom45ResponsesAt = observedAt + return + } + s.PlanFrom45Responses = "supergrok" + s.PlanFrom45ResponsesAt = observedAt + return + } + if prev != nil && strings.TrimSpace(prev.PlanFrom45Responses) != "" { + s.PlanFrom45Responses = prev.PlanFrom45Responses + s.PlanFrom45ResponsesAt = prev.PlanFrom45ResponsesAt + } +} + +// QuotaSnapshotObservedAt prefers LastHeadersSeenAt over UpdatedAt so a later +// snapshot rewrite cannot refresh a stale Heavy window. +func QuotaSnapshotObservedAt(snapshot *QuotaSnapshot) (time.Time, bool) { + if snapshot == nil { + return time.Time{}, false + } + return parseQuotaTimestamp(firstNonEmptyQuotaTime(snapshot.LastHeadersSeenAt, snapshot.UpdatedAt)) +} + +// IsQuotaSnapshotFresh reports whether a quota signal is recent enough to +// distinguish SuperGrok from Heavy. +func IsQuotaSnapshotFresh(snapshot *QuotaSnapshot, now time.Time) bool { + observedAt, ok := QuotaSnapshotObservedAt(snapshot) + if !ok { + return false + } + return isTimeFresh(observedAt, now) +} + +func isQuotaTimestampFresh(raw string, now time.Time) bool { + parsed, ok := parseQuotaTimestamp(raw) + if !ok { + return false + } + return isTimeFresh(parsed, now) +} + +func parseQuotaTimestamp(raw string) (time.Time, bool) { + raw = strings.TrimSpace(raw) + if raw == "" { + return time.Time{}, false + } + parsed, err := time.Parse(time.RFC3339, raw) + if err != nil { + return time.Time{}, false + } + return parsed, true +} + +func isTimeFresh(observedAt, now time.Time) bool { + if now.IsZero() { + now = time.Now() + } + age := now.Sub(observedAt) + return age <= GrokQuotaSignalMaxAge && age >= -5*time.Minute +} + +func quotaHasLimitWindow(quota *QuotaSnapshot) bool { + if quota == nil { + return false + } + if quota.Requests != nil && quota.Requests.Limit != nil { + return true + } + return quota.Tokens != nil && quota.Tokens.Limit != nil +} + +func quotaLooksLikeGrokHeavy(quota *QuotaSnapshot) bool { + if quota == nil { + return false + } + var requestLimit, tokenLimit int64 + if quota.Requests != nil && quota.Requests.Limit != nil { + requestLimit = *quota.Requests.Limit + } + if quota.Tokens != nil && quota.Tokens.Limit != nil { + tokenLimit = *quota.Tokens.Limit + } + return requestLimit >= grokHeavyQuotaRequestLimit || tokenLimit >= grokHeavyQuotaTokenLimit +} + +func firstNonEmptyQuotaTime(values ...string) string { + for _, value := range values { + if strings.TrimSpace(value) != "" { + return strings.TrimSpace(value) + } + } + return "" +} diff --git a/backend/internal/pkg/xai/subscription_tier_test.go b/backend/internal/pkg/xai/subscription_tier_test.go new file mode 100644 index 0000000000..ef9236d732 --- /dev/null +++ b/backend/internal/pkg/xai/subscription_tier_test.go @@ -0,0 +1,144 @@ +//go:build unit + +package xai + +import ( + "encoding/base64" + "encoding/json" + "testing" + "time" + + "github.com/stretchr/testify/require" +) + +func TestMapJWTSubscriptionTierNumber(t *testing.T) { + t.Parallel() + + require.Equal(t, "free", MapJWTSubscriptionTier(0)) + require.Equal(t, "supergrok", MapJWTSubscriptionTier(1)) + require.Equal(t, "x_basic", MapJWTSubscriptionTier(2)) + require.Equal(t, "x_premium", MapJWTSubscriptionTier(3)) + require.Equal(t, "x_premium_plus", MapJWTSubscriptionTier(4)) + require.Equal(t, "supergrok_heavy", MapJWTSubscriptionTier(5)) + require.Equal(t, "supergrok_lite", MapJWTSubscriptionTier(6)) + require.Equal(t, "supergrok_plus", MapJWTSubscriptionTier(7)) + require.Equal(t, "9", MapJWTSubscriptionTier(9)) +} + +func TestNormalizeSubscriptionTierAliases(t *testing.T) { + t.Parallel() + + require.Equal(t, "free", NormalizeSubscriptionTier("Free")) + require.Equal(t, "free", NormalizeSubscriptionTier(" FREE ")) + require.Equal(t, "supergrok", NormalizeSubscriptionTier("SuperGrok")) + require.Equal(t, "supergrok_heavy", NormalizeSubscriptionTier("SuperGrok Heavy")) + require.Equal(t, "supergrok_pro", NormalizeSubscriptionTier("SuperGrokPro")) + require.Equal(t, "supergrok_lite", NormalizeSubscriptionTier("SuperGrok Lite")) + require.Equal(t, "supergrok_lite", NormalizeSubscriptionTier("SuperGrokLite")) + require.Equal(t, "x_basic", NormalizeSubscriptionTier("X Basic")) + require.Equal(t, "free", NormalizeSubscriptionTier("free-tier")) + require.Equal(t, "free", NormalizeSubscriptionTier("free_tier")) + require.Equal(t, "free", NormalizeSubscriptionTier("grok-basic")) + require.Equal(t, "free", NormalizeSubscriptionTier("grok_basic")) + require.Equal(t, "supergrok_lite", NormalizeSubscriptionTier("supergrok_lite")) +} + +func TestSubscriptionTierFromJWTUsesNumericClaim(t *testing.T) { + t.Parallel() + + require.Equal(t, "supergrok_heavy", SubscriptionTierFromJWT(jwtWithClaims(t, map[string]any{"tier": 5}))) + require.Equal(t, "free", SubscriptionTierFromJWT(jwtWithClaims(t, map[string]any{"tier": 0}))) + require.Equal(t, "supergrok_lite", SubscriptionTierFromJWT(jwtWithClaims(t, map[string]any{"tier": 6}))) + require.Empty(t, SubscriptionTierFromJWT(jwtWithClaims(t, map[string]any{"sub": "user"}))) + require.Empty(t, SubscriptionTierFromJWT("not-a-jwt")) +} + +func TestCanonicalGrokPlanUsesOnlyGrok45ResponsesWindow(t *testing.T) { + t.Parallel() + + zero := float64(0) + heavyReq, heavyTok := int64(8300), int64(53_000_000) + superReq, superTok := int64(900), int64(15_000_000) + fresh := time.Now().UTC().Format(time.RFC3339) + stale := time.Now().Add(-GrokQuotaSignalMaxAge - time.Hour).UTC().Format(time.RFC3339) + + require.Equal(t, "supergrok", CanonicalGrokPlan(&zero, "SuperGrokPro", nil)) + require.Equal(t, "supergrok", CanonicalGrokPlan(&zero, "SuperGrok", nil)) + require.Equal(t, "supergrok_heavy", CanonicalGrokPlan(&zero, "SuperGrok Heavy", nil)) + require.Empty(t, CanonicalGrokPlan(&zero, "", nil)) + require.Equal(t, "free", CanonicalGrokPlan(&zero, "free", nil)) + + require.Equal(t, "supergrok_heavy", CanonicalGrokPlan(&zero, "SuperGrokPro", &QuotaSnapshot{ + Model: "grok-4.5", + Requests: &QuotaWindow{Limit: &heavyReq}, + Tokens: &QuotaWindow{Limit: &heavyTok}, + LastHeadersSeenAt: fresh, + })) + require.Equal(t, "supergrok", CanonicalGrokPlan(&zero, "SuperGrokPro", &QuotaSnapshot{ + Model: "grok-4.6", + Requests: &QuotaWindow{Limit: &heavyReq}, + Tokens: &QuotaWindow{Limit: &heavyTok}, + LastHeadersSeenAt: fresh, + })) + require.Equal(t, "supergrok", CanonicalGrokPlan(&zero, "SuperGrokPro", &QuotaSnapshot{ + Requests: &QuotaWindow{Limit: &heavyReq}, + Tokens: &QuotaWindow{Limit: &heavyTok}, + LastHeadersSeenAt: fresh, + })) + require.Equal(t, "supergrok", CanonicalGrokPlan(&zero, "SuperGrokPro", &QuotaSnapshot{ + Model: "grok-4.5", + Requests: &QuotaWindow{Limit: &superReq}, + Tokens: &QuotaWindow{Limit: &superTok}, + LastHeadersSeenAt: fresh, + })) + require.Equal(t, "supergrok", CanonicalGrokPlan(&zero, "SuperGrokPro", &QuotaSnapshot{ + Model: "grok-4.5", + Requests: &QuotaWindow{Limit: &heavyReq}, + LastHeadersSeenAt: stale, + })) + require.Equal(t, "supergrok_heavy", CanonicalGrokPlan(&zero, "SuperGrokPro", &QuotaSnapshot{ + Model: "grok-4.6", + Requests: &QuotaWindow{Limit: &superReq}, + PlanFrom45Responses: "supergrok_heavy", + PlanFrom45ResponsesAt: fresh, + })) + require.Equal(t, "free", CanonicalGrokPlan(&zero, "free", &QuotaSnapshot{ + Model: "grok-4.5", + Requests: &QuotaWindow{Limit: &heavyReq}, + LastHeadersSeenAt: fresh, + })) + heavyCents := float64(SuperGrokHeavyLimitCents) + require.Equal(t, "supergrok_heavy", CanonicalGrokPlan(&heavyCents, "SuperGrokPro", nil)) +} + +func TestApplyGrok45ResponsesPlanSignalCarriesHint(t *testing.T) { + t.Parallel() + + heavyReq := int64(8300) + fresh := time.Now().UTC().Format(time.RFC3339) + prev := &QuotaSnapshot{ + Model: "grok-4.5", + Requests: &QuotaWindow{Limit: &heavyReq}, + LastHeadersSeenAt: fresh, + } + prev.ApplyGrok45ResponsesPlanSignal(nil) + require.Equal(t, "supergrok_heavy", prev.PlanFrom45Responses) + + next := &QuotaSnapshot{ + Model: "grok-4.6", + Requests: &QuotaWindow{Limit: int64Ptr(100)}, + LastHeadersSeenAt: fresh, + } + next.ApplyGrok45ResponsesPlanSignal(prev) + require.Equal(t, "supergrok_heavy", next.PlanFrom45Responses) + require.Equal(t, prev.PlanFrom45ResponsesAt, next.PlanFrom45ResponsesAt) +} + +func int64Ptr(v int64) *int64 { return &v } + +func jwtWithClaims(t *testing.T, claims map[string]any) string { + t.Helper() + payload, err := json.Marshal(claims) + require.NoError(t, err) + return "header." + base64.RawURLEncoding.EncodeToString(payload) + ".sig" +} diff --git a/backend/internal/repository/api_key_repo.go b/backend/internal/repository/api_key_repo.go index a1c7cad801..388562faeb 100644 --- a/backend/internal/repository/api_key_repo.go +++ b/backend/internal/repository/api_key_repo.go @@ -3,8 +3,10 @@ package repository import ( "context" "database/sql" + "encoding/json" "errors" "fmt" + "log/slog" "strings" "time" @@ -196,6 +198,8 @@ func (r *apiKeyRepository) GetByKeyForAuth(ctx context.Context, key string) (*se group.FieldAudioRealtimePricePerMin, group.FieldAudioTtsPricePerMillionChars, group.FieldAudioSttPricePerHour, + group.FieldLongContextPricingEnabled, + group.FieldModelPricing, group.FieldClaudeCodeOnly, group.FieldFallbackGroupID, group.FieldFallbackGroupIDOnInvalidRequest, @@ -946,6 +950,14 @@ func groupEntityToService(g *dbent.Group) *service.Group { if g == nil { return nil } + var modelPricing []service.ChannelModelPricing + if len(g.ModelPricing) > 0 { + if err := json.Unmarshal(g.ModelPricing, &modelPricing); err != nil { + slog.Warn("group model_pricing unmarshal failed; falling back to channel/builtin pricing", + "group_id", g.ID, "error", err) + modelPricing = nil + } + } return &service.Group{ ID: g.ID, Name: g.Name, @@ -980,6 +992,8 @@ func groupEntityToService(g *dbent.Group) *service.Group { AudioRealtimePricePerMin: g.AudioRealtimePricePerMin, AudioTTSPricePerMillionChars: g.AudioTtsPricePerMillionChars, AudioSTTPricePerHour: g.AudioSttPricePerHour, + LongContextPricingEnabled: g.LongContextPricingEnabled, + ModelPricing: modelPricing, DefaultValidityDays: g.DefaultValidityDays, ClaudeCodeOnly: g.ClaudeCodeOnly, FallbackGroupID: g.FallbackGroupID, diff --git a/backend/internal/repository/group_repo.go b/backend/internal/repository/group_repo.go index 4cbc27909d..b5de46353d 100644 --- a/backend/internal/repository/group_repo.go +++ b/backend/internal/repository/group_repo.go @@ -3,6 +3,7 @@ package repository import ( "context" "database/sql" + "encoding/json" "errors" "fmt" "sort" @@ -56,6 +57,10 @@ func createGroupRecord(ctx context.Context, client *dbent.Client, groupIn *servi if groupIn == nil { return errors.New("group is nil") } + modelPricing, err := json.Marshal(groupIn.ModelPricing) + if err != nil { + return fmt.Errorf("marshal group model pricing: %w", err) + } builder := client.Group.Create(). SetName(groupIn.Name). SetDescription(groupIn.Description). @@ -88,6 +93,8 @@ func createGroupRecord(ctx context.Context, client *dbent.Client, groupIn *servi SetNillableAudioRealtimePricePerMin(groupIn.AudioRealtimePricePerMin). SetNillableAudioTtsPricePerMillionChars(groupIn.AudioTTSPricePerMillionChars). SetNillableAudioSttPricePerHour(groupIn.AudioSTTPricePerHour). + SetLongContextPricingEnabled(groupIn.LongContextPricingEnabled). + SetModelPricing(modelPricing). SetDefaultValidityDays(groupIn.DefaultValidityDays). SetClaudeCodeOnly(groupIn.ClaudeCodeOnly). SetNillableFallbackGroupID(groupIn.FallbackGroupID). @@ -234,6 +241,10 @@ func (r *groupRepository) GetByIDLite(ctx context.Context, id int64) (*service.G } func (r *groupRepository) Update(ctx context.Context, groupIn *service.Group) error { + modelPricing, err := json.Marshal(groupIn.ModelPricing) + if err != nil { + return fmt.Errorf("marshal group model pricing: %w", err) + } builder := r.client.Group.UpdateOneID(groupIn.ID). SetName(groupIn.Name). SetDescription(groupIn.Description). @@ -260,6 +271,8 @@ func (r *groupRepository) Update(ctx context.Context, groupIn *service.Group) er SetNillableVideoPrice720p(groupIn.VideoPrice720P). SetNillableVideoPrice1080p(groupIn.VideoPrice1080P). SetVideoModelPrices(service.NormalizeVideoModelPrices(groupIn.VideoModelPrices)). + SetLongContextPricingEnabled(groupIn.LongContextPricingEnabled). + SetModelPricing(modelPricing). SetDefaultValidityDays(groupIn.DefaultValidityDays). SetClaudeCodeOnly(groupIn.ClaudeCodeOnly). SetModelRoutingEnabled(groupIn.ModelRoutingEnabled). diff --git a/backend/internal/server/api_contract_test.go b/backend/internal/server/api_contract_test.go index 9ac40fb556..91bdd6a360 100644 --- a/backend/internal/server/api_contract_test.go +++ b/backend/internal/server/api_contract_test.go @@ -364,6 +364,7 @@ func TestAPIContracts(t *testing.T) { "daily_limit_usd": null, "weekly_limit_usd": null, "monthly_limit_usd": null, + "long_context_pricing_enabled": false, "image_price_1k": null, "image_price_2k": null, "image_price_4k": null, diff --git a/backend/internal/server/routes/gateway.go b/backend/internal/server/routes/gateway.go index 8d5fcd6a53..6d99d2c7ea 100644 --- a/backend/internal/server/routes/gateway.go +++ b/backend/internal/server/routes/gateway.go @@ -315,6 +315,14 @@ func RegisterGatewayRoutes( } h.Gateway.WebSearch(c) }) + gateway.POST("/x_search", func(c *gin.Context) { + if getGroupPlatform(c) != service.PlatformGrok { + service.MarkOpsClientBusinessLimited(c, service.OpsClientBusinessLimitedReasonLocalFeatureGate) + c.JSON(http.StatusNotFound, gin.H{"error": gin.H{"type": "not_found_error", "message": "X Search API is not supported for this platform"}}) + return + } + h.Gateway.XSearch(c) + }) } // Gemini 原生 API 兼容层(Gemini SDK/CLI 直连) @@ -443,6 +451,14 @@ func RegisterGatewayRoutes( } h.Gateway.WebSearch(c) }) + r.POST("/x_search", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), compositeTarget, requireGroupAnthropic, func(c *gin.Context) { + if getGroupPlatform(c) != service.PlatformGrok { + service.MarkOpsClientBusinessLimited(c, service.OpsClientBusinessLimitedReasonLocalFeatureGate) + c.JSON(http.StatusNotFound, gin.H{"error": gin.H{"type": "not_found_error", "message": "X Search API is not supported for this platform"}}) + return + } + h.Gateway.XSearch(c) + }) // Antigravity 模型列表 r.GET("/antigravity/models", gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, h.Gateway.AntigravityModels) diff --git a/backend/internal/server/routes/prompt_audit_route_coverage_test.go b/backend/internal/server/routes/prompt_audit_route_coverage_test.go index f03cf5b314..af4272cb89 100644 --- a/backend/internal/server/routes/prompt_audit_route_coverage_test.go +++ b/backend/internal/server/routes/prompt_audit_route_coverage_test.go @@ -48,6 +48,7 @@ func TestEveryGatewayPOSTRouteIsClassifiedForPromptAuditCoverage(t *testing.T) { "/models/*modelAction": {"gemini_v1beta_handler.go"}, "/tts": {"grok_audio.go"}, "/web_search": {"gateway_web_search.go"}, + "/x_search": {"gateway_web_search.go"}, } excluded := map[string]string{ "/messages/count_tokens": "tokenization only; it does not execute a model request", diff --git a/backend/internal/service/account_test_service.go b/backend/internal/service/account_test_service.go index 5a62598db2..9e1d29c2f9 100644 --- a/backend/internal/service/account_test_service.go +++ b/backend/internal/service/account_test_service.go @@ -969,6 +969,7 @@ func (s *AccountTestService) observeGrokTestResponse(ctx context.Context, accoun resp.Body = io.NopCloser(bytes.NewReader(responseBody)) } snapshot := parseGrokQuotaSnapshot(resp.Header, resp.StatusCode, now) + stampGrokQuotaSnapshotForPlan(account, snapshot, grokRequestedModelFromCtx(ctx)) if snapshot != nil && s.accountRepo != nil { resetAt, limited := grokRateLimitResetAtForAccount(account, snapshot, now) if limited { @@ -1073,7 +1074,7 @@ func (s *AccountTestService) testGrokResponsesConnection(c *gin.Context, ctx con } defer func() { _ = resp.Body.Close() }() - s.observeGrokTestResponse(ctx, account, resp) + s.observeGrokTestResponse(withGrokTeamRateLimitModel(ctx, testModelID), account, resp) if resp.StatusCode != http.StatusOK { body, _ := io.ReadAll(resp.Body) @@ -1421,7 +1422,7 @@ User query: return s.sendErrorAndEnd(c, fmt.Sprintf("standalone web_search probe failed: %s", err.Error())) } defer func() { _ = resp.Body.Close() }() - s.observeGrokTestResponse(ctx, account, resp) + s.observeGrokTestResponse(withGrokTeamRateLimitModel(ctx, grokDefaultResponsesModel), account, resp) body, _ := io.ReadAll(resp.Body) if resp.StatusCode != http.StatusOK { diff --git a/backend/internal/service/admin_group.go b/backend/internal/service/admin_group.go index 5744f40cfe..98f4ab6d31 100644 --- a/backend/internal/service/admin_group.go +++ b/backend/internal/service/admin_group.go @@ -301,6 +301,10 @@ func (s *adminServiceImpl) CreateGroup(ctx context.Context, input *CreateGroupIn } platform := NormalizeGroupPlatform(input.Platform) + modelPricing, err := normalizeGroupModelPricing(platform, input.ModelPricing) + if err != nil { + return nil, err + } maxReasoningEffort, err := normalizeMaxReasoningEffortForPlatform(platform, input.MaxReasoningEffort) if err != nil { return nil, infraerrors.Newf(http.StatusBadRequest, "INVALID_MAX_REASONING_EFFORT", "%v", err) @@ -459,6 +463,8 @@ func (s *adminServiceImpl) CreateGroup(ctx context.Context, input *CreateGroupIn DailyLimitUSD: dailyLimit, WeeklyLimitUSD: weeklyLimit, MonthlyLimitUSD: monthlyLimit, + LongContextPricingEnabled: input.LongContextPricingEnabled, + ModelPricing: modelPricing, AllowImageGeneration: allowImageGeneration, AllowBatchImageGeneration: allowBatchImageGeneration, ImageRateIndependent: input.ImageRateIndependent, @@ -656,6 +662,16 @@ func (s *adminServiceImpl) UpdateGroup(ctx context.Context, id int64, input *Upd if input.Status != "" { group.Status = input.Status } + if input.LongContextPricingEnabled != nil { + group.LongContextPricingEnabled = *input.LongContextPricingEnabled + } + if input.ModelPricing != nil { + modelPricing, normalizeErr := normalizeGroupModelPricing(group.Platform, *input.ModelPricing) + if normalizeErr != nil { + return nil, normalizeErr + } + group.ModelPricing = modelPricing + } // 订阅相关字段 if input.SubscriptionType != "" { @@ -955,6 +971,28 @@ func (s *adminServiceImpl) UpdateGroup(ctx context.Context, id int64, input *Upd return group, nil } +func normalizeGroupModelPricing(platform string, pricing []ChannelModelPricing) ([]ChannelModelPricing, error) { + out := make([]ChannelModelPricing, len(pricing)) + for i := range pricing { + out[i] = pricing[i].Clone() + out[i].ID = 0 + out[i].ChannelID = 0 + if strings.TrimSpace(out[i].Platform) == "" { + out[i].Platform = platform + } + for j := range out[i].Models { + out[i].Models[j] = strings.TrimSpace(out[i].Models[j]) + } + if len(out[i].Models) == 0 { + return nil, infraerrors.New(http.StatusBadRequest, "GROUP_MODEL_PRICING_MODELS_REQUIRED", "group model pricing entry requires at least one model") + } + } + if err := validatePricingEntries(out); err != nil { + return nil, err + } + return out, nil +} + func (s *adminServiceImpl) DeleteGroup(ctx context.Context, id int64) error { var groupKeys []string if s.authCacheInvalidator != nil { diff --git a/backend/internal/service/admin_service.go b/backend/internal/service/admin_service.go index d9c4060c9e..31d0e48649 100644 --- a/backend/internal/service/admin_service.go +++ b/backend/internal/service/admin_service.go @@ -209,15 +209,17 @@ type AdminBoundAuthIdentityChannel struct { } type CreateGroupInput struct { - Name string - Description string - Platform string - RateMultiplier float64 - IsExclusive bool - SubscriptionType string // standard/subscription - DailyLimitUSD *float64 // 日限额 (USD) - WeeklyLimitUSD *float64 // 周限额 (USD) - MonthlyLimitUSD *float64 // 月限额 (USD) + Name string + Description string + Platform string + RateMultiplier float64 + IsExclusive bool + SubscriptionType string // standard/subscription + DailyLimitUSD *float64 // 日限额 (USD) + WeeklyLimitUSD *float64 // 周限额 (USD) + MonthlyLimitUSD *float64 // 月限额 (USD) + LongContextPricingEnabled bool + ModelPricing []ChannelModelPricing // 图片生成计费配置(仅 antigravity 平台使用) AllowImageGeneration bool AllowBatchImageGeneration bool @@ -281,16 +283,18 @@ type CreateGroupInput struct { } type UpdateGroupInput struct { - Name string - Description *string - Platform string - RateMultiplier *float64 // 使用指针以支持设置为0 - IsExclusive *bool - Status string - SubscriptionType string // standard/subscription - DailyLimitUSD *float64 // 日限额 (USD) - WeeklyLimitUSD *float64 // 周限额 (USD) - MonthlyLimitUSD *float64 // 月限额 (USD) + Name string + Description *string + Platform string + RateMultiplier *float64 // 使用指针以支持设置为0 + IsExclusive *bool + Status string + SubscriptionType string // standard/subscription + DailyLimitUSD *float64 // 日限额 (USD) + WeeklyLimitUSD *float64 // 周限额 (USD) + MonthlyLimitUSD *float64 // 月限额 (USD) + LongContextPricingEnabled *bool + ModelPricing *[]ChannelModelPricing // 图片生成计费配置(仅 antigravity 平台使用) AllowImageGeneration *bool AllowBatchImageGeneration *bool diff --git a/backend/internal/service/billing_search_audio_cost_test.go b/backend/internal/service/billing_search_audio_cost_test.go index 94d2e5c659..ad3111f0dd 100644 --- a/backend/internal/service/billing_search_audio_cost_test.go +++ b/backend/internal/service/billing_search_audio_cost_test.go @@ -10,8 +10,8 @@ func TestCalculateSearchCost(t *testing.T) { t.Parallel() s := &BillingService{} require.Equal(t, 0.0, s.CalculateSearchCost(0, floatPtr(10), 1).ActualCost) - // nil price → default $10/1k: 5 calls = 0.05 - require.InDelta(t, 0.05, s.CalculateSearchCost(5, nil, 1).ActualCost, 1e-9) + // nil price -> official xAI default $5/1k: 5 calls = 0.025 + require.InDelta(t, 0.025, s.CalculateSearchCost(5, nil, 1).ActualCost, 1e-9) // explicit 0 → free require.Equal(t, 0.0, s.CalculateSearchCost(5, floatPtr(0), 1).ActualCost) price := 10.0 @@ -30,10 +30,10 @@ func TestCalculateAudioCost(t *testing.T) { require.InDelta(t, 1.5, s.CalculateAudioCost("tts", 0.1, cfg, 1).ActualCost, 1e-9) require.InDelta(t, 0.25, s.CalculateAudioCost("stt", 0.5, cfg, 1).ActualCost, 1e-9) require.Equal(t, 0.0, s.CalculateAudioCost("unknown", 1, cfg, 1).ActualCost) - // nil config → defaults (realtime $0.10/min, tts $15/M, stt $0.36/hr) - require.InDelta(t, 0.10, s.CalculateAudioCost("realtime", 1, nil, 1).ActualCost, 1e-9) + // nil config -> official defaults (think-fast-1 $0.05/min, TTS $15/M, REST STT $0.10/hr) + require.InDelta(t, 0.05, s.CalculateAudioCost("realtime", 1, nil, 1).ActualCost, 1e-9) require.InDelta(t, 15.0, s.CalculateAudioCost("tts", 1, nil, 1).ActualCost, 1e-9) - require.InDelta(t, 0.36, s.CalculateAudioCost("stt", 1, nil, 1).ActualCost, 1e-9) + require.InDelta(t, 0.10, s.CalculateAudioCost("stt", 1, nil, 1).ActualCost, 1e-9) // explicit 0 → free zero := 0.0 require.Equal(t, 0.0, s.CalculateAudioCost("realtime", 1, &audioPriceConfig{RealtimePerMin: &zero}, 1).ActualCost) diff --git a/backend/internal/service/billing_service.go b/backend/internal/service/billing_service.go index 96233b7ae9..c13908443e 100644 --- a/backend/internal/service/billing_service.go +++ b/backend/internal/service/billing_service.go @@ -10,6 +10,7 @@ import ( "time" "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/Wei-Shaw/sub2api/internal/pkg/xai" ) // APIKeyRateLimitCacheData holds rate limit usage data cached in Redis. @@ -104,6 +105,7 @@ type ModelPricing struct { CacheCreation1hPrice float64 // 1小时缓存创建每token价格 (USD) SupportsCacheBreakdown bool // 是否支持详细的缓存分类 LongContextInputThreshold int // 超过阈值后按整次会话提升输入价格 + LongContextThresholdInclusive bool // 达到阈值即应用(xAI);默认保持严格大于以兼容既有模型 LongContextInputMultiplier float64 // 长上下文整次会话输入倍率 LongContextOutputMultiplier float64 // 长上下文整次会话输出倍率 ImageOutputPricePerToken float64 // 图片输出 token 价格 (USD) @@ -585,32 +587,56 @@ func (s *BillingService) initFallbackPricing() { SupportsCacheBreakdown: false, } - // xAI Grok 4.5 (official docs: $2 input / $0.50 cached input / $6 output per MTok) + // xAI Grok 4.5: $2 input / $0.30 cached input / $6 output below 200k. s.fallbackPrices["grok-4.5"] = &ModelPricing{ - InputPricePerToken: 2e-6, - OutputPricePerToken: 6e-6, - CacheReadPricePerToken: 0.5e-6, - SupportsCacheBreakdown: false, + InputPricePerToken: 2e-6, + OutputPricePerToken: 6e-6, + CacheReadPricePerToken: 0.3e-6, + SupportsCacheBreakdown: false, + LongContextInputThreshold: 200000, + LongContextThresholdInclusive: true, + LongContextInputMultiplier: 2, + LongContextOutputMultiplier: 2, } - // xAI Grok 4.3 (official docs: $1.25 input / $2.50 output per MTok) + // xAI Grok 4.6 (docs.x.ai/developers/models: $2 input / $0.50 cached input / + // $6 output per MTok under 200k prompt tokens; ≥200k is 2× on input, + // cached input, and output). + s.fallbackPrices["grok-4.6"] = &ModelPricing{ + InputPricePerToken: 2e-6, + OutputPricePerToken: 6e-6, + CacheReadPricePerToken: 0.5e-6, + SupportsCacheBreakdown: false, + LongContextInputThreshold: 200000, + LongContextThresholdInclusive: true, + LongContextInputMultiplier: 2, + LongContextOutputMultiplier: 2, + } + + // xAI Grok 4.3: $1.25 input / $0.20 cached / $2.50 output below 200k. s.fallbackPrices["grok-4.3"] = &ModelPricing{ - InputPricePerToken: 1.25e-6, - OutputPricePerToken: 2.5e-6, - CacheReadPricePerToken: 0.2e-6, - SupportsCacheBreakdown: false, - LongContextInputThreshold: 1000000, - LongContextInputMultiplier: 1, + InputPricePerToken: 1.25e-6, + OutputPricePerToken: 2.5e-6, + CacheReadPricePerToken: 0.2e-6, + SupportsCacheBreakdown: false, + LongContextInputThreshold: 200000, + LongContextThresholdInclusive: true, + LongContextInputMultiplier: 2, + LongContextOutputMultiplier: 2, } // xAI Grok Build 0.1 (official docs: $1 input / $0.20 cached input / // $2 output per MTok). Composer is available only through Grok Build and // has no standalone public API rate card, so its aliases use this coding // model rate instead of silently billing at zero. s.fallbackPrices["grok-build-0.1"] = &ModelPricing{ - InputPricePerToken: 1e-6, - OutputPricePerToken: 2e-6, - CacheReadPricePerToken: 0.2e-6, - SupportsCacheBreakdown: false, + InputPricePerToken: 1e-6, + OutputPricePerToken: 2e-6, + CacheReadPricePerToken: 0.2e-6, + SupportsCacheBreakdown: false, + LongContextInputThreshold: 200000, + LongContextThresholdInclusive: true, + LongContextInputMultiplier: 2, + LongContextOutputMultiplier: 2, } } @@ -803,8 +829,10 @@ func (s *BillingService) getFallbackPricing(model string) *ModelPricing { } switch modelLower { - case "grok", "grok-latest", "grok-4.5", "grok-4.5-latest", "grok-build-latest": + case "grok", "grok-latest", "grok-4.5", "grok-4.5-latest": return s.fallbackPrices["grok-4.5"] + case "grok-4.6", "grok-4.6-latest": + return s.fallbackPrices["grok-4.6"] case "grok-4.3", "grok-4.20-0309-reasoning", "grok-4.20-0309-non-reasoning", @@ -812,13 +840,43 @@ func (s *BillingService) getFallbackPricing(model string) *ModelPricing { "grok-4.20-reasoning", "grok-4.20-non-reasoning": return s.fallbackPrices["grok-4.3"] - case "grok-build", "grok-build-0.1", "grok-composer", "grok-composer-2.5-fast", "composer-2.5": + case "grok-build", "grok-build-latest", "grok-build-0.1", "grok-composer", "grok-composer-2.5-fast", "composer-2.5": return s.fallbackPrices["grok-build-0.1"] } + // Unknown Grok text IDs (grok-5, dated snapshots, provider-prefixed) inherit + // the current default text card so a new model cannot ship unbilled. + if pricing := s.grokUnknownTextFamilyFallback(modelLower); pricing != nil { + return pricing + } + return nil } +func (s *BillingService) grokUnknownTextFamilyFallback(model string) *ModelPricing { + if s == nil || !isGrokUnknownTextFamilyModel(model) { + return nil + } + return s.fallbackPrices["grok-4.5"] +} + +func isGrokUnknownTextFamilyModel(model string) bool { + native := strings.ToLower(strings.TrimSpace(xai.StripGrokProviderPrefix(model))) + switch { + case native == "grok", native == "grok-latest": + return true + case strings.HasPrefix(native, "grok-build"), + strings.HasPrefix(native, "grok-composer"), + strings.HasPrefix(native, "composer-"): + return true + case len(native) > 5 && strings.HasPrefix(native, "grok-"): + rest := native[len("grok-"):] + return rest[0] >= '0' && rest[0] <= '9' + default: + return false + } +} + // HasIdentifiedTokenPricing 判断模型能否在价格表中被"确定性识别"出 token 价格。 // // 与 GetModelPricing 的关键区别:本函数拒绝按子串猜系列的兜底。GetModelPricing 会 @@ -950,9 +1008,11 @@ type CostInput struct { Ctx context.Context Model string GroupID *int64 // 用于渠道定价查找 + Group *Group Tokens UsageTokens - RequestCount int // 按次计费时使用 - SizeTier string // 按次/图片模式的层级标签("1K","2K","4K","HD" 等) + RequestCount int // 按次计费时使用 + UsageUnits float64 // 音频等连续计量单位(分钟/小时/百万字符) + SizeTier string // 按次/图片模式的层级标签("1K","2K","4K","HD" 等) RateMultiplier float64 ServiceTier string // "priority","flex","" 等 Resolver *ModelPricingResolver // 定价解析器 @@ -985,6 +1045,7 @@ func (s *BillingService) CalculateCostUnified(input CostInput) (*CostBreakdown, resolved = input.Resolver.Resolve(input.Ctx, PricingInput{ Model: input.Model, GroupID: input.GroupID, + Group: input.Group, }) } @@ -996,7 +1057,7 @@ func (s *BillingService) CalculateCostUnified(input CostInput) (*CostBreakdown, var breakdown *CostBreakdown var err error switch resolved.Mode { - case BillingModePerRequest, BillingModeImage: + case BillingModePerRequest, BillingModeImage, BillingModeVideo: breakdown, err = s.calculatePerRequestCost(resolved, input) default: // BillingModeToken breakdown, err = s.calculateTokenCost(resolved, input) @@ -1022,7 +1083,7 @@ func (s *BillingService) calculateTokenCost(resolved *ResolvedPricing, input Cos pricing = s.applyModelSpecificPricingPolicy(input.Model, pricing) // 长上下文定价仅在无区间定价时应用(区间定价已包含上下文分层) - applyLongCtx := len(resolved.Intervals) == 0 + applyLongCtx := len(resolved.Intervals) == 0 && resolved.longContextPricingEnabled if input.LongContextBillingEnabled != nil { applyLongCtx = applyLongCtx && *input.LongContextBillingEnabled } @@ -1157,9 +1218,13 @@ func (s *BillingService) computeCacheCreationCost(pricing *ModelPricing, tokens // calculatePerRequestCost 按次/图片计费 func (s *BillingService) calculatePerRequestCost(resolved *ResolvedPricing, input CostInput) (*CostBreakdown, error) { - count := input.RequestCount - if count <= 0 { - count = 1 + units := input.UsageUnits + if units <= 0 { + count := input.RequestCount + if count <= 0 { + count = 1 + } + units = float64(count) } var unitPrice float64 @@ -1178,7 +1243,7 @@ func (s *BillingService) calculatePerRequestCost(resolved *ResolvedPricing, inpu unitPrice = resolved.DefaultPerRequestPrice } - totalCost := unitPrice * float64(count) + totalCost := unitPrice * units actualCost := totalCost * input.RateMultiplier return &CostBreakdown{ @@ -1280,6 +1345,9 @@ func (s *BillingService) shouldApplySessionLongContextPricing(tokens UsageTokens return false } totalInputTokens := tokens.InputTokens + tokens.CacheCreationTokens + tokens.CacheReadTokens + if pricing.LongContextThresholdInclusive { + return totalInputTokens >= pricing.LongContextInputThreshold + } return totalInputTokens > pricing.LongContextInputThreshold } @@ -1451,6 +1519,8 @@ const ( defaultGrokImagineImagePrice2K = 0.02 defaultGrokImagineImageQualityPrice1K = 0.05 defaultGrokImagineImageQualityPrice2K = 0.07 + defaultGrokImagineImage20Price1K = 0.06 // default quality is Medium + defaultGrokImagineImage20Price2K = 0.08 // 视频默认价为 xAI 官方**每秒**输出价格(USD/s),总价 = 每秒价 × 时长(秒)。 defaultGrokImagineVideoPrice480P = 0.05 @@ -1462,14 +1532,14 @@ const ( // Codex alpha/search 网页搜索单次默认价:OpenAI 官方 web search 定价 $10/1000 次。 defaultWebSearchPricePerCall = 0.01 - // Grok /v1/web_search 与 SearchCount 附加费:与 Codex 对齐 $10/1000 次(按 1k 计价字段存储)。 - defaultSearchPricePer1k = 10.0 + // xAI server-side web/X search and code execution are $5/1000 calls. + defaultSearchPricePer1k = 5.0 - // Grok Voice 默认价(分组列 NULL 时使用;显式配 0 表示免费)。 - // 保守运营占位,运维可通过 groups.audio_* 覆盖。 - defaultAudioRealtimePricePerMin = 0.10 + // Generic realtime defaults to think-fast-1.0; think-fast-2.0 can be + // configured independently through per-model group/channel pricing. + defaultAudioRealtimePricePerMin = 0.05 defaultAudioTTSPricePerMillionChars = 15.0 - defaultAudioSTTPricePerHour = 0.36 + defaultAudioSTTPricePerHour = 0.10 ) // CalculateWebSearchCost 计算 Codex alpha/search 网页搜索按次费用。 @@ -1727,6 +1797,12 @@ func (s *BillingService) getDefaultVideoPrice(model string, resolution string) f func getDefaultGrokImagineImagePrice(model string, imageSize string) (float64, bool) { model = strings.ToLower(strings.TrimSpace(model)) switch model { + case "grok-imagine-image-2.0": + return getGrokImagineImageTierPrice( + imageSize, + defaultGrokImagineImage20Price1K, + defaultGrokImagineImage20Price2K, + ), true case "grok-imagine-image-quality": return getGrokImagineImageTierPrice( imageSize, diff --git a/backend/internal/service/billing_service_test.go b/backend/internal/service/billing_service_test.go index 8f322b54e9..d1ea67645c 100644 --- a/backend/internal/service/billing_service_test.go +++ b/backend/internal/service/billing_service_test.go @@ -871,8 +871,8 @@ func TestComputeTokenBreakdown_GptImage2ImageEditIssue4386(t *testing.T) { cost := svc.computeTokenBreakdown(pricing, tokens, 1.0, "", false) - wantTextInput := float64(19) * 5e-6 // 0.000095 - wantImageInput := float64(352) * 8e-6 // 0.002816 + wantTextInput := float64(19) * 5e-6 // 0.000095 + wantImageInput := float64(352) * 8e-6 // 0.002816 wantImageOutput := float64(439) * 30e-6 // 0.013170 require.InDelta(t, wantTextInput, cost.InputCost, 1e-15, "InputCost 仅含文本输入") require.InDelta(t, wantImageInput, cost.ImageInputCost, 1e-15, "图片输入按 $8/1M 独立计费") @@ -1153,7 +1153,23 @@ func TestCalculateCostWithLongContext_PropagatesError(t *testing.T) { func TestGetModelPricing_Grok45OfficialFallback(t *testing.T) { svc := newTestBillingService() - for _, model := range []string{"grok", "grok-latest", "grok-4.5", "grok-4.5-latest", "grok-build-latest"} { + for _, model := range []string{"grok", "grok-latest", "grok-4.5", "grok-4.5-latest"} { + model := model + t.Run(model, func(t *testing.T) { + pricing, err := svc.GetModelPricing(model) + require.NoError(t, err) + require.InDelta(t, 2e-6, pricing.InputPricePerToken, 1e-12) + require.InDelta(t, 6e-6, pricing.OutputPricePerToken, 1e-12) + require.InDelta(t, 0.3e-6, pricing.CacheReadPricePerToken, 1e-12) + require.False(t, pricing.SupportsCacheBreakdown) + }) + } +} + +func TestGetModelPricing_Grok46OfficialFallback(t *testing.T) { + svc := newTestBillingService() + + for _, model := range []string{"grok-4.6", "grok-4.6-latest"} { model := model t.Run(model, func(t *testing.T) { pricing, err := svc.GetModelPricing(model) @@ -1161,11 +1177,69 @@ func TestGetModelPricing_Grok45OfficialFallback(t *testing.T) { require.InDelta(t, 2e-6, pricing.InputPricePerToken, 1e-12) require.InDelta(t, 6e-6, pricing.OutputPricePerToken, 1e-12) require.InDelta(t, 0.5e-6, pricing.CacheReadPricePerToken, 1e-12) + require.Equal(t, 200000, pricing.LongContextInputThreshold) + require.InDelta(t, 2.0, pricing.LongContextInputMultiplier, 1e-12) + require.InDelta(t, 2.0, pricing.LongContextOutputMultiplier, 1e-12) require.False(t, pricing.SupportsCacheBreakdown) }) } } +func TestCalculateCostUnified_GroupLongContextToggleUsesPresetLadder(t *testing.T) { + svc := newTestBillingService() + resolver := NewModelPricingResolver(nil, svc) + tokens := UsageTokens{InputTokens: 250000, OutputTokens: 1000} + + off := &Group{LongContextPricingEnabled: false} + disabled, err := svc.CalculateCostUnified(CostInput{ + Model: "grok-4.5", Group: off, Tokens: tokens, RateMultiplier: 1, Resolver: resolver, + }) + require.NoError(t, err) + + on := &Group{LongContextPricingEnabled: true} + enabled, err := svc.CalculateCostUnified(CostInput{ + Model: "grok-4.5", Group: on, Tokens: tokens, RateMultiplier: 1, Resolver: resolver, + }) + require.NoError(t, err) + + require.False(t, disabled.LongContextBillingApplied) + require.True(t, enabled.LongContextBillingApplied) + require.InDelta(t, disabled.InputCost*2, enabled.InputCost, 1e-12) + require.InDelta(t, disabled.OutputCost*2, enabled.OutputCost, 1e-12) +} + +func TestGetModelPricing_UnknownGrokTextFallsBackToGrok45(t *testing.T) { + svc := newTestBillingService() + baseline, err := svc.GetModelPricing("grok-4.5") + require.NoError(t, err) + + for _, model := range []string{"grok-5", "grok-5-latest", "x-ai/grok-7", "grok-4.7-beta"} { + pricing, err := svc.GetModelPricing(model) + require.NoError(t, err, "model %s", model) + require.InDelta(t, baseline.InputPricePerToken, pricing.InputPricePerToken, 1e-12, model) + require.InDelta(t, baseline.OutputPricePerToken, pricing.OutputPricePerToken, 1e-12, model) + require.InDelta(t, baseline.CacheReadPricePerToken, pricing.CacheReadPricePerToken, 1e-12, model) + } + + for _, model := range []string{ + "grok-imagine-image-3.0", + "grok-imagine-video-2", + "grok-voice-latest", + "grok-web-search", + "grok-x-search", + "grok-speech-1", + } { + _, err := svc.GetModelPricing(model) + require.Error(t, err, "non-text grok family %s must not inherit grok-4.5 token rates", model) + require.ErrorIs(t, err, ErrModelPricingUnavailable) + } + + // Known cards stay on their own rate, not the 4.5 family floor. + build, err := svc.GetModelPricing("grok-build-0.1") + require.NoError(t, err) + require.InDelta(t, 1e-6, build.InputPricePerToken, 1e-12) +} + func TestGetModelPricing_GrokCatalogFallbacks(t *testing.T) { svc := newTestBillingService() diff --git a/backend/internal/service/channel.go b/backend/internal/service/channel.go index 5ed834eeca..4d5523b0a3 100644 --- a/backend/internal/service/channel.go +++ b/backend/internal/service/channel.go @@ -20,7 +20,7 @@ const ( // IsValid 检查 BillingMode 是否为合法值 func (m BillingMode) IsValid() bool { switch m { - case BillingModeToken, BillingModePerRequest, BillingModeImage, "": + case BillingModeToken, BillingModePerRequest, BillingModeImage, BillingModeVideo, "": return true } return false @@ -87,38 +87,38 @@ type AccountStatsPricingRule struct { // ChannelModelPricing 渠道模型定价条目 type ChannelModelPricing struct { - ID int64 - ChannelID int64 - Platform string // 所属平台(anthropic/openai/gemini/...) - Models []string // 绑定的模型列表 - BillingMode BillingMode // 计费模式 - InputPrice *float64 // 每 token 输入价格(USD)— 向后兼容 flat 定价 - OutputPrice *float64 // 每 token 输出价格(USD) - CacheWritePrice *float64 // 缓存写入价格 - CacheReadPrice *float64 // 缓存读取价格 - ImageInputPrice *float64 // 图片输入 token 价格(如 gpt-image-2 图片编辑);未配置时回退文本输入价 - ImageOutputPrice *float64 // 图片输出价格(向后兼容) - PerRequestPrice *float64 // 默认按次计费价格(USD) - Intervals []PricingInterval // 区间定价列表 - CreatedAt time.Time - UpdatedAt time.Time + ID int64 `json:"id,omitempty"` + ChannelID int64 `json:"channel_id,omitempty"` + Platform string `json:"platform"` // 所属平台(anthropic/openai/gemini/...) + Models []string `json:"models"` + BillingMode BillingMode `json:"billing_mode"` + InputPrice *float64 `json:"input_price"` + OutputPrice *float64 `json:"output_price"` + CacheWritePrice *float64 `json:"cache_write_price"` + CacheReadPrice *float64 `json:"cache_read_price"` + ImageInputPrice *float64 `json:"image_input_price"` + ImageOutputPrice *float64 `json:"image_output_price"` + PerRequestPrice *float64 `json:"per_request_price"` + Intervals []PricingInterval `json:"intervals"` + CreatedAt time.Time `json:"created_at,omitempty"` + UpdatedAt time.Time `json:"updated_at,omitempty"` } // PricingInterval 定价区间(token 区间 / 按次分层 / 图片分辨率分层) type PricingInterval struct { - ID int64 - PricingID int64 - MinTokens int // 区间下界(含) - MaxTokens *int // 区间上界(不含),nil = 无上限 - TierLabel string // 层级标签(按次/图片模式:1K, 2K, 4K, HD 等) - InputPrice *float64 // token 模式:每 token 输入价 - OutputPrice *float64 // token 模式:每 token 输出价 - CacheWritePrice *float64 // token 模式:缓存写入价 - CacheReadPrice *float64 // token 模式:缓存读取价 - PerRequestPrice *float64 // 按次/图片模式:每次请求价格 - SortOrder int - CreatedAt time.Time - UpdatedAt time.Time + ID int64 `json:"id,omitempty"` + PricingID int64 `json:"pricing_id,omitempty"` + MinTokens int `json:"min_tokens"` + MaxTokens *int `json:"max_tokens"` + TierLabel string `json:"tier_label"` + InputPrice *float64 `json:"input_price"` + OutputPrice *float64 `json:"output_price"` + CacheWritePrice *float64 `json:"cache_write_price"` + CacheReadPrice *float64 `json:"cache_read_price"` + PerRequestPrice *float64 `json:"per_request_price"` + SortOrder int `json:"sort_order"` + CreatedAt time.Time `json:"created_at,omitempty"` + UpdatedAt time.Time `json:"updated_at,omitempty"` } // IsActive 判断渠道是否启用 @@ -315,7 +315,7 @@ func ValidateIntervals(intervals []PricingInterval, mode BillingMode) error { } // per_request / image 模式按 tier_label 匹配,不做 token 区间重叠校验 - if mode == BillingModePerRequest || mode == BillingModeImage { + if mode == BillingModePerRequest || mode == BillingModeImage || mode == BillingModeVideo { return nil } return validateIntervalOverlap(sorted) diff --git a/backend/internal/service/channel_service.go b/backend/internal/service/channel_service.go index e56f699a8e..a7e9bffe31 100644 --- a/backend/internal/service/channel_service.go +++ b/backend/internal/service/channel_service.go @@ -659,7 +659,7 @@ func validatePricingBillingMode(pricing []ChannelModelPricing) error { } func checkBillingModeRequirements(p ChannelModelPricing) error { - if p.BillingMode == BillingModePerRequest || p.BillingMode == BillingModeImage { + if p.BillingMode == BillingModePerRequest || p.BillingMode == BillingModeImage || p.BillingMode == BillingModeVideo { if p.PerRequestPrice == nil && len(p.Intervals) == 0 { return infraerrors.BadRequest( "BILLING_MODE_MISSING_PRICE", diff --git a/backend/internal/service/gateway_usage_billing.go b/backend/internal/service/gateway_usage_billing.go index a5998f3595..d28a1be19a 100644 --- a/backend/internal/service/gateway_usage_billing.go +++ b/backend/internal/service/gateway_usage_billing.go @@ -951,6 +951,18 @@ func (s *GatewayService) calculateRecordUsageCost( // Voice audio (TTS / STT / realtime) when present on the forward result. if result.AudioUsage != nil { + if resolved := s.resolveChannelPricing(ctx, billingModel, apiKey); resolved != nil && + resolved.Mode == BillingModePerRequest { + gid := apiKey.Group.ID + cost, err := s.billingService.CalculateCostUnified(CostInput{ + Ctx: ctx, Model: billingModel, GroupID: &gid, Group: apiKey.Group, + UsageUnits: result.AudioUsage.DurationOrUnits, SizeTier: result.AudioUsage.Mode, + RateMultiplier: multiplier, Resolver: s.resolver, Resolved: resolved, + }) + if err == nil { + return cost + } + } cfg := groupAudioPriceConfigFromAPIKey(apiKey) return s.billingService.CalculateAudioCost(result.AudioUsage.Mode, result.AudioUsage.DurationOrUnits, cfg, multiplier) } @@ -1047,8 +1059,8 @@ func (s *GatewayService) resolveChannelPricing(ctx context.Context, billingModel return nil } gid := apiKey.Group.ID - resolved := s.resolver.Resolve(ctx, PricingInput{Model: billingModel, GroupID: &gid}) - if resolved.Source == PricingSourceChannel { + resolved := s.resolver.Resolve(ctx, PricingInput{Model: billingModel, GroupID: &gid, Group: apiKey.Group}) + if resolved.Source == PricingSourceGroup || resolved.Source == PricingSourceChannel { return resolved } return nil @@ -1063,11 +1075,23 @@ func (s *GatewayService) calculateImageCost( multiplier float64, ) *CostBreakdown { sizeTier := NormalizeImageBillingTierOrDefault(result.ImageSize) + resolved := s.resolveChannelPricing(ctx, billingModel, apiKey) + if resolved != nil && resolved.Source == PricingSourceGroup { + gid := apiKey.Group.ID + cost, err := s.billingService.CalculateCostUnified(CostInput{ + Ctx: ctx, Model: billingModel, GroupID: &gid, Group: apiKey.Group, + RequestCount: result.ImageCount, SizeTier: sizeTier, + RateMultiplier: multiplier, Resolver: s.resolver, Resolved: resolved, + }) + if err == nil { + return cost + } + } groupConfig := imagePriceConfigFromAPIKey(apiKey) if apiKeyHasConfiguredImagePrice(apiKey, sizeTier) { return s.billingService.CalculateImageCost(billingModel, sizeTier, result.ImageCount, groupConfig, multiplier) } - if resolved := s.resolveChannelPricing(ctx, billingModel, apiKey); resolved != nil { + if resolved != nil && resolved.Source == PricingSourceChannel { tokens := UsageTokens{ InputTokens: result.Usage.InputTokens, OutputTokens: result.Usage.OutputTokens, @@ -1078,6 +1102,7 @@ func (s *GatewayService) calculateImageCost( Ctx: ctx, Model: billingModel, GroupID: &gid, + Group: apiKey.Group, Tokens: tokens, RequestCount: result.ImageCount, SizeTier: sizeTier, @@ -1117,22 +1142,30 @@ func (s *GatewayService) calculateTokenCost( var cost *CostBreakdown var err error - // 优先尝试渠道定价 → CalculateCostUnified + // Explicit group/channel pricing wins. Built-in pricing also uses the unified + // resolver so the group long-context toggle can veto model-native tiers. if resolved := s.resolveChannelPricing(ctx, billingModel, apiKey); resolved != nil { gid := apiKey.Group.ID cost, err = s.billingService.CalculateCostUnified(CostInput{ Ctx: ctx, Model: billingModel, GroupID: &gid, + Group: apiKey.Group, Tokens: tokens, RequestCount: 1, RateMultiplier: multiplier, Resolver: s.resolver, Resolved: resolved, }) - } else if opts.LongContextThreshold > 0 { + } else if opts.LongContextThreshold > 0 && (apiKey.Group == nil || apiKey.Group.LongContextPricingEnabled) { // 长上下文双倍计费(如 Gemini 200K 阈值) cost, err = s.billingService.CalculateCostWithLongContext(billingModel, tokens, multiplier, opts.LongContextThreshold, opts.LongContextMultiplier) + } else if s.resolver != nil && apiKey.Group != nil { + gid := apiKey.Group.ID + cost, err = s.billingService.CalculateCostUnified(CostInput{ + Ctx: ctx, Model: billingModel, GroupID: &gid, Group: apiKey.Group, + Tokens: tokens, RequestCount: 1, RateMultiplier: multiplier, Resolver: s.resolver, + }) } else { cost, err = s.billingService.CalculateCost(billingModel, tokens, multiplier) } diff --git a/backend/internal/service/grok_audio.go b/backend/internal/service/grok_audio.go index 0b16047966..06c8410b8d 100644 --- a/backend/internal/service/grok_audio.go +++ b/backend/internal/service/grok_audio.go @@ -8,6 +8,7 @@ import ( "net/http" "net/url" "strings" + "sync/atomic" "time" coderws "github.com/coder/websocket" @@ -117,20 +118,20 @@ func (s *OpenAIGatewayService) ForwardGrokVoice(ctx context.Context, c *gin.Cont // ProxyGrokRealtime relays JSON Realtime events to xAI's native Voice WS. // Audio is carried as base64 inside JSON events, so preserving the JSON bytes // is sufficient and avoids translating protocol event types. -func (s *OpenAIGatewayService) ProxyGrokRealtime(ctx context.Context, c *gin.Context, client *coderws.Conn, account *Account, token, model string) error { +func (s *OpenAIGatewayService) ProxyGrokRealtime(ctx context.Context, c *gin.Context, client *coderws.Conn, account *Account, token, model string) (bool, error) { if s == nil || client == nil || account == nil { - return fmt.Errorf("realtime service, client, and account are required") + return false, fmt.Errorf("realtime service, client, and account are required") } if account.Platform != PlatformGrok { - return fmt.Errorf("account platform %s is not supported for grok realtime", account.Platform) + return false, fmt.Errorf("account platform %s is not supported for grok realtime", account.Platform) } base, err := buildGrokVoiceURL(account, s.cfg, "realtime") if err != nil { - return err + return false, err } u, err := url.Parse(base) if err != nil { - return err + return false, err } u.Scheme = "wss" u.RawQuery = "model=" + url.QueryEscape(firstNonEmpty(model, "grok-voice-latest")) @@ -150,13 +151,14 @@ func (s *OpenAIGatewayService) ProxyGrokRealtime(ctx context.Context, c *gin.Con } upstream, _, _, err := dialer.Dial(ctx, u.String(), headers, proxyURL) if err != nil { - return err + return false, err } defer func() { _ = upstream.Close() }() ctx, cancel := context.WithCancel(ctx) defer cancel() errCh := make(chan error, 2) + var audioObserved atomic.Bool // Upstream → client go func() { @@ -166,6 +168,9 @@ func (s *OpenAIGatewayService) ProxyGrokRealtime(ctx context.Context, c *gin.Con errCh <- readErr return } + if grokRealtimeEventHasAudio(msg) { + audioObserved.Store(true) + } if writeErr := client.Write(ctx, coderws.MessageText, msg); writeErr != nil { errCh <- writeErr return @@ -184,6 +189,9 @@ func (s *OpenAIGatewayService) ProxyGrokRealtime(ctx context.Context, c *gin.Con if kind != coderws.MessageText && kind != coderws.MessageBinary { continue } + if grokRealtimeEventHasAudio(msg) { + audioObserved.Store(true) + } var raw json.RawMessage if unmarshalErr := json.Unmarshal(msg, &raw); unmarshalErr != nil { errCh <- fmt.Errorf("invalid realtime event: %w", unmarshalErr) @@ -196,7 +204,32 @@ func (s *OpenAIGatewayService) ProxyGrokRealtime(ctx context.Context, c *gin.Con } }() - return <-errCh + return awaitGrokRealtimeAudioObserved(errCh, &audioObserved) +} + +func awaitGrokRealtimeAudioObserved(errCh <-chan error, audioObserved *atomic.Bool) (bool, error) { + err := <-errCh + if audioObserved == nil { + return false, err + } + return audioObserved.Load(), err +} + +func grokRealtimeEventHasAudio(msg []byte) bool { + if !gjson.ValidBytes(msg) { + return false + } + eventType := strings.ToLower(strings.TrimSpace(gjson.GetBytes(msg, "type").String())) + if !strings.Contains(eventType, "audio") || strings.Contains(eventType, "transcript") { + return false + } + for _, path := range []string{"audio", "delta", "data"} { + value := gjson.GetBytes(msg, path) + if value.Type == gjson.String && strings.TrimSpace(value.String()) != "" { + return true + } + } + return false } // estimateGrokVoiceAudioUsage derives billing units from the request/response. diff --git a/backend/internal/service/grok_audio_test.go b/backend/internal/service/grok_audio_test.go index 57a60431ad..410af06042 100644 --- a/backend/internal/service/grok_audio_test.go +++ b/backend/internal/service/grok_audio_test.go @@ -2,6 +2,8 @@ package service import ( "context" + "io" + "sync/atomic" "testing" "github.com/Wei-Shaw/sub2api/internal/pkg/xai" @@ -59,6 +61,26 @@ func TestForwardGrokVoice_RejectsNonGrok(t *testing.T) { require.Contains(t, err.Error(), "not supported") } +func TestAwaitGrokRealtimeAudioObservedReadsFlagAfterRelayExits(t *testing.T) { + errCh := make(chan error, 1) + var observed atomic.Bool + go func() { + observed.Store(true) + errCh <- io.EOF + }() + got, err := awaitGrokRealtimeAudioObserved(errCh, &observed) + require.ErrorIs(t, err, io.EOF) + require.True(t, got, "audioObserved must be read after the relay returns, not before <-errCh") +} + +func TestGrokRealtimeEventHasAudio(t *testing.T) { + require.False(t, grokRealtimeEventHasAudio([]byte(`{"type":"session.created"}`))) + require.False(t, grokRealtimeEventHasAudio([]byte(`{"type":"response.audio_transcript.delta","delta":"hi"}`))) + require.False(t, grokRealtimeEventHasAudio([]byte(`{"type":"response.audio.delta","delta":""}`))) + require.True(t, grokRealtimeEventHasAudio([]byte(`{"type":"response.audio.delta","delta":"abc"}`))) + require.True(t, grokRealtimeEventHasAudio([]byte(`{"type":"response.output_audio.delta","audio":"abc"}`))) +} + func TestForwardGrokVoice_RejectsUnknownEndpoint(t *testing.T) { svc := &OpenAIGatewayService{} _, err := svc.ForwardGrokVoice(context.Background(), nil, &Account{Platform: PlatformGrok}, "unknown", []byte(`{}`), "application/json") diff --git a/backend/internal/service/grok_media.go b/backend/internal/service/grok_media.go index 953d3db32a..2efb2332f6 100644 --- a/backend/internal/service/grok_media.go +++ b/backend/internal/service/grok_media.go @@ -695,7 +695,7 @@ func (s *OpenAIGatewayService) ForwardGrokMedia( return s.handleGrokMediaErrorResponse(ctx, resp, c, account, requestIDHeader, requestModel) } - s.updateGrokUsageFromResponse(ctx, account, resp.Header, resp.StatusCode) + s.updateGrokUsageFromResponse(withGrokTeamRateLimitModel(ctx, requestModel), account, resp.Header, resp.StatusCode) respBody, err := ReadUpstreamResponseBody(resp.Body, s.cfg, c, openAITooLargeError) if err != nil { return nil, err @@ -857,7 +857,7 @@ func (s *OpenAIGatewayService) forwardGrokMediaVideoContent( return s.handleGrokMediaErrorResponse(ctx, contentResp, c, account, contentRequestID, "") } - s.updateGrokUsageFromResponse(ctx, account, contentResp.Header, contentResp.StatusCode) + s.updateGrokUsageFromResponse(withGrokTeamRateLimitModel(ctx, ""), account, contentResp.Header, contentResp.StatusCode) if err := writeGrokMediaContentResponse(c, contentResp); err != nil { return nil, err } diff --git a/backend/internal/service/grok_model_quota_block.go b/backend/internal/service/grok_model_quota_block.go index 0bc9cd40b8..e6c7841da4 100644 --- a/backend/internal/service/grok_model_quota_block.go +++ b/backend/internal/service/grok_model_quota_block.go @@ -47,6 +47,32 @@ func markGrokModelQuotaBlock(accountID int64, model string, until time.Time) { if max := now.Add(grokModelQuotaBlockMaxTTL); until.After(max) { until = max } + storeGrokModelQuotaBlock(accountID, model, until, now) +} + +const ( + grokModelTransientBlockMinTTL = 500 * time.Millisecond + grokModelTransientBlockMaxTTL = 5 * time.Minute +) + +// markGrokModelTransientBlock soft-blocks a single model for a short capacity +// burst without the free-usage 20m floor (and without unscheduling the account). +func markGrokModelTransientBlock(accountID int64, model string, until time.Time) { + model = strings.TrimSpace(model) + if accountID <= 0 || model == "" || until.IsZero() { + return + } + now := time.Now() + if !until.After(now.Add(grokModelTransientBlockMinTTL)) { + until = now.Add(grokModelTransientBlockMinTTL) + } + if max := now.Add(grokModelTransientBlockMaxTTL); until.After(max) { + until = max + } + storeGrokModelQuotaBlock(accountID, model, until, now) +} + +func storeGrokModelQuotaBlock(accountID int64, model string, until, now time.Time) { key := grokModelQuotaBlockKey(accountID, model) globalGrokModelQuotaBlocks.mu.Lock() defer globalGrokModelQuotaBlocks.mu.Unlock() diff --git a/backend/internal/service/grok_oauth_service.go b/backend/internal/service/grok_oauth_service.go index 19120e84bd..2d6c38cb26 100644 --- a/backend/internal/service/grok_oauth_service.go +++ b/backend/internal/service/grok_oauth_service.go @@ -330,8 +330,14 @@ func (s *GrokOAuthService) RefreshAccountToken(ctx context.Context, account *Acc if err != nil { return nil, err } - tokenInfo.SubscriptionTier = account.GetCredential("subscription_tier") - tokenInfo.EntitlementStatus = account.GetCredential("entitlement_status") + // New access-token JWT is authoritative. Keep the stored value only when + // the refreshed token has no tier claim (opaque AT / missing field). + if strings.TrimSpace(tokenInfo.SubscriptionTier) == "" { + tokenInfo.SubscriptionTier = account.GetCredential("subscription_tier") + } + if strings.TrimSpace(tokenInfo.EntitlementStatus) == "" { + tokenInfo.EntitlementStatus = account.GetCredential("entitlement_status") + } return tokenInfo, nil } @@ -404,8 +410,8 @@ func (s *GrokOAuthService) tokenInfoFromResponse(tokenResp *xai.TokenResponse, c if info.TokenType == "" { info.TokenType = "Bearer" } - applyGrokTokenClaims(info, tokenResp.IDToken) - applyGrokTokenClaims(info, tokenResp.AccessToken) + applyGrokTokenClaims(info, tokenResp.IDToken, false) + applyGrokTokenClaims(info, tokenResp.AccessToken, true) if existing != nil { if info.Email == "" { if email, _ := existing["email"].(string); email != "" { @@ -446,7 +452,7 @@ func (s *GrokOAuthService) proxyURL(ctx context.Context, proxyID *int64) (string return proxy.URL(), nil } -func applyGrokTokenClaims(info *GrokTokenInfo, token string) { +func applyGrokTokenClaims(info *GrokTokenInfo, token string, includeTier bool) { if info == nil || strings.TrimSpace(token) == "" { return } @@ -463,4 +469,9 @@ func applyGrokTokenClaims(info *GrokTokenInfo, token string) { if info.TeamID == "" { info.TeamID = xai.JWTClaimString(claims, "team_id") } + if includeTier { + if tier := xai.SubscriptionTierFromJWT(token); tier != "" { + info.SubscriptionTier = tier + } + } } diff --git a/backend/internal/service/grok_oauth_service_test.go b/backend/internal/service/grok_oauth_service_test.go index cbbd802e01..d6fc4918fd 100644 --- a/backend/internal/service/grok_oauth_service_test.go +++ b/backend/internal/service/grok_oauth_service_test.go @@ -197,7 +197,7 @@ func TestGrokOAuthServiceBuildAccountCredentialsDefaultsToSubscriptionProxy(t *t func TestGrokOAuthServiceConvertFromSSOExtractsBuildClaims(t *testing.T) { svc := NewGrokOAuthService(nil, &grokOAuthClientStub{ ssoResponse: &xai.TokenResponse{ - AccessToken: makeGrokOAuthJWT(map[string]any{"sub": "user-sub", "team_id": "team-1"}), + AccessToken: makeGrokOAuthJWT(map[string]any{"sub": "user-sub", "team_id": "team-1", "tier": 5}), RefreshToken: "refresh-token", IDToken: makeGrokOAuthJWT(map[string]any{"email": "user@example.com"}), ExpiresIn: 3600, @@ -210,14 +210,93 @@ func TestGrokOAuthServiceConvertFromSSOExtractsBuildClaims(t *testing.T) { require.Equal(t, "user@example.com", info.Email) require.Equal(t, "user-sub", info.Subject) require.Equal(t, "team-1", info.TeamID) + require.Equal(t, "supergrok_heavy", info.SubscriptionTier) credentials := svc.BuildAccountCredentials(info) require.Equal(t, "user@example.com", credentials["email"]) require.Equal(t, "user-sub", credentials["sub"]) require.Equal(t, "team-1", credentials["team_id"]) + require.Equal(t, "supergrok_heavy", credentials["subscription_tier"]) require.NotContains(t, credentials, "sso_token") } +func TestGrokOAuthServiceRefreshAccountTokenOverwritesStaleTierFromNewJWT(t *testing.T) { + svc := NewGrokOAuthService(nil, &grokOAuthClientStub{ + refreshResponse: &xai.TokenResponse{ + AccessToken: makeGrokOAuthJWT(map[string]any{"sub": "user-sub", "tier": 0}), + TokenType: "Bearer", + ExpiresIn: 3600, + }, + }) + defer svc.Stop() + + account := &Account{ + Platform: PlatformGrok, + Type: AccountTypeOAuth, + Credentials: map[string]any{ + "refresh_token": "refresh-token", + "client_id": "client-id", + "subscription_tier": "supergrok_heavy", + }, + } + + info, err := svc.RefreshAccountToken(context.Background(), account) + require.NoError(t, err) + require.Equal(t, "free", info.SubscriptionTier) + + credentials := svc.BuildAccountCredentials(info) + require.Equal(t, "free", credentials["subscription_tier"]) +} + +func TestGrokOAuthServiceRefreshAccountTokenIgnoresIDTokenTierWhenAccessTokenHasNone(t *testing.T) { + svc := NewGrokOAuthService(nil, &grokOAuthClientStub{ + refreshResponse: &xai.TokenResponse{ + AccessToken: "opaque-access-token", + IDToken: makeGrokOAuthJWT(map[string]any{"tier": 5}), + TokenType: "Bearer", + ExpiresIn: 3600, + }, + }) + defer svc.Stop() + + account := &Account{ + Platform: PlatformGrok, + Type: AccountTypeOAuth, + Credentials: map[string]any{ + "refresh_token": "refresh-token", + "subscription_tier": "supergrok_lite", + }, + } + + info, err := svc.RefreshAccountToken(context.Background(), account) + require.NoError(t, err) + require.Equal(t, "supergrok_lite", info.SubscriptionTier) +} + +func TestGrokOAuthServiceRefreshAccountTokenKeepsStoredTierWhenJWTHasNoClaim(t *testing.T) { + svc := NewGrokOAuthService(nil, &grokOAuthClientStub{ + refreshResponse: &xai.TokenResponse{ + AccessToken: "opaque-access-token", + TokenType: "Bearer", + ExpiresIn: 3600, + }, + }) + defer svc.Stop() + + account := &Account{ + Platform: PlatformGrok, + Type: AccountTypeOAuth, + Credentials: map[string]any{ + "refresh_token": "refresh-token", + "subscription_tier": "supergrok_lite", + }, + } + + info, err := svc.RefreshAccountToken(context.Background(), account) + require.NoError(t, err) + require.Equal(t, "supergrok_lite", info.SubscriptionTier) +} + func TestGrokOAuthServiceValidateSSOTokenReturnsOAuthTokensWithoutPersistingSSO(t *testing.T) { svc := NewGrokOAuthService(nil, &grokOAuthClientStub{ ssoResponse: &xai.TokenResponse{ diff --git a/backend/internal/service/grok_quota_fetcher.go b/backend/internal/service/grok_quota_fetcher.go index 4979dfd0b9..f41fa0bd2f 100644 --- a/backend/internal/service/grok_quota_fetcher.go +++ b/backend/internal/service/grok_quota_fetcher.go @@ -71,7 +71,7 @@ func (f *GrokQuotaFetcher) BuildUsageInfo(account *Account) *UsageInfo { } if err != nil || snapshot == nil { - applyGrokCredentialUsageFallback(usage, account) + applyGrokCredentialUsageFallback(usage, account, billing, nil) 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" @@ -139,7 +139,7 @@ func (f *GrokQuotaFetcher) BuildUsageInfo(account *Account) *UsageInfo { usage.ErrorCode = "spending_limit" } } - applyGrokCredentialUsageFallback(usage, account) + applyGrokCredentialUsageFallback(usage, account, billing, snapshot) if activeProbeClearsForbidden && strings.TrimSpace(snapshot.EntitlementStatus) == "" && strings.EqualFold(strings.TrimSpace(usage.GrokEntitlementStatus), "forbidden") { usage.GrokEntitlementStatus = "" @@ -170,18 +170,47 @@ func firstGrokObservationTime(values ...string) (time.Time, bool) { return time.Time{}, false } -func applyGrokCredentialUsageFallback(usage *UsageInfo, account *Account) { +func applyGrokCredentialUsageFallback(usage *UsageInfo, account *Account, billing *xai.BillingSummary, snapshot *xai.QuotaSnapshot) { 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")) } + applyGrokResolvedSubscriptionTier(usage, account, billing, snapshot) +} + +func applyGrokResolvedSubscriptionTier(usage *UsageInfo, account *Account, billing *xai.BillingSummary, snapshot *xai.QuotaSnapshot) { + if usage == nil || account == nil { + return + } + if jwtTier := xai.SubscriptionTierFromJWT(account.GetCredential("access_token")); jwtTier != "" { + usage.SubscriptionTier = jwtTier + usage.SubscriptionTierRaw = jwtTier + return + } + signal := strings.TrimSpace(account.GetCredential("subscription_tier")) + if signal == "" && snapshot != nil { + signal = strings.TrimSpace(snapshot.SubscriptionTier) + } + if signal == "" && billing != nil { + signal = strings.TrimSpace(billing.Plan) + } + var limit *float64 + if billing != nil { + limit = billing.MonthlyLimitCents + } + if plan := xai.CanonicalGrokPlan(limit, signal, snapshot); plan != "" { + usage.SubscriptionTier = plan + if usage.SubscriptionTierRaw == "" { + usage.SubscriptionTierRaw = firstNonEmpty(signal, plan) + } + return + } + if usage.SubscriptionTier == "" && signal != "" { + usage.SubscriptionTier = signal + usage.SubscriptionTierRaw = signal + } } func grokBillingSnapshotFromExtra(extra map[string]any) (*xai.BillingSummary, error) { @@ -220,6 +249,23 @@ func grokBillingSnapshotFromExtra(extra map[string]any) (*xai.BillingSummary, er } } +func stampGrokQuotaSnapshotForPlan(account *Account, snapshot *xai.QuotaSnapshot, model string) { + if snapshot == nil { + return + } + if strings.TrimSpace(snapshot.Model) == "" { + model = strings.TrimSpace(model) + if model != "" { + snapshot.Model = xai.ResolveGrokTextResponsesModelID(model) + } + } + var prev *xai.QuotaSnapshot + if account != nil { + prev, _ = grokQuotaSnapshotFromExtra(account.Extra) + } + snapshot.ApplyGrok45ResponsesPlanSignal(prev) +} + 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 a54ecc6c81..c87a925e9e 100644 --- a/backend/internal/service/grok_quota_fetcher_test.go +++ b/backend/internal/service/grok_quota_fetcher_test.go @@ -24,6 +24,110 @@ func TestGrokQuotaFetcherBuildUsageInfoUnknownUntilFirstSnapshot(t *testing.T) { require.Contains(t, usage.Error, "unknown until billing is probed") } +func TestGrokQuotaFetcherDoesNotTreatGrok45ResponsesWindowAsHeavy(t *testing.T) { + t.Parallel() + + // 8300 / 53M is the grok-4.5 Responses rate-limit window, not a plan fingerprint. + reqLimit, tokLimit := int64(8300), int64(53_000_000) + fresh := time.Now().UTC().Format(time.RFC3339) + account := &Account{ + Platform: PlatformGrok, + Type: AccountTypeOAuth, + Credentials: map[string]any{ + "subscription_tier": "SuperGrokPro", + }, + Extra: map[string]any{ + grokQuotaSnapshotExtraKey: &xai.QuotaSnapshot{ + Requests: &xai.QuotaWindow{Limit: &reqLimit}, + Tokens: &xai.QuotaWindow{Limit: &tokLimit}, + LastHeadersSeenAt: fresh, + HeadersObserved: true, + UpdatedAt: fresh, + }, + }, + } + + usage := NewGrokQuotaFetcher().BuildUsageInfo(account) + require.Equal(t, "supergrok", usage.SubscriptionTier) +} + +func TestGrokQuotaFetcherUsesGrok45ResponsesWindowAsHeavy(t *testing.T) { + t.Parallel() + + reqLimit, tokLimit := int64(8300), int64(53_000_000) + fresh := time.Now().UTC().Format(time.RFC3339) + account := &Account{ + Platform: PlatformGrok, + Type: AccountTypeOAuth, + Credentials: map[string]any{ + "subscription_tier": "SuperGrokPro", + }, + Extra: map[string]any{ + grokQuotaSnapshotExtraKey: &xai.QuotaSnapshot{ + Model: "grok-4.5", + Requests: &xai.QuotaWindow{Limit: &reqLimit}, + Tokens: &xai.QuotaWindow{Limit: &tokLimit}, + LastHeadersSeenAt: fresh, + HeadersObserved: true, + UpdatedAt: fresh, + }, + }, + } + + usage := NewGrokQuotaFetcher().BuildUsageInfo(account) + require.Equal(t, "supergrok_heavy", usage.SubscriptionTier) +} + +func TestGrokQuotaFetcherJWTBeatsAmbiguousSuperGrokProQuota(t *testing.T) { + t.Parallel() + + heavyReq := int64(8300) + fresh := time.Now().UTC().Format(time.RFC3339) + account := &Account{ + Platform: PlatformGrok, + Type: AccountTypeOAuth, + Credentials: map[string]any{ + "access_token": makeGrokOAuthJWT(map[string]any{"tier": 1}), + "subscription_tier": "SuperGrokPro", + }, + Extra: map[string]any{ + grokQuotaSnapshotExtraKey: &xai.QuotaSnapshot{ + Requests: &xai.QuotaWindow{Limit: &heavyReq}, + LastHeadersSeenAt: fresh, + HeadersObserved: true, + UpdatedAt: fresh, + }, + }, + } + + usage := NewGrokQuotaFetcher().BuildUsageInfo(account) + require.Equal(t, "supergrok", usage.SubscriptionTier) +} + +func TestGrokQuotaFetcherPrefersLiveJWTTierOverStaleBillingPlan(t *testing.T) { + t.Parallel() + + account := &Account{ + Platform: PlatformGrok, + Type: AccountTypeOAuth, + Credentials: map[string]any{ + "access_token": makeGrokOAuthJWT(map[string]any{"tier": 0}), + "subscription_tier": "supergrok_heavy", + }, + Extra: map[string]any{ + grokBillingExtraKey: &xai.BillingSummary{ + Plan: "SuperGrok Heavy", + StatusCode: http.StatusOK, + UpdatedAt: "2030-01-01T00:00:00Z", + }, + }, + } + + usage := NewGrokQuotaFetcher().BuildUsageInfo(account) + require.Equal(t, "free", usage.SubscriptionTier) + require.Equal(t, "free", usage.SubscriptionTierRaw) +} + func TestGrokQuotaFetcherUsesCredentialTierWhenBillingHasNoPlan(t *testing.T) { t.Parallel() @@ -46,7 +150,7 @@ func TestGrokQuotaFetcherUsesCredentialTierWhenBillingHasNoPlan(t *testing.T) { usage := NewGrokQuotaFetcher().BuildUsageInfo(account) require.NotNil(t, usage.GrokBilling) - require.Equal(t, "FREE", usage.SubscriptionTier) + require.Equal(t, "free", usage.SubscriptionTier) require.Equal(t, "FREE", usage.SubscriptionTierRaw) require.Equal(t, "active", usage.GrokEntitlementStatus) } diff --git a/backend/internal/service/grok_quota_service.go b/backend/internal/service/grok_quota_service.go index 6658888a59..293dbcceab 100644 --- a/backend/internal/service/grok_quota_service.go +++ b/backend/internal/service/grok_quota_service.go @@ -181,6 +181,7 @@ func (s *GrokQuotaService) probeUsage(ctx context.Context, accountID int64) (*Gr defer func() { _ = resp.Body.Close() }() snapshot := xai.ObserveQuotaHeaders(resp.Header, resp.StatusCode, "active_probe") + stampGrokQuotaSnapshotForPlan(account, snapshot, probeModel) resetAt, limited := grokRateLimitResetAtForAccount(account, snapshot, time.Now()) if limited { normalizeGrokExhaustedWindowResets(snapshot, resetAt, time.Now()) diff --git a/backend/internal/service/grok_upstream_failure.go b/backend/internal/service/grok_upstream_failure.go index 1348d068c5..e7bed08bf0 100644 --- a/backend/internal/service/grok_upstream_failure.go +++ b/backend/internal/service/grok_upstream_failure.go @@ -489,6 +489,9 @@ func (s *OpenAIGatewayService) applyGrokUpstreamFailureDecision( case GrokFailureEmptyUpstream: reason = "grok empty model output" case GrokFailureModelCapacity: + if persistGrokTransientModelCooldown(account, decision) { + return true + } reason = "grok model capacity" case GrokFailureRateLimit: // Pure 429 without free-usage language keeps the existing rate-limit diff --git a/backend/internal/service/grok_upstream_failure_test.go b/backend/internal/service/grok_upstream_failure_test.go index 6f0a916a78..0bd1063577 100644 --- a/backend/internal/service/grok_upstream_failure_test.go +++ b/backend/internal/service/grok_upstream_failure_test.go @@ -133,6 +133,22 @@ func TestHandleGrokAccountUpstreamError_EmptyOutputCoolsAccount(t *testing.T) { require.WithinDuration(t, before.Add(4*time.Minute), repo.lastTempUnschedUntil, time.Second) } +func TestHandleGrokAccountUpstreamError_MultiAgentCapacityBlocksOnlyThatModel(t *testing.T) { + repo := &grokQuotaAccountRepo{} + svc := &OpenAIGatewayService{accountRepo: repo} + account := &Account{ID: 9120, Platform: PlatformGrok, Type: AccountTypeOAuth} + ctx := withGrokTeamRateLimitModel(context.Background(), "grok-4.20-multi-agent-0309") + + svc.handleGrokAccountUpstreamError( + ctx, account, http.StatusBadGateway, nil, + []byte(`{"error":{"message":"engine_overloaded"}}`), + ) + + require.Zero(t, repo.tempUnschedCalls) + require.True(t, isGrokModelQuotaBlocked(account.ID, "grok-4.20-multi-agent-0309", time.Now())) + require.False(t, isGrokModelQuotaBlocked(account.ID, "grok-4.5", time.Now())) +} + func TestHandleGrokAccountUpstreamError_FreeUsageDoesNotCoolPoolMode(t *testing.T) { repo := &grokQuotaAccountRepo{} svc := &OpenAIGatewayService{accountRepo: repo} diff --git a/backend/internal/service/group.go b/backend/internal/service/group.go index ddde276e22..bf0440afc3 100644 --- a/backend/internal/service/group.go +++ b/backend/internal/service/group.go @@ -70,6 +70,11 @@ type Group struct { AudioTTSPricePerMillionChars *float64 AudioSTTPricePerHour *float64 + // ModelPricing overrides channel and built-in prices for matching models. + // Token intervals are selected only when LongContextPricingEnabled is true. + LongContextPricingEnabled bool + ModelPricing []ChannelModelPricing + // Claude Code 客户端限制 ClaudeCodeOnly bool FallbackGroupID *int64 diff --git a/backend/internal/service/model_pricing_resolver.go b/backend/internal/service/model_pricing_resolver.go index ea5f4a29d1..52c6b930c5 100644 --- a/backend/internal/service/model_pricing_resolver.go +++ b/backend/internal/service/model_pricing_resolver.go @@ -3,10 +3,12 @@ package service import ( "context" "log/slog" + "strings" ) // PricingSource 定价来源标识 const ( + PricingSourceGroup = "group" PricingSourceChannel = "channel" PricingSourceLiteLLM = "litellm" PricingSourceFallback = "fallback" @@ -37,10 +39,12 @@ type ResolvedPricing struct { // 渠道定价原始配置(用于区间模式下获取 ImageOutputPrice) channelPricing *ChannelModelPricing + + longContextPricingEnabled bool } // ModelPricingResolver 统一模型定价解析器。 -// 解析链:Channel → LiteLLM → Fallback。 +// 解析链:Group → Channel → LiteLLM → Fallback。 type ModelPricingResolver struct { channelService *ChannelService billingService *BillingService @@ -58,12 +62,27 @@ func NewModelPricingResolver(channelService *ChannelService, billingService *Bil type PricingInput struct { Model string GroupID *int64 // nil 表示不检查渠道 + Group *Group } // Resolve 解析模型定价。 // 1. 获取基础定价(LiteLLM → Fallback) // 2. 如果指定了 GroupID,查找渠道定价并覆盖 func (r *ModelPricingResolver) Resolve(ctx context.Context, input PricingInput) *ResolvedPricing { + longContextPricingEnabled := input.Group == nil || input.Group.LongContextPricingEnabled + if groupPricing := matchGroupModelPricing(input.Group, input.Model); groupPricing != nil { + // Group token cards only override the first-tier / flat rates. + // Long-context ladders come from official presets, gated by the checkbox. + if groupPricing.BillingMode == "" || groupPricing.BillingMode == BillingModeToken { + stripped := groupPricing.Clone() + stripped.Intervals = nil + groupPricing = &stripped + } + resolved := r.resolveConfiguredPricing(groupPricing, input.Model, PricingSourceGroup) + resolved.longContextPricingEnabled = longContextPricingEnabled + return resolved + } + var chPricing *ChannelModelPricing if input.GroupID != nil && r.channelService != nil { chPricing = r.channelService.GetChannelModelPricing(ctx, *input.GroupID, input.Model) @@ -72,12 +91,13 @@ func (r *ModelPricingResolver) Resolve(ctx context.Context, input PricingInput) if mode == "" { mode = BillingModeToken } - if mode == BillingModePerRequest || mode == BillingModeImage { + if mode == BillingModePerRequest || mode == BillingModeImage || mode == BillingModeVideo { resolved := &ResolvedPricing{ Mode: mode, Source: PricingSourceChannel, channelPricing: chPricing, } + resolved.longContextPricingEnabled = longContextPricingEnabled r.applyRequestTierOverrides(chPricing, resolved) return resolved } @@ -93,19 +113,79 @@ func (r *ModelPricingResolver) Resolve(ctx context.Context, input PricingInput) Source: source, SupportsCacheBreakdown: basePricing != nil && basePricing.SupportsCacheBreakdown, } + resolved.longContextPricingEnabled = longContextPricingEnabled // 2. 如果有 GroupID,尝试渠道覆盖 if chPricing != nil { resolved.Source = PricingSourceChannel resolved.channelPricing = chPricing r.applyTokenOverrides(chPricing, resolved) - } else if input.GroupID != nil { + if !longContextPricingEnabled { + r.applyFirstTokenTier(resolved, chPricing) + } + } else if input.GroupID != nil && r.channelService != nil { r.applyChannelOverrides(ctx, *input.GroupID, input.Model, resolved) + if resolved.Source == PricingSourceChannel && !longContextPricingEnabled { + r.applyFirstTokenTier(resolved, resolved.channelPricing) + } } return resolved } +func (r *ModelPricingResolver) resolveConfiguredPricing(config *ChannelModelPricing, model, source string) *ResolvedPricing { + mode := config.BillingMode + if mode == "" { + mode = BillingModeToken + } + resolved := &ResolvedPricing{Mode: mode, Source: source, channelPricing: config} + if mode == BillingModePerRequest || mode == BillingModeImage || mode == BillingModeVideo { + r.applyRequestTierOverrides(config, resolved) + return resolved + } + resolved.BasePricing, _ = r.resolveBasePricing(model) + resolved.SupportsCacheBreakdown = resolved.BasePricing != nil && resolved.BasePricing.SupportsCacheBreakdown + r.applyTokenOverrides(config, resolved) + return resolved +} + +func matchGroupModelPricing(group *Group, model string) *ChannelModelPricing { + if group == nil { + return nil + } + model = normalizeChannelPricingModelName(model) + var wildcard *ChannelModelPricing + for i := range group.ModelPricing { + entry := &group.ModelPricing[i] + for _, pattern := range entry.Models { + normalized := normalizeChannelPricingModelName(pattern) + if normalized == model { + cp := entry.Clone() + return &cp + } + if strings.HasSuffix(normalized, "*") && strings.HasPrefix(model, strings.TrimSuffix(normalized, "*")) && wildcard == nil { + cp := entry.Clone() + wildcard = &cp + } + } + } + return wildcard +} + +func (r *ModelPricingResolver) applyFirstTokenTier(resolved *ResolvedPricing, config *ChannelModelPricing) { + if resolved == nil || len(resolved.Intervals) == 0 { + return + } + first := resolved.Intervals[0] + for _, interval := range resolved.Intervals[1:] { + if interval.MinTokens < first.MinTokens { + first = interval + } + } + resolved.BasePricing = intervalToModelPricing(&first, resolved.SupportsCacheBreakdown, config) + resolved.Intervals = nil +} + // resolveBasePricing 从 LiteLLM 或 Fallback 获取基础定价 func (r *ModelPricingResolver) resolveBasePricing(model string) (*ModelPricing, string) { pricing, err := r.billingService.GetModelPricing(model) @@ -134,7 +214,7 @@ func (r *ModelPricingResolver) applyChannelOverrides(ctx context.Context, groupI switch resolved.Mode { case BillingModeToken: r.applyTokenOverrides(chPricing, resolved) - case BillingModePerRequest, BillingModeImage: + case BillingModePerRequest, BillingModeImage, BillingModeVideo: r.applyRequestTierOverrides(chPricing, resolved) } } diff --git a/backend/internal/service/model_pricing_resolver_test.go b/backend/internal/service/model_pricing_resolver_test.go index 71ba9904c2..11613fa9eb 100644 --- a/backend/internal/service/model_pricing_resolver_test.go +++ b/backend/internal/service/model_pricing_resolver_test.go @@ -833,3 +833,66 @@ func TestApplyTokenOverrides_IntervalDoesNotPolluteFallbackPrices(t *testing.T) require.InDelta(t, 15e-6, fp.OutputPricePerToken, 1e-12, "fallback OutputPricePerToken polluted") require.False(t, fp.ImageOutputPriceExplicit, "fallback ImageOutputPriceExplicit polluted") } + +func TestResolve_GroupPricingOverridesChannel(t *testing.T) { + r := newResolverWithChannel(t, []ChannelModelPricing{{ + Platform: "anthropic", Models: []string{"claude-sonnet-4"}, BillingMode: BillingModeToken, + InputPrice: testPtrFloat64(10e-6), OutputPrice: testPtrFloat64(20e-6), + }}) + group := &Group{ID: 100, ModelPricing: []ChannelModelPricing{{ + Models: []string{"claude-sonnet-*"}, BillingMode: BillingModeToken, + InputPrice: testPtrFloat64(1e-6), OutputPrice: testPtrFloat64(2e-6), + }}} + resolved := r.Resolve(context.Background(), PricingInput{Model: "claude-sonnet-4", GroupID: groupIDPtr(), Group: group}) + + require.Equal(t, PricingSourceGroup, resolved.Source) + require.InDelta(t, 1e-6, resolved.BasePricing.InputPricePerToken, 1e-12) + require.InDelta(t, 2e-6, resolved.BasePricing.OutputPricePerToken, 1e-12) +} + +func TestResolve_GroupLongContextUsesPresetNotCustomIntervals(t *testing.T) { + bs := newTestBillingServiceForResolver() + bs.fallbackPrices["claude-sonnet-4"].LongContextInputThreshold = 200000 + bs.fallbackPrices["claude-sonnet-4"].LongContextThresholdInclusive = true + bs.fallbackPrices["claude-sonnet-4"].LongContextInputMultiplier = 2 + bs.fallbackPrices["claude-sonnet-4"].LongContextOutputMultiplier = 2 + r := NewModelPricingResolver(nil, bs) + group := &Group{ID: 100, ModelPricing: []ChannelModelPricing{{ + Models: []string{"claude-sonnet-4"}, BillingMode: BillingModeToken, + InputPrice: testPtrFloat64(1e-6), OutputPrice: testPtrFloat64(2e-6), + Intervals: []PricingInterval{ + {MinTokens: 0, MaxTokens: testPtrInt(200000), InputPrice: testPtrFloat64(9e-6)}, + {MinTokens: 200000, InputPrice: testPtrFloat64(18e-6)}, + }, + }}} + + resolved := r.Resolve(context.Background(), PricingInput{Model: "claude-sonnet-4", Group: group}) + require.False(t, resolved.longContextPricingEnabled) + require.Empty(t, resolved.Intervals, "group token intervals are not a user-facing long-context ladder") + require.InDelta(t, 1e-6, r.GetIntervalPricing(resolved, 300000).InputPricePerToken, 1e-12) + require.Equal(t, 200000, resolved.BasePricing.LongContextInputThreshold) + + group.LongContextPricingEnabled = true + resolved = r.Resolve(context.Background(), PricingInput{Model: "claude-sonnet-4", Group: group}) + require.True(t, resolved.longContextPricingEnabled) + require.Empty(t, resolved.Intervals) + require.InDelta(t, 1e-6, r.GetIntervalPricing(resolved, 300000).InputPricePerToken, 1e-12) + require.Equal(t, 200000, resolved.BasePricing.LongContextInputThreshold) + require.InDelta(t, 2.0, resolved.BasePricing.LongContextInputMultiplier, 1e-12) +} + +func TestCalculateCostUnified_UsesContinuousMediaUnits(t *testing.T) { + bs := newTestBillingServiceForResolver() + r := NewModelPricingResolver(nil, bs) + price := 0.08 + group := &Group{ModelPricing: []ChannelModelPricing{{ + Models: []string{"grok-voice-think-fast-2.0"}, BillingMode: BillingModePerRequest, + PerRequestPrice: &price, + }}} + cost, err := bs.CalculateCostUnified(CostInput{ + Ctx: context.Background(), Model: "grok-voice-think-fast-2.0", Group: group, + UsageUnits: 1.5, RateMultiplier: 1, Resolver: r, + }) + require.NoError(t, err) + require.InDelta(t, 0.12, cost.TotalCost, 1e-12) +} diff --git a/backend/internal/service/openai_gateway_chat_completions_raw.go b/backend/internal/service/openai_gateway_chat_completions_raw.go index 010cd4df2a..34c281b28e 100644 --- a/backend/internal/service/openai_gateway_chat_completions_raw.go +++ b/backend/internal/service/openai_gateway_chat_completions_raw.go @@ -191,7 +191,7 @@ func (s *OpenAIGatewayService) forwardAsRawChatCompletions( Kind: kind, Message: upstreamMsg, }) - s.handleGrokAccountUpstreamError(ctx, account, resp.StatusCode, resp.Header, respBody) + s.handleGrokAccountUpstreamError(withGrokTeamRateLimitModel(ctx, upstreamModel), account, resp.StatusCode, resp.Header, respBody) if s.shouldFailoverGrokUpstreamError(resp.StatusCode, respBody) { return nil, &UpstreamFailoverError{ StatusCode: resp.StatusCode, @@ -209,7 +209,7 @@ func (s *OpenAIGatewayService) forwardAsRawChatCompletions( } if account.Platform == PlatformGrok { - s.updateGrokUsageFromResponse(ctx, account, resp.Header, resp.StatusCode) + s.updateGrokUsageFromResponse(withGrokTeamRateLimitModel(ctx, upstreamModel), account, resp.Header, resp.StatusCode) } // 8. Forward response diff --git a/backend/internal/service/openai_gateway_grok.go b/backend/internal/service/openai_gateway_grok.go index 3a737f72d5..adb5f8fbc3 100644 --- a/backend/internal/service/openai_gateway_grok.go +++ b/backend/internal/service/openai_gateway_grok.go @@ -632,7 +632,8 @@ func normalizeGrokReasoningEffortValue(raw string) (string, bool) { func grokSupportsReasoningEffort(model string) bool { model = strings.ToLower(xai.StripGrokProviderPrefix(strings.TrimSpace(model))) switch model { - case xai.DefaultTextModel, "grok-4.5-latest", "grok-4.3", "grok-4.3-latest", + case xai.DefaultTextModel, "grok-4.5-latest", "grok-4.6", "grok-4.6-latest", + "grok-4.3", "grok-4.3-latest", "grok-3-mini", "grok-3-mini-fast", "grok-4.20-0309-reasoning", "grok-4.20-reasoning", "grok-4.20-multi-agent-0309": return true @@ -1093,7 +1094,7 @@ func (s *OpenAIGatewayService) describeGrokComposerImage( Kind: kind, Message: upstreamMsg, }) - s.handleGrokAccountUpstreamError(ctx, account, resp.StatusCode, resp.Header, respBody) + s.handleGrokAccountUpstreamError(withGrokTeamRateLimitModel(ctx, grokComposerImageBridgeVisionModel), account, resp.StatusCode, resp.Header, respBody) if s.shouldFailoverGrokUpstreamError(resp.StatusCode, respBody) { return "", OpenAIUsage{}, &UpstreamFailoverError{ StatusCode: resp.StatusCode, @@ -1105,7 +1106,7 @@ func (s *OpenAIGatewayService) describeGrokComposerImage( return "", OpenAIUsage{}, fmt.Errorf("grok composer image bridge upstream error: %s", upstreamMsg) } - s.updateGrokUsageFromResponse(ctx, account, resp.Header, resp.StatusCode) + s.updateGrokUsageFromResponse(withGrokTeamRateLimitModel(ctx, grokComposerImageBridgeVisionModel), account, resp.Header, resp.StatusCode) respBody, err := ReadUpstreamResponseBody(resp.Body, s.cfg, c, nil) if err != nil { return "", OpenAIUsage{}, fmt.Errorf("read grok composer image bridge response: %w", err) @@ -1305,6 +1306,13 @@ func (s *OpenAIGatewayService) updateGrokUsageSnapshot(ctx context.Context, acco stateCtx, cancel = openAIAccountStateContext(ctx) defer cancel() } + // Account pointers on the request path are per-request copies (Redis/DB decode), + // not a shared in-process cache. Mutating Extra here matches token refresh / + // rate-limit writers; do not reuse the same *Account across goroutines. + if account.Extra == nil { + account.Extra = map[string]any{} + } + account.Extra[grokQuotaSnapshotExtraKey] = snapshot if s.accountRepo != nil { _ = s.accountRepo.UpdateExtra(stateCtx, accountID, updates) } @@ -1322,6 +1330,7 @@ func (s *OpenAIGatewayService) updateGrokUsageSnapshot(ctx context.Context, acco func (s *OpenAIGatewayService) updateGrokUsageFromResponse(ctx context.Context, account *Account, headers http.Header, statusCode int) { snapshot := parseGrokQuotaSnapshot(headers, statusCode, time.Now()) if snapshot != nil { + stampGrokQuotaSnapshotForPlan(account, snapshot, grokRequestedModelFromCtx(ctx)) s.updateGrokUsageSnapshot(ctx, account, snapshot) return } @@ -1626,6 +1635,35 @@ func withGrokTeamRateLimitModel(ctx context.Context, model string) context.Conte return context.WithValue(ctx, grokTeamRateLimitModelContextKey{}, model) } +func grokRequestedModelFromCtx(ctx context.Context) string { + if ctx == nil { + return "" + } + model, _ := ctx.Value(grokTeamRateLimitModelContextKey{}).(string) + return strings.TrimSpace(model) +} + +func isGrokHeavyTransientModel(requestedModel string) bool { + model := strings.ToLower(strings.TrimSpace(xai.ResolveGrokTextResponsesModelID(requestedModel))) + return strings.Contains(model, "multi-agent") +} + +func persistGrokTransientModelCooldown(account *Account, decision GrokUpstreamFailureDecision) bool { + if account == nil { + return false + } + model := strings.TrimSpace(decision.Model) + if model == "" || !isGrokHeavyTransientModel(model) { + return false + } + cooldown := decision.Cooldown + if cooldown <= 0 { + cooldown = 3 * time.Minute + } + markGrokModelTransientBlock(account.ID, model, time.Now().Add(cooldown)) + return true +} + func (s *OpenAIGatewayService) handleGrokAccountUpstreamError(ctx context.Context, account *Account, statusCode int, headers http.Header, responseBody []byte) { if s == nil || account == nil { return @@ -1634,12 +1672,14 @@ func (s *OpenAIGatewayService) handleGrokAccountUpstreamError(ctx context.Contex return } now := time.Now() - s.updateGrokUsageSnapshot(ctx, account, parseGrokQuotaSnapshot(headers, statusCode, now)) + snapshot := parseGrokQuotaSnapshot(headers, statusCode, now) + stampGrokQuotaSnapshotForPlan(account, snapshot, grokRequestedModelFromCtx(ctx)) + s.updateGrokUsageSnapshot(ctx, account, snapshot) // Body-first free-usage / empty / billing / capacity must run before the // status switch so non-429 free-usage bodies still cool the account. // Pool-mode still skips durable mutation unless an explicit temp rule matches. - decision := classifyGrokUpstreamFailure(statusCode, responseBody, "") + decision := classifyGrokUpstreamFailure(statusCode, responseBody, grokRequestedModelFromCtx(ctx)) if decision.ShouldCooldown && decision.Class != GrokFailureNone && decision.Class != GrokFailureRateLimit { if account.IsPoolMode() { // Allow configured temp rules (403) below; skip default body cools. diff --git a/backend/internal/service/openai_gateway_grok_cache.go b/backend/internal/service/openai_gateway_grok_cache.go index 82fa39438e..321f2d0faf 100644 --- a/backend/internal/service/openai_gateway_grok_cache.go +++ b/backend/internal/service/openai_gateway_grok_cache.go @@ -311,6 +311,11 @@ func isKnownGrokFreeAccount(account *Account) bool { if account == nil || !account.IsGrokOAuth() { return false } + // Live access-token JWT wins over stale billing/credential snapshots + // so a downgrade to free is visible as soon as the AT is refreshed. + if jwtTier := xai.SubscriptionTierFromJWT(account.GetCredential("access_token")); jwtTier != "" { + return isGrokFreeSubscriptionTier(jwtTier) + } freeSignal := false paidSignal := false inferredFreeSignal := false @@ -360,8 +365,8 @@ func isKnownGrokFreeAccount(account *Account) bool { } func isGrokFreeSubscriptionTier(tier string) bool { - switch strings.ToLower(strings.TrimSpace(tier)) { - case "free", "grok-free", "grok_free", "free-tier", "free_tier", "basic", "grok-basic", "grok_basic": + switch xai.NormalizeSubscriptionTier(tier) { + case "free", "x_basic": return true default: return false diff --git a/backend/internal/service/openai_gateway_grok_chat_bridge.go b/backend/internal/service/openai_gateway_grok_chat_bridge.go index ed910eab43..e2b8f9a477 100644 --- a/backend/internal/service/openai_gateway_grok_chat_bridge.go +++ b/backend/internal/service/openai_gateway_grok_chat_bridge.go @@ -10,6 +10,7 @@ import ( "time" "github.com/Wei-Shaw/sub2api/internal/pkg/apicompat" + "github.com/Wei-Shaw/sub2api/internal/pkg/xai" "github.com/gin-gonic/gin" "github.com/tidwall/gjson" ) @@ -495,8 +496,17 @@ func grokChatResponsesCacheIntentBody(body []byte) ([]byte, error) { return json.Marshal(root) } +func grokChatResponsesBridgeModel(model string) bool { + switch strings.ToLower(xai.StripGrokProviderPrefix(strings.TrimSpace(model))) { + case "grok-4.5", "grok-4.6", "grok-4.6-latest": + return true + default: + return false + } +} + func grokChatResponsesRuntimeEligible(upstreamModel, cacheIdentity string) bool { - return strings.TrimSpace(upstreamModel) == "grok-4.5" && strings.TrimSpace(cacheIdentity) != "" + return grokChatResponsesBridgeModel(upstreamModel) && strings.TrimSpace(cacheIdentity) != "" } // forwardGrokChatCompletionsViaResponses converts a strictly compatible Chat @@ -527,7 +537,7 @@ func (s *OpenAIGatewayService) forwardGrokChatCompletionsViaResponses( // for non-composer models, so they would be silently dropped. Route them to // Responses even when no prompt-cache identity is available. hasImageInput := openAIJSONValueMayContainImageInput(gjson.GetBytes(body, "messages")) - if !grokChatResponsesRuntimeEligible(upstreamModel, cacheIdentity) && (!hasImageInput || strings.TrimSpace(upstreamModel) != "grok-4.5") { + if !grokChatResponsesRuntimeEligible(upstreamModel, cacheIdentity) && (!hasImageInput || !grokChatResponsesBridgeModel(upstreamModel)) { return s.forwardAsRawChatCompletions(ctx, c, account, body, defaultMappedModel) } @@ -621,7 +631,7 @@ func (s *OpenAIGatewayService) forwardGrokChatCompletionsViaResponses( Kind: kind, Message: upstreamMsg, }) - s.handleGrokAccountUpstreamError(ctx, account, resp.StatusCode, resp.Header, respBody) + s.handleGrokAccountUpstreamError(withGrokTeamRateLimitModel(ctx, upstreamModel), account, resp.StatusCode, resp.Header, respBody) if s.shouldFailoverGrokUpstreamError(resp.StatusCode, respBody) { return nil, &UpstreamFailoverError{ StatusCode: resp.StatusCode, @@ -633,7 +643,7 @@ func (s *OpenAIGatewayService) forwardGrokChatCompletionsViaResponses( return s.handleChatCompletionsErrorResponse(resp, c, account, billingModel) } - s.updateGrokUsageFromResponse(ctx, account, resp.Header, resp.StatusCode) + s.updateGrokUsageFromResponse(withGrokTeamRateLimitModel(ctx, upstreamModel), account, resp.Header, resp.StatusCode) var result *OpenAIForwardResult if clientStream { diff --git a/backend/internal/service/openai_gateway_grok_chat_bridge_test.go b/backend/internal/service/openai_gateway_grok_chat_bridge_test.go index 6c4137d9d8..373907f0d7 100644 --- a/backend/internal/service/openai_gateway_grok_chat_bridge_test.go +++ b/backend/internal/service/openai_gateway_grok_chat_bridge_test.go @@ -218,9 +218,12 @@ func TestGrokChatResponsesBridgeEligibility(t *testing.T) { func TestGrokChatResponsesRuntimeEligibility(t *testing.T) { t.Parallel() require.True(t, grokChatResponsesRuntimeEligible("grok-4.5", "isolated-id")) + require.True(t, grokChatResponsesRuntimeEligible("grok-4.6", "isolated-id")) + require.True(t, grokChatResponsesRuntimeEligible("grok-4.6-latest", "isolated-id")) require.False(t, grokChatResponsesRuntimeEligible("grok-4.3", "isolated-id")) require.False(t, grokChatResponsesRuntimeEligible("grok-4.5-build-free", "isolated-id")) require.False(t, grokChatResponsesRuntimeEligible("grok-4.5", "")) + require.False(t, grokChatResponsesRuntimeEligible("grok-4.6", "")) } func TestForwardGrokChatViaResponsesNonStreamingCachesAndReturnsChat(t *testing.T) { diff --git a/backend/internal/service/openai_gateway_grok_test.go b/backend/internal/service/openai_gateway_grok_test.go index 04edd0a4a6..243736b1c8 100644 --- a/backend/internal/service/openai_gateway_grok_test.go +++ b/backend/internal/service/openai_gateway_grok_test.go @@ -56,6 +56,8 @@ func TestPatchGrokResponsesBodySanitizesComposerReasoningParameters(t *testing.T {name: "composer legacy alias", upstreamModel: "composer-2.5"}, {name: "provider-prefixed composer", upstreamModel: "xai/grok-composer-2.5-fast"}, {name: "grok 4.5", upstreamModel: "grok-4.5", wantReasoning: true}, + {name: "grok 4.6", upstreamModel: "grok-4.6", wantReasoning: true}, + {name: "grok 4.6 latest", upstreamModel: "grok-4.6-latest", wantReasoning: true}, } bodyTemplate := []byte(`{ diff --git a/backend/internal/service/openai_gateway_messages.go b/backend/internal/service/openai_gateway_messages.go index 2aa226dcd4..268aec7b5a 100644 --- a/backend/internal/service/openai_gateway_messages.go +++ b/backend/internal/service/openai_gateway_messages.go @@ -443,7 +443,7 @@ func (s *OpenAIGatewayService) ForwardAsAnthropic( return s.handleAnthropicErrorResponse(resp, c, account, billingModel) } if account.Platform == PlatformGrok && account.Type == AccountTypeOAuth && !account.IsShadow() { - s.updateGrokUsageFromResponse(ctx, account, resp.Header, resp.StatusCode) + s.updateGrokUsageFromResponse(withGrokTeamRateLimitModel(ctx, upstreamModel), account, resp.Header, resp.StatusCode) } if account.Type == AccountTypeOAuth && promptCacheKey != "" { diff --git a/backend/internal/service/openai_gateway_record_usage_test.go b/backend/internal/service/openai_gateway_record_usage_test.go index 4ae09575e3..97ce3d185d 100644 --- a/backend/internal/service/openai_gateway_record_usage_test.go +++ b/backend/internal/service/openai_gateway_record_usage_test.go @@ -256,6 +256,17 @@ func newOpenAIRecordUsageServiceForTest(usageRepo UsageLogRepository, userRepo U return svc } +func openAIRecordUsageAPIKeyWithGroup(svc *OpenAIGatewayService, id int64, groupLongContext bool) *APIKey { + svc.resolver = NewModelPricingResolver(nil, svc.billingService) + return &APIKey{ + ID: id, + Group: &Group{ + ID: 1, + LongContextPricingEnabled: groupLongContext, + }, + } +} + func newOpenAIRecordUsageServiceWithBillingRepoForTest(usageRepo UsageLogRepository, billingRepo UsageBillingRepository, userRepo UserRepository, subRepo UserSubscriptionRepository, rateRepo UserGroupRateRepository) *OpenAIGatewayService { svc := newOpenAIRecordUsageServiceForTest(usageRepo, userRepo, subRepo, rateRepo) svc.usageBillingRepo = billingRepo @@ -1088,7 +1099,7 @@ func TestOpenAIGatewayServiceRecordUsage_Gpt54LongContextBillingDisabledByDefaul Model: "gpt-5.4-2026-03-05", Duration: time.Second, }, - APIKey: &APIKey{ID: 1014}, + APIKey: openAIRecordUsageAPIKeyWithGroup(svc, 1014, true), User: &User{ID: 2014}, Account: &Account{ID: 3014, Platform: PlatformOpenAI}, }) @@ -1122,7 +1133,7 @@ func TestOpenAIGatewayServiceRecordUsage_Gpt54LongContextBillingEnabledPerAccoun Model: "gpt-5.4-2026-03-05", Duration: time.Second, }, - APIKey: &APIKey{ID: 1015}, + APIKey: openAIRecordUsageAPIKeyWithGroup(svc, 1015, true), User: &User{ID: 2015}, Account: &Account{ ID: 3015, @@ -1143,6 +1154,62 @@ func TestOpenAIGatewayServiceRecordUsage_Gpt54LongContextBillingEnabledPerAccoun require.True(t, usageRepo.lastLog.LongContextBillingApplied) } +func TestOpenAIGatewayServiceRecordUsage_GroupAndAccountLongContextMustBothAllow(t *testing.T) { + tokens := OpenAIUsage{InputTokens: 300000, OutputTokens: 2000} + baseInput := 300000 * 2.5e-6 + baseOutput := 2000 * 15e-6 + + t.Run("group on account off", func(t *testing.T) { + usageRepo := &openAIRecordUsageLogRepoStub{inserted: true} + svc := newOpenAIRecordUsageServiceForTest(usageRepo, &openAIRecordUsageUserRepoStub{}, &openAIRecordUsageSubRepoStub{}, nil) + err := svc.RecordUsage(context.Background(), &OpenAIRecordUsageInput{ + Result: &OpenAIForwardResult{RequestID: "resp_and_off", Usage: tokens, Model: "gpt-5.4-2026-03-05", Duration: time.Second}, + APIKey: openAIRecordUsageAPIKeyWithGroup(svc, 1020, true), + User: &User{ID: 2020}, + Account: &Account{ID: 3020, Platform: PlatformOpenAI}, + }) + require.NoError(t, err) + require.False(t, usageRepo.lastLog.LongContextBillingApplied) + require.InDelta(t, baseInput, usageRepo.lastLog.InputCost, 1e-10) + require.InDelta(t, baseOutput, usageRepo.lastLog.OutputCost, 1e-10) + }) + + t.Run("group off account on", func(t *testing.T) { + usageRepo := &openAIRecordUsageLogRepoStub{inserted: true} + svc := newOpenAIRecordUsageServiceForTest(usageRepo, &openAIRecordUsageUserRepoStub{}, &openAIRecordUsageSubRepoStub{}, nil) + err := svc.RecordUsage(context.Background(), &OpenAIRecordUsageInput{ + Result: &OpenAIForwardResult{RequestID: "resp_and_group_off", Usage: tokens, Model: "gpt-5.4-2026-03-05", Duration: time.Second}, + APIKey: openAIRecordUsageAPIKeyWithGroup(svc, 1021, false), + User: &User{ID: 2021}, + Account: &Account{ + ID: 3021, Platform: PlatformOpenAI, + Extra: map[string]any{"openai_long_context_billing_enabled": true}, + }, + }) + require.NoError(t, err) + require.False(t, usageRepo.lastLog.LongContextBillingApplied) + require.InDelta(t, baseInput, usageRepo.lastLog.InputCost, 1e-10) + }) + + t.Run("group on account on", func(t *testing.T) { + usageRepo := &openAIRecordUsageLogRepoStub{inserted: true} + svc := newOpenAIRecordUsageServiceForTest(usageRepo, &openAIRecordUsageUserRepoStub{}, &openAIRecordUsageSubRepoStub{}, nil) + err := svc.RecordUsage(context.Background(), &OpenAIRecordUsageInput{ + Result: &OpenAIForwardResult{RequestID: "resp_and_on", Usage: tokens, Model: "gpt-5.4-2026-03-05", Duration: time.Second}, + APIKey: openAIRecordUsageAPIKeyWithGroup(svc, 1022, true), + User: &User{ID: 2022}, + Account: &Account{ + ID: 3022, Platform: PlatformOpenAI, + Extra: map[string]any{"openai_long_context_billing_enabled": true}, + }, + }) + require.NoError(t, err) + require.True(t, usageRepo.lastLog.LongContextBillingApplied) + require.InDelta(t, baseInput*2, usageRepo.lastLog.InputCost, 1e-10) + require.InDelta(t, baseOutput*1.5, usageRepo.lastLog.OutputCost, 1e-10) + }) +} + func TestOpenAIGatewayServiceRecordUsage_SparkShadowUsesCurrentParentBillingSetting(t *testing.T) { tests := []struct { name string @@ -1177,7 +1244,7 @@ func TestOpenAIGatewayServiceRecordUsage_SparkShadowUsesCurrentParentBillingSett Model: "gpt-5.4-2026-03-05", Duration: time.Second, }, - APIKey: &APIKey{ID: 1016}, + APIKey: openAIRecordUsageAPIKeyWithGroup(svc, 1016, true), User: &User{ID: 2016}, Account: &Account{ ID: 3016, diff --git a/backend/internal/service/openai_gateway_usage.go b/backend/internal/service/openai_gateway_usage.go index c2de3c6472..0b88439437 100644 --- a/backend/internal/service/openai_gateway_usage.go +++ b/backend/internal/service/openai_gateway_usage.go @@ -507,6 +507,15 @@ func (s *OpenAIGatewayService) calculateOpenAIRecordUsageCost( } } if result != nil && result.AudioUsage != nil { + if resolved := s.resolveOpenAIChannelPricing(ctx, billingModel, apiKey); resolved != nil && + (resolved.Mode == BillingModePerRequest) { + gid := apiKey.Group.ID + return s.billingService.CalculateCostUnified(CostInput{ + Ctx: ctx, Model: billingModel, GroupID: &gid, Group: apiKey.Group, + UsageUnits: result.AudioUsage.DurationOrUnits, SizeTier: result.AudioUsage.Mode, + RateMultiplier: webSearchMultiplier, Resolver: s.resolver, Resolved: resolved, + }) + } cfg := groupAudioPriceConfigFromAPIKey(apiKey) return s.billingService.CalculateAudioCost(result.AudioUsage.Mode, result.AudioUsage.DurationOrUnits, cfg, webSearchMultiplier), nil } @@ -628,14 +637,9 @@ func (s *OpenAIGatewayService) calculateOpenAIRecordUsageTokenCost( 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, Group: apiKey.Group, + Tokens: tokens, RequestCount: 1, RateMultiplier: multiplier, + ServiceTier: serviceTier, Resolver: s.resolver, LongContextBillingEnabled: &longContextBillingEnabled, }) } @@ -656,6 +660,19 @@ func (s *OpenAIGatewayService) calculateOpenAIImageCost( multiplier float64, ) *CostBreakdown { sizeTier := NormalizeImageBillingTierOrDefault(result.ImageSize) + resolved := s.resolveOpenAIChannelPricing(ctx, billingModel, apiKey) + if resolved != nil && resolved.Source == PricingSourceGroup && + (resolved.Mode == BillingModePerRequest || resolved.Mode == BillingModeImage) { + gid := apiKey.Group.ID + cost, err := s.billingService.CalculateCostUnified(CostInput{ + Ctx: ctx, Model: billingModel, GroupID: &gid, Group: apiKey.Group, + RequestCount: result.ImageCount, SizeTier: sizeTier, + RateMultiplier: multiplier, Resolver: s.resolver, Resolved: resolved, + }) + if err == nil { + return cost + } + } groupConfig := imagePriceConfigFromAPIKey(apiKey) if apiKeyHasConfiguredImagePrice(apiKey, sizeTier) { return s.billingService.CalculateImageCost(billingModel, sizeTier, result.ImageCount, groupConfig, multiplier) @@ -667,13 +684,14 @@ func (s *OpenAIGatewayService) calculateOpenAIImageCost( return s.billingService.CalculateImageCost(billingModel, sizeTier, result.ImageCount, groupConfig, multiplier) } } - if resolved := s.resolveOpenAIChannelPricing(ctx, billingModel, apiKey); resolved != nil && + if resolved != nil && resolved.Source == PricingSourceChannel && (resolved.Mode == BillingModePerRequest || resolved.Mode == BillingModeImage) { gid := apiKey.Group.ID cost, err := s.billingService.CalculateCostUnified(CostInput{ Ctx: ctx, Model: billingModel, GroupID: &gid, + Group: apiKey.Group, RequestCount: result.ImageCount, SizeTier: sizeTier, RateMultiplier: multiplier, @@ -702,6 +720,18 @@ func (s *OpenAIGatewayService) calculateOpenAIVideoCost( } resolution := NormalizeVideoBillingResolutionOrDefault(result.VideoResolution) durationSeconds := NormalizeVideoBillingDurationSecondsOrDefault(result.VideoDurationSeconds) + resolved := s.resolveOpenAIChannelPricing(ctx, billingModel, apiKey) + if resolved != nil && resolved.Source == PricingSourceGroup && resolved.Mode == BillingModeVideo { + gid := apiKey.Group.ID + cost, err := s.billingService.CalculateCostUnified(CostInput{ + Ctx: ctx, Model: billingModel, GroupID: &gid, Group: apiKey.Group, + UsageUnits: float64(videoCount * durationSeconds), SizeTier: resolution, + RateMultiplier: multiplier, Resolver: s.resolver, Resolved: resolved, + }) + if err == nil { + return cost + } + } groupConfig := videoPriceConfigFromAPIKey(apiKey) if apiKeyHasConfiguredVideoPrice(apiKey, billingModel, resolution) { return s.billingService.CalculateVideoCost(billingModel, resolution, videoCount, durationSeconds, groupConfig, multiplier) @@ -713,15 +743,21 @@ func (s *OpenAIGatewayService) calculateOpenAIVideoCost( return s.billingService.CalculateVideoCost(billingModel, resolution, videoCount, durationSeconds, groupConfig, multiplier) } } - if resolved := s.resolveOpenAIChannelPricing(ctx, billingModel, apiKey); resolved != nil && - (resolved.Mode == BillingModePerRequest || resolved.Mode == BillingModeImage) { + if resolved != nil && resolved.Source == PricingSourceChannel && + (resolved.Mode == BillingModePerRequest || resolved.Mode == BillingModeImage || resolved.Mode == BillingModeVideo) { // 渠道 per_request/image 定价保持"按请求次数"口径(价格由管理员按次配置),不乘视频时长。 gid := apiKey.Group.ID + units := float64(videoCount) + if resolved.Mode == BillingModeVideo { + units = float64(videoCount * durationSeconds) + } cost, err := s.billingService.CalculateCostUnified(CostInput{ Ctx: ctx, Model: billingModel, GroupID: &gid, + Group: apiKey.Group, RequestCount: videoCount, + UsageUnits: units, SizeTier: resolution, RateMultiplier: multiplier, Resolver: s.resolver, @@ -778,6 +814,9 @@ func groupMediaPricingLooksIncomplete(group *Group) bool { if len(group.VideoModelPrices) > 0 { return false } + if len(group.ModelPricing) > 0 || group.LongContextPricingEnabled { + return false + } if group.SearchPricePer1k != nil || group.AudioRealtimePricePerMin != nil || group.AudioTTSPricePerMillionChars != nil || @@ -794,8 +833,8 @@ func (s *OpenAIGatewayService) resolveOpenAIChannelPricing(ctx context.Context, return nil } gid := apiKey.Group.ID - resolved := s.resolver.Resolve(ctx, PricingInput{Model: billingModel, GroupID: &gid}) - if resolved.Source == PricingSourceChannel { + resolved := s.resolver.Resolve(ctx, PricingInput{Model: billingModel, GroupID: &gid, Group: apiKey.Group}) + if resolved.Source == PricingSourceGroup || resolved.Source == PricingSourceChannel { return resolved } return nil diff --git a/backend/internal/service/openai_ws_http_bridge.go b/backend/internal/service/openai_ws_http_bridge.go index ac6a24e78a..87abf1ef04 100644 --- a/backend/internal/service/openai_ws_http_bridge.go +++ b/backend/internal/service/openai_ws_http_bridge.go @@ -250,7 +250,7 @@ func (s *OpenAIGatewayService) proxyOpenAIWSHTTPBridgeTurn( shouldFailover := s.shouldFailoverOpenAIUpstreamResponse(resp.StatusCode, upstreamMsg, respBody) if account.Platform == PlatformGrok { shouldFailover = s.shouldFailoverGrokUpstreamError(resp.StatusCode, respBody) - s.handleGrokAccountUpstreamError(ctx, account, resp.StatusCode, resp.Header, respBody) + s.handleGrokAccountUpstreamError(withGrokTeamRateLimitModel(ctx, resolveGrokWSUpstreamModel(account, body, originalModel)), account, resp.StatusCode, resp.Header, respBody) if turn == 1 && shouldFailover { return nil, newOpenAIUpstreamFailoverError(resp.StatusCode, resp.Header, respBody, upstreamMsg, false) } @@ -265,7 +265,7 @@ func (s *OpenAIGatewayService) proxyOpenAIWSHTTPBridgeTurn( return nil, fmt.Errorf("upstream http bridge error: status=%d message=%s", resp.StatusCode, upstreamMsg) } if account.Platform == PlatformGrok { - s.updateGrokUsageFromResponse(ctx, account, resp.Header, resp.StatusCode) + s.updateGrokUsageFromResponse(withGrokTeamRateLimitModel(ctx, resolveGrokWSUpstreamModel(account, body, originalModel)), account, resp.Header, resp.StatusCode) } responseID := "" diff --git a/backend/migrations/221_group_model_pricing.sql b/backend/migrations/221_group_model_pricing.sql new file mode 100644 index 0000000000..52cf002891 --- /dev/null +++ b/backend/migrations/221_group_model_pricing.sql @@ -0,0 +1,12 @@ +ALTER TABLE groups + ADD COLUMN IF NOT EXISTS long_context_pricing_enabled BOOLEAN NOT NULL DEFAULT TRUE, + ADD COLUMN IF NOT EXISTS model_pricing JSONB; + +UPDATE groups +SET long_context_pricing_enabled = TRUE +WHERE long_context_pricing_enabled IS DISTINCT FROM TRUE; + +COMMENT ON COLUMN groups.long_context_pricing_enabled IS + 'Whether token pricing selects official/preset long-context tiers; default true preserves existing long-context billing'; +COMMENT ON COLUMN groups.model_pricing IS + 'Per-model group pricing overrides channel and built-in model pricing'; diff --git a/frontend/src/components/account/AccountUsageCell.vue b/frontend/src/components/account/AccountUsageCell.vue index e7e0a298fc..485b4b8701 100644 --- a/frontend/src/components/account/AccountUsageCell.vue +++ b/frontend/src/components/account/AccountUsageCell.vue @@ -1148,19 +1148,19 @@ const grokPlanLabelIsPaid = (value: string) => { const grokIsFree = computed(() => { if (props.account.platform !== 'grok' || props.account.type !== 'oauth') return false const billing = grokBilling.value + const plan = (billing?.plan || '').trim().toLowerCase() + const tier = (usageInfo.value?.subscription_tier || '').trim().toLowerCase() + const entitlement = (usageInfo.value?.grok_entitlement_status || '').toLowerCase() + if (grokPlanLabelIsFree(tier)) return true + if (grokPlanLabelIsPaid(tier)) return false 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 (grokPlanLabelIsPaid(plan)) return false if ( grokPlanLabelIsFree(plan) || - grokPlanLabelIsFree(tier) || grokPlanLabelIsFree(entitlement) ) return true return billing != null diff --git a/frontend/src/components/account/__tests__/AccountUsageCell.spec.ts b/frontend/src/components/account/__tests__/AccountUsageCell.spec.ts index 46f3908254..ff0e94d664 100644 --- a/frontend/src/components/account/__tests__/AccountUsageCell.spec.ts +++ b/frontend/src/components/account/__tests__/AccountUsageCell.spec.ts @@ -903,6 +903,81 @@ describe('AccountUsageCell', () => { expect(wrapper.text()).not.toContain('250.0K') }) + it('Grok JWT free tier shows 24h bar even when leftover Heavy billing metrics remain', async () => { + getUsage.mockResolvedValue({ + grok_free_token_limit: 500_000, + subscription_tier: 'free', + grok_billing: { + plan: 'SuperGrok Heavy', + monthly_limit_cents: 150_000, + usage_percent: 10, + used_percent: 5 + }, + grok_local_usage_24h: { + requests: 2, + tokens: 250_000, + cost: 0, + standard_cost: 0 + } + }) + + const wrapper = mount(AccountUsageCell, { + props: { + account: makeAccount({ id: 4404, platform: 'grok', type: 'oauth', extra: {} }) + }, + global: { + stubs: { + UsageProgressBar: { + props: ['label', 'utilization'], + template: '
{{ label }}|{{ utilization }}
' + }, + AccountQuotaInfo: true + } + } + }) + + await flushPromises() + expect(wrapper.text()).toContain('24h|') + expect(wrapper.text()).not.toContain('7d|') + expect(wrapper.text()).not.toContain('30d|') + }) + + it('Grok SuperGrok Lite stays on paid 7d bar, not free 24h', async () => { + getUsage.mockResolvedValue({ + subscription_tier: 'supergrok_lite', + grok_billing: { + period_type: 'weekly', + plan: 'SuperGrok', + usage_percent: 20 + }, + grok_local_usage_24h: { + requests: 1, + tokens: 100, + cost: 0, + standard_cost: 0 + } + }) + + const wrapper = mount(AccountUsageCell, { + props: { + account: makeAccount({ id: 4405, platform: 'grok', type: 'oauth', extra: {} }) + }, + global: { + stubs: { + UsageProgressBar: { + props: ['label', 'utilization'], + template: '
{{ label }}|{{ utilization }}
' + }, + AccountQuotaInfo: true + } + } + }) + + await flushPromises() + expect(wrapper.text()).toContain('7d|') + expect(wrapper.text()).not.toContain('24h|') + }) + it('Grok credential Free tier keeps the 1M fallback when billing is unavailable', async () => { getUsage.mockResolvedValue({ grok_free_token_limit: 1_000_000, diff --git a/frontend/src/components/admin/channel/PricingEntryCard.vue b/frontend/src/components/admin/channel/PricingEntryCard.vue index 086a9cd307..d09e19feb2 100644 --- a/frontend/src/components/admin/channel/PricingEntryCard.vue +++ b/frontend/src/components/admin/channel/PricingEntryCard.vue @@ -134,8 +134,8 @@ - -
+ +
- -
+ +
@@ -209,9 +209,9 @@
-
@@ -245,10 +245,13 @@ import channelsAPI from '@/api/admin/channels' const { t } = useI18n() -const props = defineProps<{ +const props = withDefaults(defineProps<{ entry: PricingFormEntry platform?: string -}>() + hideTokenIntervals?: boolean +}>(), { + hideTokenIntervals: false, +}) const emit = defineEmits<{ update: [entry: PricingFormEntry] @@ -261,7 +264,8 @@ const collapsed = ref(props.entry.models.length > 0) const billingModeOptions = computed(() => [ { value: 'token', label: t('admin.channels.billingMode.token') }, { value: 'per_request', label: t('admin.channels.billingMode.perRequest') }, - { value: 'image', label: t('admin.channels.billingMode.image') } + { value: 'image', label: t('admin.channels.billingMode.image') }, + { value: 'video', label: t('admin.channels.billingMode.video') } ]) const billingModeLabel = computed(() => { @@ -284,9 +288,11 @@ function addInterval() { emit('update', { ...props.entry, intervals }) } -function addImageTier() { +function addMediaTier() { const intervals = [...(props.entry.intervals || [])] - const labels = ['1K', '2K', '4K', 'HD'] + const labels = props.entry.billing_mode === 'video' + ? ['480p', '720p', '1080p'] + : ['1K', '2K', '4K', 'HD'] intervals.push({ min_tokens: 0, max_tokens: null, tier_label: labels[intervals.length] || '', input_price: null, output_price: null, cache_write_price: null, diff --git a/frontend/src/components/common/PlatformTypeBadge.vue b/frontend/src/components/common/PlatformTypeBadge.vue index d5af11f5d8..ef86879ae0 100644 --- a/frontend/src/components/common/PlatformTypeBadge.vue +++ b/frontend/src/components/common/PlatformTypeBadge.vue @@ -137,10 +137,16 @@ const planLabel = computed(() => { return props.platform === 'grok' ? 'Grok Free' : 'Free' case 'supergrok': return 'SuperGrok' + case 'supergroklite': + return 'SuperGrok Lite' + case 'supergrokplus': + return 'SuperGrok Plus' case 'supergrokheavy': return 'SuperGrok Heavy' case 'heavy': return 'Heavy' + case 'xbasic': + return 'X Basic' case 'abnormal': return t('admin.accounts.subscriptionAbnormal') default: @@ -150,7 +156,9 @@ const planLabel = computed(() => { const isGrokFreePlan = computed(() => props.platform === 'grok' && - (normalizedPlanType.value === 'free' || normalizedPlanType.value === 'basic') + (normalizedPlanType.value === 'free' || + normalizedPlanType.value === 'basic' || + normalizedPlanType.value === 'xbasic') ) const planIconName = computed<'bolt' | null>(() => { @@ -204,7 +212,11 @@ const planBadgeClass = computed(() => { return 'bg-red-100 text-red-600 dark:bg-red-900/30 dark:text-red-400' } // Free stays muted gray; paid Grok tiers get distinct colors. - if (normalizedPlanType.value === 'free' || normalizedPlanType.value === 'basic') { + if ( + normalizedPlanType.value === 'free' || + normalizedPlanType.value === 'basic' || + normalizedPlanType.value === 'xbasic' + ) { return 'bg-gray-100 text-gray-600 dark:bg-gray-700 dark:text-gray-300' } if (props.platform === 'grok' && normalizedPlanType.value) { @@ -235,7 +247,11 @@ const planBadgeClass = computed(() => { // Subscription expiration label (non-free only) const expiresLabel = computed(() => { if (!props.subscriptionExpiresAt || !props.planType) return '' - if (normalizedPlanType.value === 'free' || normalizedPlanType.value === 'basic') return '' + if ( + normalizedPlanType.value === 'free' || + normalizedPlanType.value === 'basic' || + normalizedPlanType.value === 'xbasic' + ) return '' try { const d = new Date(props.subscriptionExpiresAt) if (isNaN(d.getTime())) return '' diff --git a/frontend/src/components/common/__tests__/PlatformTypeBadge.grok.spec.ts b/frontend/src/components/common/__tests__/PlatformTypeBadge.grok.spec.ts index fc665316a8..e65c7f81de 100644 --- a/frontend/src/components/common/__tests__/PlatformTypeBadge.grok.spec.ts +++ b/frontend/src/components/common/__tests__/PlatformTypeBadge.grok.spec.ts @@ -76,6 +76,12 @@ describe('PlatformTypeBadge Grok plans', () => { expect(heavy.text()).toContain('Heavy') expect(heavy.html()).toContain('bg-purple-100') expect(heavy.find('[data-testid="grok-plan-icon"]').exists()).toBe(true) + + const lite = mount(PlatformTypeBadge, { + props: { platform: 'grok', type: 'oauth', planType: 'supergrok_lite' }, + }) + expect(lite.text()).toContain('SuperGrok Lite') + expect(lite.html()).toContain('bg-cyan-100') }) it('uses a dedicated 12px currentColor Grok mark with a Free sparkle', () => { diff --git a/frontend/src/composables/__tests__/useModelWhitelist.spec.ts b/frontend/src/composables/__tests__/useModelWhitelist.spec.ts index 714164324d..9fd93e6573 100644 --- a/frontend/src/composables/__tests__/useModelWhitelist.spec.ts +++ b/frontend/src/composables/__tests__/useModelWhitelist.spec.ts @@ -46,6 +46,8 @@ describe('useModelWhitelist', () => { it('xAI 模型列表包含 Grok 4.5 官方模型和别名', () => { const models = getModelsByPlatform('grok') + expect(models).toContain('grok-4.6') + expect(models).toContain('grok-4.6-latest') expect(models).toContain('grok-4.5') expect(models).toContain('grok-4.5-latest') expect(models).toContain('grok-build-latest') diff --git a/frontend/src/composables/useModelWhitelist.ts b/frontend/src/composables/useModelWhitelist.ts index 3d4c42a29e..e2c5924592 100644 --- a/frontend/src/composables/useModelWhitelist.ts +++ b/frontend/src/composables/useModelWhitelist.ts @@ -136,6 +136,7 @@ const metaModels = [ // xAI Grok const xaiModels = [ + 'grok-4.6', 'grok-4.5', 'grok-4.3', 'grok-build-0.1', @@ -147,6 +148,7 @@ const xaiModels = [ 'grok-4.20-multi-agent-latest', 'grok-4.3-latest', 'grok-latest', + 'grok-4.6-latest', 'grok-4.5-latest', 'grok-build-latest', 'composer-2.5', @@ -303,6 +305,7 @@ const geminiPresetMappings = [ ] const grokPresetMappings = [ + { label: 'Grok 4.6', from: 'grok-4.6', to: 'grok-4.6', color: 'bg-slate-100 text-slate-700 hover:bg-slate-200 dark:bg-slate-800/50 dark:text-slate-300' }, { label: 'Grok 4.5', from: 'grok-4.5', to: 'grok-4.5', color: 'bg-slate-100 text-slate-700 hover:bg-slate-200 dark:bg-slate-800/50 dark:text-slate-300' }, { label: 'Grok 4.3', from: 'grok-4.3', to: 'grok-4.3', color: 'bg-slate-100 text-slate-700 hover:bg-slate-200 dark:bg-slate-800/50 dark:text-slate-300' }, { label: 'Grok Latest', from: 'grok-latest', to: 'grok-4.5', color: 'bg-emerald-100 text-emerald-700 hover:bg-emerald-200 dark:bg-emerald-900/30 dark:text-emerald-400' }, diff --git a/frontend/src/constants/channel.ts b/frontend/src/constants/channel.ts index 6b54b47df0..1e1b9205de 100644 --- a/frontend/src/constants/channel.ts +++ b/frontend/src/constants/channel.ts @@ -7,10 +7,12 @@ export type ChannelStatus = typeof CHANNEL_STATUS_ACTIVE | typeof CHANNEL_STATUS export const BILLING_MODE_TOKEN = 'token' as const export const BILLING_MODE_PER_REQUEST = 'per_request' as const export const BILLING_MODE_IMAGE = 'image' as const +export const BILLING_MODE_VIDEO = 'video' as const export type BillingMode = | typeof BILLING_MODE_TOKEN | typeof BILLING_MODE_PER_REQUEST | typeof BILLING_MODE_IMAGE + | typeof BILLING_MODE_VIDEO /** Billing-model-source values (must match service.BillingModelSource* constants in Go). */ export const BILLING_MODEL_SOURCE_REQUESTED = 'requested' as const diff --git a/frontend/src/i18n/locales/en/admin/channels.ts b/frontend/src/i18n/locales/en/admin/channels.ts index 6b26e21da9..d68ea5060d 100644 --- a/frontend/src/i18n/locales/en/admin/channels.ts +++ b/frontend/src/i18n/locales/en/admin/channels.ts @@ -93,7 +93,8 @@ export default { billingMode: { token: 'Token', perRequest: 'Per Request', - image: 'Image (Per Request)' + image: 'Image (Per Request)', + video: 'Video (Per Second)' }, form: { name: 'Name', @@ -127,6 +128,7 @@ export default { addInterval: 'Add Interval', requestTiers: 'Request Tiers', imageTiers: 'Image Tiers (Per Request)', + videoTiers: 'Video Resolution Tiers (Per Second)', addTier: 'Add Tier', noTiersYet: 'No tiers yet. Click add to configure per-request pricing.', noPricingRules: 'No pricing rules yet. Click "Add" to create one.', @@ -152,6 +154,7 @@ export default { restrictModelsHint: 'When enabled, only models in the pricing list are allowed. Others will be rejected.', defaultPerRequestPrice: 'Default per-request price (fallback when no tier matches)', defaultImagePrice: 'Default image price (fallback when no tier matches)', + defaultVideoPrice: 'Default video price per second (fallback when no tier matches)', platformConfig: 'Platform Configuration', webSearchEmulation: 'Web Search Emulation', webSearchEmulationHint: '⚠️ When enabled, all accounts in this channel\'s Anthropic groups will intercept web_search requests. Use with caution.', diff --git a/frontend/src/i18n/locales/en/admin/overview.ts b/frontend/src/i18n/locales/en/admin/overview.ts index f626b6edc6..1fcc7bf0e8 100644 --- a/frontend/src/i18n/locales/en/admin/overview.ts +++ b/frontend/src/i18n/locales/en/admin/overview.ts @@ -1005,6 +1005,13 @@ export default { searchPricePer1k: 'Search price per 1k calls (USD)', pricePlaceholder: 'optional' }, + modelPricing: { + title: 'Per-model group pricing', + description: 'Overrides channel and built-in prices for matching models. Long-context tiers come from official presets — do not enter custom intervals. Use per-request tiers such as realtime, tts, and stt for audio.', + longContext: 'Enable long-context tier pricing', + longContextHint: 'When checked, official/preset long-context tiers apply. When unchecked, token models stay on the first-tier base rate.', + add: 'Add model price' + }, voicePricing: { title: 'Grok Voice Pricing', description: 'Optional per-group prices for Voice realtime / TTS / STT (USD). Leave empty to leave unpriced.', diff --git a/frontend/src/i18n/locales/zh/admin/channels.ts b/frontend/src/i18n/locales/zh/admin/channels.ts index 3da00a3120..808e78fa88 100644 --- a/frontend/src/i18n/locales/zh/admin/channels.ts +++ b/frontend/src/i18n/locales/zh/admin/channels.ts @@ -93,7 +93,8 @@ export default { billingMode: { token: 'Token', perRequest: '按次', - image: '图片(按次)' + image: '图片(按次)', + video: '视频(按秒)' }, form: { name: '名称', @@ -127,6 +128,7 @@ export default { addInterval: '添加区间', requestTiers: '按次计费层级', imageTiers: '图片计费层级(按次)', + videoTiers: '视频分辨率层级(按秒)', addTier: '添加层级', noTiersYet: '暂无层级,点击添加配置按次计费价格', noPricingRules: '暂无定价规则,点击"添加"创建', @@ -152,6 +154,7 @@ export default { restrictModelsHint: '开启后,仅允许模型定价列表中的模型。不在列表中的模型请求将被拒绝。', defaultPerRequestPrice: '默认单次价格(未命中层级时使用)', defaultImagePrice: '默认图片价格(未命中层级时使用)', + defaultVideoPrice: '默认视频每秒价格(未命中层级时使用)', platformConfig: '平台配置', webSearchEmulation: 'Web Search 模拟', webSearchEmulationHint: '⚠️ 开启后该渠道下所有 Anthropic 分组的账号将自动拦截 web_search 请求,请谨慎操作', diff --git a/frontend/src/i18n/locales/zh/admin/overview.ts b/frontend/src/i18n/locales/zh/admin/overview.ts index c4ad9a542b..f4b02a9cab 100644 --- a/frontend/src/i18n/locales/zh/admin/overview.ts +++ b/frontend/src/i18n/locales/zh/admin/overview.ts @@ -1002,6 +1002,13 @@ export default { searchPricePer1k: '搜索每千次价格(USD)', pricePlaceholder: '可选' }, + modelPricing: { + title: '分组逐模型定价', + description: '匹配模型后覆盖渠道和内置价格。长上下文阶梯沿用官方/预设价卡,无需再手填区间。音频可用按次层级配置 realtime、tts、stt。', + longContext: '启用长上下文阶梯定价', + longContextHint: '勾选后按官方/预设阶梯计费;关闭则始终按第一档基础价。', + add: '添加模型价格' + }, voicePricing: { title: 'Grok Voice 定价', description: '分组级 Voice realtime / TTS / STT 单价(USD)。留空表示未配置。', diff --git a/frontend/src/types/index.ts b/frontend/src/types/index.ts index 24498cf7ca..19ef541bc6 100644 --- a/frontend/src/types/index.ts +++ b/frontend/src/types/index.ts @@ -558,6 +558,7 @@ export interface Group { daily_limit_usd: number | null weekly_limit_usd: number | null monthly_limit_usd: number | null + long_context_pricing_enabled: boolean // 图片生成计费配置 allow_image_generation: boolean allow_batch_image_generation: boolean @@ -604,6 +605,7 @@ export interface Group { } export interface AdminGroup extends Group { + model_pricing: import('@/api/admin/channels').ChannelModelPricing[] // 分组利润控制(openai/anthropic/gemini/grok/antigravity 分组可启用;margin/buffer 为小数存储)。 // 仅管理员可见:与 rate_multiplier 相乘即可反推上游成本上限,不得下放到 Group。 profit_control_enabled: boolean @@ -766,6 +768,8 @@ export interface CreateGroupRequest { daily_limit_usd?: number | null weekly_limit_usd?: number | null monthly_limit_usd?: number | null + long_context_pricing_enabled?: boolean + model_pricing?: import('@/api/admin/channels').ChannelModelPricing[] allow_image_generation?: boolean allow_batch_image_generation?: boolean image_rate_independent?: boolean @@ -826,6 +830,8 @@ export interface UpdateGroupRequest { daily_limit_usd?: number | null weekly_limit_usd?: number | null monthly_limit_usd?: number | null + long_context_pricing_enabled?: boolean + model_pricing?: import('@/api/admin/channels').ChannelModelPricing[] allow_image_generation?: boolean allow_batch_image_generation?: boolean image_rate_independent?: boolean diff --git a/frontend/src/utils/__tests__/accountUsageRefresh.spec.ts b/frontend/src/utils/__tests__/accountUsageRefresh.spec.ts index aef73b0fa4..f9559697e1 100644 --- a/frontend/src/utils/__tests__/accountUsageRefresh.spec.ts +++ b/frontend/src/utils/__tests__/accountUsageRefresh.spec.ts @@ -1,5 +1,5 @@ import { describe, expect, it } from 'vitest' -import { buildOpenAIUsageRefreshKey } from '../accountUsageRefresh' +import { buildGrokUsageRefreshKey, buildOpenAIUsageRefreshKey } from '../accountUsageRefresh' describe('buildOpenAIUsageRefreshKey', () => { it('会在 codex 快照变化时生成不同 key', () => { @@ -61,3 +61,101 @@ describe('buildOpenAIUsageRefreshKey', () => { } as any)).toBe('') }) }) + +describe('buildGrokUsageRefreshKey', () => { + it('changes when a canonical Grok billing or usage snapshot changes', () => { + const base = { + platform: 'grok', + extra: { + grok_billing_snapshot: { plan: 'Free', usage_percent: 0 }, + grok_usage_snapshot: { subscription_tier: 'Free', status_code: 200 } + } + } as any + + expect(buildGrokUsageRefreshKey(base)).not.toBe(buildGrokUsageRefreshKey({ + ...base, + extra: { + ...base.extra, + grok_billing_snapshot: { plan: 'SuperGrok', usage_percent: 0 } + } + })) + expect(buildGrokUsageRefreshKey(base)).not.toBe(buildGrokUsageRefreshKey({ + ...base, + extra: { + ...base.extra, + grok_usage_snapshot: { subscription_tier: 'SuperGrok', status_code: 200 } + } + })) + }) + + it('ignores object key order and a legacy alias shadowed by canonical usage', () => { + const first = { + platform: 'grok', + extra: { + grok_billing_snapshot: { + plan: 'SuperGrok', + limits: { monthly: 100, weekly: 25 } + }, + grok_usage_snapshot: { status_code: 200, subscription_tier: 'SuperGrok' }, + grok_quota_snapshot: { subscription_tier: 'Free' } + } + } as any + const reordered = { + platform: 'grok', + extra: { + grok_quota_snapshot: { subscription_tier: 'SuperGrok Heavy' }, + grok_usage_snapshot: { subscription_tier: 'SuperGrok', status_code: 200 }, + grok_billing_snapshot: { + limits: { weekly: 25, monthly: 100 }, + plan: 'SuperGrok' + } + } + } as any + + expect(buildGrokUsageRefreshKey(first)).toBe(buildGrokUsageRefreshKey(reordered)) + }) + + it('uses the legacy quota alias only when the canonical snapshot is absent', () => { + const base = { + platform: 'grok', + extra: { grok_quota_snapshot: { subscription_tier: 'Free' } } + } as any + const next = { + platform: 'grok', + extra: { grok_quota_snapshot: { subscription_tier: 'SuperGrok' } } + } as any + + expect(buildGrokUsageRefreshKey(base)).not.toBe(buildGrokUsageRefreshKey(next)) + }) + + it('tracks the legacy tier when the canonical snapshot has no usable tier', () => { + for (const canonicalSnapshot of [ + { status_code: 200 }, + { status_code: 200, subscription_tier: ' ' }, + ]) { + const base = { + platform: 'grok', + extra: { + grok_usage_snapshot: canonicalSnapshot, + grok_quota_snapshot: { subscription_tier: 'Free' }, + }, + } as any + const next = { + ...base, + extra: { + ...base.extra, + grok_quota_snapshot: { subscription_tier: 'SuperGrok' }, + }, + } + + expect(buildGrokUsageRefreshKey(base)).not.toBe(buildGrokUsageRefreshKey(next)) + } + }) + + it('returns an empty key for non-Grok accounts', () => { + expect(buildGrokUsageRefreshKey({ + platform: 'openai', + extra: { grok_usage_snapshot: { subscription_tier: 'SuperGrok' } } + } as any)).toBe('') + }) +}) diff --git a/frontend/src/utils/accountUsageRefresh.ts b/frontend/src/utils/accountUsageRefresh.ts index 3406c7a504..121dae253a 100644 --- a/frontend/src/utils/accountUsageRefresh.ts +++ b/frontend/src/utils/accountUsageRefresh.ts @@ -5,6 +5,30 @@ const normalizeUsageRefreshValue = (value: unknown): string => { return String(value) } +const normalizeSnapshotRefreshValue = (value: unknown): unknown => { + if (Array.isArray(value)) { + return value.map(normalizeSnapshotRefreshValue) + } + if (value && typeof value === 'object') { + return Object.fromEntries( + Object.entries(value as Record) + .filter(([, entry]) => entry !== undefined) + .sort(([left], [right]) => left.localeCompare(right)) + .map(([key, entry]) => [key, normalizeSnapshotRefreshValue(entry)]) + ) + } + return value +} + +const serializeSnapshotRefreshValue = (value: unknown): string => { + if (value == null) return '' + return JSON.stringify(normalizeSnapshotRefreshValue(value)) ?? '' +} + +const isNonBlankString = (value: unknown): value is string => ( + typeof value === 'string' && value.trim().length > 0 +) + export const buildOpenAIUsageRefreshKey = (account: Pick): string => { if (account.platform !== 'openai' || account.type !== 'oauth') { return '' @@ -27,3 +51,21 @@ export const buildOpenAIUsageRefreshKey = (account: Pick): string => { + if (account.platform !== 'grok') { + return '' + } + + const extra = account.extra ?? {} + const usageSnapshot = extra.grok_usage_snapshot + const canonicalTier = (usageSnapshot as Record | null | undefined)?.subscription_tier + const legacyQuotaFallback = isNonBlankString(canonicalTier) + ? undefined + : extra.grok_quota_snapshot + return [ + serializeSnapshotRefreshValue(extra.grok_billing_snapshot), + serializeSnapshotRefreshValue(usageSnapshot), + serializeSnapshotRefreshValue(legacyQuotaFallback) + ].join('|') +} diff --git a/frontend/src/views/admin/AccountsView.vue b/frontend/src/views/admin/AccountsView.vue index 3dfe2c2797..3bfc9357ec 100644 --- a/frontend/src/views/admin/AccountsView.vue +++ b/frontend/src/views/admin/AccountsView.vue @@ -525,7 +525,7 @@ import Icon from '@/components/icons/Icon.vue' import ErrorPassthroughRulesModal from '@/components/admin/ErrorPassthroughRulesModal.vue' import TLSFingerprintProfilesModal from '@/components/admin/TLSFingerprintProfilesModal.vue' import { fetchAllAccountIds } from '@/utils/accountSelection' -import { buildOpenAIUsageRefreshKey } from '@/utils/accountUsageRefresh' +import { buildGrokUsageRefreshKey, buildOpenAIUsageRefreshKey } from '@/utils/accountUsageRefresh' import { formatDateTime, formatRelativeTime } from '@/utils/format' import { proxyExpiryBadgeClass, proxyExpiryLabelKey } from '@/utils/proxyExpiry' import { extractApiErrorMessage } from '@/utils/apiError' @@ -1300,7 +1300,8 @@ const shouldReplaceAutoRefreshRow = (current: Account, next: Account) => { current.rate_limit_reset_at !== next.rate_limit_reset_at || current.overload_until !== next.overload_until || current.temp_unschedulable_until !== next.temp_unschedulable_until || - buildOpenAIUsageRefreshKey(current) !== buildOpenAIUsageRefreshKey(next) + buildOpenAIUsageRefreshKey(current) !== buildOpenAIUsageRefreshKey(next) || + buildGrokUsageRefreshKey(current) !== buildGrokUsageRefreshKey(next) ) } @@ -1481,25 +1482,109 @@ const { pause: pauseAutoRefresh, resume: resumeAutoRefresh } = useIntervalFn( { immediate: false } ) -// Fresh billing/quota snapshots are authoritative. Imported credential tiers -// can be stale, so they remain fallbacks together with legacy plan_type fields. +const GROK_QUOTA_SIGNAL_MAX_AGE_MS = 24 * 60 * 60 * 1000 +const GROK_QUOTA_SIGNAL_MAX_FUTURE_SKEW_MS = 5 * 60 * 1000 + +function firstNonBlankString(...values: unknown[]): string | undefined { + return values.find((value): value is string => ( + typeof value === 'string' && value.trim().length > 0 + )) +} + +function normalizeGrokPlanKey(value: unknown): string { + if (typeof value !== 'string') return '' + return value + .trim() + .toLowerCase() + .replace(/[\s_-]+/g, '') +} + +function grokPersistedQuotaSnapshot(extra: Record): Record | undefined { + const usage = extra.grok_usage_snapshot + if (usage && typeof usage === 'object' && !Array.isArray(usage)) { + return usage as Record + } + const legacy = extra.grok_quota_snapshot + if (legacy && typeof legacy === 'object' && !Array.isArray(legacy)) { + return legacy as Record + } + return undefined +} + +function isGrokQuotaTimestampFresh(raw: unknown): boolean { + const value = String(raw || '').trim() + if (!value) return false + const observedAt = Date.parse(value) + if (!Number.isFinite(observedAt)) return false + const age = Date.now() - observedAt + return age <= GROK_QUOTA_SIGNAL_MAX_AGE_MS && age >= -GROK_QUOTA_SIGNAL_MAX_FUTURE_SKEW_MS +} + +function isGrok45ResponsesQuotaModel(model: unknown): boolean { + const value = String(model || '') + .trim() + .toLowerCase() + .replace(/^(x-ai|xai)\//, '') + return value === 'grok-4.5' || value.startsWith('grok-4.5-') +} + +function grokQuotaLooksHeavy(snapshot: Record | undefined): boolean { + const req = Number(snapshot?.requests?.limit ?? 0) + const tok = Number(snapshot?.tokens?.limit ?? 0) + return req >= 8300 || tok >= 53_000_000 +} + +function grok45ResponsesPlanIsHeavy(snapshot: Record | undefined): boolean { + if (!snapshot) return false + const hint = normalizeGrokPlanKey(snapshot.plan_from_45_responses) + if (hint === 'supergrokheavy' && isGrokQuotaTimestampFresh(snapshot.plan_from_45_responses_at)) { + return true + } + const observedAt = snapshot.last_headers_seen_at || snapshot.updated_at + return ( + isGrok45ResponsesQuotaModel(snapshot.model) && + isGrokQuotaTimestampFresh(observedAt) && + grokQuotaLooksHeavy(snapshot) + ) +} + +// JWT / unambiguous credentials outrank snapshots. SuperGrokPro is ambiguous +// (covers SuperGrok and Heavy). 8300/53M only upgrades when the window came +// from grok-4.5 Responses (or a carried 4.5 hint). function getAccountPlanType(row: any): string | undefined { if (!row) return undefined if (row.platform === 'grok') { const extra = (row.extra || {}) as Record const billing = extra.grok_billing_snapshot as Record | undefined - const quota = extra.grok_quota_snapshot as Record | undefined - return ( - billing?.plan || - quota?.subscription_tier || - row.credentials?.subscription_tier || - extra.subscription_tier || - row.credentials?.plan_type || - row.parent_plan_type || - undefined + const usage = extra.grok_usage_snapshot as Record | undefined + const legacyQuota = extra.grok_quota_snapshot as Record | undefined + const quota = grokPersistedQuotaSnapshot(extra) + const cred = firstNonBlankString(row.credentials?.subscription_tier) + const credKey = normalizeGrokPlanKey(cred) + if (credKey && credKey !== 'supergrokpro') { + return cred + } + if ( + grok45ResponsesPlanIsHeavy(quota) && + (credKey === 'supergrokpro' || + normalizeGrokPlanKey(billing?.plan) === 'supergrok' || + normalizeGrokPlanKey(billing?.plan) === 'supergrokpro') + ) { + return 'SuperGrok Heavy' + } + if (credKey === 'supergrokpro') { + return firstNonBlankString(billing?.plan) || 'SuperGrok' + } + return firstNonBlankString( + billing?.plan, + usage?.subscription_tier, + legacyQuota?.subscription_tier, + extra.subscription_tier, + row.credentials?.plan_type, + row.parent_plan_type ) } - return row.credentials?.plan_type || row.parent_plan_type || undefined + return firstNonBlankString(row.credentials?.plan_type, row.parent_plan_type) } function getOpenAIAuthMode(row: any): string | undefined { @@ -2355,6 +2440,7 @@ const handleClickOutside = (event: MouseEvent) => { onMounted(async () => { if (typeof window !== 'undefined') { + loadSavedAutoRefresh() desktopViewportMediaQuery = window.matchMedia(desktopViewportQuery) isDesktopViewport.value = desktopViewportMediaQuery.matches desktopViewportListener = (event: MediaQueryListEvent) => { diff --git a/frontend/src/views/admin/GroupsView.vue b/frontend/src/views/admin/GroupsView.vue index bc9ac2d19b..5cb036f7bc 100644 --- a/frontend/src/views/admin/GroupsView.vue +++ b/frontend/src/views/admin/GroupsView.vue @@ -1484,6 +1484,25 @@
+
+
+
+

{{ t("admin.groups.modelPricing.title") }}

+

{{ t("admin.groups.modelPricing.description") }}

+
+ +
+ +
+ +
+
+
+
+
+
+

{{ t("admin.groups.modelPricing.title") }}

+

{{ t("admin.groups.modelPricing.description") }}

+
+ +
+ +
+ +
+
+
({ + models: [], + billing_mode: "token", + input_price: null, + output_price: null, + cache_write_price: null, + cache_read_price: null, + image_input_price: null, + image_output_price: null, + per_request_price: null, + intervals: [], +}); + +const addGroupPricing = (entries: PricingFormEntry[]) => + entries.push(emptyGroupPricing()); + +const groupPricingFromAPI = ( + pricing: ChannelModelPricing[] | undefined, +): PricingFormEntry[] => + (pricing || []).map((entry) => ({ + models: entry.models || [], + billing_mode: entry.billing_mode || "token", + input_price: perTokenToMTok(entry.input_price), + output_price: perTokenToMTok(entry.output_price), + cache_write_price: perTokenToMTok(entry.cache_write_price), + cache_read_price: perTokenToMTok(entry.cache_read_price), + image_input_price: perTokenToMTok(entry.image_input_price), + image_output_price: perTokenToMTok(entry.image_output_price), + per_request_price: entry.per_request_price, + intervals: apiIntervalsToForm(entry.intervals || []), + })); + +const groupPricingToAPI = ( + pricing: PricingFormEntry[], + platform: string, +): ChannelModelPricing[] => + pricing + .filter((entry) => entry.models.length > 0) + .map((entry) => ({ + platform, + models: entry.models, + billing_mode: entry.billing_mode, + input_price: mTokToPerToken(entry.input_price), + output_price: mTokToPerToken(entry.output_price), + cache_write_price: mTokToPerToken(entry.cache_write_price), + cache_read_price: mTokToPerToken(entry.cache_read_price), + image_input_price: mTokToPerToken(entry.image_input_price), + image_output_price: mTokToPerToken(entry.image_output_price), + per_request_price: toNullableNumber(entry.per_request_price), + intervals: + entry.billing_mode === "token" + ? [] + : formIntervalsToAPI(entry.intervals || []), + })); + const { t } = useI18n(); const appStore = useAppStore(); const onboardingStore = useOnboardingStore(); @@ -4913,6 +5016,8 @@ const createForm = reactive({ daily_limit_usd: null as number | null, weekly_limit_usd: null as number | null, monthly_limit_usd: null as number | null, + long_context_pricing_enabled: true, + model_pricing: [] as PricingFormEntry[], // 图片生成计费配置 allow_image_generation: false, allow_batch_image_generation: false, @@ -5272,6 +5377,8 @@ const editForm = reactive({ daily_limit_usd: null as number | null, weekly_limit_usd: null as number | null, monthly_limit_usd: null as number | null, + long_context_pricing_enabled: true, + model_pricing: [] as PricingFormEntry[], // 图片生成计费配置 allow_image_generation: false, allow_batch_image_generation: false, @@ -5744,6 +5851,8 @@ const closeCreateModal = () => { createForm.video_price_720p = null; createForm.video_price_1080p = null; createForm.video_model_prices = createVideoModelPricesForm(); + createForm.long_context_pricing_enabled = true; + createForm.model_pricing = []; createForm.web_search_price_per_call = null; createForm.search_price_per_1k = null; createForm.audio_realtime_price_per_min = null; @@ -5843,6 +5952,10 @@ const handleCreateGroup = async () => { // 构建请求数据,包含模型路由配置 const requestData = { ...createGroupForm, + model_pricing: groupPricingToAPI( + createForm.model_pricing, + createForm.platform, + ), daily_limit_usd: normalizeOptionalLimit( createForm.daily_limit_usd as number | string | null, ), @@ -5965,6 +6078,9 @@ const handleEdit = async (group: AdminGroup) => { editForm.daily_limit_usd = group.daily_limit_usd; editForm.weekly_limit_usd = group.weekly_limit_usd; editForm.monthly_limit_usd = group.monthly_limit_usd; + editForm.long_context_pricing_enabled = + group.long_context_pricing_enabled ?? true; + editForm.model_pricing = groupPricingFromAPI(group.model_pricing); editForm.allow_image_generation = group.allow_image_generation ?? false; editForm.allow_batch_image_generation = group.allow_batch_image_generation ?? false; @@ -6069,6 +6185,8 @@ const closeEditModal = () => { editForm.video_price_720p = null; editForm.video_price_1080p = null; editForm.video_model_prices = createVideoModelPricesForm(); + editForm.long_context_pricing_enabled = true; + editForm.model_pricing = []; editForm.web_search_price_per_call = null; editForm.search_price_per_1k = null; editForm.audio_realtime_price_per_min = null; @@ -6101,6 +6219,10 @@ const handleUpdateGroup = async () => { // 转换 fallback_group_id: null -> 0 (后端使用 0 表示清除) const payload = { ...editForm, + model_pricing: groupPricingToAPI( + editForm.model_pricing, + editForm.platform, + ), daily_limit_usd: normalizeOptionalLimit( editForm.daily_limit_usd as number | string | null, ), diff --git a/frontend/src/views/admin/__tests__/AccountsView.sparkShadow.spec.ts b/frontend/src/views/admin/__tests__/AccountsView.sparkShadow.spec.ts index b6a1dc4650..08191b511b 100644 --- a/frontend/src/views/admin/__tests__/AccountsView.sparkShadow.spec.ts +++ b/frontend/src/views/admin/__tests__/AccountsView.sparkShadow.spec.ts @@ -270,6 +270,8 @@ describe('admin AccountsView — 账号行展示', () => { }) afterEach(() => { + vi.useRealTimers() + vi.restoreAllMocks() vi.unstubAllGlobals() }) @@ -346,7 +348,7 @@ describe('admin AccountsView — 账号行展示', () => { wrapper.unmount() }) - it('passes fresh Grok billing and quota snapshots before stale credential fallbacks', async () => { + it('prefers persisted Grok JWT tier over lagging billing/quota snapshots', async () => { const grokAccounts = [ { id: 201, @@ -396,6 +398,63 @@ describe('admin AccountsView — 账号行展示', () => { type: 'oauth', credentials: { plan_type: 'SuperGrok' }, }, + { + id: 206, + name: 'supergrokpro-responses-quota', + platform: 'grok', + type: 'oauth', + credentials: { subscription_tier: 'SuperGrokPro' }, + extra: { + grok_billing_snapshot: { plan: 'SuperGrok' }, + grok_usage_snapshot: { + model: 'grok-4.5', + last_headers_seen_at: new Date().toISOString(), + requests: { limit: 8300 }, + tokens: { limit: 53_000_000 }, + }, + grok_quota_snapshot: { + model: 'grok-4.6', + last_headers_seen_at: new Date().toISOString(), + requests: { limit: 8300 }, + tokens: { limit: 53_000_000 }, + }, + }, + }, + { + id: 207, + name: 'supergrokpro-other-model-quota', + platform: 'grok', + type: 'oauth', + credentials: { subscription_tier: 'SuperGrokPro' }, + extra: { + grok_billing_snapshot: { plan: 'SuperGrok' }, + grok_usage_snapshot: { + model: 'grok-4.6', + last_headers_seen_at: new Date().toISOString(), + requests: { limit: 8300 }, + tokens: { limit: 53_000_000 }, + }, + }, + }, + { + id: 208, + name: 'usage-over-legacy-quota', + platform: 'grok', + type: 'oauth', + credentials: {}, + extra: { + grok_usage_snapshot: { subscription_tier: 'SuperGrok' }, + grok_quota_snapshot: { subscription_tier: 'Free' }, + }, + }, + { + id: 209, + name: 'legacy-quota-alias', + platform: 'grok', + type: 'oauth', + credentials: {}, + extra: { grok_quota_snapshot: { subscription_tier: 'SuperGrok' } }, + }, ] listAccounts.mockResolvedValue({ @@ -411,13 +470,113 @@ describe('admin AccountsView — 账号行展示', () => { const badges = wrapper.findAllComponents(PlatformTypeBadge) expect(badges.map((badge) => badge.props('planType'))).toEqual([ + 'FREE', + 'SuperGrok Heavy', + 'FREE', + 'BASIC', 'SuperGrok', 'SuperGrok Heavy', 'SuperGrok', - 'BASIC', + 'SuperGrok', 'SuperGrok', ]) wrapper.unmount() }) + + it('skips malformed Grok plan fields and safely uses the next valid fallback', async () => { + const grokAccounts = [ + { + id: 210, + name: 'legacy-fallback', + platform: 'grok', + type: 'oauth', + credentials: {}, + extra: { + grok_usage_snapshot: { subscription_tier: { name: 'SuperGrok Heavy' } }, + grok_quota_snapshot: { subscription_tier: 'SuperGrok' }, + }, + }, + { + id: 211, + name: 'credential-plan-fallback', + platform: 'grok', + type: 'oauth', + credentials: { subscription_tier: 0, plan_type: 'SuperGrok Heavy' }, + extra: { + grok_billing_snapshot: { plan: {} }, + grok_usage_snapshot: { subscription_tier: 1 }, + grok_quota_snapshot: { subscription_tier: [] }, + subscription_tier: ' ', + }, + }, + { + id: 212, + name: 'no-valid-plan', + platform: 'grok', + type: 'oauth', + credentials: { subscription_tier: {}, plan_type: 2 }, + parent_plan_type: [], + extra: { + grok_billing_snapshot: { plan: [] }, + grok_usage_snapshot: { subscription_tier: 1 }, + grok_quota_snapshot: { subscription_tier: {} }, + subscription_tier: null, + }, + }, + ] + + listAccounts.mockResolvedValue({ + items: grokAccounts, + total: grokAccounts.length, + page: 1, + page_size: 20, + pages: 1, + }) + + const wrapper = mountViewWithRow() + await flushPromises() + + expect(wrapper.findAllComponents(PlatformTypeBadge).map((badge) => badge.props('planType'))).toEqual([ + 'SuperGrok', + 'SuperGrok Heavy', + undefined, + ]) + wrapper.unmount() + }) + + it('replaces a Grok row when auto refresh returns a changed canonical usage snapshot', async () => { + vi.useFakeTimers() + vi.spyOn(document, 'hidden', 'get').mockReturnValue(false) + localStorage.setItem('account-auto-refresh', JSON.stringify({ enabled: true, interval_seconds: 5 })) + + const initialAccount = { + id: 213, + name: 'refresh-tier', + platform: 'grok', + type: 'oauth', + extra: { grok_usage_snapshot: { subscription_tier: 'Free', status_code: 200 } }, + } + const refreshedAccount = { + ...initialAccount, + extra: { grok_usage_snapshot: { subscription_tier: 'SuperGrok', status_code: 200 } }, + } + listAccounts.mockResolvedValue({ items: [initialAccount], total: 1, page: 1, page_size: 20, pages: 1 }) + listWithEtag.mockResolvedValueOnce({ + notModified: false, + etag: 'grok-snapshot-2', + data: { items: [refreshedAccount], total: 1, page: 1, page_size: 20, pages: 1 }, + }) + + const wrapper = mountViewWithRow() + await flushPromises() + expect(wrapper.findComponent(PlatformTypeBadge).props('planType')).toBe('Free') + + await vi.advanceTimersByTimeAsync(6000) + await flushPromises() + + expect(listWithEtag).toHaveBeenCalledTimes(1) + expect(wrapper.findComponent(PlatformTypeBadge).props('planType')).toBe('SuperGrok') + wrapper.unmount() + }) })