Merge pull request #5571 from IanShaw027/feat/grok-jwt-tier-and-4.6

feat(grok): 以 JWT tier 识别订阅档位并接入 grok-4.6,修正徽章滞后、未知模型零计费与搜索/Voice 边界
This commit is contained in:
Wesley Liddick
2026-08-13 09:08:35 +08:00
committed by GitHub
86 changed files with 3111 additions and 276 deletions
+26 -2
View File
@@ -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(", ")
+13
View File
@@ -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()
+25
View File
@@ -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))
+136
View File
@@ -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) {
+93
View File
@@ -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)
}
+3 -1
View File
@@ -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",
+145 -1
View File
@@ -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
+22 -18
View File
@@ -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()
+9
View File
@@ -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").
+27 -19
View File
@@ -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,
+2
View File
@@ -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,
+9 -7
View File
@@ -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"`
+1 -1
View File
@@ -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
+56 -17
View File
@@ -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 {
+14 -11
View File
@@ -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
@@ -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)
}
}
@@ -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)
}
@@ -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())
}
@@ -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 选择,
@@ -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",
@@ -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
}
@@ -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))
}
+18 -2
View File
@@ -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.
+3
View File
@@ -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",
+9
View File
@@ -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(""))
+2
View File
@@ -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"])
+6
View File
@@ -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 {
@@ -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 ""
}
@@ -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"
}
@@ -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,
+13
View File
@@ -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).
@@ -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,
+16
View File
@@ -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)
@@ -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",
@@ -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 {
+38
View File
@@ -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 {
+23 -19
View File
@@ -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
@@ -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)
+108 -32
View File
@@ -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,
@@ -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()
+30 -30
View File
@@ -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)
+1 -1
View File
@@ -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",
@@ -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)
}
+40 -7
View File
@@ -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.
@@ -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")
+2 -2
View File
@@ -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
}
@@ -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()
+16 -5
View File
@@ -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
}
}
}
@@ -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{
+54 -8
View File
@@ -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
@@ -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)
}
@@ -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())
@@ -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
@@ -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}
+5
View File
@@ -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
@@ -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)
}
}
@@ -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)
}
@@ -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
@@ -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.
@@ -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
@@ -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 {
@@ -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) {
@@ -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(`{
@@ -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 != "" {
@@ -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,
@@ -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
@@ -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 := ""
@@ -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';
@@ -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
@@ -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: '<div class="usage-bar">{{ label }}|{{ utilization }}</div>'
},
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: '<div class="usage-bar">{{ label }}|{{ utilization }}</div>'
},
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,
@@ -134,8 +134,8 @@
</div>
</div>
<!-- Token intervals -->
<div class="mt-3">
<!-- Token intervals (channel-only; group long-context uses official presets) -->
<div v-if="!hideTokenIntervals" class="mt-3">
<div class="flex items-center justify-between">
<label class="text-xs font-medium text-gray-500 dark:text-gray-400">
{{ t('admin.channels.form.intervals') }}
@@ -194,11 +194,11 @@
</div>
</div>
<!-- Image mode -->
<div v-else-if="entry.billing_mode === 'image'">
<!-- Image/video mode -->
<div v-else-if="entry.billing_mode === 'image' || entry.billing_mode === 'video'">
<!-- Default image price (per-request, same as per_request mode) -->
<label class="mt-3 block text-xs font-medium text-gray-500 dark:text-gray-400">
{{ t('admin.channels.form.defaultImagePrice') }}
{{ entry.billing_mode === 'video' ? t('admin.channels.form.defaultVideoPrice') : t('admin.channels.form.defaultImagePrice') }}
<span class="ml-1 font-normal text-gray-400">$</span>
</label>
<div class="mt-1 w-48">
@@ -209,9 +209,9 @@
<!-- Image tiers -->
<div class="mt-3 flex items-center justify-between">
<label class="text-xs font-medium text-gray-500 dark:text-gray-400">
{{ t('admin.channels.form.imageTiers') }}
{{ entry.billing_mode === 'video' ? t('admin.channels.form.videoTiers') : t('admin.channels.form.imageTiers') }}
</label>
<button type="button" @click="addImageTier" class="text-xs text-primary-600 hover:text-primary-700">
<button type="button" @click="addMediaTier" class="text-xs text-primary-600 hover:text-primary-700">
+ {{ t('admin.channels.form.addTier') }}
</button>
</div>
@@ -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,
@@ -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 ''
@@ -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', () => {
@@ -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')
@@ -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' },
+2
View File
@@ -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
@@ -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.',
@@ -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.',
@@ -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 请求,请谨慎操作',
@@ -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)。留空表示未配置。',
+6
View File
@@ -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
@@ -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('')
})
})
+42
View File
@@ -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<string, unknown>)
.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<Account, 'id' | 'platform' | 'type' | 'updated_at' | 'last_used_at' | 'rate_limit_reset_at' | 'extra'>): string => {
if (account.platform !== 'openai' || account.type !== 'oauth') {
return ''
@@ -27,3 +51,21 @@ export const buildOpenAIUsageRefreshKey = (account: Pick<Account, 'id' | 'platfo
extra.codex_7d_window_minutes
].map(normalizeUsageRefreshValue).join('|')
}
export const buildGrokUsageRefreshKey = (account: Pick<Account, 'platform' | 'extra'>): string => {
if (account.platform !== 'grok') {
return ''
}
const extra = account.extra ?? {}
const usageSnapshot = extra.grok_usage_snapshot
const canonicalTier = (usageSnapshot as Record<string, unknown> | null | undefined)?.subscription_tier
const legacyQuotaFallback = isNonBlankString(canonicalTier)
? undefined
: extra.grok_quota_snapshot
return [
serializeSnapshotRefreshValue(extra.grok_billing_snapshot),
serializeSnapshotRefreshValue(usageSnapshot),
serializeSnapshotRefreshValue(legacyQuotaFallback)
].join('|')
}
+100 -14
View File
@@ -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<string, any>): Record<string, any> | undefined {
const usage = extra.grok_usage_snapshot
if (usage && typeof usage === 'object' && !Array.isArray(usage)) {
return usage as Record<string, any>
}
const legacy = extra.grok_quota_snapshot
if (legacy && typeof legacy === 'object' && !Array.isArray(legacy)) {
return legacy as Record<string, any>
}
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<string, any> | 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<string, any> | 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<string, any>
const billing = extra.grok_billing_snapshot as Record<string, any> | undefined
const quota = extra.grok_quota_snapshot as Record<string, any> | 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<string, any> | undefined
const legacyQuota = extra.grok_quota_snapshot as Record<string, any> | 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) => {
+122
View File
@@ -1484,6 +1484,25 @@
</div>
<div class="border-t border-gray-200 pt-4 mt-4 dark:border-dark-400">
<div class="flex items-start justify-between gap-4">
<div>
<h4 class="text-sm font-medium text-gray-700 dark:text-gray-300">{{ t("admin.groups.modelPricing.title") }}</h4>
<p class="mt-1 text-xs text-gray-500 dark:text-gray-400">{{ t("admin.groups.modelPricing.description") }}</p>
</div>
<button type="button" class="btn btn-secondary" @click="addGroupPricing(createForm.model_pricing)">
<Icon name="plus" size="sm" class="mr-1" />{{ t("admin.groups.modelPricing.add") }}
</button>
</div>
<label class="mt-3 flex items-start gap-2">
<input v-model="createForm.long_context_pricing_enabled" type="checkbox" class="mt-0.5" />
<span><span class="block text-sm text-gray-700 dark:text-gray-300">{{ t("admin.groups.modelPricing.longContext") }}</span><span class="block text-xs text-gray-500">{{ t("admin.groups.modelPricing.longContextHint") }}</span></span>
</label>
<div class="mt-3 space-y-2">
<PricingEntryCard v-for="(entry, index) in createForm.model_pricing" :key="index" :entry="entry" :platform="createForm.platform" hide-token-intervals @update="createForm.model_pricing[index] = $event" @remove="createForm.model_pricing.splice(index, 1)" />
</div>
</div>
<!-- Grok Voice 显式定价(仅 grok 平台) -->
<div
v-if="createForm.platform === 'grok'"
@@ -3187,6 +3206,25 @@
</div>
<div class="border-t border-gray-200 pt-4 mt-4 dark:border-dark-400">
<div class="flex items-start justify-between gap-4">
<div>
<h4 class="text-sm font-medium text-gray-700 dark:text-gray-300">{{ t("admin.groups.modelPricing.title") }}</h4>
<p class="mt-1 text-xs text-gray-500 dark:text-gray-400">{{ t("admin.groups.modelPricing.description") }}</p>
</div>
<button type="button" class="btn btn-secondary" @click="addGroupPricing(editForm.model_pricing)">
<Icon name="plus" size="sm" class="mr-1" />{{ t("admin.groups.modelPricing.add") }}
</button>
</div>
<label class="mt-3 flex items-start gap-2">
<input v-model="editForm.long_context_pricing_enabled" type="checkbox" class="mt-0.5" />
<span><span class="block text-sm text-gray-700 dark:text-gray-300">{{ t("admin.groups.modelPricing.longContext") }}</span><span class="block text-xs text-gray-500">{{ t("admin.groups.modelPricing.longContextHint") }}</span></span>
</label>
<div class="mt-3 space-y-2">
<PricingEntryCard v-for="(entry, index) in editForm.model_pricing" :key="index" :entry="entry" :platform="editForm.platform" hide-token-intervals @update="editForm.model_pricing[index] = $event" @remove="editForm.model_pricing.splice(index, 1)" />
</div>
</div>
<!-- Grok Voice 显式定价(仅 grok 平台) -->
<div
v-if="editForm.platform === 'grok'"
@@ -4377,6 +4415,16 @@ import GroupRateMultipliersModal from "@/components/admin/group/GroupRateMultipl
import GroupRPMOverridesModal from "@/components/admin/group/GroupRPMOverridesModal.vue";
import GroupCapacityBadge from "@/components/common/GroupCapacityBadge.vue";
import ReasoningEffortPolicyFields from "@/components/admin/group/ReasoningEffortPolicyFields.vue";
import PricingEntryCard from "@/components/admin/channel/PricingEntryCard.vue";
import type { PricingFormEntry } from "@/components/admin/channel/types";
import {
apiIntervalsToForm,
formIntervalsToAPI,
mTokToPerToken,
perTokenToMTok,
toNullableNumber,
} from "@/components/admin/channel/types";
import type { ChannelModelPricing } from "@/api/admin/channels";
import { VueDraggable } from "vue-draggable-plus";
import { createStableObjectKeyResolver } from "@/utils/stableObjectKey";
import { extractApiErrorMessage } from "@/utils/apiError";
@@ -4430,6 +4478,61 @@ import {
videoModelPriceFamilyRows,
} from "./groupsVideoModelPricing";
const emptyGroupPricing = (): PricingFormEntry => ({
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,
),
@@ -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()
})
})