mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-10-07 11:33:18 +08:00
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:
+26
-2
@@ -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(", ")
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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))
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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").
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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"`
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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))
|
||||
}
|
||||
@@ -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.
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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(""))
|
||||
|
||||
@@ -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"])
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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{
|
||||
|
||||
@@ -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}
|
||||
|
||||
@@ -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' },
|
||||
|
||||
@@ -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)。留空表示未配置。',
|
||||
|
||||
@@ -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('')
|
||||
})
|
||||
})
|
||||
|
||||
@@ -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('|')
|
||||
}
|
||||
|
||||
@@ -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) => {
|
||||
|
||||
@@ -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()
|
||||
})
|
||||
})
|
||||
|
||||
Reference in New Issue
Block a user