diff --git a/.gitignore b/.gitignore index 1b783b31cb87..6d496fe7e292 100644 --- a/.gitignore +++ b/.gitignore @@ -183,3 +183,5 @@ issues/ !docs/credential-encryption.md !docs/screenshots/credential-encryption/ !docs/screenshots/pelican-showcase-api-ui/ + +!docs/typesafe-jev.md diff --git a/backend/ent/migrate/schema.go b/backend/ent/migrate/schema.go index f0213edf1fcd..2a759e73d5ef 100644 --- a/backend/ent/migrate/schema.go +++ b/backend/ent/migrate/schema.go @@ -1132,6 +1132,7 @@ var ( {Name: "amount", Type: field.TypeFloat64, SchemaType: map[string]string{"postgres": "decimal(20,2)"}}, {Name: "pay_amount", Type: field.TypeFloat64, SchemaType: map[string]string{"postgres": "decimal(20,2)"}}, {Name: "fee_rate", Type: field.TypeFloat64, Default: 0, SchemaType: map[string]string{"postgres": "decimal(10,4)"}}, + {Name: "bonus_amount", Type: field.TypeFloat64, Default: 0, SchemaType: map[string]string{"postgres": "decimal(20,2)"}}, {Name: "recharge_code", Type: field.TypeString, Size: 64}, {Name: "out_trade_no", Type: field.TypeString, Size: 64, Default: ""}, {Name: "payment_type", Type: field.TypeString, Size: 30}, @@ -1174,7 +1175,7 @@ var ( ForeignKeys: []*schema.ForeignKey{ { Symbol: "payment_orders_users_payment_orders", - Columns: []*schema.Column{PaymentOrdersColumns[39]}, + Columns: []*schema.Column{PaymentOrdersColumns[40]}, RefColumns: []*schema.Column{UsersColumns[0]}, OnDelete: schema.NoAction, }, @@ -1183,7 +1184,7 @@ var ( { Name: "paymentorder_out_trade_no", Unique: true, - Columns: []*schema.Column{PaymentOrdersColumns[8]}, + Columns: []*schema.Column{PaymentOrdersColumns[9]}, Annotation: &entsql.IndexAnnotation{ Where: "out_trade_no <> ''", }, @@ -1191,37 +1192,37 @@ var ( { Name: "paymentorder_user_id", Unique: false, - Columns: []*schema.Column{PaymentOrdersColumns[39]}, + Columns: []*schema.Column{PaymentOrdersColumns[40]}, }, { Name: "paymentorder_status", Unique: false, - Columns: []*schema.Column{PaymentOrdersColumns[21]}, + Columns: []*schema.Column{PaymentOrdersColumns[22]}, }, { Name: "paymentorder_expires_at", Unique: false, - Columns: []*schema.Column{PaymentOrdersColumns[29]}, + Columns: []*schema.Column{PaymentOrdersColumns[30]}, }, { Name: "paymentorder_created_at", Unique: false, - Columns: []*schema.Column{PaymentOrdersColumns[37]}, + Columns: []*schema.Column{PaymentOrdersColumns[38]}, }, { Name: "paymentorder_paid_at", Unique: false, - Columns: []*schema.Column{PaymentOrdersColumns[30]}, + Columns: []*schema.Column{PaymentOrdersColumns[31]}, }, { Name: "paymentorder_payment_type_paid_at", Unique: false, - Columns: []*schema.Column{PaymentOrdersColumns[9], PaymentOrdersColumns[30]}, + Columns: []*schema.Column{PaymentOrdersColumns[10], PaymentOrdersColumns[31]}, }, { Name: "paymentorder_order_type", Unique: false, - Columns: []*schema.Column{PaymentOrdersColumns[14]}, + Columns: []*schema.Column{PaymentOrdersColumns[15]}, }, }, } diff --git a/backend/ent/mutation.go b/backend/ent/mutation.go index 647a2ffde4a4..df0feb32c417 100644 --- a/backend/ent/mutation.go +++ b/backend/ent/mutation.go @@ -30415,6 +30415,8 @@ type PaymentOrderMutation struct { addpay_amount *float64 fee_rate *float64 addfee_rate *float64 + bonus_amount *float64 + addbonus_amount *float64 recharge_code *string out_trade_no *string payment_type *string @@ -30882,6 +30884,62 @@ func (m *PaymentOrderMutation) ResetFeeRate() { m.addfee_rate = nil } +// SetBonusAmount sets the "bonus_amount" field. +func (m *PaymentOrderMutation) SetBonusAmount(f float64) { + m.bonus_amount = &f + m.addbonus_amount = nil +} + +// BonusAmount returns the value of the "bonus_amount" field in the mutation. +func (m *PaymentOrderMutation) BonusAmount() (r float64, exists bool) { + v := m.bonus_amount + if v == nil { + return + } + return *v, true +} + +// OldBonusAmount returns the old "bonus_amount" field's value of the PaymentOrder entity. +// If the PaymentOrder 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 *PaymentOrderMutation) OldBonusAmount(ctx context.Context) (v float64, err error) { + if !m.op.Is(OpUpdateOne) { + return v, errors.New("OldBonusAmount is only allowed on UpdateOne operations") + } + if m.id == nil || m.oldValue == nil { + return v, errors.New("OldBonusAmount requires an ID field in the mutation") + } + oldValue, err := m.oldValue(ctx) + if err != nil { + return v, fmt.Errorf("querying old value for OldBonusAmount: %w", err) + } + return oldValue.BonusAmount, nil +} + +// AddBonusAmount adds f to the "bonus_amount" field. +func (m *PaymentOrderMutation) AddBonusAmount(f float64) { + if m.addbonus_amount != nil { + *m.addbonus_amount += f + } else { + m.addbonus_amount = &f + } +} + +// AddedBonusAmount returns the value that was added to the "bonus_amount" field in this mutation. +func (m *PaymentOrderMutation) AddedBonusAmount() (r float64, exists bool) { + v := m.addbonus_amount + if v == nil { + return + } + return *v, true +} + +// ResetBonusAmount resets all changes to the "bonus_amount" field. +func (m *PaymentOrderMutation) ResetBonusAmount() { + m.bonus_amount = nil + m.addbonus_amount = nil +} + // SetRechargeCode sets the "recharge_code" field. func (m *PaymentOrderMutation) SetRechargeCode(s string) { m.recharge_code = &s @@ -32425,7 +32483,7 @@ func (m *PaymentOrderMutation) Type() string { // order to get all numeric fields that were incremented/decremented, call // AddedFields(). func (m *PaymentOrderMutation) Fields() []string { - fields := make([]string, 0, 39) + fields := make([]string, 0, 40) if m.user != nil { fields = append(fields, paymentorder.FieldUserID) } @@ -32447,6 +32505,9 @@ func (m *PaymentOrderMutation) Fields() []string { if m.fee_rate != nil { fields = append(fields, paymentorder.FieldFeeRate) } + if m.bonus_amount != nil { + fields = append(fields, paymentorder.FieldBonusAmount) + } if m.recharge_code != nil { fields = append(fields, paymentorder.FieldRechargeCode) } @@ -32565,6 +32626,8 @@ func (m *PaymentOrderMutation) Field(name string) (ent.Value, bool) { return m.PayAmount() case paymentorder.FieldFeeRate: return m.FeeRate() + case paymentorder.FieldBonusAmount: + return m.BonusAmount() case paymentorder.FieldRechargeCode: return m.RechargeCode() case paymentorder.FieldOutTradeNo: @@ -32652,6 +32715,8 @@ func (m *PaymentOrderMutation) OldField(ctx context.Context, name string) (ent.V return m.OldPayAmount(ctx) case paymentorder.FieldFeeRate: return m.OldFeeRate(ctx) + case paymentorder.FieldBonusAmount: + return m.OldBonusAmount(ctx) case paymentorder.FieldRechargeCode: return m.OldRechargeCode(ctx) case paymentorder.FieldOutTradeNo: @@ -32774,6 +32839,13 @@ func (m *PaymentOrderMutation) SetField(name string, value ent.Value) error { } m.SetFeeRate(v) return nil + case paymentorder.FieldBonusAmount: + v, ok := value.(float64) + if !ok { + return fmt.Errorf("unexpected type %T for field %s", value, name) + } + m.SetBonusAmount(v) + return nil case paymentorder.FieldRechargeCode: v, ok := value.(string) if !ok { @@ -33015,6 +33087,9 @@ func (m *PaymentOrderMutation) AddedFields() []string { if m.addfee_rate != nil { fields = append(fields, paymentorder.FieldFeeRate) } + if m.addbonus_amount != nil { + fields = append(fields, paymentorder.FieldBonusAmount) + } if m.addplan_id != nil { fields = append(fields, paymentorder.FieldPlanID) } @@ -33041,6 +33116,8 @@ func (m *PaymentOrderMutation) AddedField(name string) (ent.Value, bool) { return m.AddedPayAmount() case paymentorder.FieldFeeRate: return m.AddedFeeRate() + case paymentorder.FieldBonusAmount: + return m.AddedBonusAmount() case paymentorder.FieldPlanID: return m.AddedPlanID() case paymentorder.FieldSubscriptionGroupID: @@ -33079,6 +33156,13 @@ func (m *PaymentOrderMutation) AddField(name string, value ent.Value) error { } m.AddFeeRate(v) return nil + case paymentorder.FieldBonusAmount: + v, ok := value.(float64) + if !ok { + return fmt.Errorf("unexpected type %T for field %s", value, name) + } + m.AddBonusAmount(v) + return nil case paymentorder.FieldPlanID: v, ok := value.(int64) if !ok { @@ -33278,6 +33362,9 @@ func (m *PaymentOrderMutation) ResetField(name string) error { case paymentorder.FieldFeeRate: m.ResetFeeRate() return nil + case paymentorder.FieldBonusAmount: + m.ResetBonusAmount() + return nil case paymentorder.FieldRechargeCode: m.ResetRechargeCode() return nil diff --git a/backend/ent/paymentorder.go b/backend/ent/paymentorder.go index b131b8c88045..db5699499075 100644 --- a/backend/ent/paymentorder.go +++ b/backend/ent/paymentorder.go @@ -33,6 +33,8 @@ type PaymentOrder struct { PayAmount float64 `json:"pay_amount,omitempty"` // FeeRate holds the value of the "fee_rate" field. FeeRate float64 `json:"fee_rate,omitempty"` + // BonusAmount holds the value of the "bonus_amount" field. + BonusAmount float64 `json:"bonus_amount,omitempty"` // RechargeCode holds the value of the "recharge_code" field. RechargeCode string `json:"recharge_code,omitempty"` // OutTradeNo holds the value of the "out_trade_no" field. @@ -132,7 +134,7 @@ func (*PaymentOrder) scanValues(columns []string) ([]any, error) { values[i] = new([]byte) case paymentorder.FieldForceRefund: values[i] = new(sql.NullBool) - case paymentorder.FieldAmount, paymentorder.FieldPayAmount, paymentorder.FieldFeeRate, paymentorder.FieldRefundAmount: + case paymentorder.FieldAmount, paymentorder.FieldPayAmount, paymentorder.FieldFeeRate, paymentorder.FieldBonusAmount, paymentorder.FieldRefundAmount: values[i] = new(sql.NullFloat64) case paymentorder.FieldID, paymentorder.FieldUserID, paymentorder.FieldPlanID, paymentorder.FieldSubscriptionGroupID, paymentorder.FieldSubscriptionDays: values[i] = new(sql.NullInt64) @@ -204,6 +206,12 @@ func (_m *PaymentOrder) assignValues(columns []string, values []any) error { } else if value.Valid { _m.FeeRate = value.Float64 } + case paymentorder.FieldBonusAmount: + if value, ok := values[i].(*sql.NullFloat64); !ok { + return fmt.Errorf("unexpected type %T for field bonus_amount", values[i]) + } else if value.Valid { + _m.BonusAmount = value.Float64 + } case paymentorder.FieldRechargeCode: if value, ok := values[i].(*sql.NullString); !ok { return fmt.Errorf("unexpected type %T for field recharge_code", values[i]) @@ -480,6 +488,9 @@ func (_m *PaymentOrder) String() string { builder.WriteString("fee_rate=") builder.WriteString(fmt.Sprintf("%v", _m.FeeRate)) builder.WriteString(", ") + builder.WriteString("bonus_amount=") + builder.WriteString(fmt.Sprintf("%v", _m.BonusAmount)) + builder.WriteString(", ") builder.WriteString("recharge_code=") builder.WriteString(_m.RechargeCode) builder.WriteString(", ") diff --git a/backend/ent/paymentorder/paymentorder.go b/backend/ent/paymentorder/paymentorder.go index 628837943428..391389b159e4 100644 --- a/backend/ent/paymentorder/paymentorder.go +++ b/backend/ent/paymentorder/paymentorder.go @@ -28,6 +28,8 @@ const ( FieldPayAmount = "pay_amount" // FieldFeeRate holds the string denoting the fee_rate field in the database. FieldFeeRate = "fee_rate" + // FieldBonusAmount holds the string denoting the bonus_amount field in the database. + FieldBonusAmount = "bonus_amount" // FieldRechargeCode holds the string denoting the recharge_code field in the database. FieldRechargeCode = "recharge_code" // FieldOutTradeNo holds the string denoting the out_trade_no field in the database. @@ -115,6 +117,7 @@ var Columns = []string{ FieldAmount, FieldPayAmount, FieldFeeRate, + FieldBonusAmount, FieldRechargeCode, FieldOutTradeNo, FieldPaymentType, @@ -166,6 +169,8 @@ var ( UserNameValidator func(string) error // DefaultFeeRate holds the default value on creation for the "fee_rate" field. DefaultFeeRate float64 + // DefaultBonusAmount holds the default value on creation for the "bonus_amount" field. + DefaultBonusAmount float64 // RechargeCodeValidator is a validator for the "recharge_code" field. It is called by the builders before save. RechargeCodeValidator func(string) error // DefaultOutTradeNo holds the default value on creation for the "out_trade_no" field. @@ -249,6 +254,11 @@ func ByFeeRate(opts ...sql.OrderTermOption) OrderOption { return sql.OrderByField(FieldFeeRate, opts...).ToFunc() } +// ByBonusAmount orders the results by the bonus_amount field. +func ByBonusAmount(opts ...sql.OrderTermOption) OrderOption { + return sql.OrderByField(FieldBonusAmount, opts...).ToFunc() +} + // ByRechargeCode orders the results by the recharge_code field. func ByRechargeCode(opts ...sql.OrderTermOption) OrderOption { return sql.OrderByField(FieldRechargeCode, opts...).ToFunc() diff --git a/backend/ent/paymentorder/where.go b/backend/ent/paymentorder/where.go index e96bf51ebd09..33e5267791ef 100644 --- a/backend/ent/paymentorder/where.go +++ b/backend/ent/paymentorder/where.go @@ -90,6 +90,11 @@ func FeeRate(v float64) predicate.PaymentOrder { return predicate.PaymentOrder(sql.FieldEQ(FieldFeeRate, v)) } +// BonusAmount applies equality check predicate on the "bonus_amount" field. It's identical to BonusAmountEQ. +func BonusAmount(v float64) predicate.PaymentOrder { + return predicate.PaymentOrder(sql.FieldEQ(FieldBonusAmount, v)) +} + // RechargeCode applies equality check predicate on the "recharge_code" field. It's identical to RechargeCodeEQ. func RechargeCode(v string) predicate.PaymentOrder { return predicate.PaymentOrder(sql.FieldEQ(FieldRechargeCode, v)) @@ -590,6 +595,46 @@ func FeeRateLTE(v float64) predicate.PaymentOrder { return predicate.PaymentOrder(sql.FieldLTE(FieldFeeRate, v)) } +// BonusAmountEQ applies the EQ predicate on the "bonus_amount" field. +func BonusAmountEQ(v float64) predicate.PaymentOrder { + return predicate.PaymentOrder(sql.FieldEQ(FieldBonusAmount, v)) +} + +// BonusAmountNEQ applies the NEQ predicate on the "bonus_amount" field. +func BonusAmountNEQ(v float64) predicate.PaymentOrder { + return predicate.PaymentOrder(sql.FieldNEQ(FieldBonusAmount, v)) +} + +// BonusAmountIn applies the In predicate on the "bonus_amount" field. +func BonusAmountIn(vs ...float64) predicate.PaymentOrder { + return predicate.PaymentOrder(sql.FieldIn(FieldBonusAmount, vs...)) +} + +// BonusAmountNotIn applies the NotIn predicate on the "bonus_amount" field. +func BonusAmountNotIn(vs ...float64) predicate.PaymentOrder { + return predicate.PaymentOrder(sql.FieldNotIn(FieldBonusAmount, vs...)) +} + +// BonusAmountGT applies the GT predicate on the "bonus_amount" field. +func BonusAmountGT(v float64) predicate.PaymentOrder { + return predicate.PaymentOrder(sql.FieldGT(FieldBonusAmount, v)) +} + +// BonusAmountGTE applies the GTE predicate on the "bonus_amount" field. +func BonusAmountGTE(v float64) predicate.PaymentOrder { + return predicate.PaymentOrder(sql.FieldGTE(FieldBonusAmount, v)) +} + +// BonusAmountLT applies the LT predicate on the "bonus_amount" field. +func BonusAmountLT(v float64) predicate.PaymentOrder { + return predicate.PaymentOrder(sql.FieldLT(FieldBonusAmount, v)) +} + +// BonusAmountLTE applies the LTE predicate on the "bonus_amount" field. +func BonusAmountLTE(v float64) predicate.PaymentOrder { + return predicate.PaymentOrder(sql.FieldLTE(FieldBonusAmount, v)) +} + // RechargeCodeEQ applies the EQ predicate on the "recharge_code" field. func RechargeCodeEQ(v string) predicate.PaymentOrder { return predicate.PaymentOrder(sql.FieldEQ(FieldRechargeCode, v)) diff --git a/backend/ent/paymentorder_create.go b/backend/ent/paymentorder_create.go index 3ee24f8e918d..a8693d091446 100644 --- a/backend/ent/paymentorder_create.go +++ b/backend/ent/paymentorder_create.go @@ -81,6 +81,20 @@ func (_c *PaymentOrderCreate) SetNillableFeeRate(v *float64) *PaymentOrderCreate return _c } +// SetBonusAmount sets the "bonus_amount" field. +func (_c *PaymentOrderCreate) SetBonusAmount(v float64) *PaymentOrderCreate { + _c.mutation.SetBonusAmount(v) + return _c +} + +// SetNillableBonusAmount sets the "bonus_amount" field if the given value is not nil. +func (_c *PaymentOrderCreate) SetNillableBonusAmount(v *float64) *PaymentOrderCreate { + if v != nil { + _c.SetBonusAmount(*v) + } + return _c +} + // SetRechargeCode sets the "recharge_code" field. func (_c *PaymentOrderCreate) SetRechargeCode(v string) *PaymentOrderCreate { _c.mutation.SetRechargeCode(v) @@ -517,6 +531,10 @@ func (_c *PaymentOrderCreate) defaults() { v := paymentorder.DefaultFeeRate _c.mutation.SetFeeRate(v) } + if _, ok := _c.mutation.BonusAmount(); !ok { + v := paymentorder.DefaultBonusAmount + _c.mutation.SetBonusAmount(v) + } if _, ok := _c.mutation.OutTradeNo(); !ok { v := paymentorder.DefaultOutTradeNo _c.mutation.SetOutTradeNo(v) @@ -577,6 +595,9 @@ func (_c *PaymentOrderCreate) check() error { if _, ok := _c.mutation.FeeRate(); !ok { return &ValidationError{Name: "fee_rate", err: errors.New(`ent: missing required field "PaymentOrder.fee_rate"`)} } + if _, ok := _c.mutation.BonusAmount(); !ok { + return &ValidationError{Name: "bonus_amount", err: errors.New(`ent: missing required field "PaymentOrder.bonus_amount"`)} + } if _, ok := _c.mutation.RechargeCode(); !ok { return &ValidationError{Name: "recharge_code", err: errors.New(`ent: missing required field "PaymentOrder.recharge_code"`)} } @@ -725,6 +746,10 @@ func (_c *PaymentOrderCreate) createSpec() (*PaymentOrder, *sqlgraph.CreateSpec) _spec.SetField(paymentorder.FieldFeeRate, field.TypeFloat64, value) _node.FeeRate = value } + if value, ok := _c.mutation.BonusAmount(); ok { + _spec.SetField(paymentorder.FieldBonusAmount, field.TypeFloat64, value) + _node.BonusAmount = value + } if value, ok := _c.mutation.RechargeCode(); ok { _spec.SetField(paymentorder.FieldRechargeCode, field.TypeString, value) _node.RechargeCode = value @@ -1030,6 +1055,24 @@ func (u *PaymentOrderUpsert) AddFeeRate(v float64) *PaymentOrderUpsert { return u } +// SetBonusAmount sets the "bonus_amount" field. +func (u *PaymentOrderUpsert) SetBonusAmount(v float64) *PaymentOrderUpsert { + u.Set(paymentorder.FieldBonusAmount, v) + return u +} + +// UpdateBonusAmount sets the "bonus_amount" field to the value that was provided on create. +func (u *PaymentOrderUpsert) UpdateBonusAmount() *PaymentOrderUpsert { + u.SetExcluded(paymentorder.FieldBonusAmount) + return u +} + +// AddBonusAmount adds v to the "bonus_amount" field. +func (u *PaymentOrderUpsert) AddBonusAmount(v float64) *PaymentOrderUpsert { + u.Add(paymentorder.FieldBonusAmount, v) + return u +} + // SetRechargeCode sets the "recharge_code" field. func (u *PaymentOrderUpsert) SetRechargeCode(v string) *PaymentOrderUpsert { u.Set(paymentorder.FieldRechargeCode, v) @@ -1711,6 +1754,27 @@ func (u *PaymentOrderUpsertOne) UpdateFeeRate() *PaymentOrderUpsertOne { }) } +// SetBonusAmount sets the "bonus_amount" field. +func (u *PaymentOrderUpsertOne) SetBonusAmount(v float64) *PaymentOrderUpsertOne { + return u.Update(func(s *PaymentOrderUpsert) { + s.SetBonusAmount(v) + }) +} + +// AddBonusAmount adds v to the "bonus_amount" field. +func (u *PaymentOrderUpsertOne) AddBonusAmount(v float64) *PaymentOrderUpsertOne { + return u.Update(func(s *PaymentOrderUpsert) { + s.AddBonusAmount(v) + }) +} + +// UpdateBonusAmount sets the "bonus_amount" field to the value that was provided on create. +func (u *PaymentOrderUpsertOne) UpdateBonusAmount() *PaymentOrderUpsertOne { + return u.Update(func(s *PaymentOrderUpsert) { + s.UpdateBonusAmount() + }) +} + // SetRechargeCode sets the "recharge_code" field. func (u *PaymentOrderUpsertOne) SetRechargeCode(v string) *PaymentOrderUpsertOne { return u.Update(func(s *PaymentOrderUpsert) { @@ -2643,6 +2707,27 @@ func (u *PaymentOrderUpsertBulk) UpdateFeeRate() *PaymentOrderUpsertBulk { }) } +// SetBonusAmount sets the "bonus_amount" field. +func (u *PaymentOrderUpsertBulk) SetBonusAmount(v float64) *PaymentOrderUpsertBulk { + return u.Update(func(s *PaymentOrderUpsert) { + s.SetBonusAmount(v) + }) +} + +// AddBonusAmount adds v to the "bonus_amount" field. +func (u *PaymentOrderUpsertBulk) AddBonusAmount(v float64) *PaymentOrderUpsertBulk { + return u.Update(func(s *PaymentOrderUpsert) { + s.AddBonusAmount(v) + }) +} + +// UpdateBonusAmount sets the "bonus_amount" field to the value that was provided on create. +func (u *PaymentOrderUpsertBulk) UpdateBonusAmount() *PaymentOrderUpsertBulk { + return u.Update(func(s *PaymentOrderUpsert) { + s.UpdateBonusAmount() + }) +} + // SetRechargeCode sets the "recharge_code" field. func (u *PaymentOrderUpsertBulk) SetRechargeCode(v string) *PaymentOrderUpsertBulk { return u.Update(func(s *PaymentOrderUpsert) { diff --git a/backend/ent/paymentorder_update.go b/backend/ent/paymentorder_update.go index 378e0dad2f90..3ab893cbe6ee 100644 --- a/backend/ent/paymentorder_update.go +++ b/backend/ent/paymentorder_update.go @@ -154,6 +154,27 @@ func (_u *PaymentOrderUpdate) AddFeeRate(v float64) *PaymentOrderUpdate { return _u } +// SetBonusAmount sets the "bonus_amount" field. +func (_u *PaymentOrderUpdate) SetBonusAmount(v float64) *PaymentOrderUpdate { + _u.mutation.ResetBonusAmount() + _u.mutation.SetBonusAmount(v) + return _u +} + +// SetNillableBonusAmount sets the "bonus_amount" field if the given value is not nil. +func (_u *PaymentOrderUpdate) SetNillableBonusAmount(v *float64) *PaymentOrderUpdate { + if v != nil { + _u.SetBonusAmount(*v) + } + return _u +} + +// AddBonusAmount adds value to the "bonus_amount" field. +func (_u *PaymentOrderUpdate) AddBonusAmount(v float64) *PaymentOrderUpdate { + _u.mutation.AddBonusAmount(v) + return _u +} + // SetRechargeCode sets the "recharge_code" field. func (_u *PaymentOrderUpdate) SetRechargeCode(v string) *PaymentOrderUpdate { _u.mutation.SetRechargeCode(v) @@ -881,6 +902,12 @@ func (_u *PaymentOrderUpdate) sqlSave(ctx context.Context) (_node int, err error if value, ok := _u.mutation.AddedFeeRate(); ok { _spec.AddField(paymentorder.FieldFeeRate, field.TypeFloat64, value) } + if value, ok := _u.mutation.BonusAmount(); ok { + _spec.SetField(paymentorder.FieldBonusAmount, field.TypeFloat64, value) + } + if value, ok := _u.mutation.AddedBonusAmount(); ok { + _spec.AddField(paymentorder.FieldBonusAmount, field.TypeFloat64, value) + } if value, ok := _u.mutation.RechargeCode(); ok { _spec.SetField(paymentorder.FieldRechargeCode, field.TypeString, value) } @@ -1217,6 +1244,27 @@ func (_u *PaymentOrderUpdateOne) AddFeeRate(v float64) *PaymentOrderUpdateOne { return _u } +// SetBonusAmount sets the "bonus_amount" field. +func (_u *PaymentOrderUpdateOne) SetBonusAmount(v float64) *PaymentOrderUpdateOne { + _u.mutation.ResetBonusAmount() + _u.mutation.SetBonusAmount(v) + return _u +} + +// SetNillableBonusAmount sets the "bonus_amount" field if the given value is not nil. +func (_u *PaymentOrderUpdateOne) SetNillableBonusAmount(v *float64) *PaymentOrderUpdateOne { + if v != nil { + _u.SetBonusAmount(*v) + } + return _u +} + +// AddBonusAmount adds value to the "bonus_amount" field. +func (_u *PaymentOrderUpdateOne) AddBonusAmount(v float64) *PaymentOrderUpdateOne { + _u.mutation.AddBonusAmount(v) + return _u +} + // SetRechargeCode sets the "recharge_code" field. func (_u *PaymentOrderUpdateOne) SetRechargeCode(v string) *PaymentOrderUpdateOne { _u.mutation.SetRechargeCode(v) @@ -1974,6 +2022,12 @@ func (_u *PaymentOrderUpdateOne) sqlSave(ctx context.Context) (_node *PaymentOrd if value, ok := _u.mutation.AddedFeeRate(); ok { _spec.AddField(paymentorder.FieldFeeRate, field.TypeFloat64, value) } + if value, ok := _u.mutation.BonusAmount(); ok { + _spec.SetField(paymentorder.FieldBonusAmount, field.TypeFloat64, value) + } + if value, ok := _u.mutation.AddedBonusAmount(); ok { + _spec.AddField(paymentorder.FieldBonusAmount, field.TypeFloat64, value) + } if value, ok := _u.mutation.RechargeCode(); ok { _spec.SetField(paymentorder.FieldRechargeCode, field.TypeString, value) } diff --git a/backend/ent/runtime/runtime.go b/backend/ent/runtime/runtime.go index d34ba4151e35..2a1497f79a62 100644 --- a/backend/ent/runtime/runtime.go +++ b/backend/ent/runtime/runtime.go @@ -1337,70 +1337,74 @@ func init() { paymentorderDescFeeRate := paymentorderFields[6].Descriptor() // paymentorder.DefaultFeeRate holds the default value on creation for the fee_rate field. paymentorder.DefaultFeeRate = paymentorderDescFeeRate.Default.(float64) + // paymentorderDescBonusAmount is the schema descriptor for bonus_amount field. + paymentorderDescBonusAmount := paymentorderFields[7].Descriptor() + // paymentorder.DefaultBonusAmount holds the default value on creation for the bonus_amount field. + paymentorder.DefaultBonusAmount = paymentorderDescBonusAmount.Default.(float64) // paymentorderDescRechargeCode is the schema descriptor for recharge_code field. - paymentorderDescRechargeCode := paymentorderFields[7].Descriptor() + paymentorderDescRechargeCode := paymentorderFields[8].Descriptor() // paymentorder.RechargeCodeValidator is a validator for the "recharge_code" field. It is called by the builders before save. paymentorder.RechargeCodeValidator = paymentorderDescRechargeCode.Validators[0].(func(string) error) // paymentorderDescOutTradeNo is the schema descriptor for out_trade_no field. - paymentorderDescOutTradeNo := paymentorderFields[8].Descriptor() + paymentorderDescOutTradeNo := paymentorderFields[9].Descriptor() // paymentorder.DefaultOutTradeNo holds the default value on creation for the out_trade_no field. paymentorder.DefaultOutTradeNo = paymentorderDescOutTradeNo.Default.(string) // paymentorder.OutTradeNoValidator is a validator for the "out_trade_no" field. It is called by the builders before save. paymentorder.OutTradeNoValidator = paymentorderDescOutTradeNo.Validators[0].(func(string) error) // paymentorderDescPaymentType is the schema descriptor for payment_type field. - paymentorderDescPaymentType := paymentorderFields[9].Descriptor() + paymentorderDescPaymentType := paymentorderFields[10].Descriptor() // paymentorder.PaymentTypeValidator is a validator for the "payment_type" field. It is called by the builders before save. paymentorder.PaymentTypeValidator = paymentorderDescPaymentType.Validators[0].(func(string) error) // paymentorderDescPaymentTradeNo is the schema descriptor for payment_trade_no field. - paymentorderDescPaymentTradeNo := paymentorderFields[10].Descriptor() + paymentorderDescPaymentTradeNo := paymentorderFields[11].Descriptor() // paymentorder.PaymentTradeNoValidator is a validator for the "payment_trade_no" field. It is called by the builders before save. paymentorder.PaymentTradeNoValidator = paymentorderDescPaymentTradeNo.Validators[0].(func(string) error) // paymentorderDescOrderType is the schema descriptor for order_type field. - paymentorderDescOrderType := paymentorderFields[14].Descriptor() + paymentorderDescOrderType := paymentorderFields[15].Descriptor() // paymentorder.DefaultOrderType holds the default value on creation for the order_type field. paymentorder.DefaultOrderType = paymentorderDescOrderType.Default.(string) // paymentorder.OrderTypeValidator is a validator for the "order_type" field. It is called by the builders before save. paymentorder.OrderTypeValidator = paymentorderDescOrderType.Validators[0].(func(string) error) // paymentorderDescProviderInstanceID is the schema descriptor for provider_instance_id field. - paymentorderDescProviderInstanceID := paymentorderFields[18].Descriptor() + paymentorderDescProviderInstanceID := paymentorderFields[19].Descriptor() // paymentorder.ProviderInstanceIDValidator is a validator for the "provider_instance_id" field. It is called by the builders before save. paymentorder.ProviderInstanceIDValidator = paymentorderDescProviderInstanceID.Validators[0].(func(string) error) // paymentorderDescProviderKey is the schema descriptor for provider_key field. - paymentorderDescProviderKey := paymentorderFields[19].Descriptor() + paymentorderDescProviderKey := paymentorderFields[20].Descriptor() // paymentorder.ProviderKeyValidator is a validator for the "provider_key" field. It is called by the builders before save. paymentorder.ProviderKeyValidator = paymentorderDescProviderKey.Validators[0].(func(string) error) // paymentorderDescStatus is the schema descriptor for status field. - paymentorderDescStatus := paymentorderFields[21].Descriptor() + paymentorderDescStatus := paymentorderFields[22].Descriptor() // paymentorder.DefaultStatus holds the default value on creation for the status field. paymentorder.DefaultStatus = paymentorderDescStatus.Default.(string) // paymentorder.StatusValidator is a validator for the "status" field. It is called by the builders before save. paymentorder.StatusValidator = paymentorderDescStatus.Validators[0].(func(string) error) // paymentorderDescRefundAmount is the schema descriptor for refund_amount field. - paymentorderDescRefundAmount := paymentorderFields[22].Descriptor() + paymentorderDescRefundAmount := paymentorderFields[23].Descriptor() // paymentorder.DefaultRefundAmount holds the default value on creation for the refund_amount field. paymentorder.DefaultRefundAmount = paymentorderDescRefundAmount.Default.(float64) // paymentorderDescForceRefund is the schema descriptor for force_refund field. - paymentorderDescForceRefund := paymentorderFields[25].Descriptor() + paymentorderDescForceRefund := paymentorderFields[26].Descriptor() // paymentorder.DefaultForceRefund holds the default value on creation for the force_refund field. paymentorder.DefaultForceRefund = paymentorderDescForceRefund.Default.(bool) // paymentorderDescRefundRequestedBy is the schema descriptor for refund_requested_by field. - paymentorderDescRefundRequestedBy := paymentorderFields[28].Descriptor() + paymentorderDescRefundRequestedBy := paymentorderFields[29].Descriptor() // paymentorder.RefundRequestedByValidator is a validator for the "refund_requested_by" field. It is called by the builders before save. paymentorder.RefundRequestedByValidator = paymentorderDescRefundRequestedBy.Validators[0].(func(string) error) // paymentorderDescClientIP is the schema descriptor for client_ip field. - paymentorderDescClientIP := paymentorderFields[34].Descriptor() + paymentorderDescClientIP := paymentorderFields[35].Descriptor() // paymentorder.ClientIPValidator is a validator for the "client_ip" field. It is called by the builders before save. paymentorder.ClientIPValidator = paymentorderDescClientIP.Validators[0].(func(string) error) // paymentorderDescSrcHost is the schema descriptor for src_host field. - paymentorderDescSrcHost := paymentorderFields[35].Descriptor() + paymentorderDescSrcHost := paymentorderFields[36].Descriptor() // paymentorder.SrcHostValidator is a validator for the "src_host" field. It is called by the builders before save. paymentorder.SrcHostValidator = paymentorderDescSrcHost.Validators[0].(func(string) error) // paymentorderDescCreatedAt is the schema descriptor for created_at field. - paymentorderDescCreatedAt := paymentorderFields[37].Descriptor() + paymentorderDescCreatedAt := paymentorderFields[38].Descriptor() // paymentorder.DefaultCreatedAt holds the default value on creation for the created_at field. paymentorder.DefaultCreatedAt = paymentorderDescCreatedAt.Default.(func() time.Time) // paymentorderDescUpdatedAt is the schema descriptor for updated_at field. - paymentorderDescUpdatedAt := paymentorderFields[38].Descriptor() + paymentorderDescUpdatedAt := paymentorderFields[39].Descriptor() // paymentorder.DefaultUpdatedAt holds the default value on creation for the updated_at field. paymentorder.DefaultUpdatedAt = paymentorderDescUpdatedAt.Default.(func() time.Time) // paymentorder.UpdateDefaultUpdatedAt holds the default value on update for the updated_at field. diff --git a/backend/ent/schema/payment_order.go b/backend/ent/schema/payment_order.go index d25d1e5e17c3..c8ce5f1ed3eb 100644 --- a/backend/ent/schema/payment_order.go +++ b/backend/ent/schema/payment_order.go @@ -50,6 +50,10 @@ func (PaymentOrder) Fields() []ent.Field { field.Float("fee_rate"). SchemaType(map[string]string{dialect.Postgres: "decimal(10,4)"}). Default(0), + // 充值赠送额度(USD)。已计入 amount;单独记录用于订单展示与推广返利基数剔除。 + field.Float("bonus_amount"). + SchemaType(map[string]string{dialect.Postgres: "decimal(20,2)"}). + Default(0), field.String("recharge_code"). MaxLen(64), diff --git a/backend/ent/schema/user_platform_quota.go b/backend/ent/schema/user_platform_quota.go index 20c0b3aa14ad..663be0a6ecbc 100644 --- a/backend/ent/schema/user_platform_quota.go +++ b/backend/ent/schema/user_platform_quota.go @@ -42,7 +42,7 @@ func (UserPlatformQuota) Fields() []ent.Field { // 此处为 ent 构建期约束,需与 service.AllowedQuotaPlatforms 保持同步。 switch s { case "anthropic", "openai", "gemini", "antigravity", "grok", - "kimi", "zhipu", "deepseek", "minimax", "opencode_go": + "kimi", "zhipu", "deepseek", "minimax", "opencode_go", "typesafe": return nil default: return fmt.Errorf("platform %q is not allowed", s) diff --git a/backend/internal/domain/constants.go b/backend/internal/domain/constants.go index 99b11a28a36a..ffa70b218e4b 100644 --- a/backend/internal/domain/constants.go +++ b/backend/internal/domain/constants.go @@ -29,6 +29,7 @@ const ( PlatformZhipu = "zhipu" // 智谱 GLM (bigmodel) PlatformDeepseek = "deepseek" // DeepSeek PlatformMiniMax = "minimax" // MiniMax (M 系列) + PlatformTypeSafe = "typesafe" // TypeSafe AI System One (Jev) // PlatformOpenCodeGo 是 OpenCode 平台(账号类型 Zen 按量 / Go 订阅)。 // 值保持 opencode_go 以兼容已落库的分组、配额与 Composite 路由 CHECK。 PlatformOpenCodeGo = "opencode_go" diff --git a/backend/internal/handler/admin/account_handler.go b/backend/internal/handler/admin/account_handler.go index 98f8d2eb6d9d..62e98b958ecb 100644 --- a/backend/internal/handler/admin/account_handler.go +++ b/backend/internal/handler/admin/account_handler.go @@ -27,6 +27,7 @@ import ( "github.com/Wei-Shaw/sub2api/internal/pkg/openai" "github.com/Wei-Shaw/sub2api/internal/pkg/response" "github.com/Wei-Shaw/sub2api/internal/pkg/timezone" + "github.com/Wei-Shaw/sub2api/internal/pkg/typesafe" "github.com/Wei-Shaw/sub2api/internal/pkg/xai" "github.com/Wei-Shaw/sub2api/internal/service" @@ -3085,6 +3086,12 @@ func (h *AccountHandler) GetAvailableModels(c *gin.Context) { return } + // TypeSafe accounts serve only the native System One model. + if account.IsTypeSafe() { + response.Success(c, []claude.Model{{ID: typesafe.JevLatestModel, Type: "model", DisplayName: typesafe.JevLatestModel}}) + return + } + // Handle Claude/Anthropic accounts // For OAuth and Setup-Token accounts: return default models if account.IsOAuth() { diff --git a/backend/internal/handler/admin/account_handler_available_models_test.go b/backend/internal/handler/admin/account_handler_available_models_test.go index 28dc6d1492a4..3bcfce11d02b 100644 --- a/backend/internal/handler/admin/account_handler_available_models_test.go +++ b/backend/internal/handler/admin/account_handler_available_models_test.go @@ -526,3 +526,36 @@ func TestAccountHandlerSyncUpstreamModels_MetadataEnrichmentFailureReturnsWarnin require.Len(t, resp.Data.Warnings, 1) require.Equal(t, "upstream_model_metadata_incomplete", resp.Data.Warnings[0].Code) } + +func TestAccountHandlerGetAvailableModels_TypeSafeOnlyReturnsJev(t *testing.T) { + for _, credentials := range []map[string]any{ + {"api_key": "ts-secret"}, + {"api_key": "ts-secret", "model_mapping": map[string]any{"jev-latest": "jev-latest"}}, + } { + svc := &availableModelsAdminService{ + stubAdminService: newStubAdminService(), + account: service.Account{ + ID: 46, + Name: "typesafe", + Platform: service.PlatformTypeSafe, + Type: service.AccountTypeAPIKey, + Status: service.StatusActive, + Credentials: credentials, + }, + } + router := setupAvailableModelsRouter(svc) + + rec := httptest.NewRecorder() + router.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/api/v1/admin/accounts/46/models", nil)) + require.Equal(t, http.StatusOK, rec.Code) + + var resp struct { + Data []struct { + ID string `json:"id"` + } `json:"data"` + } + require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &resp)) + require.Len(t, resp.Data, 1) + require.Equal(t, "jev-latest", resp.Data[0].ID) + } +} diff --git a/backend/internal/handler/admin/channel_handler.go b/backend/internal/handler/admin/channel_handler.go index 376e08f88ecb..065aedb2f497 100644 --- a/backend/internal/handler/admin/channel_handler.go +++ b/backend/internal/handler/admin/channel_handler.go @@ -644,6 +644,7 @@ var platformToLiteLLMProvider = map[string]string{ service.PlatformDeepseek: "deepseek", service.PlatformMiniMax: "minimax", service.PlatformOpenCodeGo: "opencode-go", + service.PlatformTypeSafe: "typesafe", } // SyncPricingModels 返回 LiteLLM 定价目录中指定平台的最新模型列表 diff --git a/backend/internal/handler/admin/channel_handler_test.go b/backend/internal/handler/admin/channel_handler_test.go index 5678a9499be5..0c01631ecc87 100644 --- a/backend/internal/handler/admin/channel_handler_test.go +++ b/backend/internal/handler/admin/channel_handler_test.go @@ -565,7 +565,7 @@ func TestSyncPricingModels_ValidPlatform_EmptyService(t *testing.T) { svc := service.NewPricingService(nil, nil) router := setupSyncPricingModelsRouter(svc) - for _, platform := range []string{"anthropic", "openai", "gemini", "antigravity", "grok", "kimi", "zhipu", "deepseek", "minimax"} { + for _, platform := range []string{"anthropic", "openai", "gemini", "antigravity", "grok", "kimi", "zhipu", "deepseek", "minimax", "typesafe"} { req := httptest.NewRequest(http.MethodGet, "/channels/pricing/sync-models?platform="+platform, nil) w := httptest.NewRecorder() router.ServeHTTP(w, req) diff --git a/backend/internal/handler/admin/group_handler.go b/backend/internal/handler/admin/group_handler.go index 253addc6823d..15d5253f8669 100644 --- a/backend/internal/handler/admin/group_handler.go +++ b/backend/internal/handler/admin/group_handler.go @@ -184,7 +184,7 @@ func sanitizeUpdateGroupRequestForSimpleMode(req *UpdateGroupRequest) { 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 kimi zhipu deepseek minimax opencode_go composite"` + Platform string `json:"platform" binding:"omitempty,oneof=anthropic openai gemini antigravity grok kimi zhipu deepseek minimax opencode_go typesafe composite"` RateMultiplier float64 `json:"rate_multiplier"` IsExclusive bool `json:"is_exclusive"` SubscriptionType string `json:"subscription_type" binding:"omitempty,oneof=standard subscription"` @@ -259,7 +259,7 @@ type CreateGroupRequest struct { type UpdateGroupRequest struct { Name string `json:"name"` Description *string `json:"description"` - Platform string `json:"platform" binding:"omitempty,oneof=anthropic openai gemini antigravity grok kimi zhipu deepseek minimax opencode_go composite"` + Platform string `json:"platform" binding:"omitempty,oneof=anthropic openai gemini antigravity grok kimi zhipu deepseek minimax opencode_go typesafe composite"` RateMultiplier *float64 `json:"rate_multiplier"` IsExclusive *bool `json:"is_exclusive"` Status string `json:"status" binding:"omitempty,oneof=active inactive"` @@ -334,7 +334,7 @@ type UpdateGroupRequest struct { type CompositeRouteRequest struct { PublicModel string `json:"public_model" binding:"required"` MatchType string `json:"match_type" binding:"omitempty,oneof=exact prefix"` - TargetPlatform string `json:"target_platform" binding:"required,oneof=anthropic openai gemini antigravity grok kimi zhipu deepseek minimax opencode_go"` + TargetPlatform string `json:"target_platform" binding:"required,oneof=anthropic openai gemini antigravity grok kimi zhipu deepseek minimax opencode_go typesafe"` UpstreamModel string `json:"upstream_model"` Endpoint string `json:"endpoint" binding:"omitempty,oneof=any messages count_tokens responses chat_completions embeddings images gemini"` Priority int `json:"priority"` diff --git a/backend/internal/handler/admin/group_handler_platform_test.go b/backend/internal/handler/admin/group_handler_platform_test.go index 0c820868b482..d5d48845e1d3 100644 --- a/backend/internal/handler/admin/group_handler_platform_test.go +++ b/backend/internal/handler/admin/group_handler_platform_test.go @@ -28,6 +28,7 @@ func TestGroupPlatformBinding_AllowedPlatforms(t *testing.T) { allowed := []string{ "anthropic", "openai", "gemini", "antigravity", "grok", "kimi", "zhipu", "deepseek", "minimax", "opencode_go", "composite", + "typesafe", } for _, platform := range allowed { t.Run("create_"+platform, func(t *testing.T) { @@ -71,8 +72,8 @@ func TestGroupPlatformBinding_RejectsInvalidPlatforms(t *testing.T) { } } -func TestCompositeRouteTargetPlatform_AllowsCNProviders(t *testing.T) { - for _, platform := range []string{"kimi", "zhipu", "deepseek", "minimax", "opencode_go"} { +func TestCompositeRouteTargetPlatform_AllowsConcreteProviders(t *testing.T) { + for _, platform := range []string{"kimi", "zhipu", "deepseek", "minimax", "opencode_go", "typesafe"} { var req CompositeRouteRequest body := fmt.Sprintf(`{"public_model":"m","target_platform":%q}`, platform) require.NoError(t, bindGroupPlatformJSON(t, &req, body)) diff --git a/backend/internal/handler/admin/payment_handler.go b/backend/internal/handler/admin/payment_handler.go index 1749d196705a..3598635cb50f 100644 --- a/backend/internal/handler/admin/payment_handler.go +++ b/backend/internal/handler/admin/payment_handler.go @@ -125,6 +125,7 @@ type AdminPaymentOrderResult struct { Amount float64 `json:"amount"` PayAmount float64 `json:"pay_amount"` FeeRate float64 `json:"fee_rate"` + BonusAmount float64 `json:"bonus_amount"` Currency string `json:"currency"` RechargeCode string `json:"recharge_code,omitempty"` OutTradeNo string `json:"out_trade_no"` @@ -182,6 +183,7 @@ func sanitizeAdminPaymentOrderForResponse(order *dbent.PaymentOrder) *AdminPayme Amount: order.Amount, PayAmount: order.PayAmount, FeeRate: order.FeeRate, + BonusAmount: order.BonusAmount, Currency: service.PaymentOrderCurrency(order), RechargeCode: order.RechargeCode, OutTradeNo: order.OutTradeNo, diff --git a/backend/internal/handler/admin/setting_handler.go b/backend/internal/handler/admin/setting_handler.go index 5ac33ad3dbae..8f44ae1d8bfc 100644 --- a/backend/internal/handler/admin/setting_handler.go +++ b/backend/internal/handler/admin/setting_handler.go @@ -371,6 +371,9 @@ func (h *SettingHandler) GetSettings(c *gin.Context) { PaymentBalanceRechargeMultiplier: paymentCfg.BalanceRechargeMultiplier, PaymentSubscriptionUSDToCNYRate: paymentCfg.SubscriptionUSDToCNYRate, PaymentRechargeFeeRate: paymentCfg.RechargeFeeRate, + PaymentRechargeBonusTiers: rechargeBonusTiersToDTO(paymentCfg.RechargeBonusTiers), + PaymentRechargeBonusMode: rechargeBonusModeToDTO(paymentCfg.RechargeBonusMode), + PaymentRechargeBonusNotice: paymentCfg.RechargeBonusNotice, PaymentLoadBalanceStrat: paymentCfg.LoadBalanceStrategy, PaymentProductNamePrefix: paymentCfg.ProductNamePrefix, PaymentProductNameSuffix: paymentCfg.ProductNameSuffix, diff --git a/backend/internal/handler/admin/setting_handler_recharge_bonus.go b/backend/internal/handler/admin/setting_handler_recharge_bonus.go new file mode 100644 index 000000000000..334aedcef2bd --- /dev/null +++ b/backend/internal/handler/admin/setting_handler_recharge_bonus.go @@ -0,0 +1,33 @@ +package admin + +import ( + "github.com/Wei-Shaw/sub2api/internal/handler/dto" + "github.com/Wei-Shaw/sub2api/internal/service" +) + +// rechargeBonusTiersFromDTO 请求 nil 表示未携带该字段(保持现值);空数组表示清空阶梯。 +func rechargeBonusTiersFromDTO(items *[]dto.RechargeBonusTier) *[]service.RechargeBonusTier { + if items == nil { + return nil + } + out := make([]service.RechargeBonusTier, 0, len(*items)) + for _, item := range *items { + out = append(out, service.RechargeBonusTier{MinAmount: item.MinAmount, BonusPercent: item.BonusPercent}) + } + return &out +} + +// rechargeBonusModeToDTO 输出已归一化的模式(空/非法按 bonus)。 +func rechargeBonusModeToDTO(mode string) string { + normalized, _ := service.NormalizeRechargeBonusMode(mode) + return normalized +} + +// rechargeBonusTiersToDTO 始终返回非 nil 切片,空配置输出 []。 +func rechargeBonusTiersToDTO(items []service.RechargeBonusTier) []dto.RechargeBonusTier { + out := make([]dto.RechargeBonusTier, 0, len(items)) + for _, item := range items { + out = append(out, dto.RechargeBonusTier{MinAmount: item.MinAmount, BonusPercent: item.BonusPercent}) + } + return out +} diff --git a/backend/internal/handler/admin/setting_handler_update.go b/backend/internal/handler/admin/setting_handler_update.go index 520533771a4c..36878a185439 100644 --- a/backend/internal/handler/admin/setting_handler_update.go +++ b/backend/internal/handler/admin/setting_handler_update.go @@ -322,11 +322,15 @@ type UpdateSettingsRequest struct { PaymentBalanceRechargeMultiplier *float64 `json:"payment_balance_recharge_multiplier"` PaymentSubscriptionUSDToCNYRate *float64 `json:"payment_subscription_usd_to_cny_rate"` PaymentRechargeFeeRate *float64 `json:"payment_recharge_fee_rate"` - PaymentLoadBalanceStrat *string `json:"payment_load_balance_strategy"` - PaymentProductNamePrefix *string `json:"payment_product_name_prefix"` - PaymentProductNameSuffix *string `json:"payment_product_name_suffix"` - PaymentHelpImageURL *string `json:"payment_help_image_url"` - PaymentHelpText *string `json:"payment_help_text"` + // nil 表示不更新;空数组表示清空阶梯 + PaymentRechargeBonusTiers *[]dto.RechargeBonusTier `json:"payment_recharge_bonus_tiers"` + PaymentRechargeBonusMode *string `json:"payment_recharge_bonus_mode"` + PaymentRechargeBonusNotice *string `json:"payment_recharge_bonus_notice"` + PaymentLoadBalanceStrat *string `json:"payment_load_balance_strategy"` + PaymentProductNamePrefix *string `json:"payment_product_name_prefix"` + PaymentProductNameSuffix *string `json:"payment_product_name_suffix"` + PaymentHelpImageURL *string `json:"payment_help_image_url"` + PaymentHelpText *string `json:"payment_help_text"` // Cancel rate limit PaymentCancelRateLimitEnabled *bool `json:"payment_cancel_rate_limit_enabled"` @@ -2384,6 +2388,9 @@ func (h *SettingHandler) UpdateSettings(c *gin.Context) { BalanceRechargeMultiplier: req.PaymentBalanceRechargeMultiplier, SubscriptionUSDToCNYRate: req.PaymentSubscriptionUSDToCNYRate, RechargeFeeRate: req.PaymentRechargeFeeRate, + RechargeBonusTiers: rechargeBonusTiersFromDTO(req.PaymentRechargeBonusTiers), + RechargeBonusMode: req.PaymentRechargeBonusMode, + RechargeBonusNotice: req.PaymentRechargeBonusNotice, LoadBalanceStrategy: req.PaymentLoadBalanceStrat, ProductNamePrefix: req.PaymentProductNamePrefix, ProductNameSuffix: req.PaymentProductNameSuffix, @@ -2672,6 +2679,9 @@ func (h *SettingHandler) UpdateSettings(c *gin.Context) { PaymentBalanceRechargeMultiplier: updatedPaymentCfg.BalanceRechargeMultiplier, PaymentSubscriptionUSDToCNYRate: updatedPaymentCfg.SubscriptionUSDToCNYRate, PaymentRechargeFeeRate: updatedPaymentCfg.RechargeFeeRate, + PaymentRechargeBonusTiers: rechargeBonusTiersToDTO(updatedPaymentCfg.RechargeBonusTiers), + PaymentRechargeBonusMode: rechargeBonusModeToDTO(updatedPaymentCfg.RechargeBonusMode), + PaymentRechargeBonusNotice: updatedPaymentCfg.RechargeBonusNotice, PaymentLoadBalanceStrat: updatedPaymentCfg.LoadBalanceStrategy, PaymentProductNamePrefix: updatedPaymentCfg.ProductNamePrefix, PaymentProductNameSuffix: updatedPaymentCfg.ProductNameSuffix, @@ -2771,6 +2781,7 @@ func hasPaymentFields(req UpdateSettingsRequest) bool { req.PaymentEnabledTypes != nil || req.PaymentBalanceDisabled != nil || req.PaymentBalanceRechargeMultiplier != nil || req.PaymentSubscriptionUSDToCNYRate != nil || req.PaymentRechargeFeeRate != nil || + req.PaymentRechargeBonusTiers != nil || req.PaymentRechargeBonusMode != nil || req.PaymentRechargeBonusNotice != nil || req.PaymentLoadBalanceStrat != nil || req.PaymentProductNamePrefix != nil || req.PaymentProductNameSuffix != nil || req.PaymentHelpImageURL != nil || req.PaymentHelpText != nil || req.PaymentCancelRateLimitEnabled != nil || diff --git a/backend/internal/handler/auth_oauth_pending_flow_test.go b/backend/internal/handler/auth_oauth_pending_flow_test.go index 91a65d5f1190..dd8e1c03fe35 100644 --- a/backend/internal/handler/auth_oauth_pending_flow_test.go +++ b/backend/internal/handler/auth_oauth_pending_flow_test.go @@ -3593,3 +3593,20 @@ func (oauthPendingFlowTotpEncryptorStub) Encrypt(plaintext string) (string, erro func (oauthPendingFlowTotpEncryptorStub) Decrypt(ciphertext string) (string, error) { return ciphertext, nil } + +func (s *oauthPendingFlowEmailCacheStub) IncrVerificationCodeAttempts(_ context.Context, email string) (int, error) { + data := s.verificationCodes[email] + if data == nil { + return 0, errors.New("verification code not found") + } + data.Attempts++ + return data.Attempts, nil +} + +func (s *oauthPendingFlowEmailCacheStub) IncrNotifyVerifyCodeAttempts(context.Context, string) (int, error) { + return 0, errors.New("notify verification code not found") +} + +func (s *oauthPendingFlowEmailCacheStub) ConsumePasswordResetToken(context.Context, string, string) (bool, error) { + return false, nil +} diff --git a/backend/internal/handler/dto/recharge_bonus_tiers.go b/backend/internal/handler/dto/recharge_bonus_tiers.go new file mode 100644 index 000000000000..21731566a88f --- /dev/null +++ b/backend/internal/handler/dto/recharge_bonus_tiers.go @@ -0,0 +1,7 @@ +package dto + +// RechargeBonusTier 充值赠送档位:余额充值支付金额 ≥ MinAmount 时,在到账基数上赠送 BonusPercent%。 +type RechargeBonusTier struct { + MinAmount float64 `json:"min_amount"` + BonusPercent float64 `json:"bonus_percent"` +} diff --git a/backend/internal/handler/dto/settings.go b/backend/internal/handler/dto/settings.go index 7fd50d5f0d11..8a36be8d4f99 100644 --- a/backend/internal/handler/dto/settings.go +++ b/backend/internal/handler/dto/settings.go @@ -288,11 +288,15 @@ type SystemSettings struct { PaymentBalanceRechargeMultiplier float64 `json:"payment_balance_recharge_multiplier"` PaymentSubscriptionUSDToCNYRate float64 `json:"payment_subscription_usd_to_cny_rate"` PaymentRechargeFeeRate float64 `json:"payment_recharge_fee_rate"` - PaymentLoadBalanceStrat string `json:"payment_load_balance_strategy"` - PaymentProductNamePrefix string `json:"payment_product_name_prefix"` - PaymentProductNameSuffix string `json:"payment_product_name_suffix"` - PaymentHelpImageURL string `json:"payment_help_image_url"` - PaymentHelpText string `json:"payment_help_text"` + // 充值赠送阶梯与活动文案 + PaymentRechargeBonusTiers []RechargeBonusTier `json:"payment_recharge_bonus_tiers"` + PaymentRechargeBonusMode string `json:"payment_recharge_bonus_mode"` + PaymentRechargeBonusNotice string `json:"payment_recharge_bonus_notice"` + PaymentLoadBalanceStrat string `json:"payment_load_balance_strategy"` + PaymentProductNamePrefix string `json:"payment_product_name_prefix"` + PaymentProductNameSuffix string `json:"payment_product_name_suffix"` + PaymentHelpImageURL string `json:"payment_help_image_url"` + PaymentHelpText string `json:"payment_help_text"` // Cancel rate limit PaymentCancelRateLimitEnabled bool `json:"payment_cancel_rate_limit_enabled"` diff --git a/backend/internal/handler/endpoint.go b/backend/internal/handler/endpoint.go index 2115c2d7c010..82d8279e9046 100644 --- a/backend/internal/handler/endpoint.go +++ b/backend/internal/handler/endpoint.go @@ -16,6 +16,7 @@ import ( const ( EndpointMessages = "/v1/messages" + EndpointSystemOne = "/v1/systemone" EndpointChatCompletions = "/v1/chat/completions" EndpointEmbeddings = "/v1/embeddings" EndpointAlphaSearch = "/v1/alpha/search" @@ -94,6 +95,8 @@ func NormalizeInboundEndpoint(path string) string { return EndpointChatCompletions case strings.Contains(path, EndpointMessages): return EndpointMessages + case strings.Contains(path, EndpointSystemOne): + return EndpointSystemOne case strings.Contains(path, EndpointImagesGenerations) || strings.Contains(path, "/images/generations"): return EndpointImagesGenerations case strings.Contains(path, EndpointImagesEdits) || strings.Contains(path, "/images/edits"): @@ -223,6 +226,9 @@ func DeriveUpstreamEndpoint(inbound, rawRequestPath, platform string) string { case service.PlatformAnthropic: return EndpointMessages + case service.PlatformTypeSafe: + return EndpointSystemOne + case service.PlatformGemini: return EndpointGeminiModels diff --git a/backend/internal/handler/endpoint_test.go b/backend/internal/handler/endpoint_test.go index aaa7e9f8fd9d..6941e5637411 100644 --- a/backend/internal/handler/endpoint_test.go +++ b/backend/internal/handler/endpoint_test.go @@ -23,6 +23,7 @@ func TestNormalizeInboundEndpoint(t *testing.T) { }{ // Direct canonical paths. {"/v1/messages", EndpointMessages}, + {"/v1/systemone", EndpointSystemOne}, {"/v1/chat/completions", EndpointChatCompletions}, {"/v1/embeddings", EndpointEmbeddings}, {"/v1/alpha/search", EndpointAlphaSearch}, @@ -97,6 +98,7 @@ func TestDeriveUpstreamEndpoint(t *testing.T) { }{ // Anthropic. {"anthropic messages", EndpointMessages, "/v1/messages", service.PlatformAnthropic, EndpointMessages}, + {"typesafe system one", EndpointSystemOne, "/v1/systemone", service.PlatformTypeSafe, EndpointSystemOne}, // Gemini. {"gemini models", EndpointGeminiModels, "/v1beta/models/gemini:gen", service.PlatformGemini, EndpointGeminiModels}, diff --git a/backend/internal/handler/gateway_handler.go b/backend/internal/handler/gateway_handler.go index a09f47911014..177e14b26daf 100644 --- a/backend/internal/handler/gateway_handler.go +++ b/backend/internal/handler/gateway_handler.go @@ -24,6 +24,7 @@ import ( "github.com/Wei-Shaw/sub2api/internal/pkg/logger" "github.com/Wei-Shaw/sub2api/internal/pkg/openai" "github.com/Wei-Shaw/sub2api/internal/pkg/timezone" + "github.com/Wei-Shaw/sub2api/internal/pkg/typesafe" "github.com/Wei-Shaw/sub2api/internal/pkg/xai" "github.com/Wei-Shaw/sub2api/internal/securityaudit" middleware2 "github.com/Wei-Shaw/sub2api/internal/server/middleware" @@ -222,6 +223,9 @@ func (h *GatewayHandler) Messages(c *gin.Context) { h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "Model is not supported by composite groups") return } + if rejectSystemOneOnlyPlatform(c, apiKey, h.errorResponse) { + return + } if decision := h.checkSecurityAudit(c, reqLog, apiKey, subject, service.ContentModerationProtocolAnthropicMessages, reqModel, body); decision != nil && !decision.AllowNextStage { h.anthropicSecurityAuditError(c, decision) @@ -1157,7 +1161,7 @@ func (h *GatewayHandler) Models(c *gin.Context) { } if platform == service.PlatformComposite { - availableModels := h.compositeAvailableModels(c.Request.Context(), groupID) + availableModels := h.compositeAvailableModels(c.Request.Context(), groupID, true) if apiKey != nil && apiKey.Group != nil && apiKey.Group.ModelAllowlistEnabled() { source := availableModels if len(source) == 0 { @@ -1201,6 +1205,10 @@ func (h *GatewayHandler) Models(c *gin.Context) { writeGrokModelsList(c, xai.DefaultModelIDs()) return } + if platform == service.PlatformTypeSafe { + writeModelsList(c, platform, []string{typesafe.JevLatestModel}) + return + } writeModelsListResponse(c, claude.DefaultModels) } @@ -1253,7 +1261,7 @@ func (h *GatewayHandler) codexModelIDsForGroup(ctx context.Context, group *servi platform = group.Platform } if platform == service.PlatformComposite { - availableModels := h.compositeAvailableModels(ctx, groupID) + availableModels := h.compositeAvailableModels(ctx, groupID, false) fallbackModels := defaultCodexModelIDsForPlatform(service.PlatformComposite) if group.ModelAllowlistEnabled() { source := availableModels @@ -1279,14 +1287,20 @@ func (h *GatewayHandler) codexModelIDsForGroup(ctx context.Context, group *servi return fallbackModels } -func (h *GatewayHandler) compositeAvailableModels(ctx context.Context, groupID *int64) []string { +// compositeAvailableModels lists the models the composite group can serve. +// includeSystemOne adds TypeSafe models, which only work through /v1/systemone; +// LLM client catalogs (Codex) must exclude them. +func (h *GatewayHandler) compositeAvailableModels(ctx context.Context, groupID *int64, includeSystemOne bool) []string { if h == nil || h.gatewayService == nil { return nil } seen := make(map[string]struct{}) models := make([]string, 0) schedulablePlatforms := h.gatewayService.GetSchedulablePlatforms(ctx, groupID) - for _, platform := range []string{service.PlatformAnthropic, service.PlatformGemini, service.PlatformOpenAI, service.PlatformAntigravity, service.PlatformGrok, service.PlatformKimi, service.PlatformZhipu, service.PlatformDeepseek, service.PlatformMiniMax, service.PlatformOpenCodeGo} { + for _, platform := range []string{service.PlatformAnthropic, service.PlatformGemini, service.PlatformOpenAI, service.PlatformAntigravity, service.PlatformGrok, service.PlatformKimi, service.PlatformZhipu, service.PlatformDeepseek, service.PlatformMiniMax, service.PlatformOpenCodeGo, service.PlatformTypeSafe} { + if platform == service.PlatformTypeSafe && !includeSystemOne { + continue + } platformModels := h.gatewayService.GetAvailableModels(ctx, groupID, platform) if len(platformModels) == 0 { // CN 供应商没有静态默认模型列表(defaultModelIDsForPlatform 的 @@ -1470,9 +1484,14 @@ func defaultModelIDsForPlatform(platform string) []string { return xai.DefaultModelIDs() case service.PlatformOpenCodeGo: return service.DefaultOpenCodeGoModelIDs() + case service.PlatformTypeSafe: + return []string{"jev-latest"} case service.PlatformComposite: ids := make([]string, 0) seen := make(map[string]struct{}) + // TypeSafe is deliberately absent: jev-latest only works through + // /v1/systemone, so the static fallback never advertises it to LLM + // clients. compositeAvailableModels lists it when the group can serve it. for _, concretePlatform := range []string{service.PlatformAnthropic, service.PlatformGemini, service.PlatformOpenAI, service.PlatformAntigravity, service.PlatformGrok, service.PlatformKimi, service.PlatformZhipu, service.PlatformDeepseek, service.PlatformMiniMax, service.PlatformOpenCodeGo} { for _, id := range defaultModelIDsForPlatform(concretePlatform) { if _, ok := seen[id]; ok { @@ -2161,6 +2180,9 @@ func (h *GatewayHandler) CountTokens(c *gin.Context) { h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "Model is not supported by composite groups") return } + if rejectSystemOneOnlyPlatform(c, apiKey, h.errorResponse) { + return + } setOpsRequestContext(c, parsedReq.Model, parsedReq.Stream) setOpsEndpointContext(c, "", int16(service.RequestTypeFromLegacy(parsedReq.Stream, false))) diff --git a/backend/internal/handler/gateway_handler_chat_completions.go b/backend/internal/handler/gateway_handler_chat_completions.go index 868317d4e505..d956ed680d76 100644 --- a/backend/internal/handler/gateway_handler_chat_completions.go +++ b/backend/internal/handler/gateway_handler_chat_completions.go @@ -82,6 +82,9 @@ func (h *GatewayHandler) ChatCompletions(c *gin.Context) { h.chatCompletionsErrorResponse(c, http.StatusBadRequest, "invalid_request_error", "Model is not supported by composite groups") return } + if rejectSystemOneOnlyPlatform(c, apiKey, h.chatCompletionsErrorResponse) { + return + } reqStream, ok := parseOpenAICompatibleStream(body) if !ok { h.chatCompletionsErrorResponse(c, http.StatusBadRequest, "invalid_request_error", invalidStreamFieldTypeMessage) diff --git a/backend/internal/handler/gateway_handler_responses.go b/backend/internal/handler/gateway_handler_responses.go index 54961e6c8b94..88da0009517a 100644 --- a/backend/internal/handler/gateway_handler_responses.go +++ b/backend/internal/handler/gateway_handler_responses.go @@ -82,6 +82,9 @@ func (h *GatewayHandler) Responses(c *gin.Context) { h.responsesErrorResponse(c, http.StatusBadRequest, "invalid_request_error", "Model is not supported by composite groups") return } + if rejectSystemOneOnlyPlatform(c, apiKey, h.responsesErrorResponse) { + return + } reqStream, ok := parseOpenAICompatibleStream(body) if !ok { h.responsesErrorResponse(c, http.StatusBadRequest, "invalid_request_error", invalidStreamFieldTypeMessage) diff --git a/backend/internal/handler/gateway_models_retrieve_test.go b/backend/internal/handler/gateway_models_retrieve_test.go index 26a5b6107ac8..fc590a35dfd2 100644 --- a/backend/internal/handler/gateway_models_retrieve_test.go +++ b/backend/internal/handler/gateway_models_retrieve_test.go @@ -27,7 +27,7 @@ func requestModelForTest(h *GatewayHandler, group *service.Group, modelID, etag func TestRetrieveModelMatchesVisibleCatalogue(t *testing.T) { gin.SetMode(gin.TestMode) - for _, platform := range []string{service.PlatformOpenAI, service.PlatformAnthropic, service.PlatformGemini, service.PlatformGrok, service.PlatformComposite} { + for _, platform := range []string{service.PlatformOpenAI, service.PlatformAnthropic, service.PlatformGemini, service.PlatformGrok, service.PlatformTypeSafe, service.PlatformComposite} { for _, mapped := range []bool{false, true} { name := platform + "/fallback" if mapped { diff --git a/backend/internal/handler/gateway_models_test.go b/backend/internal/handler/gateway_models_test.go index 0b49ca5f3586..d2fdb5cfa084 100644 --- a/backend/internal/handler/gateway_models_test.go +++ b/backend/internal/handler/gateway_models_test.go @@ -932,6 +932,10 @@ func TestDefaultCodexModelIDsForPlatform_DeepSeekUsesDeepSeekModels(t *testing.T require.Equal(t, defaultModelIDsForPlatform(service.PlatformAnthropic), defaultCodexModelIDsForPlatform(service.PlatformAnthropic)) } +func TestDefaultModelIDsForPlatform_TypeSafeUsesJev(t *testing.T) { + require.Equal(t, []string{"jev-latest"}, defaultModelIDsForPlatform(service.PlatformTypeSafe)) +} + func TestGatewayCodexModels_DeepSeekWithoutMappingUsesDeepSeekDefaults(t *testing.T) { gin.SetMode(gin.TestMode) const groupID int64 = 130 @@ -1522,3 +1526,47 @@ func TestGatewayModels_GPT6SolLunaDiscoveryRespectsGroupAndAccountRestrictions(t }) } } + +// Scenario: jev-latest only works through /v1/systemone, so Composite groups list +// it in /v1/models only when they can serve it, and never in the Codex manifest. +func TestGatewayModels_CompositeTypeSafeListingScope(t *testing.T) { + gin.SetMode(gin.TestMode) + + groupID := int64(66) + h := newGatewayModelsHandlerForTest(&gatewayModelsAccountRepoStub{ + byGroup: map[int64][]service.Account{ + groupID: {{ID: 1, Platform: service.PlatformAnthropic}, {ID: 2, Platform: service.PlatformTypeSafe, Type: service.AccountTypeAPIKey}}, + }, + }) + newContext := func(path string) (*gin.Context, *httptest.ResponseRecorder) { + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodGet, path, nil) + c.Set(string(middleware2.ContextKeyAPIKey), &service.APIKey{ + Group: &service.Group{ID: groupID, Platform: service.PlatformComposite}, + }) + return c, rec + } + + c, rec := newContext("/v1/models") + h.Models(c) + require.Equal(t, http.StatusOK, rec.Code) + var models gatewayModelsResponseForTest + require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &models)) + require.Contains(t, modelIDsForTest(models.Data), "jev-latest") + require.Contains(t, modelIDsForTest(models.Data), "claude-opus-4-6") + + c, rec = newContext("/models?client_version=0.147.0") + h.CodexModels(c) + require.Equal(t, http.StatusOK, rec.Code) + var manifest codexModelsResponseForTest + require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &manifest)) + slugs := codexModelSlugsForTest(manifest.Models) + require.Contains(t, slugs, "claude-opus-4-6") + require.NotContains(t, slugs, "jev-latest") +} + +func TestDefaultModelIDsForPlatform_CompositeFallbackExcludesTypeSafe(t *testing.T) { + require.NotContains(t, defaultModelIDsForPlatform(service.PlatformComposite), "jev-latest") + require.NotContains(t, defaultCodexModelIDsForPlatform(service.PlatformComposite), "jev-latest") +} diff --git a/backend/internal/handler/gateway_systemone.go b/backend/internal/handler/gateway_systemone.go new file mode 100644 index 000000000000..a366a4b1119c --- /dev/null +++ b/backend/internal/handler/gateway_systemone.go @@ -0,0 +1,275 @@ +package handler + +import ( + "context" + "errors" + "net/http" + "strconv" + "time" + + "github.com/Wei-Shaw/sub2api/internal/pkg/ip" + "github.com/Wei-Shaw/sub2api/internal/pkg/logger" + "github.com/Wei-Shaw/sub2api/internal/pkg/typesafe" + middleware2 "github.com/Wei-Shaw/sub2api/internal/server/middleware" + "github.com/Wei-Shaw/sub2api/internal/service" + "github.com/gin-gonic/gin" + "go.uber.org/zap" +) + +const systemOneOnlyPlatformMessage = "TypeSafe models are only available through POST /v1/systemone" + +// rejectSystemOneOnlyPlatform stops TypeSafe traffic from entering a non System +// One protocol chain. TypeSafe accounts only speak the native System One +// protocol; letting them reach the Anthropic/OpenAI converters would send +// foreign payloads (and the account key) to the wrong upstream path and feed +// the resulting auth failures back into account state. +func rejectSystemOneOnlyPlatform(c *gin.Context, apiKey *service.APIKey, writeError func(*gin.Context, int, string, string)) bool { + if _, forced := middleware2.GetForcePlatformFromContext(c); forced { + return false + } + if effectiveAPIKeyPlatform(c, apiKey) != service.PlatformTypeSafe { + return false + } + service.MarkOpsClientBusinessLimited(c, service.OpsClientBusinessLimitedReasonLocalFeatureGate) + writeError(c, http.StatusNotFound, "not_found_error", systemOneOnlyPlatformMessage) + return true +} + +// SystemOne proxies TypeSafe's native, non-streaming System One protocol. +func (h *GatewayHandler) SystemOne(c *gin.Context) { + requestStart := time.Now() + apiKey, ok := middleware2.GetAPIKeyFromContext(c) + if !ok || apiKey == nil || apiKey.Group == nil { + h.errorResponse(c, http.StatusUnauthorized, "authentication_error", "Invalid API key") + return + } + subject, ok := middleware2.GetAuthSubjectFromContext(c) + if !ok { + h.errorResponse(c, http.StatusInternalServerError, "api_error", "User context not found") + return + } + reqLog := requestLogger(c, "handler.gateway.systemone", + zap.Int64("user_id", subject.UserID), + zap.Int64("api_key_id", apiKey.ID), + zap.Any("group_id", apiKey.GroupID), + ) + + body, err := readLenientJSONRequestBodyWithPrealloc(c.Request, h.cfg) + if err != nil { + if maxErr, ok := extractMaxBytesError(err); ok { + h.errorResponse(c, http.StatusRequestEntityTooLarge, "invalid_request_error", buildBodyTooLargeMessage(maxErr.Limit)) + return + } + h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "Failed to read request body") + return + } + model, err := typesafe.ValidateSystemOneRequest(body) + if err != nil { + h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", err.Error()) + return + } + ensureCompositeTargetPlatform(c, apiKey, model) + if apiKey.Group.Platform != service.PlatformTypeSafe && + (apiKey.Group.Platform != service.PlatformComposite || !compositeTargetPlatformAllowed(c, apiKey, model, service.PlatformTypeSafe)) { + h.errorResponse(c, http.StatusNotFound, "not_found_error", "System One is only available for TypeSafe and compatible Composite groups") + return + } + if decision := h.checkSecurityAudit(c, reqLog, apiKey, subject, service.ContentModerationProtocolTypeSafeSystemOne, model, body); decision != nil && !decision.AllowNextStage { + h.anthropicSecurityAuditError(c, decision) + return + } + + setOpsRequestContext(c, model, false) + setOpsEndpointContext(c, "", int16(service.RequestTypeSync)) + service.SetOpsLatencyMs(c, service.OpsAuthLatencyMsKey, time.Since(requestStart).Milliseconds()) + pricingCtx, pricingAt := service.WithGatewayTokenRequestPricing(c.Request.Context()) + c.Request = c.Request.WithContext(pricingCtx) + channelMapping, _ := h.gatewayService.ResolveChannelMappingAndRestrict(c.Request.Context(), apiKey.GroupID, model) + subscription, _ := middleware2.GetSubscriptionFromContext(c) + + streamStarted := false + userRelease, err := h.concurrencyHelper.AcquireUserSlotWithWait(c, subject.UserID, subject.Concurrency, apiKey.ID, apiKey.ConcurrencyLimit, false, &streamStarted) + if err != nil { + reqLog.Warn("systemone.user_slot_acquire_failed", zap.Error(err)) + h.handleConcurrencyError(c, err, "user", false) + return + } + // 在请求结束或 Context 取消时确保释放槽位,避免客户端断开造成泄漏。 + userRelease = wrapReleaseOnDone(c.Request.Context(), userRelease) + if userRelease != nil { + defer userRelease() + } + if err := h.billingCacheService.CheckBillingEligibility(c.Request.Context(), apiKey.User, apiKey, apiKey.Group, subscription, service.QuotaPlatform(c.Request.Context(), apiKey)); err != nil { + reqLog.Info("systemone.billing_eligibility_check_failed", zap.Error(err)) + status, code, message, retryAfter := billingErrorDetails(err) + if retryAfter > 0 { + c.Header("Retry-After", strconv.Itoa(retryAfter)) + } + h.errorResponse(c, status, code, message) + return + } + + fs := NewFailoverState(h.maxAccountSwitches, false) + for { + if failoverClientGone(c) { + return + } + selection, err := h.gatewayService.SelectAccountWithLoadAwareness(c.Request.Context(), apiKey.GroupID, "", model, fs.FailedAccountIDs, "", subject.UserID) + if err == nil && (selection == nil || selection.Account == nil) { + err = service.ErrNoAvailableAccounts + } + if err != nil { + if failoverClientGone(c) { + reqLog.Info("systemone.account_select_aborted_client_disconnected", zap.Error(err)) + return + } + if len(fs.FailedAccountIDs) == 0 { + cls := classifyNoAccountErrorFromGin(c, h.gatewayService, apiKey, model, model, service.PlatformTypeSafe) + cls = classifySelectionFailureError(err, cls) + if !cls.ModelNotFound { + markOpsRoutingCapacityLimitedIfNoAvailable(c, err) + } + reqLog.Warn("systemone.select_account_no_available", zap.Bool("model_not_found", cls.ModelNotFound), zap.Error(err)) + h.errorResponse(c, cls.Status, cls.ErrType, cls.Message) + return + } + switch fs.HandleSelectionExhausted(c.Request.Context()) { + case FailoverContinue: + continue + case FailoverCanceled: + failoverClientGone(c) + return + default: + if fs.LastFailoverErr != nil { + h.handleFailoverExhausted(c, fs.LastFailoverErr, service.PlatformTypeSafe, false) + } else { + h.handleFailoverExhaustedSimple(c, http.StatusBadGateway, false) + } + return + } + } + account := selection.Account + + accountRelease := selection.ReleaseFunc + if !selection.Acquired { + if selection.WaitPlan == nil { + markOpsRoutingCapacityLimited(c) + h.errorResponse(c, http.StatusServiceUnavailable, "api_error", "No available accounts") + return + } + accountRelease, err = h.concurrencyHelper.AcquireAccountSlotWithWaitTimeout(c, account.ID, selection.WaitPlan.MaxConcurrency, selection.WaitPlan.Timeout, false, &streamStarted) + if err != nil { + reqLog.Warn("systemone.account_slot_acquire_failed", zap.Int64("account_id", account.ID), zap.Error(err)) + h.handleConcurrencyError(c, err, "account", false) + return + } + } + // 准入终检:与其他网关入口一致,利润控制否决的账号不得承接本次请求。 + admissionCtx := service.ContextWithSelectionProfitGate(c.Request.Context(), selection) + latest, vetoed, reason := h.gatewayService.GatewayProfitControlVetoLatest(admissionCtx, account) + if vetoed { + if accountRelease != nil { + accountRelease() + } + reqLog.Debug("systemone.account_slot_profit_vetoed", zap.Int64("account_id", account.ID), zap.String("reason", reason)) + if fs.RecordProfitVeto(account.ID) == FailoverExhausted { + reqLog.Warn("systemone.profit_veto_attempts_exhausted", zap.Int("profit_veto_count", fs.ProfitVetoCount())) + h.errorResponse(c, http.StatusServiceUnavailable, "api_error", profitVetoExhaustedMessage) + return + } + continue + } + account = latest + accountRelease = wrapReleaseOnDone(c.Request.Context(), accountRelease) + setOpsSelectedAccount(c, account.ID, account.Platform) + service.SetOpsUpstreamModel(c, model) + + forwardStart := time.Now() + result, forwardErr := h.gatewayService.ForwardSystemOne(c.Request.Context(), c, account, body) + if accountRelease != nil { + accountRelease() + } + service.SetOpsLatencyMs(c, service.OpsResponseLatencyMsKey, time.Since(forwardStart).Milliseconds()) + + if forwardErr != nil { + var failoverErr *service.UpstreamFailoverError + if errors.As(forwardErr, &failoverErr) { + switch fs.HandleFailoverError(c.Request.Context(), h.gatewayService, account.ID, account.Platform, account.GetPoolModeRetryCount(), failoverErr) { + case FailoverContinue: + reqLog.Warn("systemone.upstream_failover_switching", + zap.Int64("account_id", account.ID), + zap.Int("upstream_status", failoverErr.StatusCode), + zap.Int("switch_count", fs.SwitchCount), + ) + continue + case FailoverExhausted: + h.handleFailoverExhausted(c, fs.LastFailoverErr, service.PlatformTypeSafe, false) + return + case FailoverCanceled: + failoverClientGone(c) + return + } + } + if failoverClientGone(c) { + return + } + var upstreamErr *service.SystemOneUpstreamError + if errors.As(forwardErr, &upstreamErr) { + status := upstreamErr.StatusCode + if !service.IsSystemOneRequestErrorStatus(status) { + status = http.StatusBadGateway + } + h.errorResponse(c, status, "upstream_error", "TypeSafe rejected the request") + return + } + reqLog.Warn("systemone.forward_failed", zap.Int64("account_id", account.ID), zap.Error(forwardErr)) + if errors.Is(forwardErr, typesafe.ErrSystemOneResponseTooLarge) { + h.errorResponse(c, http.StatusBadGateway, "upstream_error", "TypeSafe response exceeds the gateway size limit") + return + } + h.errorResponse(c, http.StatusBadGateway, "upstream_error", "TypeSafe upstream request failed") + return + } + + c.Data(result.StatusCode, result.ContentType, result.Body) + h.recordSystemOneUsage(c, apiKey, account, subscription, channelMapping, model, body, result, subject.UserID, pricingAt) + return + } +} + +func (h *GatewayHandler) recordSystemOneUsage(c *gin.Context, apiKey *service.APIKey, account *service.Account, subscription *service.UserSubscription, mapping service.ChannelMappingResult, model string, body []byte, result *service.SystemOneForwardResult, userID int64, pricingAt time.Time) { + userAgent := c.GetHeader("User-Agent") + clientIP := ip.GetClientIP(c) + inboundEndpoint := GetInboundEndpoint(c) + upstreamEndpoint := GetUpstreamEndpoint(c, account.Platform) + quotaPlatform := service.QuotaPlatform(c.Request.Context(), apiKey) + sessionID := service.ExtractClientSessionID(c) + requestPayloadHash := service.HashUsageRequestPayload(body) + + h.submitMandatoryUsageRecordTask(c.Request.Context(), func(ctx context.Context) { + if err := h.gatewayService.RecordUsage(ctx, &service.RecordUsageInput{ + Result: &result.ForwardResult, + APIKey: apiKey, + User: apiKey.User, + Account: account, + Subscription: subscription, + PricingAt: pricingAt, + InboundEndpoint: inboundEndpoint, + UpstreamEndpoint: upstreamEndpoint, + UserAgent: userAgent, + IPAddress: clientIP, + SessionID: sessionID, + RequestPayloadHash: requestPayloadHash, + APIKeyService: h.apiKeyService, + QuotaPlatform: quotaPlatform, + ChannelUsageFields: clientRequestedUsageFields(c, mapping, model, result.UpstreamModel), + }); err != nil { + logger.L().With( + zap.String("component", "handler.gateway.systemone"), + zap.Int64("user_id", userID), + zap.Int64("api_key_id", apiKey.ID), + zap.Int64("account_id", account.ID), + ).Error("systemone.record_usage_failed", zap.Error(err)) + } + }) +} diff --git a/backend/internal/handler/gateway_systemone_test.go b/backend/internal/handler/gateway_systemone_test.go new file mode 100644 index 000000000000..084b37ed9fe9 --- /dev/null +++ b/backend/internal/handler/gateway_systemone_test.go @@ -0,0 +1,109 @@ +package handler + +import ( + "net/http" + "net/http/httptest" + "strings" + "testing" + + "github.com/Wei-Shaw/sub2api/internal/config" + middleware2 "github.com/Wei-Shaw/sub2api/internal/server/middleware" + "github.com/Wei-Shaw/sub2api/internal/service" + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" +) + +const validSystemOneHandlerBody = `{"model":"jev-latest","state":"sample","questions":{"q":{"type":"noul","instructions":"Evaluate"}}}` + +func newSystemOneHandlerContext(body string) (*gin.Context, *httptest.ResponseRecorder) { + gin.SetMode(gin.TestMode) + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/systemone", strings.NewReader(body)) + c.Request.Header.Set("Content-Type", "application/json") + return c, recorder +} + +func TestSystemOneRequiresAuthentication(t *testing.T) { + c, recorder := newSystemOneHandlerContext(validSystemOneHandlerBody) + (&GatewayHandler{}).SystemOne(c) + require.Equal(t, http.StatusUnauthorized, recorder.Code) + require.Contains(t, recorder.Body.String(), "authentication_error") +} + +func TestSystemOneRejectsNonTypeSafeGroupBeforeScheduling(t *testing.T) { + c, recorder := newSystemOneHandlerContext(validSystemOneHandlerBody) + groupID := int64(3) + c.Set(string(middleware2.ContextKeyAPIKey), &service.APIKey{ID: 4, UserID: 5, GroupID: &groupID, Group: &service.Group{ID: groupID, Platform: service.PlatformOpenAI}}) + c.Set(string(middleware2.ContextKeyUser), middleware2.AuthSubject{UserID: 5, Concurrency: 1}) + + (&GatewayHandler{cfg: &config.Config{Gateway: config.GatewayConfig{MaxBodySize: 1 << 20}}}).SystemOne(c) + require.Equal(t, http.StatusNotFound, recorder.Code) + require.Contains(t, recorder.Body.String(), "only available for TypeSafe") +} + +func newTypeSafeGroupContext(t *testing.T, path, body, groupPlatform string) (*gin.Context, *httptest.ResponseRecorder) { + t.Helper() + c, recorder := newSystemOneHandlerContext(body) + c.Request = httptest.NewRequest(http.MethodPost, path, strings.NewReader(body)) + c.Request.Header.Set("Content-Type", "application/json") + groupID := int64(9) + c.Set(string(middleware2.ContextKeyAPIKey), &service.APIKey{ID: 4, UserID: 5, GroupID: &groupID, Group: &service.Group{ID: groupID, Platform: groupPlatform}}) + c.Set(string(middleware2.ContextKeyUser), middleware2.AuthSubject{UserID: 5, Concurrency: 1}) + return c, recorder +} + +func TestRejectSystemOneOnlyPlatform(t *testing.T) { + write := func(c *gin.Context, status int, errType, message string) { + c.JSON(status, gin.H{"type": errType, "message": message}) + } + + c, recorder := newTypeSafeGroupContext(t, "/v1/messages", `{}`, service.PlatformTypeSafe) + require.True(t, rejectSystemOneOnlyPlatform(c, mustAPIKey(t, c), write)) + require.Equal(t, http.StatusNotFound, recorder.Code) + require.Contains(t, recorder.Body.String(), systemOneOnlyPlatformMessage) + + c, _ = newTypeSafeGroupContext(t, "/v1/messages", `{}`, service.PlatformComposite) + require.False(t, rejectSystemOneOnlyPlatform(c, mustAPIKey(t, c), write)) + c.Request = c.Request.WithContext(service.WithResolvedTargetPlatform(c.Request.Context(), service.PlatformTypeSafe)) + require.True(t, rejectSystemOneOnlyPlatform(c, mustAPIKey(t, c), write)) + + c, _ = newTypeSafeGroupContext(t, "/v1/messages", `{}`, service.PlatformAnthropic) + require.False(t, rejectSystemOneOnlyPlatform(c, mustAPIKey(t, c), write)) + + c, _ = newTypeSafeGroupContext(t, "/antigravity/v1/messages", `{}`, service.PlatformTypeSafe) + c.Set(string(middleware2.ContextKeyForcePlatform), service.PlatformAntigravity) + require.False(t, rejectSystemOneOnlyPlatform(c, mustAPIKey(t, c), write)) +} + +func mustAPIKey(t *testing.T, c *gin.Context) *service.APIKey { + t.Helper() + apiKey, ok := middleware2.GetAPIKeyFromContext(c) + require.True(t, ok) + return apiKey +} + +func TestTypeSafeGroupsRejectNonSystemOneProtocolsBeforeScheduling(t *testing.T) { + h := &GatewayHandler{ + cfg: &config.Config{Gateway: config.GatewayConfig{MaxBodySize: 1 << 20}}, + gatewayService: &service.GatewayService{}, + } + for _, tc := range []struct { + name string + path string + body string + handler func(*gin.Context) + }{ + {"messages", "/v1/messages", `{"model":"jev-latest","max_tokens":1,"messages":[{"role":"user","content":"hi"}]}`, h.Messages}, + {"count_tokens", "/v1/messages/count_tokens", `{"model":"jev-latest","messages":[{"role":"user","content":"hi"}]}`, h.CountTokens}, + {"chat_completions", "/v1/chat/completions", `{"model":"jev-latest","messages":[{"role":"user","content":"hi"}]}`, h.ChatCompletions}, + {"responses", "/v1/responses", `{"model":"jev-latest","input":"hi"}`, h.Responses}, + } { + t.Run(tc.name, func(t *testing.T) { + c, recorder := newTypeSafeGroupContext(t, tc.path, tc.body, service.PlatformTypeSafe) + tc.handler(c) + require.Equal(t, http.StatusNotFound, recorder.Code, recorder.Body.String()) + require.Contains(t, recorder.Body.String(), systemOneOnlyPlatformMessage) + }) + } +} diff --git a/backend/internal/handler/payment_handler.go b/backend/internal/handler/payment_handler.go index 9aab3255044e..24e318cc5a04 100644 --- a/backend/internal/handler/payment_handler.go +++ b/backend/internal/handler/payment_handler.go @@ -149,6 +149,9 @@ func (h *PaymentHandler) GetCheckoutInfo(c *gin.Context) { BalanceRechargeMultiplier: cfg.BalanceRechargeMultiplier, SubscriptionUSDToCNYRate: cfg.SubscriptionUSDToCNYRate, RechargeFeeRate: cfg.RechargeFeeRate, + RechargeBonusTiers: cfg.RechargeBonusTiers, + RechargeBonusMode: cfg.RechargeBonusMode, + RechargeBonusNotice: cfg.RechargeBonusNotice, HelpText: cfg.HelpText, HelpImageURL: cfg.HelpImageURL, StripePublishableKey: cfg.StripePublishableKey, @@ -166,6 +169,9 @@ type checkoutInfoResponse struct { BalanceRechargeMultiplier float64 `json:"balance_recharge_multiplier"` SubscriptionUSDToCNYRate float64 `json:"subscription_usd_to_cny_rate"` RechargeFeeRate float64 `json:"recharge_fee_rate"` + RechargeBonusTiers []service.RechargeBonusTier `json:"recharge_bonus_tiers"` + RechargeBonusMode string `json:"recharge_bonus_mode"` + RechargeBonusNotice string `json:"recharge_bonus_notice"` HelpText string `json:"help_text"` HelpImageURL string `json:"help_image_url"` StripePublishableKey string `json:"stripe_publishable_key"` @@ -482,6 +488,7 @@ type PublicOrderResult struct { Amount float64 `json:"amount"` PayAmount float64 `json:"pay_amount"` FeeRate float64 `json:"fee_rate"` + BonusAmount float64 `json:"bonus_amount"` Currency string `json:"currency"` PaymentType string `json:"payment_type"` OrderType string `json:"order_type"` @@ -517,6 +524,7 @@ func buildPublicOrderResult(order *dbent.PaymentOrder) PublicOrderResult { Amount: order.Amount, PayAmount: order.PayAmount, FeeRate: order.FeeRate, + BonusAmount: order.BonusAmount, Currency: service.PaymentOrderCurrency(order), PaymentType: order.PaymentType, OrderType: order.OrderType, @@ -625,6 +633,7 @@ type PaymentOrderResult struct { Amount float64 `json:"amount"` PayAmount float64 `json:"pay_amount"` FeeRate float64 `json:"fee_rate"` + BonusAmount float64 `json:"bonus_amount"` Currency string `json:"currency"` PaymentType string `json:"payment_type"` OutTradeNo string `json:"out_trade_no"` @@ -663,6 +672,7 @@ func sanitizePaymentOrderForResponse(order *dbent.PaymentOrder) *PaymentOrderRes Amount: order.Amount, PayAmount: order.PayAmount, FeeRate: order.FeeRate, + BonusAmount: order.BonusAmount, Currency: service.PaymentOrderCurrency(order), PaymentType: order.PaymentType, OutTradeNo: order.OutTradeNo, diff --git a/backend/internal/handler/user_handler_test.go b/backend/internal/handler/user_handler_test.go index a40755c92092..796c64f33981 100644 --- a/backend/internal/handler/user_handler_test.go +++ b/backend/internal/handler/user_handler_test.go @@ -6,6 +6,7 @@ import ( "bytes" "context" "encoding/json" + "errors" "net/http" "net/http/httptest" "testing" @@ -811,3 +812,19 @@ func TestUserHandlerStartIdentityBindingReturnsAuthorizeURL(t *testing.T) { require.Contains(t, resp.Data.AuthorizeURL, "intent=bind_current_user") require.Contains(t, resp.Data.AuthorizeURL, "redirect=%2Fsettings%2Fprofile") } + +func (s *userHandlerEmailCacheStub) IncrVerificationCodeAttempts(context.Context, string) (int, error) { + if s.data == nil { + return 0, errors.New("verification code not found") + } + s.data.Attempts++ + return s.data.Attempts, nil +} + +func (s *userHandlerEmailCacheStub) IncrNotifyVerifyCodeAttempts(context.Context, string) (int, error) { + return 0, errors.New("notify verification code not found") +} + +func (s *userHandlerEmailCacheStub) ConsumePasswordResetToken(context.Context, string, string) (bool, error) { + return false, nil +} diff --git a/backend/internal/model/error_passthrough_rule.go b/backend/internal/model/error_passthrough_rule.go index a0c81934d7d7..da928125f29b 100644 --- a/backend/internal/model/error_passthrough_rule.go +++ b/backend/internal/model/error_passthrough_rule.go @@ -46,6 +46,7 @@ const ( PlatformDeepseek = domain.PlatformDeepseek PlatformMiniMax = domain.PlatformMiniMax PlatformOpenCodeGo = domain.PlatformOpenCodeGo + PlatformTypeSafe = domain.PlatformTypeSafe ) // AllPlatforms 返回所有支持的平台列表 @@ -61,6 +62,7 @@ func AllPlatforms() []string { PlatformDeepseek, PlatformMiniMax, PlatformOpenCodeGo, + PlatformTypeSafe, } } diff --git a/backend/internal/model/error_passthrough_rule_test.go b/backend/internal/model/error_passthrough_rule_test.go index 34ceb414e8fd..17e6fdcbf281 100644 --- a/backend/internal/model/error_passthrough_rule_test.go +++ b/backend/internal/model/error_passthrough_rule_test.go @@ -18,5 +18,6 @@ func TestAllPlatformsIncludesEveryConcretePlatform(t *testing.T) { "deepseek", "minimax", "opencode_go", + "typesafe", }, AllPlatforms()) } diff --git a/backend/internal/pkg/typesafe/client.go b/backend/internal/pkg/typesafe/client.go index 35783249e0ef..0974204e2a4f 100644 --- a/backend/internal/pkg/typesafe/client.go +++ b/backend/internal/pkg/typesafe/client.go @@ -13,6 +13,12 @@ import ( "strings" ) +const ( + DefaultBaseURL = "https://api.typesafe.ai" + SystemOnePath = "/v1/systemone" + JevLatestModel = "jev-latest" +) + type Question struct { Type string `json:"type"` Instructions string `json:"instructions"` @@ -35,22 +41,103 @@ type Usage struct { OutputTokens int `json:"output_tokens"` } -// Evaluate performs one attempt. The caller owns timeouts, retries and key rotation. -func Evaluate(ctx context.Context, client *http.Client, baseURL, key string, input Request) (*Result, int, error) { - endpoint, err := url.JoinPath(strings.TrimRight(baseURL, "/"), "/v1/systemone") +type SystemOneResponse struct { + Body []byte + Model string + Usage Usage +} + +func NewSystemOneRequest(ctx context.Context, baseURL, key string, body []byte) (*http.Request, error) { + endpoint, err := url.JoinPath(strings.TrimRight(baseURL, "/"), SystemOnePath) if err != nil { - return nil, 0, errors.New("typesafe invalid endpoint") + return nil, errors.New("typesafe invalid endpoint") + } + req, err := http.NewRequestWithContext(ctx, http.MethodPost, endpoint, bytes.NewReader(body)) + if err != nil { + return nil, errors.New("typesafe invalid request") + } + req.Header.Set("Authorization", "Bearer "+key) + req.Header.Set("Content-Type", "application/json") + return req, nil +} + +// MaxSystemOneResponseBytes bounds a buffered System One response body. +const MaxSystemOneResponseBytes = 4 << 20 + +var ErrSystemOneResponseTooLarge = errors.New("typesafe response exceeds size limit") + +func DecodeSystemOneResponse(r io.Reader) (*SystemOneResponse, error) { + body, err := io.ReadAll(io.LimitReader(r, MaxSystemOneResponseBytes+1)) + if err != nil { + return nil, errors.New("typesafe invalid response") + } + // A truncated body would otherwise surface as a misleading "invalid JSON". + if len(body) > MaxSystemOneResponseBytes { + return nil, ErrSystemOneResponseTooLarge + } + if !json.Valid(body) { + return nil, errors.New("typesafe invalid response") } + var envelope map[string]json.RawMessage + if err := json.Unmarshal(body, &envelope); err != nil || envelope == nil { + return nil, errors.New("typesafe invalid response") + } + // The upstream already answered (and charged); an unexpected model or usage + // shape must not discard the answer, so both are decoded leniently. + var model string + _ = json.Unmarshal(envelope["model"], &model) + var usage map[string]json.RawMessage + _ = json.Unmarshal(envelope["usage"], &usage) + return &SystemOneResponse{ + Body: body, + Model: model, + Usage: Usage{ + InputTokens: systemOneTokenCount(usage["input_tokens"]), + OutputTokens: systemOneTokenCount(usage["output_tokens"]), + }, + }, nil +} + +// maxSystemOneTokenCount bounds a reported token count before int conversion. +const maxSystemOneTokenCount = 1 << 40 + +// systemOneTokenCount accepts integer, float, or numeric-string token counts. +func systemOneTokenCount(raw json.RawMessage) int { + raw = bytes.TrimSpace(raw) + if len(raw) == 0 { + return 0 + } + if raw[0] == '"' { + var text string + if json.Unmarshal(raw, &text) != nil { + return 0 + } + raw = json.RawMessage(strings.TrimSpace(text)) + } + var number json.Number + if json.Unmarshal(raw, &number) != nil { + return 0 + } + value, err := number.Float64() + if err != nil || math.IsNaN(value) || value <= 0 { + return 0 + } + if value > maxSystemOneTokenCount { + value = maxSystemOneTokenCount + } + return int(math.Round(value)) +} + +// Evaluate performs one attempt. The caller owns timeouts, retries and key rotation. +func Evaluate(ctx context.Context, client *http.Client, baseURL, key string, input Request) (*Result, int, error) { body, err := json.Marshal(input) if err != nil { return nil, 0, errors.New("typesafe invalid request") } - req, err := http.NewRequestWithContext(ctx, http.MethodPost, endpoint, bytes.NewReader(body)) + req, err := NewSystemOneRequest(ctx, baseURL, key, body) if err != nil { - return nil, 0, errors.New("typesafe invalid request") + return nil, 0, err } - req.Header.Set("Authorization", "Bearer "+key) - req.Header.Set("Content-Type", "application/json") resp, err := client.Do(req) if err != nil { if ctx.Err() != nil { diff --git a/backend/internal/pkg/typesafe/systemone.go b/backend/internal/pkg/typesafe/systemone.go new file mode 100644 index 000000000000..4fab1f1d695e --- /dev/null +++ b/backend/internal/pkg/typesafe/systemone.go @@ -0,0 +1,202 @@ +package typesafe + +import ( + "bytes" + "encoding/json" + "errors" + "fmt" + "strings" +) + +var ErrStreamingUnsupported = errors.New("typesafe system one does not support streaming") + +type systemOneEnvelope struct { + Model json.RawMessage `json:"model"` + State json.RawMessage `json:"state"` + Questions json.RawMessage `json:"questions"` + Stream json.RawMessage `json:"stream"` +} + +var ( + systemOneRequestFields = []string{"model", "state", "questions", "stream"} + systemOneQuestionFields = []string{"type", "instructions", "criteria"} +) + +func ValidateSystemOneRequest(body []byte) (string, error) { + var envelope systemOneEnvelope + if err := json.Unmarshal(body, &envelope); err != nil { + return "", errors.New("invalid JSON request") + } + // encoding/json matches struct fields case-insensitively and keeps the last + // duplicate, while the raw body is forwarded upstream unchanged. Reject any + // spelling the upstream could read differently from this validator. + if err := checkSystemOneObjectKeys(body, "request", systemOneRequestFields); err != nil { + return "", err + } + + model, err := requiredString(envelope.Model, "model") + if err != nil { + return "", err + } + if model != JevLatestModel { + return "", fmt.Errorf("model must be %s", JevLatestModel) + } + if err := validateStringObjectOrArray(envelope.State, "state"); err != nil { + return "", err + } + questionsRaw := bytes.TrimSpace(envelope.Questions) + var questions map[string]json.RawMessage + if len(questionsRaw) == 0 || questionsRaw[0] != '{' || json.Unmarshal(questionsRaw, &questions) != nil || len(questions) == 0 { + return "", errors.New("questions must be a non-empty object") + } + if err := checkSystemOneObjectKeys(questionsRaw, "questions", nil); err != nil { + return "", err + } + if len(envelope.Stream) > 0 && string(envelope.Stream) != "null" { + var stream bool + if err := json.Unmarshal(envelope.Stream, &stream); err != nil { + return "", errors.New("stream must be a boolean") + } + if stream { + return "", ErrStreamingUnsupported + } + } + for id, raw := range questions { + if err := validateQuestion(id, raw); err != nil { + return "", err + } + } + return model, nil +} + +func validateQuestion(id string, raw json.RawMessage) error { + raw = bytes.TrimSpace(raw) + var question struct { + Type json.RawMessage `json:"type"` + Instructions json.RawMessage `json:"instructions"` + Criteria json.RawMessage `json:"criteria"` + } + if len(raw) == 0 || raw[0] != '{' || json.Unmarshal(raw, &question) != nil { + return fmt.Errorf("question %q must be an object", id) + } + if err := checkSystemOneObjectKeys(raw, fmt.Sprintf("question %q", id), systemOneQuestionFields); err != nil { + return err + } + typ, err := requiredString(question.Type, "question type") + if err != nil { + return fmt.Errorf("question %q: %w", id, err) + } + // Mirror the wire schema generated from https://api.typesafe.ai/openapi.json + // in TypeSafe SDK v0.5.7: instructions are optional and nullable for all types. + if err := validateOptionalDescription(question.Instructions, "instructions"); err != nil { + return fmt.Errorf("question %q: %w", id, err) + } + + switch typ { + case "noul": + criteria := bytes.TrimSpace(question.Criteria) + if len(criteria) == 0 || bytes.Equal(criteria, []byte("null")) { + return nil + } + var descriptions map[string]json.RawMessage + if criteria[0] != '{' || json.Unmarshal(criteria, &descriptions) != nil { + return fmt.Errorf("question %q: noul criteria must be an object", id) + } + for _, outcome := range []string{"true", "false"} { + if err := validateOptionalDescription(descriptions[outcome], "noul criteria "+outcome); err != nil { + return fmt.Errorf("question %q: %w", id, err) + } + } + case "choice": + var criteria map[string]json.RawMessage + if json.Unmarshal(question.Criteria, &criteria) != nil || criteria == nil { + return fmt.Errorf("question %q: choice criteria must be an object", id) + } + for _, value := range criteria { + if err := validateOptionalDescription(value, "choice criteria value"); err != nil { + return fmt.Errorf("question %q: %w", id, err) + } + } + case "score": + var criteria []json.RawMessage + if json.Unmarshal(question.Criteria, &criteria) != nil || len(criteria) == 0 { + return fmt.Errorf("question %q: score criteria must contain at least one level", id) + } + for _, value := range criteria { + if err := validateStringObjectOrArray(value, "score criteria value"); err != nil { + return fmt.Errorf("question %q: %w", id, err) + } + } + default: + return fmt.Errorf("question %q: unsupported type %q", id, typ) + } + return nil +} + +func validateOptionalDescription(raw json.RawMessage, name string) error { + raw = bytes.TrimSpace(raw) + if len(raw) == 0 || bytes.Equal(raw, []byte("null")) { + return nil + } + return validateStringObjectOrArray(raw, name) +} + +func requiredString(raw json.RawMessage, name string) (string, error) { + var value string + if len(raw) == 0 || json.Unmarshal(raw, &value) != nil || strings.TrimSpace(value) == "" { + return "", fmt.Errorf("%s must be a non-empty string", name) + } + return value, nil +} + +func validateStringObjectOrArray(raw json.RawMessage, name string) error { + raw = bytes.TrimSpace(raw) + if len(raw) == 0 || string(raw) == "null" { + return fmt.Errorf("%s is required", name) + } + switch raw[0] { + case '{', '[': + return nil + case '"': + if rawString(raw) { + return nil + } + } + return fmt.Errorf("%s must be a string, object, or array", name) +} + +func rawString(raw json.RawMessage) bool { + var value string + return json.Unmarshal(raw, &value) == nil +} + +// checkSystemOneObjectKeys rejects duplicate keys and non-canonical spellings +// (any case variant) of the known fields in one JSON object level. +func checkSystemOneObjectKeys(raw []byte, scope string, canonical []string) error { + decoder := json.NewDecoder(bytes.NewReader(raw)) + if token, err := decoder.Token(); err != nil || token != json.Delim('{') { + return fmt.Errorf("%s must be an object", scope) + } + seen := make(map[string]struct{}) + for decoder.More() { + token, err := decoder.Token() + if err != nil { + return errors.New("invalid JSON request") + } + key, _ := token.(string) + if _, duplicate := seen[key]; duplicate { + return fmt.Errorf("%s contains duplicate field %q", scope, key) + } + seen[key] = struct{}{} + for _, field := range canonical { + if key != field && strings.EqualFold(key, field) { + return fmt.Errorf("%s field %q must be written as %q", scope, key, field) + } + } + var value json.RawMessage + if err := decoder.Decode(&value); err != nil { + return errors.New("invalid JSON request") + } + } + return nil +} diff --git a/backend/internal/pkg/typesafe/systemone_test.go b/backend/internal/pkg/typesafe/systemone_test.go new file mode 100644 index 000000000000..65991cf14367 --- /dev/null +++ b/backend/internal/pkg/typesafe/systemone_test.go @@ -0,0 +1,124 @@ +package typesafe + +import ( + "errors" + "strings" + "testing" + + "github.com/stretchr/testify/require" +) + +func TestValidateSystemOneRequestValidQuestionTypes(t *testing.T) { + for _, tc := range []struct { + name string + body string + }{ + {"noul string state", `{"model":"jev-latest","state":"sample","questions":{"safety":{"type":"noul","instructions":"Evaluate safety","criteria":{"safe":"No harm"}}}}`}, + {"choice object state", `{"model":"jev-latest","state":{"text":"sample"},"questions":{"label":{"type":"choice","instructions":{"task":"Classify"},"criteria":{"safe":"Allowed","unsafe":null}}},"stream":false}`}, + {"score array state", `{"model":"jev-latest","state":["sample"],"questions":{"quality":{"type":"score","instructions":["Rate quality"],"criteria":["poor","good"]}}}`}, + {"noul omitted instructions", `{"model":"jev-latest","state":"sample","questions":{"q":{"type":"noul"}}}`}, + {"noul nullable instructions and criteria", `{"model":"jev-latest","state":"sample","questions":{"q":{"type":"noul","instructions":null,"criteria":null}}}`}, + {"noul structured descriptions", `{"model":"jev-latest","state":"sample","questions":{"q":{"type":"noul","criteria":{"true":{"examples":["yes",null]},"false":["no",null],"extension":42}}}}`}, + {"choice structured descriptions", `{"model":"jev-latest","state":"sample","questions":{"q":{"type":"choice","instructions":null,"criteria":{"a":{"description":"A","extra":null},"b":["B",null],"c":null}}}}`}, + {"choice empty criteria", `{"model":"jev-latest","state":"sample","questions":{"q":{"type":"choice","criteria":{}}}}`}, + {"score one level", `{"model":"jev-latest","state":"sample","questions":{"q":{"type":"score","criteria":["only"]}}}`}, + {"score object level", `{"model":"jev-latest","state":"sample","questions":{"q":{"type":"score","instructions":null,"criteria":[{"description":"only","extra":null}]}}}`}, + {"score array level", `{"model":"jev-latest","state":"sample","questions":{"q":{"type":"score","criteria":[["only",null]]}}}`}, + {"native extensions preserved", `{"model":"jev-latest","state":"sample","questions":{"q":{"type":"noul","extension":{"kept":true}}},"extension":[1,null]}`}, + } { + t.Run(tc.name, func(t *testing.T) { + model, err := ValidateSystemOneRequest([]byte(tc.body)) + require.NoError(t, err) + require.Equal(t, JevLatestModel, model) + }) + } +} + +func TestValidateSystemOneRequestRejectsInvalidRequests(t *testing.T) { + for _, tc := range []struct { + name string + body string + want string + }{ + {"invalid json", `{`, "invalid JSON"}, + {"missing model", `{"state":"x","questions":{"q":{"type":"noul","instructions":"x"}}}`, "model"}, + {"illegal model", `{"model":"jev-old","state":"x","questions":{"q":{"type":"noul","instructions":"x"}}}`, "jev-latest"}, + {"model with whitespace", `{"model":" jev-latest ","state":"x","questions":{"q":{"type":"noul","instructions":"x"}}}`, "jev-latest"}, + {"missing state", `{"model":"jev-latest","questions":{"q":{"type":"noul","instructions":"x"}}}`, "state"}, + {"scalar state", `{"model":"jev-latest","state":42,"questions":{"q":{"type":"noul","instructions":"x"}}}`, "state"}, + {"empty questions", `{"model":"jev-latest","state":"x","questions":{}}`, "questions"}, + {"array questions", `{"model":"jev-latest","state":"x","questions":[{"type":"noul"}]}`, "questions must be a non-empty object"}, + {"null questions", `{"model":"jev-latest","state":"x","questions":null}`, "questions must be a non-empty object"}, + {"unknown question type", `{"model":"jev-latest","state":"x","questions":{"q":{"type":"boolean","instructions":"x"}}}`, "unsupported type"}, + {"question type with whitespace", `{"model":"jev-latest","state":"x","questions":{"q":{"type":" noul ","instructions":"x"}}}`, "unsupported type"}, + {"noul criteria array", `{"model":"jev-latest","state":"x","questions":{"q":{"type":"noul","instructions":"x","criteria":[]}}}`, "noul criteria"}, + {"noul numeric description", `{"model":"jev-latest","state":"x","questions":{"q":{"type":"noul","criteria":{"true":1}}}}`, "noul criteria"}, + {"noul boolean description", `{"model":"jev-latest","state":"x","questions":{"q":{"type":"noul","criteria":{"false":false}}}}`, "noul criteria"}, + {"missing choice criteria", `{"model":"jev-latest","state":"x","questions":{"q":{"type":"choice"}}}`, "choice criteria"}, + {"null choice criteria", `{"model":"jev-latest","state":"x","questions":{"q":{"type":"choice","criteria":null}}}`, "choice criteria"}, + {"choice criteria array", `{"model":"jev-latest","state":"x","questions":{"q":{"type":"choice","criteria":[]}}}`, "choice criteria"}, + {"choice numeric value", `{"model":"jev-latest","state":"x","questions":{"q":{"type":"choice","instructions":"x","criteria":{"one":1}}}}`, "choice criteria"}, + {"missing score criteria", `{"model":"jev-latest","state":"x","questions":{"q":{"type":"score"}}}`, "score criteria"}, + {"empty score criteria", `{"model":"jev-latest","state":"x","questions":{"q":{"type":"score","criteria":[]}}}`, "score criteria"}, + {"null score criteria", `{"model":"jev-latest","state":"x","questions":{"q":{"type":"score","criteria":null}}}`, "score criteria"}, + {"score map is not wire protocol", `{"model":"jev-latest","state":"x","questions":{"q":{"type":"score","criteria":{"0":"low"}}}}`, "score criteria"}, + {"numeric score level", `{"model":"jev-latest","state":"x","questions":{"q":{"type":"score","criteria":["low",2]}}}`, "score criteria"}, + {"null score level", `{"model":"jev-latest","state":"x","questions":{"q":{"type":"score","criteria":[null]}}}`, "score criteria"}, + {"numeric instructions", `{"model":"jev-latest","state":"x","questions":{"q":{"type":"noul","instructions":1}}}`, "instructions"}, + {"boolean instructions", `{"model":"jev-latest","state":"x","questions":{"q":{"type":"choice","instructions":true,"criteria":{}}}}`, "instructions"}, + {"stream true", `{"model":"jev-latest","state":"x","questions":{"q":{"type":"noul","instructions":"x"}},"stream":true}`, "streaming"}, + {"case variant model smuggles upstream model", `{"model":"jev-pro","MODEL":"jev-latest","state":"x","questions":{"q":{"type":"noul"}}}`, `must be written as "model"`}, + {"case variant stream hides streaming", `{"model":"jev-latest","state":"x","questions":{"q":{"type":"noul"}},"stream":true,"Stream":false}`, `must be written as "stream"`}, + {"unicode fold variant state", `{"model":"jev-latest","state":"x","ſtate":"y","questions":{"q":{"type":"noul"}}}`, `must be written as "state"`}, + {"duplicate model", `{"model":"jev-pro","model":"jev-latest","state":"x","questions":{"q":{"type":"noul"}}}`, `duplicate field "model"`}, + {"escaped duplicate model", `{"\u006dodel":"jev-pro","model":"jev-latest","state":"x","questions":{"q":{"type":"noul"}}}`, `duplicate field "model"`}, + {"duplicate state", `{"model":"jev-latest","state":"benign","state":"payload","questions":{"q":{"type":"noul"}}}`, `duplicate field "state"`}, + {"duplicate question id", `{"model":"jev-latest","state":"x","questions":{"q":{"type":"noul"},"q":{"type":"noul"}}}`, `duplicate field "q"`}, + {"duplicate question instructions", `{"model":"jev-latest","state":"x","questions":{"q":{"type":"noul","instructions":"a","instructions":"b"}}}`, `duplicate field "instructions"`}, + {"case variant question type", `{"model":"jev-latest","state":"x","questions":{"q":{"type":"noul","Type":"choice"}}}`, `must be written as "type"`}, + } { + t.Run(tc.name, func(t *testing.T) { + _, err := ValidateSystemOneRequest([]byte(tc.body)) + require.Error(t, err) + require.Contains(t, err.Error(), tc.want) + if tc.name == "stream true" { + require.True(t, errors.Is(err, ErrStreamingUnsupported)) + } + }) + } +} + +func TestDecodeSystemOneResponseToleratesUsageShapes(t *testing.T) { + for _, tc := range []struct { + name string + body string + model string + input, output int + }{ + {"integers", `{"model":"jev-1","usage":{"input_tokens":12,"output_tokens":3}}`, "jev-1", 12, 3}, + {"floats", `{"model":"jev-1","usage":{"input_tokens":12.0,"output_tokens":2.6}}`, "jev-1", 12, 3}, + {"numeric strings", `{"usage":{"input_tokens":"15","output_tokens":" 4 "}}`, "", 15, 4}, + {"non numeric usage", `{"model":7,"usage":{"input_tokens":"many","output_tokens":null}}`, "", 0, 0}, + {"negative usage", `{"usage":{"input_tokens":-5}}`, "", 0, 0}, + {"usage not object", `{"model":"jev-1","usage":"none","answers":{}}`, "jev-1", 0, 0}, + {"missing usage", `{"answers":{}}`, "", 0, 0}, + } { + t.Run(tc.name, func(t *testing.T) { + decoded, err := DecodeSystemOneResponse(strings.NewReader(tc.body)) + require.NoError(t, err) + require.Equal(t, []byte(tc.body), decoded.Body) + require.Equal(t, tc.model, decoded.Model) + require.Equal(t, tc.input, decoded.Usage.InputTokens) + require.Equal(t, tc.output, decoded.Usage.OutputTokens) + }) + } +} + +func TestDecodeSystemOneResponseRejectsNonObjects(t *testing.T) { + for _, body := range []string{`[]`, `"text"`, `null`, `{`, `data: {}`} { + _, err := DecodeSystemOneResponse(strings.NewReader(body)) + require.Error(t, err, body) + } + _, err := DecodeSystemOneResponse(strings.NewReader(`{"pad":"` + strings.Repeat("a", MaxSystemOneResponseBytes) + `"}`)) + require.ErrorIs(t, err, ErrSystemOneResponseTooLarge) +} diff --git a/backend/internal/pkg/xai/billing.go b/backend/internal/pkg/xai/billing.go index c4247b76cecf..d1386bd2398d 100644 --- a/backend/internal/pkg/xai/billing.go +++ b/backend/internal/pkg/xai/billing.go @@ -19,10 +19,7 @@ const ( // repository and service layers build their own client identity from it, so // one bump here covers OAuth traffic and billing probes together. // Keep in sync with the latest stable @xai-official/grok release. - CLIClientVersion = "1.0.44" - // billingCLIUserAgent is the legacy pager/shell UA used by billing probes. - // Distinct from CLIUserAgent() in cli_identity.go (workspace-style UA). - billingCLIUserAgent = "grok-pager/" + CLIClientVersion + " grok-shell/" + CLIClientVersion + " (macos; aarch64)" + CLIClientVersion = "1.0.46" BillingWeeklyPath = "/billing?format=credits" BillingMonthlyPath = "/billing" @@ -148,7 +145,8 @@ func ApplyCLIBillingHeaders(req *http.Request, accessToken string) { req.Header.Set("Content-Type", "application/json") req.Header.Set(CLITokenAuthHeader, CLITokenAuthValue) req.Header.Set(CLIClientVersionHeader, CLIClientVersion) - req.Header.Set("User-Agent", billingCLIUserAgent) + req.Header.Set("User-Agent", CLIUserAgent(CLIClientVersion)) + req.Header.Set("x-grok-client-mode", CLIClientMode) } // ParseBillingPayload unmarshals a billing API response body. diff --git a/backend/internal/pkg/xai/billing_test.go b/backend/internal/pkg/xai/billing_test.go index 3dbfc89bf2b6..51e3f2b1b4f9 100644 --- a/backend/internal/pkg/xai/billing_test.go +++ b/backend/internal/pkg/xai/billing_test.go @@ -38,7 +38,8 @@ func TestApplyCLIBillingHeaders(t *testing.T) { require.Equal(t, "Bearer token", req.Header.Get("Authorization")) require.Equal(t, CLITokenAuthValue, req.Header.Get(CLITokenAuthHeader)) require.Equal(t, CLIClientVersion, req.Header.Get(CLIClientVersionHeader)) - require.Equal(t, "grok-pager/"+CLIClientVersion+" grok-shell/"+CLIClientVersion+" (macos; aarch64)", req.UserAgent()) + require.Equal(t, CLIUserAgent(CLIClientVersion), req.UserAgent()) + require.Equal(t, "interactive", req.Header.Get("x-grok-client-mode")) } func TestBuildBillingSummaryWeeklyAndMonthly(t *testing.T) { diff --git a/backend/internal/pkg/xai/cli_identity.go b/backend/internal/pkg/xai/cli_identity.go index c34b7a0c084e..febb03f10351 100644 --- a/backend/internal/pkg/xai/cli_identity.go +++ b/backend/internal/pkg/xai/cli_identity.go @@ -3,6 +3,7 @@ package xai import ( "net/http" "os" + "runtime" "strings" "golang.org/x/mod/semver" @@ -26,11 +27,11 @@ const ( // CLITokenAuth is required by cli-chat-proxy for Grok Build OAuth tokens. CLITokenAuth = "xai-grok-cli" - // CLIClientIdentifier is the x-grok-client-identifier value used by Grok shell/CLI. - CLIClientIdentifier = "grok-shell" + // CLIClientIdentifier 对齐官方交互式 CLI 主请求的客户端标识。 + CLIClientIdentifier = "grok-pager" - // CLIClientMode is used by billing / quota probes on the CLI surface. - CLIClientMode = "cli" + // CLIClientMode 对齐官方 CLI 正常交互模式。 + CLIClientMode = "interactive" ) // ResolveCLIVersion returns a supported CLI client version. @@ -56,12 +57,24 @@ func IsSupportedCLIVersion(version string) bool { semver.Compare(canonical, minimum) >= 0 } -// CLIUserAgent builds the workspace-style User-Agent for a CLI client version. +// CLIUserAgent 对齐官方交互式 CLI 的 UA,平台名称使用 Rust 的格式。 func CLIUserAgent(version string) string { if strings.TrimSpace(version) == "" { version = CLIClientVersion } - return "xai-grok-workspace/" + version + platform, arch := runtime.GOOS, runtime.GOARCH + if platform == "darwin" { + platform = "macos" + } + switch arch { + case "amd64": + arch = "x86_64" + case "arm64": + arch = "aarch64" + case "386": + arch = "x86" + } + return "grok-pager/" + version + " grok-shell/" + version + " (" + platform + "; " + arch + ")" } // ApplyCLIProxyHeaders stamps the fixed Grok CLI identity when the request @@ -77,5 +90,7 @@ func ApplyCLIProxyHeaders(req *http.Request) { req.Header.Set("X-XAI-Token-Auth", CLITokenAuth) req.Header.Set("x-grok-client-version", version) req.Header.Set("x-grok-client-identifier", CLIClientIdentifier) + req.Header.Set("x-grok-client-mode", CLIClientMode) + req.Header.Set("x-authenticateresponse", "authenticate-response") req.Header.Set("User-Agent", CLIUserAgent(version)) } diff --git a/backend/internal/pkg/xai/cli_identity_test.go b/backend/internal/pkg/xai/cli_identity_test.go index 996c5cb8c05d..ce21d2a5cdc7 100644 --- a/backend/internal/pkg/xai/cli_identity_test.go +++ b/backend/internal/pkg/xai/cli_identity_test.go @@ -2,11 +2,40 @@ package xai import ( "net/http" + "runtime" "testing" "github.com/stretchr/testify/require" + "golang.org/x/mod/semver" ) +func TestApplyCLIProxyHeadersMeetsUpstreamMinimumVersion(t *testing.T) { + // 上游 426 明确要求至少 1.0.13;断言独立于生产版本常量,防止旧覆盖值绕过下限。 + for _, override := range []string{"", "0.2.93", "0.2.120", "1.0.12", "1.0.13-beta.1", "1.0.13", "1.0.14-alpha.1"} { + t.Run("override="+override, func(t *testing.T) { + t.Setenv(CLIVersionEnv, override) + req, err := http.NewRequest(http.MethodPost, "https://cli-chat-proxy.grok.com/v1/responses", nil) + require.NoError(t, err) + + ApplyCLIProxyHeaders(req) + + version := req.Header.Get("x-grok-client-version") + require.True(t, semver.IsValid("v"+version)) + require.GreaterOrEqual(t, semver.Compare("v"+version, "v1.0.13"), 0, + "Grok Responses 会以 426 拒绝版本 %s,最低要求为 1.0.13", version) + require.Contains(t, req.Header.Get("User-Agent"), "grok-shell/"+version+" (") + }) + } +} + +func TestCLIUserAgentMatchesOfficialInteractiveCapture(t *testing.T) { + // 来源:官方 CLI 1.0.46 交互式界面在 Linux x86_64 的主 Responses 请求抓包。 + if runtime.GOOS != "linux" || runtime.GOARCH != "amd64" { + t.Skip("该抓包样本来自 Linux x86_64") + } + require.Equal(t, "grok-pager/1.0.46 grok-shell/1.0.46 (linux; x86_64)", CLIUserAgent("1.0.46")) +} + func TestResolveCLIVersionDefaultsToPinnedClientVersion(t *testing.T) { t.Setenv(CLIVersionEnv, "") // The default pin and minimum accepted version stay aligned. @@ -16,16 +45,16 @@ func TestResolveCLIVersionDefaultsToPinnedClientVersion(t *testing.T) { } func TestResolveCLIVersionAcceptsValidOverride(t *testing.T) { - t.Setenv(CLIVersionEnv, "1.0.45-alpha.1") - require.Equal(t, "1.0.45-alpha.1", ResolveCLIVersion()) + t.Setenv(CLIVersionEnv, "1.0.47-alpha.1") + require.Equal(t, "1.0.47-alpha.1", ResolveCLIVersion()) } func TestResolveCLIVersionRejectsUnsafeOrTooOld(t *testing.T) { for _, version := range []string{ - "1.0.43", - "1.0.44-beta.1", - "1.0.45\r\nX-Injected: true", - "1.0.044", + "1.0.45", + "1.0.46-beta.1", + "1.0.47\r\nX-Injected: true", + "1.0.046", "1.1", "1", } { @@ -48,11 +77,13 @@ func TestApplyCLIProxyHeaders(t *testing.T) { require.Equal(t, CLIClientVersion, req.Header.Get("x-grok-client-version")) require.Equal(t, CLIClientIdentifier, req.Header.Get("x-grok-client-identifier")) require.Equal(t, CLITokenAuth, req.Header.Get("X-XAI-Token-Auth")) + require.Equal(t, "interactive", req.Header.Get("x-grok-client-mode")) + require.Equal(t, "authenticate-response", req.Header.Get("x-authenticateresponse")) require.Equal(t, CLIUserAgent(CLIClientVersion), req.Header.Get("User-Agent")) } func TestApplyCLIProxyHeadersLeavesAPIHostUnchanged(t *testing.T) { - t.Setenv(CLIVersionEnv, "1.0.44") + t.Setenv(CLIVersionEnv, "1.0.46") req, err := http.NewRequest(http.MethodPost, "https://api.x.ai/v1/responses", nil) require.NoError(t, err) @@ -63,5 +94,7 @@ func TestApplyCLIProxyHeadersLeavesAPIHostUnchanged(t *testing.T) { require.Empty(t, req.Header.Get("x-grok-client-version")) require.Empty(t, req.Header.Get("x-grok-client-identifier")) require.Empty(t, req.Header.Get("X-XAI-Token-Auth")) + require.Empty(t, req.Header.Get("x-grok-client-mode")) + require.Empty(t, req.Header.Get("x-authenticateresponse")) require.Equal(t, "direct-api-client/1.0", req.Header.Get("User-Agent")) } diff --git a/backend/internal/repository/api_key_repo.go b/backend/internal/repository/api_key_repo.go index 01d97fa5d696..94e5bbadf755 100644 --- a/backend/internal/repository/api_key_repo.go +++ b/backend/internal/repository/api_key_repo.go @@ -658,6 +658,17 @@ func apiKeyListOrder(params pagination.PaginationParams) []func(*entsql.Selector sortBy := strings.ToLower(strings.TrimSpace(params.SortBy)) sortOrder := params.NormalizedSortOrder(pagination.SortOrderDesc) + if sortBy == "group" { + // Sort before pagination, keeping ungrouped keys last in either direction. + opts := []entsql.OrderTermOption{entsql.OrderNullsLast()} + tieOrder := dbent.Asc(apikey.FieldID) + if sortOrder == pagination.SortOrderDesc { + opts = append(opts, entsql.OrderDesc()) + tieOrder = dbent.Desc(apikey.FieldID) + } + return []func(*entsql.Selector){apikey.ByGroupField(group.FieldName, opts...), tieOrder} + } + var field string switch sortBy { case "name": diff --git a/backend/internal/repository/api_key_repo_sort_test.go b/backend/internal/repository/api_key_repo_sort_test.go new file mode 100644 index 000000000000..369b6a59003e --- /dev/null +++ b/backend/internal/repository/api_key_repo_sort_test.go @@ -0,0 +1,78 @@ +package repository + +import ( + "context" + "testing" + + dbent "github.com/Wei-Shaw/sub2api/ent" + "github.com/Wei-Shaw/sub2api/internal/pkg/pagination" + "github.com/Wei-Shaw/sub2api/internal/service" + "github.com/stretchr/testify/require" +) + +func TestAPIKeyRepositoryListByUserIDSortByGroup(t *testing.T) { + repo, client := newAPIKeyRepoSQLite(t) + ctx := context.Background() + user := mustCreateAPIKeyRepoUser(t, ctx, client, "group-sort@test.com") + otherUser := mustCreateAPIKeyRepoUser(t, ctx, client, "other-group-sort@test.com") + createGroup := func(name string) *dbent.Group { + g, err := client.Group.Create().SetName(name).Save(ctx) + require.NoError(t, err) + return g + } + // Create groups in reverse name order so sorting by group ID cannot pass. + zulu := createGroup("Zulu") + alpha := createGroup("Alpha") + createKey := func(userID int64, name string, groupID *int64, status string) int64 { + key := &service.APIKey{ + UserID: userID, Key: "sk-" + name, Name: name, GroupID: groupID, Status: status, + } + require.NoError(t, repo.Create(ctx, key)) + return key.ID + } + zuluFirst := createKey(user.ID, "match-zulu-first", &zulu.ID, service.StatusActive) + ungroupedFirst := createKey(user.ID, "match-ungrouped-first", nil, service.StatusActive) + alphaFirst := createKey(user.ID, "match-alpha-first", &alpha.ID, service.StatusActive) + zuluSecond := createKey(user.ID, "match-zulu-second", &zulu.ID, service.StatusActive) + alphaSecond := createKey(user.ID, "match-alpha-second", &alpha.ID, service.StatusActive) + ungroupedSecond := createKey(user.ID, "match-ungrouped-second", nil, service.StatusActive) + createKey(otherUser.ID, "match-other-user", &alpha.ID, service.StatusActive) + createKey(user.ID, "excluded-search", &alpha.ID, service.StatusActive) + createKey(user.ID, "match-inactive", &alpha.ID, service.StatusDisabled) + deleted := createKey(user.ID, "match-deleted", &alpha.ID, service.StatusActive) + require.NoError(t, repo.Delete(ctx, deleted)) + + ungroupedID := int64(0) + for _, tc := range []struct { + name string + order string + groupID *int64 + want []int64 + }{ + {"ascending", "asc", nil, []int64{alphaFirst, alphaSecond, zuluFirst, zuluSecond, ungroupedFirst, ungroupedSecond}}, + {"descending", "desc", nil, []int64{zuluSecond, zuluFirst, alphaSecond, alphaFirst, ungroupedSecond, ungroupedFirst}}, + {"group filter", "asc", &alpha.ID, []int64{alphaFirst, alphaSecond}}, + {"ungrouped filter", "desc", &ungroupedID, []int64{ungroupedSecond, ungroupedFirst}}, + } { + t.Run(tc.name, func(t *testing.T) { + var got []int64 + const pageSize = 3 + pages := (len(tc.want) + pageSize - 1) / pageSize + for page := 1; page <= pages; page++ { + keys, result, err := repo.ListByUserID(ctx, user.ID, pagination.PaginationParams{ + Page: page, PageSize: pageSize, SortBy: "group", SortOrder: tc.order, + }, service.APIKeyListFilters{Search: "match", Status: service.StatusActive, GroupID: tc.groupID}) + require.NoError(t, err) + require.EqualValues(t, len(tc.want), result.Total) + require.Equal(t, pages, result.Pages) + for _, key := range keys { + got = append(got, key.ID) + if key.GroupID != nil { + require.NotNil(t, key.Group, "sorting must preserve group preloading") + } + } + } + require.Equal(t, tc.want, got) + }) + } +} diff --git a/backend/internal/repository/email_cache.go b/backend/internal/repository/email_cache.go index 96a23a8ed685..9f59ff05b226 100644 --- a/backend/internal/repository/email_cache.go +++ b/backend/internal/repository/email_cache.go @@ -17,8 +17,42 @@ const ( passwordResetKeyPrefix = "password_reset:" passwordResetSentAtKeyPrefix = "password_reset_sent:" notifyCodeUserRateKeyPrefix = "notify_code_user_rate:" + + // attemptsKeySuffix stores the failed-attempt counter next to a verification code. + // Kept in a separate key so it can be incremented atomically with INCR. + attemptsKeySuffix = ":attempts" ) +// incrAttemptsScript atomically increments the attempt counter for an existing +// verification code and aligns the counter TTL with the code TTL. +// KEYS[1] = code key, KEYS[2] = attempts key. Returns -1 when the code is missing. +var incrAttemptsScript = redis.NewScript(` +if redis.call('EXISTS', KEYS[1]) == 0 then + return -1 +end +local n = redis.call('INCR', KEYS[2]) +local ttl = redis.call('PTTL', KEYS[1]) +if ttl > 0 then + redis.call('PEXPIRE', KEYS[2], ttl) +end +return n +`) + +// consumeResetTokenScript atomically compares the stored token hash and deletes it. +// KEYS[1] = reset key, ARGV[1] = expected token hash. Returns 1 on success, 0 otherwise. +var consumeResetTokenScript = redis.NewScript(` +local v = redis.call('GET', KEYS[1]) +if not v then + return 0 +end +local ok, d = pcall(cjson.decode, v) +if not ok or type(d) ~= 'table' or d['Token'] ~= ARGV[1] then + return 0 +end +redis.call('DEL', KEYS[1]) +return 1 +`) + // verifyCodeKey generates the Redis key for email verification code. // Email is lowercased for case-insensitive consistency. func verifyCodeKey(email string) string { @@ -50,8 +84,7 @@ func NewEmailCache(rdb *redis.Client) service.EmailCache { return &emailCache{rdb: rdb} } -func (c *emailCache) GetVerificationCode(ctx context.Context, email string) (*service.VerificationCodeData, error) { - key := verifyCodeKey(email) +func (c *emailCache) getCode(ctx context.Context, key string) (*service.VerificationCodeData, error) { val, err := c.rdb.Get(ctx, key).Result() if err != nil { return nil, err @@ -60,21 +93,56 @@ func (c *emailCache) GetVerificationCode(ctx context.Context, email string) (*se if err := json.Unmarshal([]byte(val), &data); err != nil { return nil, err } + if n, err := c.rdb.Get(ctx, key+attemptsKeySuffix).Int(); err == nil && n > data.Attempts { + data.Attempts = n + } return &data, nil } -func (c *emailCache) SetVerificationCode(ctx context.Context, email string, data *service.VerificationCodeData, ttl time.Duration) error { - key := verifyCodeKey(email) +func (c *emailCache) setCode(ctx context.Context, key string, data *service.VerificationCodeData, ttl time.Duration) error { val, err := json.Marshal(data) if err != nil { return err } - return c.rdb.Set(ctx, key, val, ttl).Err() + pipe := c.rdb.TxPipeline() + pipe.Set(ctx, key, val, ttl) + pipe.Del(ctx, key+attemptsKeySuffix) + if data.Attempts > 0 { + pipe.Set(ctx, key+attemptsKeySuffix, data.Attempts, ttl) + } + _, err = pipe.Exec(ctx) + return err +} + +func (c *emailCache) incrCodeAttempts(ctx context.Context, key string) (int, error) { + n, err := incrAttemptsScript.Run(ctx, c.rdb, []string{key, key + attemptsKeySuffix}).Int() + if err != nil { + return 0, err + } + if n < 0 { + return 0, redis.Nil + } + return n, nil +} + +func (c *emailCache) deleteCode(ctx context.Context, key string) error { + return c.rdb.Del(ctx, key, key+attemptsKeySuffix).Err() +} + +func (c *emailCache) GetVerificationCode(ctx context.Context, email string) (*service.VerificationCodeData, error) { + return c.getCode(ctx, verifyCodeKey(email)) +} + +func (c *emailCache) SetVerificationCode(ctx context.Context, email string, data *service.VerificationCodeData, ttl time.Duration) error { + return c.setCode(ctx, verifyCodeKey(email), data, ttl) +} + +func (c *emailCache) IncrVerificationCodeAttempts(ctx context.Context, email string) (int, error) { + return c.incrCodeAttempts(ctx, verifyCodeKey(email)) } func (c *emailCache) DeleteVerificationCode(ctx context.Context, email string) error { - key := verifyCodeKey(email) - return c.rdb.Del(ctx, key).Err() + return c.deleteCode(ctx, verifyCodeKey(email)) } // Password reset token methods @@ -101,6 +169,16 @@ func (c *emailCache) SetPasswordResetToken(ctx context.Context, email string, da return c.rdb.Set(ctx, key, val, ttl).Err() } +// ConsumePasswordResetToken atomically deletes the stored reset token when its +// stored hash equals tokenHash. Returns true only for the single winning caller. +func (c *emailCache) ConsumePasswordResetToken(ctx context.Context, email, tokenHash string) (bool, error) { + n, err := consumeResetTokenScript.Run(ctx, c.rdb, []string{passwordResetKey(email)}, tokenHash).Int() + if err != nil { + return false, err + } + return n == 1, nil +} + func (c *emailCache) DeletePasswordResetToken(ctx context.Context, email string) error { key := passwordResetKey(email) return c.rdb.Del(ctx, key).Err() @@ -122,30 +200,19 @@ func (c *emailCache) SetPasswordResetEmailCooldown(ctx context.Context, email st // Notify email verification code methods func (c *emailCache) GetNotifyVerifyCode(ctx context.Context, email string) (*service.VerificationCodeData, error) { - key := notifyVerifyKey(email) - val, err := c.rdb.Get(ctx, key).Result() - if err != nil { - return nil, err - } - var data service.VerificationCodeData - if err := json.Unmarshal([]byte(val), &data); err != nil { - return nil, err - } - return &data, nil + return c.getCode(ctx, notifyVerifyKey(email)) } func (c *emailCache) SetNotifyVerifyCode(ctx context.Context, email string, data *service.VerificationCodeData, ttl time.Duration) error { - key := notifyVerifyKey(email) - val, err := json.Marshal(data) - if err != nil { - return err - } - return c.rdb.Set(ctx, key, val, ttl).Err() + return c.setCode(ctx, notifyVerifyKey(email), data, ttl) +} + +func (c *emailCache) IncrNotifyVerifyCodeAttempts(ctx context.Context, email string) (int, error) { + return c.incrCodeAttempts(ctx, notifyVerifyKey(email)) } func (c *emailCache) DeleteNotifyVerifyCode(ctx context.Context, email string) error { - key := notifyVerifyKey(email) - return c.rdb.Del(ctx, key).Err() + return c.deleteCode(ctx, notifyVerifyKey(email)) } // User-level rate limiting for notify email verification codes diff --git a/backend/internal/repository/email_cache_atomic_test.go b/backend/internal/repository/email_cache_atomic_test.go new file mode 100644 index 000000000000..0086bbe4604a --- /dev/null +++ b/backend/internal/repository/email_cache_atomic_test.go @@ -0,0 +1,150 @@ +package repository + +import ( + "context" + "crypto/sha256" + "encoding/hex" + "errors" + "strings" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/Wei-Shaw/sub2api/internal/service" + "github.com/alicebob/miniredis/v2" + "github.com/redis/go-redis/v9" + "github.com/stretchr/testify/require" +) + +func newMiniredisEmailCache(t *testing.T) (service.EmailCache, *miniredis.Miniredis, *redis.Client) { + t.Helper() + mr := miniredis.RunT(t) + rdb := redis.NewClient(&redis.Options{Addr: mr.Addr()}) + t.Cleanup(func() { _ = rdb.Close() }) + return NewEmailCache(rdb), mr, rdb +} + +func TestEmailCache_ConcurrentWrongCodesCannotExceedAttemptCap(t *testing.T) { + cache, _, _ := newMiniredisEmailCache(t) + ctx := context.Background() + email := "user@example.com" + + svc := service.NewEmailService(nil, cache) + require.NoError(t, cache.SetVerificationCode(ctx, email, &service.VerificationCodeData{ + Code: "123456", + CreatedAt: time.Now(), + ExpiresAt: time.Now().Add(15 * time.Minute), + }, 15*time.Minute)) + + const workers = 50 + var invalid, maxed atomic.Int32 + var wg sync.WaitGroup + for i := 0; i < workers; i++ { + wg.Add(1) + go func() { + defer wg.Done() + err := svc.VerifyCode(ctx, email, "000000") + switch { + case errors.Is(err, service.ErrInvalidVerifyCode): + invalid.Add(1) + case errors.Is(err, service.ErrVerifyCodeMaxAttempts): + maxed.Add(1) + default: + t.Errorf("unexpected result: %v", err) + } + }() + } + wg.Wait() + + // Only attempts 1..4 may return "invalid"; every other guess is rejected by the cap. + require.LessOrEqual(t, int(invalid.Load()), 4) + require.Equal(t, workers, int(invalid.Load()+maxed.Load())) + + // Even the correct code is now rejected. + require.ErrorIs(t, svc.VerifyCode(ctx, email, "123456"), service.ErrVerifyCodeMaxAttempts) + + data, err := cache.GetVerificationCode(ctx, email) + require.NoError(t, err) + require.GreaterOrEqual(t, data.Attempts, 5) +} + +func TestEmailCache_AttemptsResetOnNewCodeAndTTLFollowsCode(t *testing.T) { + cache, mr, _ := newMiniredisEmailCache(t) + ctx := context.Background() + email := "User@Example.com" + + require.NoError(t, cache.SetVerificationCode(ctx, email, &service.VerificationCodeData{Code: "1"}, time.Minute)) + n, err := cache.IncrVerificationCodeAttempts(ctx, email) + require.NoError(t, err) + require.Equal(t, 1, n) + require.Greater(t, mr.TTL(verifyCodeKey(email)+attemptsKeySuffix), time.Duration(0)) + + require.NoError(t, cache.SetVerificationCode(ctx, email, &service.VerificationCodeData{Code: "2"}, time.Minute)) + data, err := cache.GetVerificationCode(ctx, email) + require.NoError(t, err) + require.Equal(t, 0, data.Attempts) + + require.NoError(t, cache.DeleteVerificationCode(ctx, email)) + _, err = cache.IncrVerificationCodeAttempts(ctx, email) + require.Error(t, err) + require.False(t, mr.Exists(verifyCodeKey(email)+attemptsKeySuffix)) +} + +func TestEmailCache_PasswordResetTokenHashedAndSingleUse(t *testing.T) { + cache, mr, _ := newMiniredisEmailCache(t) + ctx := context.Background() + email := "reset@example.com" + + svc := service.NewEmailService(nil, cache) + + // Seed the token the same way SendPasswordResetEmail does (hash only). + token, err := svc.GeneratePasswordResetToken() + require.NoError(t, err) + sum := sha256.Sum256([]byte(token)) + require.NoError(t, cache.SetPasswordResetToken(ctx, email, &service.PasswordResetTokenData{ + Token: hex.EncodeToString(sum[:]), CreatedAt: time.Now(), + }, 30*time.Minute)) + + raw, err := mr.Get(passwordResetKey(email)) + require.NoError(t, err) + require.False(t, strings.Contains(raw, token), "plaintext token must not be stored") + + require.NoError(t, svc.VerifyPasswordResetToken(ctx, email, token)) + require.ErrorIs(t, svc.ConsumePasswordResetToken(ctx, email, "wrong"), service.ErrInvalidResetToken) + + const workers = 30 + var ok atomic.Int32 + var wg sync.WaitGroup + for i := 0; i < workers; i++ { + wg.Add(1) + go func() { + defer wg.Done() + if svc.ConsumePasswordResetToken(ctx, email, token) == nil { + ok.Add(1) + } + }() + } + wg.Wait() + require.Equal(t, int32(1), ok.Load()) + require.False(t, mr.Exists(passwordResetKey(email))) +} + +func TestEmailCache_ConsumePasswordResetTokenMismatchKeepsToken(t *testing.T) { + cache, mr, _ := newMiniredisEmailCache(t) + ctx := context.Background() + email := "keep@example.com" + require.NoError(t, cache.SetPasswordResetToken(ctx, email, &service.PasswordResetTokenData{Token: "abc"}, time.Minute)) + + ok, err := cache.ConsumePasswordResetToken(ctx, email, "xyz") + require.NoError(t, err) + require.False(t, ok) + require.True(t, mr.Exists(passwordResetKey(email))) + + ok, err = cache.ConsumePasswordResetToken(ctx, email, "abc") + require.NoError(t, err) + require.True(t, ok) + ok, err = cache.ConsumePasswordResetToken(ctx, email, "abc") + require.NoError(t, err) + require.False(t, ok) +} diff --git a/backend/internal/repository/http_upstream.go b/backend/internal/repository/http_upstream.go index 2c3eb8d57bb6..30a73fe042ed 100644 --- a/backend/internal/repository/http_upstream.go +++ b/backend/internal/repository/http_upstream.go @@ -509,6 +509,8 @@ func newGrokOfficialAPIFallbackRequest(req *http.Request) (*http.Request, error) for _, header := range []string{ "X-XAI-Token-Auth", "X-Grok-Client-Version", + "X-Grok-Client-Mode", + "X-Authenticateresponse", "X-Grok-Client-Surface", "X-UserID", "X-Email", @@ -574,6 +576,8 @@ func applyGrokCLIProxyHeaders(req *http.Request) { req.Header.Set("X-XAI-Token-Auth", xai.CLITokenAuth) req.Header.Set("x-grok-client-version", version) req.Header.Set("x-grok-client-identifier", xai.CLIClientIdentifier) + req.Header.Set("x-grok-client-mode", xai.CLIClientMode) + req.Header.Set("x-authenticateresponse", "authenticate-response") req.Header.Set("User-Agent", xai.CLIUserAgent(version)) } diff --git a/backend/internal/repository/http_upstream_test.go b/backend/internal/repository/http_upstream_test.go index f9c579064877..4f67f55a369f 100644 --- a/backend/internal/repository/http_upstream_test.go +++ b/backend/internal/repository/http_upstream_test.go @@ -237,6 +237,10 @@ func TestHTTPUpstreamDoAppliesGrokCLIIdentityBeforeOAuthRoundTrip(t *testing.T) require.NoError(t, resp.Body.Close()) require.Equal(t, xai.CLIClientVersion, capturedHeaders.Get("x-grok-client-version")) + require.Equal(t, "1.0.46", capturedHeaders.Get("x-grok-client-version")) + require.Equal(t, "grok-pager", capturedHeaders.Get("x-grok-client-identifier")) + require.Equal(t, "interactive", capturedHeaders.Get("x-grok-client-mode")) + require.Equal(t, "authenticate-response", capturedHeaders.Get("x-authenticateresponse")) require.Equal(t, "xai-grok-cli", capturedHeaders.Get("X-XAI-Token-Auth")) require.Equal(t, xai.CLIUserAgent(xai.CLIClientVersion), capturedHeaders.Get("User-Agent")) }) @@ -309,6 +313,8 @@ func TestHTTPUpstreamDoFallsBackToOfficialGrokAPIOnCLIAccessDenied(t *testing.T) require.Equal(t, "Bearer oauth-token", fallbackHeaders.Get("Authorization")) require.Empty(t, fallbackHeaders.Get("X-XAI-Token-Auth")) require.Empty(t, fallbackHeaders.Get("x-grok-client-version")) + require.Empty(t, fallbackHeaders.Get("x-grok-client-mode")) + require.Empty(t, fallbackHeaders.Get("x-authenticateresponse")) require.Empty(t, fallbackHeaders.Get("User-Agent")) } @@ -465,18 +471,18 @@ func TestApplyGrokCLIProxyHeaders(t *testing.T) { }) t.Run("accepts a valid operator override", func(t *testing.T) { - t.Setenv("XAI_GROK_CLI_VERSION", "1.0.45-alpha.1") + t.Setenv("XAI_GROK_CLI_VERSION", "1.0.47-alpha.1") req, err := http.NewRequest(http.MethodPost, "https://cli-chat-proxy.grok.com/v1/chat/completions", nil) require.NoError(t, err) applyGrokCLIProxyHeaders(req) - require.Equal(t, "1.0.45-alpha.1", req.Header.Get("x-grok-client-version")) - require.Equal(t, xai.CLIUserAgent("1.0.45-alpha.1"), req.Header.Get("User-Agent")) + require.Equal(t, "1.0.47-alpha.1", req.Header.Get("x-grok-client-version")) + require.Equal(t, xai.CLIUserAgent("1.0.47-alpha.1"), req.Header.Get("User-Agent")) }) t.Run("rejects an unsafe override", func(t *testing.T) { - t.Setenv("XAI_GROK_CLI_VERSION", "1.0.45\r\nX-Injected: true") + t.Setenv("XAI_GROK_CLI_VERSION", "1.0.47\r\nX-Injected: true") req, err := http.NewRequest(http.MethodPost, "https://cli-chat-proxy.grok.com/v1/responses", nil) require.NoError(t, err) @@ -487,7 +493,7 @@ func TestApplyGrokCLIProxyHeaders(t *testing.T) { }) t.Run("rejects an override below the supported minimum", func(t *testing.T) { - t.Setenv("XAI_GROK_CLI_VERSION", "1.0.43") + t.Setenv("XAI_GROK_CLI_VERSION", "1.0.45") req, err := http.NewRequest(http.MethodPost, "https://cli-chat-proxy.grok.com/v1/responses", nil) require.NoError(t, err) @@ -498,7 +504,7 @@ func TestApplyGrokCLIProxyHeaders(t *testing.T) { }) t.Run("rejects a prerelease override at the minimum version", func(t *testing.T) { - t.Setenv("XAI_GROK_CLI_VERSION", "1.0.44-beta.1") + t.Setenv("XAI_GROK_CLI_VERSION", "1.0.46-beta.1") req, err := http.NewRequest(http.MethodPost, "https://cli-chat-proxy.grok.com/v1/responses", nil) require.NoError(t, err) @@ -511,11 +517,11 @@ func TestApplyGrokCLIProxyHeaders(t *testing.T) { // Every entry sits above the pinned minimum, so a rejection here can only be // caused by the malformed semver and never by the version being too old. for _, version := range []string{ - "1.0.045", - "1.0.45-alpha..1", + "1.0.047", + "1.0.47-alpha..1", "1.1", "1", - "1.0.45+build.1", + "1.0.47+build.1", } { t.Run("rejects invalid semver "+version, func(t *testing.T) { t.Setenv("XAI_GROK_CLI_VERSION", version) @@ -530,7 +536,7 @@ func TestApplyGrokCLIProxyHeaders(t *testing.T) { } t.Run("leaves direct xAI API requests unchanged", func(t *testing.T) { - t.Setenv("XAI_GROK_CLI_VERSION", "1.0.44") + t.Setenv("XAI_GROK_CLI_VERSION", "1.0.46") req, err := http.NewRequest(http.MethodPost, "https://api.x.ai/v1/responses", nil) require.NoError(t, err) req.Header.Set("User-Agent", "direct-api-client/1.0") diff --git a/backend/internal/repository/usage_billing_deleted_key_integration_test.go b/backend/internal/repository/usage_billing_deleted_key_integration_test.go index d2d5059b24b4..9949bfa51709 100644 --- a/backend/internal/repository/usage_billing_deleted_key_integration_test.go +++ b/backend/internal/repository/usage_billing_deleted_key_integration_test.go @@ -16,6 +16,8 @@ import ( // The request has already captured its billing command when the owner deletes // the key. Replaying that interleaving must charge once and keep the key revoked. +// As upstream #7816 does, a deleted key's own quota and rate-limit counters are +// skipped while the user's balance or subscription is still charged. func TestUsageBillingRepositoryApply_SettlesSoftDeletedAPIKey(t *testing.T) { for _, subscriptionBilling := range []bool{false, true} { for _, limits := range []struct { @@ -139,10 +141,10 @@ func TestUsageBillingRepositoryApply_SettlesSoftDeletedAPIKey(t *testing.T) { "SELECT quota_used, usage_5h, usage_1d, usage_7d, key, status, deleted_at FROM api_keys WHERE id = $1", key.ID). Scan("aUsed, &usage5h, &usage1d, &usage7d, &keyAfter, &statusAfter, &deletedAfter)) expectedQuota, expectedWindow := 0.0, 0.0 - if limits.quota { + if limits.quota && !deleted { expectedQuota = 1.25 } - if limits.window { + if limits.window && !deleted { expectedWindow = 1.25 } require.InDelta(t, expectedQuota, quotaUsed, 1e-8) diff --git a/backend/internal/repository/usage_billing_deleted_key_unit_test.go b/backend/internal/repository/usage_billing_deleted_key_unit_test.go index f483d66346b9..2e412d54f53e 100644 --- a/backend/internal/repository/usage_billing_deleted_key_unit_test.go +++ b/backend/internal/repository/usage_billing_deleted_key_unit_test.go @@ -14,59 +14,79 @@ import ( "github.com/Wei-Shaw/sub2api/internal/service" ) -// Including soft-deleted rows must not turn a genuinely missing row or a -// database failure into a successful, partially applied billing transaction. +func expectDeletedKeySettlementPrefix(mock sqlmock.Sqlmock) { + mock.ExpectBegin() + mock.ExpectQuery("INSERT INTO usage_billing_dedup"). + WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(1)) + mock.ExpectQuery("SELECT request_fingerprint.*FROM usage_billing_dedup_archive"). + WillReturnRows(sqlmock.NewRows([]string{"request_fingerprint"})) + mock.ExpectQuery(conditionalBalanceDeductSQL). + WithArgs(1.25, int64(42)). + WillReturnRows(sqlmock.NewRows([]string{"balance"}).AddRow(98.75)) +} + +// A database failure on a key counter must not turn into a successful, +// partially applied billing transaction. func TestUsageBillingRepositoryApply_KeyUpdateFailureStillRollsBack(t *testing.T) { dbFailure := errors.New("database unavailable") for _, field := range []string{"quota", "window"} { - for _, missing := range []bool{false, true} { - name := field + "/database_error" - if missing { - name = field + "/missing_row" + t.Run(field+"/database_error", func(t *testing.T) { + db, mock, err := sqlmock.New() + require.NoError(t, err) + defer func() { _ = db.Close() }() + cmd := &service.UsageBillingCommand{ + RequestID: "failed-key-settlement", APIKeyID: 7, UserID: 42, BalanceCost: 1.25, + } + expectDeletedKeySettlementPrefix(mock) + if field == "quota" { + cmd.APIKeyQuotaCost = 1.25 + mock.ExpectQuery("UPDATE api_keys"). + WithArgs(1.25, int64(7), service.StatusAPIKeyActive, service.StatusAPIKeyQuotaExhausted). + WillReturnError(dbFailure) + } else { + cmd.APIKeyRateLimitCost = 1.25 + mock.ExpectExec("UPDATE api_keys").WithArgs(1.25, int64(7)).WillReturnError(dbFailure) + } + mock.ExpectRollback() + result, err := NewUsageBillingRepository(nil, db).Apply(context.Background(), cmd) + require.Nil(t, result) + require.ErrorIs(t, err, dbFailure) + require.NoError(t, mock.ExpectationsWereMet()) + }) + } +} + +// When the key row no longer matches (deleted mid-request), upstream #7816 +// skips the key's own counters and still commits the user's charge. +func TestUsageBillingRepositoryApply_MissingKeySkipsKeyCounters(t *testing.T) { + for _, field := range []string{"quota", "window"} { + t.Run(field, func(t *testing.T) { + db, mock, err := sqlmock.New() + require.NoError(t, err) + defer func() { _ = db.Close() }() + cmd := &service.UsageBillingCommand{ + RequestID: "deleted-key-settlement", APIKeyID: 7, UserID: 42, BalanceCost: 1.25, + } + expectDeletedKeySettlementPrefix(mock) + if field == "quota" { + cmd.APIKeyQuotaCost = 1.25 + mock.ExpectQuery("UPDATE api_keys"). + WithArgs(1.25, int64(7), service.StatusAPIKeyActive, service.StatusAPIKeyQuotaExhausted). + WillReturnError(sql.ErrNoRows) + } else { + cmd.APIKeyRateLimitCost = 1.25 + mock.ExpectExec("UPDATE api_keys").WithArgs(1.25, int64(7)). + WillReturnResult(sqlmock.NewResult(0, 0)) } - t.Run(name, func(t *testing.T) { - db, mock, err := sqlmock.New() - require.NoError(t, err) - defer func() { _ = db.Close() }() - cmd := &service.UsageBillingCommand{ - RequestID: "failed-key-settlement", APIKeyID: 7, UserID: 42, BalanceCost: 1.25, - } - mock.ExpectBegin() - mock.ExpectQuery("INSERT INTO usage_billing_dedup"). - WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(1)) - mock.ExpectQuery("SELECT request_fingerprint.*FROM usage_billing_dedup_archive"). - WillReturnRows(sqlmock.NewRows([]string{"request_fingerprint"})) - mock.ExpectQuery(conditionalBalanceDeductSQL). - WithArgs(1.25, int64(42)). - WillReturnRows(sqlmock.NewRows([]string{"balance"}).AddRow(98.75)) - if field == "quota" { - cmd.APIKeyQuotaCost = 1.25 - q := mock.ExpectQuery("UPDATE api_keys"). - WithArgs(1.25, int64(7), service.StatusAPIKeyActive, service.StatusAPIKeyQuotaExhausted) - if missing { - q.WillReturnError(sql.ErrNoRows) - } else { - q.WillReturnError(dbFailure) - } - } else { - cmd.APIKeyRateLimitCost = 1.25 - q := mock.ExpectExec("UPDATE api_keys").WithArgs(1.25, int64(7)) - if missing { - q.WillReturnResult(sqlmock.NewResult(0, 0)) - } else { - q.WillReturnError(dbFailure) - } - } - mock.ExpectRollback() - result, err := NewUsageBillingRepository(nil, db).Apply(context.Background(), cmd) - require.Nil(t, result) - wantErr := dbFailure - if missing { - wantErr = service.ErrAPIKeyNotFound - } - require.ErrorIs(t, err, wantErr) - require.NoError(t, mock.ExpectationsWereMet()) - }) - } + mock.ExpectCommit() + result, err := NewUsageBillingRepository(nil, db).Apply(context.Background(), cmd) + require.NoError(t, err) + require.NotNil(t, result) + require.True(t, result.Applied) + require.NotNil(t, result.NewBalance) + require.InDelta(t, 98.75, *result.NewBalance, 1e-9) + require.False(t, result.APIKeyQuotaExhausted) + require.NoError(t, mock.ExpectationsWereMet()) + }) } } diff --git a/backend/internal/repository/usage_billing_repo.go b/backend/internal/repository/usage_billing_repo.go index e697b1e83f1b..bfb1bb0de5f4 100644 --- a/backend/internal/repository/usage_billing_repo.go +++ b/backend/internal/repository/usage_billing_repo.go @@ -187,16 +187,17 @@ func (r *usageBillingRepository) applyUsageBillingEffects(ctx context.Context, t result.BalanceOverdrafted = !sufficient } + // Key 已不存在时跳过其自身的额度/限速计数,其余结算项不受影响。 if cmd.APIKeyQuotaCost > 0 { exhausted, err := incrementUsageBillingAPIKeyQuota(ctx, tx, cmd.APIKeyID, cmd.APIKeyQuotaCost) - if err != nil { + if err != nil && !errors.Is(err, service.ErrAPIKeyNotFound) { return err } result.APIKeyQuotaExhausted = exhausted } if cmd.APIKeyRateLimitCost > 0 { - if err := incrementUsageBillingAPIKeyRateLimit(ctx, tx, cmd.APIKeyID, cmd.APIKeyRateLimitCost); err != nil { + if err := incrementUsageBillingAPIKeyRateLimit(ctx, tx, cmd.APIKeyID, cmd.APIKeyRateLimitCost); err != nil && !errors.Is(err, service.ErrAPIKeyNotFound) { return err } } @@ -414,15 +415,12 @@ func userExistsForBilling(ctx context.Context, tx *sql.Tx, userID int64) (bool, } func incrementUsageBillingAPIKeyQuota(ctx context.Context, tx *sql.Tx, apiKeyID int64, amount float64) (bool, error) { - // In-flight requests still owe their usage after the key is soft-deleted. - // Bill the retained row by ID, but never change a deleted key's status or - // report it as newly exhausted. Authentication must continue to exclude it. var exhausted bool err := tx.QueryRowContext(ctx, ` UPDATE api_keys SET quota_used = quota_used + $1, status = CASE - WHEN deleted_at IS NULL AND quota > 0 + WHEN quota > 0 AND status = $3 AND quota_used < quota AND quota_used + $1 >= quota @@ -430,8 +428,8 @@ func incrementUsageBillingAPIKeyQuota(ctx context.Context, tx *sql.Tx, apiKeyID ELSE status END, updated_at = NOW() - WHERE id = $2 - RETURNING deleted_at IS NULL AND quota > 0 AND quota_used >= quota AND quota_used - $1 < quota + WHERE id = $2 AND deleted_at IS NULL + RETURNING quota > 0 AND quota_used >= quota AND quota_used - $1 < quota `, amount, apiKeyID, service.StatusAPIKeyActive, service.StatusAPIKeyQuotaExhausted).Scan(&exhausted) if errors.Is(err, sql.ErrNoRows) { return false, service.ErrAPIKeyNotFound @@ -443,8 +441,6 @@ func incrementUsageBillingAPIKeyQuota(ctx context.Context, tx *sql.Tx, apiKeyID } func incrementUsageBillingAPIKeyRateLimit(ctx context.Context, tx *sql.Tx, apiKeyID int64, cost float64) error { - // As with lifetime quota, deletion must not roll back an accepted request's - // balance/subscription charge. Only this settlement path includes tombstones. res, err := tx.ExecContext(ctx, ` UPDATE api_keys SET usage_5h = CASE WHEN window_5h_start IS NOT NULL AND window_5h_start + INTERVAL '5 hours' <= NOW() THEN $1 ELSE usage_5h + $1 END, @@ -454,7 +450,7 @@ func incrementUsageBillingAPIKeyRateLimit(ctx context.Context, tx *sql.Tx, apiKe window_1d_start = CASE WHEN window_1d_start IS NULL OR window_1d_start + INTERVAL '24 hours' <= NOW() THEN date_trunc('day', NOW()) ELSE window_1d_start END, window_7d_start = CASE WHEN window_7d_start IS NULL OR window_7d_start + INTERVAL '7 days' <= NOW() THEN date_trunc('day', NOW()) ELSE window_7d_start END, updated_at = NOW() - WHERE id = $2 + WHERE id = $2 AND deleted_at IS NULL `, cost, apiKeyID) if err != nil { return err diff --git a/backend/internal/repository/usage_billing_repo_integration_test.go b/backend/internal/repository/usage_billing_repo_integration_test.go index e8d4d32707fb..9f2f0169dcec 100644 --- a/backend/internal/repository/usage_billing_repo_integration_test.go +++ b/backend/internal/repository/usage_billing_repo_integration_test.go @@ -162,6 +162,61 @@ func TestUsageBillingRepositoryApply_RequestFingerprintConflict(t *testing.T) { require.ErrorIs(t, err, service.ErrUsageBillingRequestConflict) } +func TestUsageBillingRepositoryApply_DeletedAPIKeyStillBillsBalance(t *testing.T) { + ctx := context.Background() + client := testEntClient(t) + repo := NewUsageBillingRepository(client, integrationDB) + + user := mustCreateUser(t, client, &service.User{ + Email: fmt.Sprintf("usage-billing-deleted-key-user-%d@example.com", time.Now().UnixNano()), + PasswordHash: "hash", + Balance: 100, + }) + apiKey := mustCreateApiKey(t, client, &service.APIKey{ + UserID: user.ID, + Key: "sk-usage-billing-deleted-key-" + uuid.NewString(), + Name: "billing-deleted-key", + Quota: 50, + RateLimit5h: 50, + }) + account := mustCreateAccount(t, client, &service.Account{ + Name: "usage-billing-deleted-key-account-" + uuid.NewString(), + Type: service.AccountTypeAPIKey, + }) + + _, err := integrationDB.ExecContext(ctx, "UPDATE api_keys SET deleted_at = NOW() WHERE id = $1", apiKey.ID) + require.NoError(t, err) + + requestID := uuid.NewString() + result, err := repo.Apply(ctx, &service.UsageBillingCommand{ + RequestID: requestID, + APIKeyID: apiKey.ID, + UserID: user.ID, + AccountID: account.ID, + AccountType: service.AccountTypeAPIKey, + BalanceCost: 1.25, + APIKeyQuotaCost: 1.25, + APIKeyRateLimitCost: 1.25, + }) + require.NoError(t, err) + require.NotNil(t, result) + require.True(t, result.Applied) + require.False(t, result.APIKeyQuotaExhausted) + + var balance float64 + require.NoError(t, integrationDB.QueryRowContext(ctx, "SELECT balance FROM users WHERE id = $1", user.ID).Scan(&balance)) + require.InDelta(t, 98.75, balance, 0.000001) + + var quotaUsed, usage5h float64 + require.NoError(t, integrationDB.QueryRowContext(ctx, "SELECT quota_used, usage_5h FROM api_keys WHERE id = $1", apiKey.ID).Scan("aUsed, &usage5h)) + require.InDelta(t, 0, quotaUsed, 0.000001) + require.InDelta(t, 0, usage5h, 0.000001) + + var dedupCount int + require.NoError(t, integrationDB.QueryRowContext(ctx, "SELECT COUNT(*) FROM usage_billing_dedup WHERE request_id = $1 AND api_key_id = $2", requestID, apiKey.ID).Scan(&dedupCount)) + require.Equal(t, 1, dedupCount) +} + func TestUsageBillingRepositoryApply_UpdatesAccountQuota(t *testing.T) { ctx := context.Background() client := testEntClient(t) diff --git a/backend/internal/repository/usage_billing_repo_unit_test.go b/backend/internal/repository/usage_billing_repo_unit_test.go index 6fe4d4fc2632..40394d819c8d 100644 --- a/backend/internal/repository/usage_billing_repo_unit_test.go +++ b/backend/internal/repository/usage_billing_repo_unit_test.go @@ -20,6 +20,8 @@ const ( captureBatchImageHoldSQL = `(?s)UPDATE users\s+SET balance = balance\s+\+ CASE WHEN \$1 > \$2 THEN \$1 - \$2 ELSE 0 END\s+- CASE WHEN \$2 > \$1 THEN \$2 - \$1 ELSE 0 END,\s+frozen_balance = COALESCE\(frozen_balance, 0\) - \$1,\s+updated_at = NOW\(\)\s+WHERE id = \$3 AND deleted_at IS NULL AND COALESCE\(frozen_balance, 0\) >= \$1\s+RETURNING balance, frozen_balance` releaseBatchImageHoldSQL = `(?s)UPDATE users\s+SET balance = balance \+ \$1,\s+frozen_balance = COALESCE\(frozen_balance, 0\) - \$1,\s+updated_at = NOW\(\)\s+WHERE id = \$2 AND deleted_at IS NULL AND COALESCE\(frozen_balance, 0\) >= \$1\s+RETURNING balance, frozen_balance` userExistsForBillingSQL = `(?s)SELECT 1\s+FROM users\s+WHERE id = \$1 AND deleted_at IS NULL` + apiKeyQuotaIncrementSQL = `(?s)UPDATE api_keys\s+SET quota_used = quota_used \+ \$1,.*WHERE id = \$2 AND deleted_at IS NULL\s+RETURNING` + apiKeyRateLimitIncrementSQL = `(?s)UPDATE api_keys SET\s+usage_5h = .*WHERE id = \$2 AND deleted_at IS NULL` ) func TestDeductUsageBillingBalance_UsesSufficientBalanceGuard(t *testing.T) { @@ -99,6 +101,66 @@ func TestApplyUsageBillingEffects_FlagsBalanceOverdraft(t *testing.T) { require.NoError(t, mock.ExpectationsWereMet()) } +func TestApplyUsageBillingEffects_DeletedAPIKeyStillBillsBalance(t *testing.T) { + ctx := context.Background() + db, mock, err := sqlmock.New() + require.NoError(t, err) + defer func() { _ = db.Close() }() + + mock.ExpectBegin() + tx, err := db.BeginTx(ctx, nil) + require.NoError(t, err) + mock.ExpectQuery(conditionalBalanceDeductSQL). + WithArgs(10.0, int64(42)). + WillReturnRows(sqlmock.NewRows([]string{"balance"}).AddRow(90.0)) + mock.ExpectQuery(apiKeyQuotaIncrementSQL). + WithArgs(10.0, int64(7), service.StatusAPIKeyActive, service.StatusAPIKeyQuotaExhausted). + WillReturnError(sql.ErrNoRows) + mock.ExpectExec(apiKeyRateLimitIncrementSQL). + WithArgs(10.0, int64(7)). + WillReturnResult(sqlmock.NewResult(0, 0)) + mock.ExpectCommit() + + result := &service.UsageBillingApplyResult{Applied: true} + err = (&usageBillingRepository{}).applyUsageBillingEffects(ctx, tx, &service.UsageBillingCommand{ + UserID: 42, + APIKeyID: 7, + BalanceCost: 10, + APIKeyQuotaCost: 10, + APIKeyRateLimitCost: 10, + }, result) + require.NoError(t, err) + require.NotNil(t, result.NewBalance) + require.InDelta(t, 90.0, *result.NewBalance, 0.000001) + require.False(t, result.APIKeyQuotaExhausted) + require.NoError(t, tx.Commit()) + require.NoError(t, mock.ExpectationsWereMet()) +} + +func TestApplyUsageBillingEffects_APIKeyCounterErrorStillFails(t *testing.T) { + ctx := context.Background() + db, mock, err := sqlmock.New() + require.NoError(t, err) + defer func() { _ = db.Close() }() + + mock.ExpectBegin() + tx, err := db.BeginTx(ctx, nil) + require.NoError(t, err) + mock.ExpectExec(apiKeyRateLimitIncrementSQL). + WithArgs(10.0, int64(7)). + WillReturnError(sql.ErrConnDone) + mock.ExpectRollback() + + err = (&usageBillingRepository{}).applyUsageBillingEffects(ctx, tx, &service.UsageBillingCommand{ + UserID: 42, + APIKeyID: 7, + APIKeyRateLimitCost: 10, + }, &service.UsageBillingApplyResult{Applied: true}) + require.ErrorIs(t, err, sql.ErrConnDone) + require.NoError(t, tx.Rollback()) + require.NoError(t, mock.ExpectationsWereMet()) +} + func TestDeductUsageBillingBalance_ReturnsUserNotFoundWhenNoUserUpdated(t *testing.T) { ctx := context.Background() db, mock, err := sqlmock.New() diff --git a/backend/internal/securityaudit/prompt_snapshot.go b/backend/internal/securityaudit/prompt_snapshot.go index b47e8467417e..74a0d2597436 100644 --- a/backend/internal/securityaudit/prompt_snapshot.go +++ b/backend/internal/securityaudit/prompt_snapshot.go @@ -106,6 +106,8 @@ func extractProtocolSegments(protocol string, document any) []promptSegment { return append(extractInstructions(root["instructions"]), extractResponses(root["input"])...) case "openai_images", "grok_media", "media", "images": return userPromptSegments(extractMediaPrompts(root)) + case "typesafe_systemone": + return extractSystemOneSegments(root) default: if segments := extractChatLikeSegments(root); len(segments) > 0 { return segments @@ -120,6 +122,79 @@ func extractProtocolSegments(protocol string, document any) []promptSegment { } } +// extractSystemOneSegments collects every client-controlled text of a TypeSafe +// System One request: question IDs, every question field except the validated +// type, unknown top-level extension fields, and the evaluated state. Object keys +// are text too (Jev reads the whole JSON), so they are collected with values. +// The state comes last so it is the prioritized segment, and keys are visited +// in sorted order to keep the prompt hash stable across requests. +func extractSystemOneSegments(root map[string]any) []promptSegment { + if root == nil { + return nil + } + texts := make([]string, 0, 4) + questions, isObject := root["questions"].(map[string]any) + if !isObject { + texts = appendJSONStringLeaves(texts, root["questions"]) + } + for _, id := range sortedJSONKeys(questions) { + texts = appendJSONStringLeaves(texts, id) + question, ok := questions[id].(map[string]any) + if !ok { + texts = appendJSONStringLeaves(texts, questions[id]) + continue + } + for _, field := range sortedJSONKeys(question) { + switch field { + case "type": + continue + case "instructions", "criteria": + default: + texts = appendJSONStringLeaves(texts, field) + } + texts = appendJSONStringLeaves(texts, question[field]) + } + } + for _, field := range sortedJSONKeys(root) { + switch field { + case "model", "stream", "state", "questions": + continue + } + texts = appendJSONStringLeaves(texts, field) + texts = appendJSONStringLeaves(texts, root[field]) + } + texts = appendJSONStringLeaves(texts, root["state"]) + return userPromptSegments(texts) +} + +func appendJSONStringLeaves(texts []string, value any) []string { + switch typed := value.(type) { + case string: + if text := strings.TrimSpace(typed); text != "" { + texts = append(texts, text) + } + case []any: + for _, item := range typed { + texts = appendJSONStringLeaves(texts, item) + } + case map[string]any: + for _, key := range sortedJSONKeys(typed) { + texts = appendJSONStringLeaves(texts, key) + texts = appendJSONStringLeaves(texts, typed[key]) + } + } + return texts +} + +func sortedJSONKeys(values map[string]any) []string { + keys := make([]string, 0, len(values)) + for key := range values { + keys = append(keys, key) + } + sort.Strings(keys) + return keys +} + // clientInstructionRoles are roles a client may freely populate. Attackers can // place jailbreak/PII text in assistant/tool turns, so blocking audit must scan // them too—not only user/system/developer instructions. diff --git a/backend/internal/securityaudit/prompt_snapshot_test.go b/backend/internal/securityaudit/prompt_snapshot_test.go index 70d26f925c66..751043f550ee 100644 --- a/backend/internal/securityaudit/prompt_snapshot_test.go +++ b/backend/internal/securityaudit/prompt_snapshot_test.go @@ -412,3 +412,53 @@ func mustJSON(t *testing.T, value string) []byte { func metadataTextForTest(scanText string) string { return strings.Replace(scanText, promptAuditPrioritySeparator, "\n\n", 1) } + +func TestPromptSnapshotTypeSafeSystemOneCollectsStateAndQuestions(t *testing.T) { + body := `{"model":"jev-latest","state":{"STATE_KEY":"STATE_TITLE","items":["STATE_ITEM",{"text":"STATE_NESTED"},3]},` + + `"questions":{"b":{"type":"choice","instructions":"CHOICE_INSTRUCTIONS","criteria":{"OPTION_LABEL":{"description":"OPTION_DESC"},"empty":null}},` + + `"a":{"type":"noul","instructions":["NOUL_INSTRUCTIONS"],"criteria":{"true":"NOUL_TRUE","false":"NOUL_FALSE"},"EXTENSION_KEY":"EXTENSION_VALUE"},` + + `"QUESTION_ID":{"type":"score","criteria":["SCORE_LOW",{"description":"SCORE_HIGH"}]}},"TOP_EXTENSION":{"x":"TOP_VALUE"}}` + + snapshot, err := ExtractPromptSnapshot(Request{Protocol: "typesafe_systemone", Body: []byte(body)}) + require.NoError(t, err) + for _, text := range []string{"STATE_KEY", "STATE_TITLE", "STATE_ITEM", "STATE_NESTED", "CHOICE_INSTRUCTIONS", "OPTION_LABEL", "OPTION_DESC", + "NOUL_INSTRUCTIONS", "NOUL_TRUE", "NOUL_FALSE", "EXTENSION_KEY", "EXTENSION_VALUE", "QUESTION_ID", "SCORE_LOW", "SCORE_HIGH", + "TOP_EXTENSION", "TOP_VALUE"} { + require.Contains(t, snapshot.ScanText, text) + } + // Canonical field names and validated enum values are not client text. + for _, text := range []string{"jev-latest", "instructions", "criteria", "questions", "choice", "score"} { + require.NotContains(t, snapshot.ScanText, text) + } + + // Map iteration order must not change the audited text or its hash. + for range 20 { + again, err := ExtractPromptSnapshot(Request{Protocol: "typesafe_systemone", Body: []byte(body)}) + require.NoError(t, err) + require.Equal(t, snapshot.PromptHash, again.PromptHash) + require.Equal(t, snapshot.ScanText, again.ScanText) + } + + blocking, err := ExtractBlockingPromptSnapshot(Request{Protocol: "typesafe_systemone", Body: []byte(body)}, true) + require.NoError(t, err) + require.Contains(t, blocking.ScanText, "STATE_TITLE") + require.Contains(t, blocking.ScanText, "CHOICE_INSTRUCTIONS") + require.Contains(t, blocking.ScanText, "QUESTION_ID") +} + +func TestPromptSnapshotTypeSafeSystemOneStringStateIsPrioritized(t *testing.T) { + snapshot, err := ExtractPromptSnapshot(Request{Protocol: "typesafe_systemone", Body: []byte(`{"model":"jev-latest","state":"plain state","questions":{"q":{"type":"noul"}}}`)}) + require.NoError(t, err) + require.True(t, strings.HasPrefix(snapshot.ScanText, "plain state")) + require.Equal(t, 2, snapshot.MessageCount) +} + +func TestPromptSnapshotTypeSafeSystemOneAuditsKeyOnlyPayloads(t *testing.T) { + body := `{"model":"jev-latest","state":{"HIDDEN_STATE_KEY":1},"questions":{"HIDDEN_QUESTION_ID":{"type":"noul"}}}` + for _, latestTurnOnly := range []bool{false, true} { + snapshot, err := ExtractBlockingPromptSnapshot(Request{Protocol: "typesafe_systemone", Body: []byte(body)}, latestTurnOnly) + require.NoError(t, err) + require.Contains(t, snapshot.ScanText, "HIDDEN_STATE_KEY") + require.Contains(t, snapshot.ScanText, "HIDDEN_QUESTION_ID") + } +} diff --git a/backend/internal/server/api_contract_test.go b/backend/internal/server/api_contract_test.go index f9fab5ed28cd..7c86a29bf767 100644 --- a/backend/internal/server/api_contract_test.go +++ b/backend/internal/server/api_contract_test.go @@ -866,7 +866,7 @@ func TestAPIContracts(t *testing.T) { "force_email_on_third_party_signup": false, "default_concurrency": 5, "default_balance": 1.25, - "default_platform_quotas": {"anthropic":{"daily":null,"weekly":null,"monthly":null},"antigravity":{"daily":null,"weekly":null,"monthly":null},"deepseek":{"daily":null,"weekly":null,"monthly":null},"gemini":{"daily":null,"weekly":null,"monthly":null},"grok":{"daily":null,"weekly":null,"monthly":null},"kimi":{"daily":null,"weekly":null,"monthly":null},"minimax":{"daily":null,"weekly":null,"monthly":null},"openai":{"daily":null,"weekly":null,"monthly":null},"opencode_go":{"daily":null,"weekly":null,"monthly":null},"zhipu":{"daily":null,"weekly":null,"monthly":null}}, + "default_platform_quotas": {"anthropic":{"daily":null,"weekly":null,"monthly":null},"antigravity":{"daily":null,"weekly":null,"monthly":null},"deepseek":{"daily":null,"weekly":null,"monthly":null},"gemini":{"daily":null,"weekly":null,"monthly":null},"grok":{"daily":null,"weekly":null,"monthly":null},"kimi":{"daily":null,"weekly":null,"monthly":null},"minimax":{"daily":null,"weekly":null,"monthly":null},"openai":{"daily":null,"weekly":null,"monthly":null},"opencode_go":{"daily":null,"weekly":null,"monthly":null},"typesafe":{"daily":null,"weekly":null,"monthly":null},"zhipu":{"daily":null,"weekly":null,"monthly":null}}, "auth_source_default_email_platform_quotas": null, "auth_source_default_github_platform_quotas": null, "auth_source_default_google_platform_quotas": null, @@ -1001,6 +1001,9 @@ func TestAPIContracts(t *testing.T) { "payment_balance_recharge_multiplier": 0, "payment_subscription_usd_to_cny_rate": 0, "payment_recharge_fee_rate": 0, + "payment_recharge_bonus_tiers": [], + "payment_recharge_bonus_mode": "bonus", + "payment_recharge_bonus_notice": "", "payment_load_balance_strategy": "", "payment_product_name_prefix": "", "payment_product_name_suffix": "", @@ -1219,7 +1222,7 @@ func TestAPIContracts(t *testing.T) { "purchase_subscription_url": "", "table_default_page_size": 20, "table_page_size_options": [10, 20, 50], - "default_platform_quotas": {"anthropic":{"daily":null,"weekly":null,"monthly":null},"antigravity":{"daily":null,"weekly":null,"monthly":null},"deepseek":{"daily":null,"weekly":null,"monthly":null},"gemini":{"daily":null,"weekly":null,"monthly":null},"grok":{"daily":null,"weekly":null,"monthly":null},"kimi":{"daily":null,"weekly":null,"monthly":null},"minimax":{"daily":null,"weekly":null,"monthly":null},"openai":{"daily":null,"weekly":null,"monthly":null},"opencode_go":{"daily":null,"weekly":null,"monthly":null},"zhipu":{"daily":null,"weekly":null,"monthly":null}}, + "default_platform_quotas": {"anthropic":{"daily":null,"weekly":null,"monthly":null},"antigravity":{"daily":null,"weekly":null,"monthly":null},"deepseek":{"daily":null,"weekly":null,"monthly":null},"gemini":{"daily":null,"weekly":null,"monthly":null},"grok":{"daily":null,"weekly":null,"monthly":null},"kimi":{"daily":null,"weekly":null,"monthly":null},"minimax":{"daily":null,"weekly":null,"monthly":null},"openai":{"daily":null,"weekly":null,"monthly":null},"opencode_go":{"daily":null,"weekly":null,"monthly":null},"typesafe":{"daily":null,"weekly":null,"monthly":null},"zhipu":{"daily":null,"weekly":null,"monthly":null}}, "auth_source_default_email_platform_quotas": null, "auth_source_default_github_platform_quotas": null, "auth_source_default_google_platform_quotas": null, @@ -1351,6 +1354,9 @@ func TestAPIContracts(t *testing.T) { "payment_balance_recharge_multiplier": 0, "payment_subscription_usd_to_cny_rate": 0, "payment_recharge_fee_rate": 0, + "payment_recharge_bonus_tiers": [], + "payment_recharge_bonus_mode": "bonus", + "payment_recharge_bonus_notice": "", "payment_load_balance_strategy": "", "payment_product_name_prefix": "", "payment_product_name_suffix": "", diff --git a/backend/internal/server/router.go b/backend/internal/server/router.go index 18310ee5ec6c..77e92254b20b 100644 --- a/backend/internal/server/router.go +++ b/backend/internal/server/router.go @@ -137,7 +137,7 @@ func registerRoutes( routes.RegisterPublicPelicanShowcaseRoutes(v1, h, apiKeyAuth, panelRateLimiter) routes.RegisterAdminRoutes(v1, h, adminAuth, auditLog, stepUpAuth, settingService, panelRateLimiter) routes.RegisterGatewayRoutes(r, h, apiKeyAuth, apiKeyService, subscriptionService, opsService, settingService, compositeResolver, cfg) - routes.RegisterPaymentRoutes(v1, h.Payment, h.PaymentWebhook, h.Admin.Payment, jwtAuth, adminAuth, auditLog, settingService, panelRateLimiter) + routes.RegisterPaymentRoutes(v1, h.Payment, h.PaymentWebhook, h.Admin.Payment, jwtAuth, adminAuth, auditLog, settingService, panelRateLimiter, redisClient) handler.RegisterPageRoutes(v1, cfg.Pricing.DataDir, gin.HandlerFunc(jwtAuth), gin.HandlerFunc(adminAuth), settingService) } diff --git a/backend/internal/server/routes/gateway.go b/backend/internal/server/routes/gateway.go index 3d8da50d1251..83790ea06c92 100644 --- a/backend/internal/server/routes/gateway.go +++ b/backend/internal/server/routes/gateway.go @@ -238,6 +238,8 @@ func RegisterGatewayRoutes( } h.Gateway.Messages(c) }) + // System One carries only JSON text, so it uses the text body limit. + gateway.POST("/systemone", textBodyLimit, h.Gateway.SystemOne) // /v1/messages/count_tokens: OpenAI bridges upstream, Grok estimates // locally, and Anthropic-compatible platforms retain their existing path. gateway.POST("/messages/count_tokens", countTokensHandler) diff --git a/backend/internal/server/routes/gateway_model_allowlist_test.go b/backend/internal/server/routes/gateway_model_allowlist_test.go index b5c8029c0f68..156333bf48be 100644 --- a/backend/internal/server/routes/gateway_model_allowlist_test.go +++ b/backend/internal/server/routes/gateway_model_allowlist_test.go @@ -159,6 +159,7 @@ func TestGatewayRoutesGroupModelAllowlistCoversRootAliasRoutes(t *testing.T) { {http.MethodGet, "/realtime?model=gpt-4.1", ""}, {http.MethodPost, "/v1/responses", `{"model":"gpt-4.1"}`}, {http.MethodPost, "/v1/messages", `{"model":"gpt-4.1"}`}, + {http.MethodPost, "/v1/systemone", `{"model":"gpt-4.1","state":"x","questions":{"q":{"type":"noul","instructions":"x"}}}`}, {http.MethodPost, "/v1/messages/count_tokens", `{"model":"gpt-4.1","messages":[]}`}, {http.MethodPost, "/v1/chat/completions", `{"model":"gpt-4.1"}`}, {http.MethodPost, "/v1/embeddings", `{"model":"gpt-4.1","input":"hi"}`}, diff --git a/backend/internal/server/routes/payment.go b/backend/internal/server/routes/payment.go index ecda25f538cc..26d0fcd66973 100644 --- a/backend/internal/server/routes/payment.go +++ b/backend/internal/server/routes/payment.go @@ -1,12 +1,25 @@ package routes import ( + "time" + "github.com/Wei-Shaw/sub2api/internal/handler" "github.com/Wei-Shaw/sub2api/internal/handler/admin" + ratelimit "github.com/Wei-Shaw/sub2api/internal/middleware" "github.com/Wei-Shaw/sub2api/internal/server/middleware" "github.com/Wei-Shaw/sub2api/internal/service" "github.com/gin-gonic/gin" + "github.com/redis/go-redis/v9" +) + +// publicOrderVerifyRateLimit caps anonymous legacy out_trade_no lookups per +// client IP. The payment result page polls at most a handful of times per +// order, so this leaves ample headroom for real users while making +// out_trade_no enumeration impractical. +const ( + publicOrderVerifyRateLimit = 20 + publicOrderVerifyRateLimitWindow = time.Minute ) // RegisterPaymentRoutes registers all payment-related routes: @@ -21,6 +34,7 @@ func RegisterPaymentRoutes( auditLog middleware.AuditLogMiddleware, settingService *service.SettingService, panelRateLimiter *middleware.PanelRateLimiter, + redisClient *redis.Client, ) { // --- User-facing payment endpoints (authenticated) --- authenticated := v1.Group("/payment") @@ -50,9 +64,15 @@ func RegisterPaymentRoutes( // Signed resume-token recovery is the preferred public lookup path. // The legacy anonymous out_trade_no verify endpoint remains available as a // persisted-state compatibility path for staggered upgrades. + // The anonymous verify endpoint is IP rate-limited to prevent out_trade_no + // enumeration. It fails open on Redis errors so an outage never blocks a + // user who is mid-payment from seeing their result. + publicRateLimiter := ratelimit.NewRateLimiter(redisClient) public := v1.Group("/payment/public") { - public.POST("/orders/verify", paymentHandler.VerifyOrderPublic) + public.POST("/orders/verify", + publicRateLimiter.Limit("payment-public-order-verify", publicOrderVerifyRateLimit, publicOrderVerifyRateLimitWindow), + paymentHandler.VerifyOrderPublic) public.POST("/orders/resolve", paymentHandler.ResolveOrderPublicByResumeToken) } diff --git a/backend/internal/server/routes/payment_public_rate_limit_test.go b/backend/internal/server/routes/payment_public_rate_limit_test.go new file mode 100644 index 000000000000..3f8c86eab312 --- /dev/null +++ b/backend/internal/server/routes/payment_public_rate_limit_test.go @@ -0,0 +1,84 @@ +package routes + +import ( + "net/http" + "net/http/httptest" + "strings" + "testing" + "time" + + "github.com/Wei-Shaw/sub2api/internal/handler" + "github.com/Wei-Shaw/sub2api/internal/handler/admin" + servermiddleware "github.com/Wei-Shaw/sub2api/internal/server/middleware" + "github.com/alicebob/miniredis/v2" + "github.com/gin-gonic/gin" + "github.com/redis/go-redis/v9" + "github.com/stretchr/testify/require" +) + +func newPaymentRoutesTestRouter(redisClient *redis.Client) *gin.Engine { + gin.SetMode(gin.TestMode) + router := gin.New() + v1 := router.Group("/api/v1") + noop := func(c *gin.Context) { c.Next() } + + RegisterPaymentRoutes( + v1, + &handler.PaymentHandler{}, + &handler.PaymentWebhookHandler{}, + &admin.PaymentHandler{}, + servermiddleware.JWTAuthMiddleware(noop), + servermiddleware.AdminAuthMiddleware(noop), + servermiddleware.AuditLogMiddleware(noop), + nil, + nil, + redisClient, + ) + return router +} + +func postPublicOrderVerify(router *gin.Engine, remoteAddr string) *httptest.ResponseRecorder { + // Empty body fails binding in the handler, so no service is touched and a + // non-429 response proves the request passed the limiter. + req := httptest.NewRequest(http.MethodPost, "/api/v1/payment/public/orders/verify", strings.NewReader(`{}`)) + req.Header.Set("Content-Type", "application/json") + req.RemoteAddr = remoteAddr + w := httptest.NewRecorder() + router.ServeHTTP(w, req) + return w +} + +func TestPublicOrderVerifyRateLimitedPerIP(t *testing.T) { + mr := miniredis.RunT(t) + rdb := redis.NewClient(&redis.Options{Addr: mr.Addr()}) + t.Cleanup(func() { _ = rdb.Close() }) + + router := newPaymentRoutesTestRouter(rdb) + + for i := 1; i <= publicOrderVerifyRateLimit; i++ { + w := postPublicOrderVerify(router, "198.51.100.20:1234") + require.Equal(t, http.StatusBadRequest, w.Code, "request %d should reach the handler", i) + } + + w := postPublicOrderVerify(router, "198.51.100.20:1234") + require.Equal(t, http.StatusTooManyRequests, w.Code) + require.Contains(t, w.Body.String(), "rate limit exceeded") + + // A different client IP has its own budget. + w = postPublicOrderVerify(router, "198.51.100.21:1234") + require.Equal(t, http.StatusBadRequest, w.Code) +} + +func TestPublicOrderVerifyRateLimitFailsOpenWhenRedisUnavailable(t *testing.T) { + rdb := redis.NewClient(&redis.Options{ + Addr: "127.0.0.1:1", + DialTimeout: 50 * time.Millisecond, + ReadTimeout: 50 * time.Millisecond, + WriteTimeout: 50 * time.Millisecond, + }) + t.Cleanup(func() { _ = rdb.Close() }) + + router := newPaymentRoutesTestRouter(rdb) + w := postPublicOrderVerify(router, "203.0.113.30:1234") + require.Equal(t, http.StatusBadRequest, w.Code, "users mid-payment must not be blocked by a Redis outage") +} diff --git a/backend/internal/server/routes/prompt_audit_route_coverage_test.go b/backend/internal/server/routes/prompt_audit_route_coverage_test.go index a6070544579b..413665625c4f 100644 --- a/backend/internal/server/routes/prompt_audit_route_coverage_test.go +++ b/backend/internal/server/routes/prompt_audit_route_coverage_test.go @@ -29,6 +29,7 @@ func TestEveryGatewayPOSTRouteIsClassifiedForPromptAuditCoverage(t *testing.T) { audited := map[string][]string{ "/messages": {"gateway_handler.go", "openai_gateway_handler.go"}, + "/systemone": {"gateway_systemone.go"}, "/responses": {"gateway_handler_responses.go", "openai_gateway_handler.go"}, "/responses/*subpath": {"gateway_handler_responses.go", "openai_gateway_handler.go"}, "/chat/completions": {"gateway_handler_chat_completions.go", "openai_chat_completions.go"}, diff --git a/backend/internal/service/account.go b/backend/internal/service/account.go index 8fb84d39ed64..e59f399c300f 100644 --- a/backend/internal/service/account.go +++ b/backend/internal/service/account.go @@ -17,6 +17,7 @@ import ( "github.com/Wei-Shaw/sub2api/internal/domain" "github.com/Wei-Shaw/sub2api/internal/pkg/geminicli" "github.com/Wei-Shaw/sub2api/internal/pkg/openai_compat" + "github.com/Wei-Shaw/sub2api/internal/pkg/typesafe" "github.com/Wei-Shaw/sub2api/internal/pkg/xai" ) @@ -294,6 +295,10 @@ func (a *Account) IsGrok() bool { return a.Platform == PlatformGrok } +func (a *Account) IsTypeSafe() bool { + return a != nil && a.Platform == PlatformTypeSafe +} + func (a *Account) IsGrokOAuth() bool { return a.IsGrok() && a.Type == AccountTypeOAuth } @@ -1049,6 +1054,10 @@ func (a *Account) GetBaseURL() string { } baseURL := a.GetCredential("base_url") if baseURL == "" { + // TypeSafe keys must never fall back to the Anthropic host. + if a.Platform == PlatformTypeSafe { + return typesafe.DefaultBaseURL + } return "https://api.anthropic.com" } if a.Platform == PlatformAntigravity { @@ -1070,6 +1079,28 @@ func (a *Account) GetGeminiBaseURL(defaultBaseURL string) string { return baseURL } +func (a *Account) GetTypeSafeBaseURL() string { + if a == nil || !a.IsTypeSafe() || a.Type != AccountTypeAPIKey { + return "" + } + baseURL := strings.TrimRight(strings.TrimSpace(a.GetCredential("base_url")), "/") + // The System One path already carries /v1; accept a base URL pasted with it. + if len(baseURL) >= 3 && strings.EqualFold(baseURL[len(baseURL)-3:], "/v1") { + baseURL = strings.TrimRight(baseURL[:len(baseURL)-3], "/") + } + if baseURL == "" { + return typesafe.DefaultBaseURL + } + return baseURL +} + +func (a *Account) GetTypeSafeAPIKey() string { + if a == nil || !a.IsTypeSafe() || a.Type != AccountTypeAPIKey { + return "" + } + return strings.TrimSpace(a.GetCredential("api_key")) +} + func (a *Account) GetExtraString(key string) string { if a.Extra == nil { return "" diff --git a/backend/internal/service/account_service.go b/backend/internal/service/account_service.go index f4390bfcfba6..086b5a4f97a2 100644 --- a/backend/internal/service/account_service.go +++ b/backend/internal/service/account_service.go @@ -2,6 +2,7 @@ package service import ( "context" + "errors" "fmt" "time" @@ -241,6 +242,9 @@ func (s *AccountService) Create(ctx context.Context, req CreateAccountRequest) ( if err := ValidateModelMappingMode(req.Credentials); err != nil { return nil, err } + if req.Platform == PlatformTypeSafe && req.Type != AccountTypeAPIKey { + return nil, errors.New("typesafe accounts only support apikey credentials") + } // 验证分组是否存在(如果指定了分组) if len(req.GroupIDs) > 0 { if err := s.validateGroupIDsExist(ctx, req.GroupIDs); err != nil { @@ -536,6 +540,9 @@ func (s *AccountService) TestCredentials(ctx context.Context, id int64) error { case PlatformGrok: // Grok OAuth credentials are validated via token exchange/refresh and request-path probes. return nil + case PlatformTypeSafe: + // TypeSafe credentials are API keys; inference failures drive health and cooldown state. + return nil case PlatformKimi, PlatformZhipu, PlatformDeepseek, PlatformMiniMax, PlatformOpenCodeGo: // 国产 OpenAI 兼容供应商与 OpenCode:凭证为 API Key,实际可用性经余额/额度探测与转发路径验证。 return nil diff --git a/backend/internal/service/account_test_service.go b/backend/internal/service/account_test_service.go index 68ed31b17828..31259fbdf90a 100644 --- a/backend/internal/service/account_test_service.go +++ b/backend/internal/service/account_test_service.go @@ -436,6 +436,10 @@ func (s *AccountTestService) TestAccountConnection(c *gin.Context, accountID int return s.testOpenCodeGoAccountConnection(c, account, modelID, prompt) } + if account.IsTypeSafe() { + return s.testTypeSafeAccountConnection(c, account, prompt) + } + return s.testClaudeAccountConnection(c, account, modelID) } diff --git a/backend/internal/service/account_test_service_typesafe.go b/backend/internal/service/account_test_service_typesafe.go new file mode 100644 index 000000000000..6c9501d89204 --- /dev/null +++ b/backend/internal/service/account_test_service_typesafe.go @@ -0,0 +1,95 @@ +package service + +import ( + "encoding/json" + "fmt" + "io" + "net/http" + "strings" + + "github.com/Wei-Shaw/sub2api/internal/pkg/typesafe" + "github.com/gin-gonic/gin" +) + +const ( + typeSafeTestDefaultState = "Sub2API connection test" + typeSafeTestQuestionID = "connection_test" + typeSafeTestMaxPreviewBytes = 2000 +) + +// testTypeSafeAccountConnection probes a TypeSafe account with a minimal native +// System One request. TypeSafe accounts never speak the Claude protocol, so +// they must not fall through to testClaudeAccountConnection (which would send +// the key to /v1/messages and could misclassify the account). +func (s *AccountTestService) testTypeSafeAccountConnection(c *gin.Context, account *Account, prompt string) error { + ctx := c.Request.Context() + if account.Type != AccountTypeAPIKey { + return s.sendErrorAndEnd(c, fmt.Sprintf("Unsupported account type: %s", account.Type)) + } + apiKey := account.GetTypeSafeAPIKey() + if apiKey == "" { + return s.sendErrorAndEnd(c, "No API key available") + } + baseURL, err := s.validateUpstreamBaseURL(account.GetTypeSafeBaseURL()) + if err != nil { + return s.sendErrorAndEnd(c, fmt.Sprintf("Invalid base URL: %s", err.Error())) + } + + state := strings.TrimSpace(prompt) + if state == "" { + state = typeSafeTestDefaultState + } + payload, err := json.Marshal(map[string]any{ + "model": typesafe.JevLatestModel, + "state": state, + "questions": map[string]any{ + typeSafeTestQuestionID: map[string]any{ + "type": "noul", + "instructions": "Is this text a connection test?", + }, + }, + }) + if err != nil { + return s.sendErrorAndEnd(c, "Failed to create test payload") + } + + c.Writer.Header().Set("Content-Type", "text/event-stream") + c.Writer.Header().Set("Cache-Control", "no-cache") + c.Writer.Header().Set("Connection", "keep-alive") + c.Writer.Header().Set("X-Accel-Buffering", "no") + c.Writer.Flush() + + s.sendEvent(c, TestEvent{Type: "test_start", Model: typesafe.JevLatestModel}) + + req, err := typesafe.NewSystemOneRequest(ctx, baseURL, apiKey, payload) + if err != nil { + return s.sendErrorAndEnd(c, "Failed to create request") + } + proxyURL := "" + if account.ProxyID != nil && account.Proxy != nil { + proxyURL = account.Proxy.URL() + } + resp, err := s.httpUpstream.Do(req, proxyURL, account.ID, account.Concurrency) + if err != nil { + return s.sendErrorAndEnd(c, fmt.Sprintf("Request failed: %s", sanitizeUpstreamErrorMessage(err.Error()))) + } + defer func() { _ = resp.Body.Close() }() + + if resp.StatusCode != http.StatusOK { + body, _ := io.ReadAll(io.LimitReader(resp.Body, 64<<10)) + errMsg := fmt.Sprintf("API returned %d: %s", resp.StatusCode, truncateString(string(body), typeSafeTestMaxPreviewBytes)) + // 401/403 表示 API Key 无效或被上游拒绝,标记为 error 状态。 + if (resp.StatusCode == http.StatusUnauthorized || resp.StatusCode == http.StatusForbidden) && s.accountRepo != nil { + _ = s.accountRepo.SetError(ctx, account.ID, errMsg) + } + return s.sendErrorAndEnd(c, errMsg) + } + + decoded, err := typesafe.DecodeSystemOneResponse(resp.Body) + if err != nil { + return s.sendErrorAndEnd(c, fmt.Sprintf("Invalid System One response: %s", err.Error())) + } + s.sendEvent(c, TestEvent{Type: "content", Text: truncateString(string(decoded.Body), typeSafeTestMaxPreviewBytes)}) + s.sendEvent(c, TestEvent{Type: "test_complete", Success: true}) + return nil +} diff --git a/backend/internal/service/account_test_service_typesafe_test.go b/backend/internal/service/account_test_service_typesafe_test.go new file mode 100644 index 000000000000..c60abf42c3c0 --- /dev/null +++ b/backend/internal/service/account_test_service_typesafe_test.go @@ -0,0 +1,92 @@ +package service + +import ( + "encoding/json" + "io" + "net/http" + "net/http/httptest" + "strings" + "testing" + + "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/Wei-Shaw/sub2api/internal/pkg/typesafe" + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" +) + +func newTypeSafeAccountTestFixture(t *testing.T, status int, responseBody string) (*AccountTestService, *systemOnePolicyAccountRepo, *[]*http.Request, *gin.Context, *httptest.ResponseRecorder) { + t.Helper() + account := &Account{ID: 31, Platform: PlatformTypeSafe, Type: AccountTypeAPIKey, Credentials: map[string]any{"api_key": "ts-secret"}} + repo := &systemOnePolicyAccountRepo{account: account} + var requests []*http.Request + upstream := &systemOneHTTPUpstream{do: func(req *http.Request) (*http.Response, error) { + requests = append(requests, req) + return &http.Response{StatusCode: status, Header: make(http.Header), Body: io.NopCloser(strings.NewReader(responseBody))}, nil + }} + svc := &AccountTestService{ + accountRepo: repo, + httpUpstream: upstream, + cfg: &config.Config{}, + } + gin.SetMode(gin.TestMode) + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/api/v1/admin/accounts/31/test", nil) + return svc, repo, &requests, c, rec +} + +func TestTypeSafeAccountTestSendsNativeSystemOneRequest(t *testing.T) { + svc, repo, requests, c, rec := newTypeSafeAccountTestFixture(t, http.StatusOK, `{"model":"jev-1.13.0","answers":{"connection_test":{"type":"noul"}},"usage":{"input_tokens":9}}`) + + // The Claude default model must be ignored: TypeSafe only serves jev-latest. + require.NoError(t, svc.TestAccountConnection(c, 31, "claude-sonnet-4-5", "", AccountTestModeDefault)) + + require.Len(t, *requests, 1) + req := (*requests)[0] + require.Equal(t, typesafe.DefaultBaseURL+typesafe.SystemOnePath, req.URL.String()) + require.Equal(t, "Bearer ts-secret", req.Header.Get("Authorization")) + require.Empty(t, req.Header.Get("x-api-key")) + body, err := io.ReadAll(req.Body) + require.NoError(t, err) + model, err := typesafe.ValidateSystemOneRequest(body) + require.NoError(t, err) + require.Equal(t, typesafe.JevLatestModel, model) + + output := rec.Body.String() + require.Contains(t, output, `"type":"test_start"`) + require.Contains(t, output, `"model":"jev-latest"`) + require.Contains(t, output, `"type":"test_complete"`) + require.Contains(t, output, "jev-1.13.0") + require.Zero(t, repo.errorCalls) +} + +func TestTypeSafeAccountTestMarksRejectedKey(t *testing.T) { + for _, status := range []int{http.StatusUnauthorized, http.StatusForbidden} { + t.Run(http.StatusText(status), func(t *testing.T) { + svc, repo, _, c, rec := newTypeSafeAccountTestFixture(t, status, `{"detail":"invalid key"}`) + require.Error(t, svc.TestAccountConnection(c, 31, "", "", AccountTestModeDefault)) + require.Equal(t, 1, repo.errorCalls) + responseText, errMsg := parseTestSSEOutput(rec.Body.String()) + require.Empty(t, responseText) + require.Contains(t, errMsg, "API returned") + }) + } +} + +func TestTypeSafeAccountTestTransientFailureKeepsAccountState(t *testing.T) { + svc, repo, _, c, _ := newTypeSafeAccountTestFixture(t, http.StatusServiceUnavailable, `{"detail":"busy"}`) + require.Error(t, svc.TestAccountConnection(c, 31, "", "", AccountTestModeDefault)) + require.Zero(t, repo.errorCalls) +} + +func TestTypeSafeAccountTestUsesPromptAsState(t *testing.T) { + svc, _, requests, c, _ := newTypeSafeAccountTestFixture(t, http.StatusOK, `{"answers":{}}`) + require.NoError(t, svc.TestAccountConnection(c, 31, "", "custom state", AccountTestModeDefault)) + body, err := io.ReadAll((*requests)[0].Body) + require.NoError(t, err) + var payload struct { + State string `json:"state"` + } + require.NoError(t, json.Unmarshal(body, &payload)) + require.Equal(t, "custom state", payload.State) +} diff --git a/backend/internal/service/admin_account.go b/backend/internal/service/admin_account.go index bbe67ee927f3..a0e485916741 100644 --- a/backend/internal/service/admin_account.go +++ b/backend/internal/service/admin_account.go @@ -427,6 +427,9 @@ func normalizeOpenAILongContextBillingUpdateExtra(account *Account, input *Updat func buildAccountForCreate(input *CreateAccountInput, accountExtra map[string]any) (*Account, error) { accountExtra = MergeOpenAICodexTicketExtra(accountExtra, nil) + if input.Platform == PlatformTypeSafe && input.Type != AccountTypeAPIKey { + return nil, errors.New("typesafe accounts only support apikey credentials") + } // Probe/session state is system-managed. New accounts always start with automatic refresh disabled. delete(accountExtra, UpstreamBillingProbeEnabledExtraKey) delete(accountExtra, UpstreamBillingRateSyncEnabledExtraKey) @@ -621,6 +624,9 @@ func (s *adminServiceImpl) UpdateAccount(ctx context.Context, id int64, input *U if err != nil { return nil, err } + if account.Platform == PlatformTypeSafe && input.Type != "" && input.Type != AccountTypeAPIKey { + return nil, errors.New("typesafe accounts only support apikey credentials") + } var normalizedExtra map[string]any if input.Extra != nil { normalizedExtra, err = normalizeOpenAILongContextBillingUpdateExtra(account, input) diff --git a/backend/internal/service/admin_group.go b/backend/internal/service/admin_group.go index 4359a7b81214..2459a076147b 100644 --- a/backend/internal/service/admin_group.go +++ b/backend/internal/service/admin_group.go @@ -17,6 +17,7 @@ import ( "github.com/Wei-Shaw/sub2api/internal/pkg/logger" "github.com/Wei-Shaw/sub2api/internal/pkg/openai" "github.com/Wei-Shaw/sub2api/internal/pkg/pagination" + "github.com/Wei-Shaw/sub2api/internal/pkg/typesafe" "github.com/Wei-Shaw/sub2api/internal/pkg/xai" ) @@ -296,6 +297,8 @@ func defaultModelsListCandidateIDs(platform string) []string { return xai.DefaultModelIDs() case PlatformOpenCodeGo: return DefaultOpenCodeGoModelIDs() + case PlatformTypeSafe: + return []string{typesafe.JevLatestModel} case PlatformComposite: return compositeDefaultModelsListCandidateIDs() default: @@ -316,6 +319,9 @@ func defaultAllowImageGenerationForPlatform(platform string) bool { func compositeDefaultModelsListCandidateIDs() []string { seen := make(map[string]struct{}) ids := make([]string, 0) + // TypeSafe stays out of the static composite candidates (jev-latest only works + // through /v1/systemone); groups with TypeSafe accounts still get it from the + // account model mappings collected by GetGroupModelsListCandidates. for _, platform := range []string{PlatformAnthropic, PlatformGemini, PlatformOpenAI, PlatformAntigravity, PlatformGrok, PlatformKimi, PlatformZhipu, PlatformDeepseek, PlatformMiniMax, PlatformOpenCodeGo} { for _, id := range defaultModelsListCandidateIDs(platform) { if _, ok := seen[id]; ok { diff --git a/backend/internal/service/antigravity_gateway_gemini.go b/backend/internal/service/antigravity_gateway_gemini.go index 3ce13e6de4f4..4a8c70dcfb3c 100644 --- a/backend/internal/service/antigravity_gateway_gemini.go +++ b/backend/internal/service/antigravity_gateway_gemini.go @@ -194,7 +194,6 @@ func (s *AntigravityGatewayService) ForwardGemini(ctx context.Context, c *gin.Co // 处理错误响应 if resp.StatusCode >= 400 { respBody := s.readUpstreamErrorBody(resp) - contentType := resp.Header.Get("Content-Type") // 尽早关闭原始响应体,释放连接;后续逻辑仍可能需要读取 body,因此用内存副本重新包装。 _ = resp.Body.Close() resp.Body = io.NopCloser(bytes.NewReader(respBody)) @@ -301,7 +300,6 @@ func (s *AntigravityGatewayService) ForwardGemini(ctx context.Context, c *gin.Co Header: retryResp.Header.Clone(), Body: io.NopCloser(bytes.NewReader(retryRespBody)), } - contentType = resp.Header.Get("Content-Type") } } else { if switchErr, ok := IsAntigravityAccountSwitchError(retryErr); ok { @@ -393,9 +391,6 @@ func (s *AntigravityGatewayService) ForwardGemini(ctx context.Context, c *gin.Co }) return nil, &UpstreamFailoverError{StatusCode: resp.StatusCode, ResponseBody: unwrappedForOps} } - if contentType == "" { - contentType = "application/json" - } appendOpsUpstreamError(c, OpsUpstreamErrorEvent{ ProxyID: opsUpstreamProxyID(account), ProxyName: opsUpstreamProxyName(account), @@ -410,7 +405,8 @@ func (s *AntigravityGatewayService) ForwardGemini(ctx context.Context, c *gin.Co }) logger.LegacyPrintf("service.antigravity_gateway", "[antigravity-Forward] upstream error status=%d body=%s", resp.StatusCode, truncateForLog(unwrappedForOps, 500)) MarkResponseCommitted(c) - c.Data(resp.StatusCode, contentType, unwrappedForOps) + // 原始上游错误体仅保留在 ops 日志中;返回客户端的错误体需脱敏,避免泄露账号池身份(项目号/服务账号等) + c.Data(resp.StatusCode, "application/json", buildAntigravityClientErrorBody(resp.StatusCode, unwrappedForOps)) return nil, fmt.Errorf("antigravity upstream error: %d", resp.StatusCode) } diff --git a/backend/internal/service/antigravity_upstream_error_sanitize.go b/backend/internal/service/antigravity_upstream_error_sanitize.go new file mode 100644 index 000000000000..b937970d6d29 --- /dev/null +++ b/backend/internal/service/antigravity_upstream_error_sanitize.go @@ -0,0 +1,97 @@ +package service + +import ( + "encoding/json" + "net/http" + "regexp" + "strings" +) + +var ( + antigravityProjectRefRegex = regexp.MustCompile(`(?i)\bprojects/[a-z0-9][a-z0-9._:-]*`) + antigravityEmailRegex = regexp.MustCompile(`[A-Za-z0-9._%+\-]+@[A-Za-z0-9.\-]+\.[A-Za-z]{2,}`) + antigravityConsumerRegex = regexp.MustCompile(`(?i)\b(consumer|project(?:[ _-]?(?:id|number))?)(\s*[:=]?\s*['"]?)[0-9]{6,}`) +) + +// sanitizeAntigravityErrorText 清除上游错误文本中的 GCP 项目号/项目 ID、服务账号邮箱及敏感查询参数。 +func sanitizeAntigravityErrorText(msg string) string { + if msg == "" { + return msg + } + msg = sanitizeUpstreamErrorMessage(msg) + msg = antigravityProjectRefRegex.ReplaceAllString(msg, "projects/***") + msg = antigravityEmailRegex.ReplaceAllString(msg, "***") + msg = antigravityConsumerRegex.ReplaceAllString(msg, "$1$2***") + return msg +} + +// buildAntigravityClientErrorBody 为客户端构造 Gemini 风格的错误体: +// 仅保留 code/status/message(message 已脱敏),丢弃 details 等可能含账号身份的字段。 +func buildAntigravityClientErrorBody(statusCode int, body []byte) []byte { + code := statusCode + status := "" + message := "" + + var parsed struct { + Error struct { + Code int `json:"code"` + Message string `json:"message"` + Status string `json:"status"` + } `json:"error"` + } + if err := json.Unmarshal(body, &parsed); err == nil { + if parsed.Error.Code != 0 { + code = parsed.Error.Code + } + status = parsed.Error.Status + message = parsed.Error.Message + } + if strings.TrimSpace(message) == "" { + message = strings.TrimSpace(extractUpstreamErrorMessage(body)) + } + if strings.TrimSpace(message) == "" { + message = http.StatusText(statusCode) + if message == "" { + message = "Upstream request failed" + } + } + if status == "" { + status = antigravityGeminiStatusFromHTTP(statusCode) + } + + out, err := json.Marshal(map[string]any{ + "error": map[string]any{ + "code": code, + "message": sanitizeAntigravityErrorText(message), + "status": status, + }, + }) + if err != nil { + return []byte(`{"error":{"code":500,"message":"Upstream request failed","status":"INTERNAL"}}`) + } + return out +} + +func antigravityGeminiStatusFromHTTP(statusCode int) string { + switch statusCode { + case http.StatusBadRequest: + return "INVALID_ARGUMENT" + case http.StatusUnauthorized: + return "UNAUTHENTICATED" + case http.StatusForbidden: + return "PERMISSION_DENIED" + case http.StatusNotFound: + return "NOT_FOUND" + case http.StatusTooManyRequests: + return "RESOURCE_EXHAUSTED" + case http.StatusServiceUnavailable: + return "UNAVAILABLE" + case http.StatusGatewayTimeout: + return "DEADLINE_EXCEEDED" + default: + if statusCode >= 500 { + return "INTERNAL" + } + return "UNKNOWN" + } +} diff --git a/backend/internal/service/antigravity_upstream_error_sanitize_test.go b/backend/internal/service/antigravity_upstream_error_sanitize_test.go new file mode 100644 index 000000000000..3f4abf04a21e --- /dev/null +++ b/backend/internal/service/antigravity_upstream_error_sanitize_test.go @@ -0,0 +1,50 @@ +//go:build unit + +package service + +import ( + "encoding/json" + "net/http" + "net/http/httptest" + "strings" + "testing" + + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" +) + +func TestBuildAntigravityClientErrorBody_ScrubsPoolIdentity(t *testing.T) { + gin.SetMode(gin.TestMode) + upstream := []byte(`{"error":{"code":403,"message":"Permission denied on resource project projects/123456789 for consumer: projects/123456789; caller pool-sa@my-gcp-proj.iam.gserviceaccount.com","status":"PERMISSION_DENIED","details":[{"@type":"type.googleapis.com/google.rpc.ErrorInfo","metadata":{"consumer":"projects/123456789","service":"cloudcode-pa.googleapis.com"}}]}}`) + + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Data(http.StatusForbidden, "application/json", buildAntigravityClientErrorBody(http.StatusForbidden, upstream)) + + require.Equal(t, http.StatusForbidden, rec.Code) + out := rec.Body.String() + require.NotContains(t, out, "123456789") + require.NotContains(t, out, "pool-sa@") + require.NotContains(t, out, "gserviceaccount.com") + require.NotContains(t, out, "details") + + var parsed struct { + Error struct { + Code int `json:"code"` + Message string `json:"message"` + Status string `json:"status"` + } `json:"error"` + } + require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &parsed)) + require.Equal(t, 403, parsed.Error.Code) + require.Equal(t, "PERMISSION_DENIED", parsed.Error.Status) + require.True(t, strings.Contains(parsed.Error.Message, "Permission denied")) +} + +func TestBuildAntigravityClientErrorBody_NonJSONBody(t *testing.T) { + out := string(buildAntigravityClientErrorBody(http.StatusTooManyRequests, []byte("quota exceeded for consumer 987654321 sa@x.iam.gserviceaccount.com"))) + require.NotContains(t, out, "987654321") + require.NotContains(t, out, "gserviceaccount") + require.Contains(t, out, `"status":"RESOURCE_EXHAUSTED"`) + require.Contains(t, out, `"code":429`) +} diff --git a/backend/internal/service/auth_service_email_bind_test.go b/backend/internal/service/auth_service_email_bind_test.go index 9bdca02f2d6b..c41734c40bbc 100644 --- a/backend/internal/service/auth_service_email_bind_test.go +++ b/backend/internal/service/auth_service_email_bind_test.go @@ -1165,3 +1165,19 @@ func cloneEmailBindUser(user *service.User) *service.User { cloned := *user return &cloned } + +func (s *emailBindCacheStub) IncrVerificationCodeAttempts(context.Context, string) (int, error) { + if s.data == nil { + return 0, errors.New("verification code not found") + } + s.data.Attempts++ + return s.data.Attempts, nil +} + +func (s *emailBindCacheStub) IncrNotifyVerifyCodeAttempts(context.Context, string) (int, error) { + return 0, errors.New("notify verification code not found") +} + +func (s *emailBindCacheStub) ConsumePasswordResetToken(context.Context, string, string) (bool, error) { + return false, nil +} diff --git a/backend/internal/service/auth_service_register_test.go b/backend/internal/service/auth_service_register_test.go index aa7519f3d1ea..076cd025d3f5 100644 --- a/backend/internal/service/auth_service_register_test.go +++ b/backend/internal/service/auth_service_register_test.go @@ -1010,3 +1010,19 @@ func TestCanBypassRegistrationDisabledForOAuth(t *testing.T) { }) } } + +func (s *emailCacheStub) IncrVerificationCodeAttempts(context.Context, string) (int, error) { + if s.data == nil { + return 0, errors.New("verification code not found") + } + s.data.Attempts++ + return s.data.Attempts, nil +} + +func (s *emailCacheStub) IncrNotifyVerifyCodeAttempts(context.Context, string) (int, error) { + return 0, errors.New("notify verification code not found") +} + +func (s *emailCacheStub) ConsumePasswordResetToken(context.Context, string, string) (bool, error) { + return false, nil +} diff --git a/backend/internal/service/billing_service.go b/backend/internal/service/billing_service.go index c8955b07a375..b46cf7be7d06 100644 --- a/backend/internal/service/billing_service.go +++ b/backend/internal/service/billing_service.go @@ -691,6 +691,12 @@ func (s *BillingService) initFallbackPricing() { SupportsCacheBreakdown: false, } + // TypeSafe Jev bills input tokens only: $0.042 per million tokens. + s.fallbackPrices["jev-latest"] = &ModelPricing{ + InputPricePerToken: 0.042 / 1_000_000, + OutputPricePerToken: 0, + } + // ---- 智谱 GLM(Z.AI)---- // Source: https://docs.z.ai/guides/overview/pricing (USD per 1M tokens) // 注意:CacheReadPricePerToken 即"缓存命中"价格,CacheCreationPricePerToken 留空(智谱未公开写入价,按 0 处理)。 @@ -987,6 +993,9 @@ func (s *BillingService) initFallbackPricing() { // getFallbackPricing 根据模型系列获取回退价格 func (s *BillingService) getFallbackPricing(model string) *ModelPricing { modelLower := strings.ToLower(model) + if modelLower == "jev-latest" { + return s.fallbackPrices["jev-latest"] + } // 按模型系列匹配 if isClaudeFable51Model(modelLower) { diff --git a/backend/internal/service/billing_service_test.go b/backend/internal/service/billing_service_test.go index 1baa039bd8fa..180be782831c 100644 --- a/backend/internal/service/billing_service_test.go +++ b/backend/internal/service/billing_service_test.go @@ -130,6 +130,20 @@ func TestGetModelPricing_CaseInsensitive(t *testing.T) { require.Equal(t, p1.InputPricePerToken, p2.InputPricePerToken) } +func TestGetModelPricing_JevLatestInputOnlyAndChannelOverride(t *testing.T) { + svc := newTestBillingService() + pricing, err := svc.GetModelPricing("jev-latest") + require.NoError(t, err) + require.InDelta(t, 0.042/1_000_000, pricing.InputPricePerToken, 1e-15) + require.Zero(t, pricing.OutputPricePerToken) + + input, output := 0.25/1_000_000, 0.5/1_000_000 + pricing, err = svc.GetModelPricingWithChannel("jev-latest", &ChannelModelPricing{InputPrice: &input, OutputPrice: &output}) + require.NoError(t, err) + require.Equal(t, input, pricing.InputPricePerToken) + require.Equal(t, output, pricing.OutputPricePerToken) +} + // issue #3394: fallback warn 应按模型名去重,每个模型每进程最多打一条, // 避免热路径每请求刷屏 ops_system_logs。 func TestGetModelPricing_FallbackWarnLoggedOncePerModel(t *testing.T) { diff --git a/backend/internal/service/channel_service.go b/backend/internal/service/channel_service.go index fb59d7f7bbe0..b2c605cd80d9 100644 --- a/backend/internal/service/channel_service.go +++ b/backend/internal/service/channel_service.go @@ -368,7 +368,7 @@ func isPlatformPricingMatch(groupPlatform, pricingPlatform string) bool { // fallback used before a request target has been resolved. func matchingPlatforms(groupPlatform string) []string { if groupPlatform == PlatformComposite { - return []string{PlatformAnthropic, PlatformGemini, PlatformOpenAI, PlatformAntigravity, PlatformGrok, PlatformKimi, PlatformZhipu, PlatformDeepseek, PlatformMiniMax, PlatformOpenCodeGo} + return []string{PlatformAnthropic, PlatformGemini, PlatformOpenAI, PlatformAntigravity, PlatformGrok, PlatformKimi, PlatformZhipu, PlatformDeepseek, PlatformMiniMax, PlatformOpenCodeGo, PlatformTypeSafe} } return []string{groupPlatform} } diff --git a/backend/internal/service/channel_service_test.go b/backend/internal/service/channel_service_test.go index 9b85aa3e00a4..62ba8e0b9a1b 100644 --- a/backend/internal/service/channel_service_test.go +++ b/backend/internal/service/channel_service_test.go @@ -2127,7 +2127,7 @@ func TestMatchingPlatforms(t *testing.T) { {"anthropic returns itself", PlatformAnthropic, []string{PlatformAnthropic}}, {"gemini returns itself", PlatformGemini, []string{PlatformGemini}}, {"openai returns itself", PlatformOpenAI, []string{PlatformOpenAI}}, - {"composite returns concrete platforms", PlatformComposite, []string{PlatformAnthropic, PlatformGemini, PlatformOpenAI, PlatformAntigravity, PlatformGrok, PlatformKimi, PlatformZhipu, PlatformDeepseek, PlatformMiniMax, PlatformOpenCodeGo}}, + {"composite returns concrete platforms", PlatformComposite, []string{PlatformAnthropic, PlatformGemini, PlatformOpenAI, PlatformAntigravity, PlatformGrok, PlatformKimi, PlatformZhipu, PlatformDeepseek, PlatformMiniMax, PlatformOpenCodeGo, PlatformTypeSafe}}, } for _, tt := range tests { diff --git a/backend/internal/service/composite_platform.go b/backend/internal/service/composite_platform.go index 6ddb9cfa9a58..c10e335cbebe 100644 --- a/backend/internal/service/composite_platform.go +++ b/backend/internal/service/composite_platform.go @@ -114,6 +114,8 @@ func DetectModelPlatform(model string) (string, bool) { return PlatformDeepseek, true case "minimax": return PlatformMiniMax, true + case "typesafe", "jev": + return PlatformTypeSafe, true } if rest != "" { normalized = strings.TrimPrefix(rest, "models/") @@ -155,6 +157,8 @@ func DetectModelPlatform(model string) (string, bool) { strings.HasPrefix(normalized, "abab6"), strings.HasPrefix(normalized, "abab7"): return PlatformMiniMax, true + case normalized == "jev-latest" || strings.HasPrefix(normalized, "jev-"): + return PlatformTypeSafe, true default: return "", false } @@ -202,7 +206,7 @@ func (s *GatewayService) resolveCompositeRouteDecision(ctx context.Context, grou func isConcreteRequestPlatform(platform string) bool { switch platform { case PlatformAnthropic, PlatformOpenAI, PlatformGemini, PlatformAntigravity, PlatformGrok, - PlatformKimi, PlatformZhipu, PlatformDeepseek, PlatformMiniMax, PlatformOpenCodeGo: + PlatformKimi, PlatformZhipu, PlatformDeepseek, PlatformMiniMax, PlatformOpenCodeGo, PlatformTypeSafe: return true default: return false diff --git a/backend/internal/service/composite_platform_test.go b/backend/internal/service/composite_platform_test.go index 20505a3985af..4b11366a2947 100644 --- a/backend/internal/service/composite_platform_test.go +++ b/backend/internal/service/composite_platform_test.go @@ -180,6 +180,8 @@ func TestDetectModelPlatform(t *testing.T) { {name: "minimax prefix", model: "minimax/MiniMax-M2.5", platform: PlatformMiniMax, ok: true}, {name: "abab legacy", model: "abab6.5-chat", platform: PlatformMiniMax, ok: true}, {name: "abab7 legacy", model: "abab7-chat-preview", platform: PlatformMiniMax, ok: true}, + {name: "jev", model: "jev-latest", platform: PlatformTypeSafe, ok: true}, + {name: "typesafe prefix", model: "typesafe/jev-latest", platform: PlatformTypeSafe, ok: true}, {name: "abab unrelated namespace", model: "abab-other", ok: false}, {name: "unknown k3 alias", model: "k3-preview", ok: false}, {name: "unknown", model: "llama-4-maverick", ok: false}, @@ -216,13 +218,13 @@ func TestCompositeGroupSchedulerHasAllCanonicalPlatformBuckets(t *testing.T) { platforms = append(platforms, platform) } require.ElementsMatch(t, - []string{PlatformAnthropic, PlatformGemini, PlatformOpenAI, PlatformAntigravity, PlatformGrok, PlatformKimi, PlatformZhipu, PlatformDeepseek, PlatformMiniMax, PlatformOpenCodeGo}, + []string{PlatformAnthropic, PlatformGemini, PlatformOpenAI, PlatformAntigravity, PlatformGrok, PlatformKimi, PlatformZhipu, PlatformDeepseek, PlatformMiniMax, PlatformOpenCodeGo, PlatformTypeSafe}, platforms, ) } func TestCompositeConcretePlatformsIncludeCNProviders(t *testing.T) { - for _, platform := range []string{PlatformKimi, PlatformZhipu, PlatformDeepseek, PlatformMiniMax, PlatformOpenCodeGo} { + for _, platform := range []string{PlatformKimi, PlatformZhipu, PlatformDeepseek, PlatformMiniMax, PlatformOpenCodeGo, PlatformTypeSafe} { require.True(t, isConcreteRequestPlatform(platform)) require.True(t, canCopyAccountsFromGroupPlatform(PlatformComposite, platform)) } diff --git a/backend/internal/service/content_moderation.go b/backend/internal/service/content_moderation.go index cb000e2a0a25..ee162a023443 100644 --- a/backend/internal/service/content_moderation.go +++ b/backend/internal/service/content_moderation.go @@ -56,6 +56,7 @@ const ( ContentModerationProtocolOpenAIChat = "openai_chat_completions" ContentModerationProtocolGemini = "gemini" ContentModerationProtocolOpenAIImages = "openai_images" + ContentModerationProtocolTypeSafeSystemOne = "typesafe_systemone" defaultContentModerationBaseURL = "https://api.openai.com" defaultContentModerationModel = "omni-moderation-latest" diff --git a/backend/internal/service/content_moderation_input.go b/backend/internal/service/content_moderation_input.go index 3886bfff8eab..b52adfaa21a1 100644 --- a/backend/internal/service/content_moderation_input.go +++ b/backend/internal/service/content_moderation_input.go @@ -46,6 +46,10 @@ func extractContentModerationInput(protocol string, body []byte, filterReminders case ContentModerationProtocolOpenAIImages: collector.addModerationText(&parts, gjson.GetBytes(body, "prompt").String()) collector.collectContentValue(gjson.GetBytes(body, "images"), &parts, &images) + case ContentModerationProtocolTypeSafeSystemOne: + // System One carries no client-harness reminder blocks, so a literal + // is ordinary user text and must never be skipped. + moderationTextCollector{}.collectSystemOneInput(body, &parts) default: collector.collectLastResponsesInput(gjson.GetBytes(body, "input"), &parts, &images) collector.collectLastRoleMessage(gjson.GetBytes(body, "messages"), "user", &parts, &images) @@ -61,6 +65,67 @@ func extractContentModerationInput(protocol string, body []byte, filterReminders return out } +// collectSystemOneInput moderates every client-controlled text of a System One +// request: question IDs, every question field except the validated type, +// unknown top-level extension fields, and the evaluated state. Object keys are +// sent to Jev as part of the JSON, so they are moderated like values. +func (collector moderationTextCollector) collectSystemOneInput(body []byte, parts *[]string) { + root := gjson.ParseBytes(body) + questions := root.Get("questions") + if !questions.IsObject() { + collector.collectSystemOneText(questions, parts) + } + questions.ForEach(func(id, question gjson.Result) bool { + collector.addModerationText(parts, id.String()) + if !question.IsObject() { + collector.collectSystemOneText(question, parts) + return true + } + question.ForEach(func(field, value gjson.Result) bool { + switch field.String() { + case "type": + return true + case "instructions", "criteria": + default: + collector.addModerationText(parts, field.String()) + } + collector.collectSystemOneText(value, parts) + return true + }) + return true + }) + root.ForEach(func(field, value gjson.Result) bool { + switch field.String() { + case "model", "stream", "state", "questions": + return true + } + collector.addModerationText(parts, field.String()) + collector.collectSystemOneText(value, parts) + return true + }) + collector.collectSystemOneText(root.Get("state"), parts) +} + +func (collector moderationTextCollector) collectSystemOneText(value gjson.Result, parts *[]string) { + switch { + case !value.Exists(): + return + case value.Type == gjson.String: + collector.addModerationText(parts, value.String()) + case value.IsArray(): + value.ForEach(func(_, child gjson.Result) bool { + collector.collectSystemOneText(child, parts) + return true + }) + case value.IsObject(): + value.ForEach(func(key, child gjson.Result) bool { + collector.addModerationText(parts, key.String()) + collector.collectSystemOneText(child, parts) + return true + }) + } +} + func (collector moderationTextCollector) collectLastRoleMessage(messages gjson.Result, role string, parts *[]string, images *[]string) { if !messages.IsArray() { return diff --git a/backend/internal/service/content_moderation_systemone_test.go b/backend/internal/service/content_moderation_systemone_test.go new file mode 100644 index 000000000000..8f2eba816fe6 --- /dev/null +++ b/backend/internal/service/content_moderation_systemone_test.go @@ -0,0 +1,34 @@ +package service + +import ( + "testing" + + "github.com/stretchr/testify/require" +) + +func TestExtractContentModerationInputTypeSafeSystemOne(t *testing.T) { + for _, tc := range []struct { + name string + body string + want string + }{ + {"string", `{"state":"plain text"}`, "plain text"}, + {"object", `{"state":{"title":"hello","nested":{"body":"world"}}}`, "title hello nested body world"}, + {"array", `{"state":["first",{"text":"second"},3]}`, "first text second"}, + {"questions", `{"state":"state text","questions":{"c":{"type":"choice","instructions":"pick one","criteria":{"label a":{"description":"desc a"},"label b":null}},"s":{"type":"score","criteria":["low",{"text":"high"}]},"n":{"type":"noul","instructions":["judge"],"criteria":{"true":"yes"},"ext":"extra"}},"top":{"k":"v"}}`, "c pick one label a description desc a label b s low text high n judge true yes ext extra top k v state text"}, + {"key only payload", `{"state":{"hidden state key":1},"questions":{"hidden question id":{"type":"noul"}}}`, "hidden question id hidden state key"}, + } { + t.Run(tc.name, func(t *testing.T) { + input := ExtractContentModerationInput(ContentModerationProtocolTypeSafeSystemOne, []byte(tc.body)) + require.Equal(t, tc.want, input.Text) + require.Empty(t, input.Images) + }) + } +} + +func TestExtractContentModerationInputTypeSafeSystemOneKeepsReminderText(t *testing.T) { + body := `{"state":"hidden payload","questions":{"q":{"type":"noul","instructions":"hidden instructions"}}}` + input := ExtractContentModerationInput(ContentModerationProtocolTypeSafeSystemOne, []byte(body)) + require.Contains(t, input.Text, "hidden payload") + require.Contains(t, input.Text, "hidden instructions") +} diff --git a/backend/internal/service/domain_constants.go b/backend/internal/service/domain_constants.go index 5f9daa13fca6..ad93e2000386 100644 --- a/backend/internal/service/domain_constants.go +++ b/backend/internal/service/domain_constants.go @@ -49,6 +49,7 @@ const ( PlatformZhipu = domain.PlatformZhipu PlatformDeepseek = domain.PlatformDeepseek PlatformMiniMax = domain.PlatformMiniMax + PlatformTypeSafe = domain.PlatformTypeSafe PlatformOpenCodeGo = domain.PlatformOpenCodeGo PlatformComposite = domain.PlatformComposite // PlatformKiro is retained for unsupported-platform threshold tests and legacy @@ -136,6 +137,7 @@ var AllowedQuotaPlatforms = []string{ PlatformDeepseek, PlatformMiniMax, PlatformOpenCodeGo, + PlatformTypeSafe, } // AllowedSchedulingThresholdPlatforms 是允许设置账号自动停调阈值的平台列表。 diff --git a/backend/internal/service/email_service.go b/backend/internal/service/email_service.go index 8e60d9e73d4d..794dad52586c 100644 --- a/backend/internal/service/email_service.go +++ b/backend/internal/service/email_service.go @@ -3,6 +3,7 @@ package service import ( "context" "crypto/rand" + "crypto/sha256" "crypto/subtle" "crypto/tls" "crypto/x509" @@ -37,16 +38,23 @@ type EmailCache interface { GetVerificationCode(ctx context.Context, email string) (*VerificationCodeData, error) SetVerificationCode(ctx context.Context, email string, data *VerificationCodeData, ttl time.Duration) error DeleteVerificationCode(ctx context.Context, email string) error + // IncrVerificationCodeAttempts atomically increments the attempt counter of an + // existing code and returns the new count. Returns an error if the code is missing. + IncrVerificationCodeAttempts(ctx context.Context, email string) (int, error) // Notify email verification code methods GetNotifyVerifyCode(ctx context.Context, email string) (*VerificationCodeData, error) SetNotifyVerifyCode(ctx context.Context, email string, data *VerificationCodeData, ttl time.Duration) error DeleteNotifyVerifyCode(ctx context.Context, email string) error + IncrNotifyVerifyCodeAttempts(ctx context.Context, email string) (int, error) // Password reset token methods GetPasswordResetToken(ctx context.Context, email string) (*PasswordResetTokenData, error) SetPasswordResetToken(ctx context.Context, email string, data *PasswordResetTokenData, ttl time.Duration) error DeletePasswordResetToken(ctx context.Context, email string) error + // ConsumePasswordResetToken atomically compares the stored token hash with + // tokenHash and deletes it on match. Only one concurrent caller can succeed. + ConsumePasswordResetToken(ctx context.Context, email, tokenHash string) (bool, error) // Password reset email cooldown methods // Returns true if in cooldown period (email was sent recently) @@ -68,6 +76,7 @@ type VerificationCodeData struct { // PasswordResetTokenData represents password reset token data type PasswordResetTokenData struct { + // Token holds the hex-encoded SHA-256 hash of the reset token (never the plaintext). Token string CreatedAt time.Time } @@ -376,35 +385,49 @@ func (s *EmailService) SendVerifyCode(ctx context.Context, email, siteName strin // VerifyCode 验证验证码 func (s *EmailService) VerifyCode(ctx context.Context, email, code string) error { - data, err := s.cache.GetVerificationCode(ctx, email) + return verifyCodeWithAttempts(ctx, email, code, + s.cache.GetVerificationCode, s.cache.IncrVerificationCodeAttempts, + func() { + if err := s.cache.DeleteVerificationCode(ctx, email); err != nil { + slog.Error("failed to delete verification code after success", "email", email, "error", err) + } + }) +} + +// verifyCodeWithAttempts checks a verification code while enforcing the attempt cap +// atomically: each check first reserves an attempt via an atomic increment, so +// concurrent guesses can never evaluate more than maxVerifyCodeAttempts codes. +func verifyCodeWithAttempts( + ctx context.Context, email, code string, + get func(context.Context, string) (*VerificationCodeData, error), + incr func(context.Context, string) (int, error), + onSuccess func(), +) error { + data, err := get(ctx, email) if err != nil || data == nil { return ErrInvalidVerifyCode } - - // 检查是否已达到最大尝试次数 if data.Attempts >= maxVerifyCodeAttempts { return ErrVerifyCodeMaxAttempts } + attempts, err := incr(ctx, email) + if err != nil { + return ErrInvalidVerifyCode + } + if attempts > maxVerifyCodeAttempts { + return ErrVerifyCodeMaxAttempts + } - // 验证码不匹配 (constant-time comparison to prevent timing attacks) + // constant-time comparison to prevent timing attacks if subtle.ConstantTimeCompare([]byte(data.Code), []byte(code)) != 1 { - data.Attempts++ - remaining := time.Until(data.ExpiresAt) - if remaining <= 0 { - return ErrInvalidVerifyCode - } - if err := s.cache.SetVerificationCode(ctx, email, data, remaining); err != nil { - slog.Error("failed to update verification attempt count", "email", email, "error", err) - } - if data.Attempts >= maxVerifyCodeAttempts { + if attempts >= maxVerifyCodeAttempts { return ErrVerifyCodeMaxAttempts } return ErrInvalidVerifyCode } - // 验证成功,删除验证码 - if err := s.cache.DeleteVerificationCode(ctx, email); err != nil { - slog.Error("failed to delete verification code after success", "email", email, "error", err) + if onSuccess != nil { + onSuccess() } return nil } @@ -480,33 +503,18 @@ func (s *EmailService) GeneratePasswordResetToken() (string, error) { // SendPasswordResetEmail sends a password reset email with a reset link func (s *EmailService) SendPasswordResetEmail(ctx context.Context, email, siteName, resetURL string, locale ...string) error { - var token string - var needSaveToken bool - - // Check if token already exists - existing, err := s.cache.GetPasswordResetToken(ctx, email) - if err == nil && existing != nil { - // Token exists, reuse it (allows resending email without generating new token) - token = existing.Token - needSaveToken = false - } else { - // Generate new token - token, err = s.GeneratePasswordResetToken() - if err != nil { - return fmt.Errorf("generate token: %w", err) - } - needSaveToken = true + // Only the SHA-256 hash of the token is stored, so an existing token cannot be + // re-sent; always issue a fresh token (the email cooldown bounds resend frequency). + token, err := s.GeneratePasswordResetToken() + if err != nil { + return fmt.Errorf("generate token: %w", err) } - - // Save token to Redis (only if new token generated) - if needSaveToken { - data := &PasswordResetTokenData{ - Token: token, - CreatedAt: time.Now(), - } - if err := s.cache.SetPasswordResetToken(ctx, email, data, passwordResetTokenTTL); err != nil { - return fmt.Errorf("save reset token: %w", err) - } + data := &PasswordResetTokenData{ + Token: hashPasswordResetToken(token), + CreatedAt: time.Now(), + } + if err := s.cache.SetPasswordResetToken(ctx, email, data, passwordResetTokenTTL); err != nil { + return fmt.Errorf("save reset token: %w", err) } // Build full reset URL with URL-encoded token and email @@ -566,31 +574,40 @@ func (s *EmailService) SendPasswordResetEmailWithCooldown(ctx context.Context, e return nil } +// hashPasswordResetToken returns the hex-encoded SHA-256 of a reset token. +func hashPasswordResetToken(token string) string { + sum := sha256.Sum256([]byte(token)) + return hex.EncodeToString(sum[:]) +} + // VerifyPasswordResetToken verifies the password reset token without consuming it func (s *EmailService) VerifyPasswordResetToken(ctx context.Context, email, token string) error { data, err := s.cache.GetPasswordResetToken(ctx, email) - if err != nil || data == nil { + if err != nil || data == nil || token == "" { return ErrInvalidResetToken } // Use constant-time comparison to prevent timing attacks - if subtle.ConstantTimeCompare([]byte(data.Token), []byte(token)) != 1 { + if subtle.ConstantTimeCompare([]byte(data.Token), []byte(hashPasswordResetToken(token))) != 1 { return ErrInvalidResetToken } return nil } -// ConsumePasswordResetToken verifies and deletes the token (one-time use) +// ConsumePasswordResetToken verifies and deletes the token atomically (one-time use). func (s *EmailService) ConsumePasswordResetToken(ctx context.Context, email, token string) error { - // Verify first + // Constant-time pre-check in Go; the atomic compare-and-delete below is authoritative. if err := s.VerifyPasswordResetToken(ctx, email, token); err != nil { return err } - - // Delete after verification (one-time use) - if err := s.cache.DeletePasswordResetToken(ctx, email); err != nil { - slog.Error("failed to delete password reset token after consumption", "email", email, "error", err) + ok, err := s.cache.ConsumePasswordResetToken(ctx, email, hashPasswordResetToken(token)) + if err != nil { + slog.Error("failed to consume password reset token", "email", email, "error", err) + return ErrInvalidResetToken + } + if !ok { + return ErrInvalidResetToken } return nil } diff --git a/backend/internal/service/email_service_reset_token_test.go b/backend/internal/service/email_service_reset_token_test.go new file mode 100644 index 000000000000..15efd50d6e69 --- /dev/null +++ b/backend/internal/service/email_service_reset_token_test.go @@ -0,0 +1,50 @@ +//go:build unit + +package service + +import ( + "context" + "crypto/sha256" + "encoding/hex" + "testing" + + "github.com/stretchr/testify/require" +) + +type resetTokenCacheStub struct { + emailCacheStub + stored *PasswordResetTokenData + consumedHash string +} + +func (s *resetTokenCacheStub) GetPasswordResetToken(context.Context, string) (*PasswordResetTokenData, error) { + return s.stored, nil +} + +func (s *resetTokenCacheStub) ConsumePasswordResetToken(_ context.Context, _ string, tokenHash string) (bool, error) { + s.consumedHash = tokenHash + if s.stored == nil || s.stored.Token != tokenHash { + return false, nil + } + s.stored = nil + return true, nil +} + +func TestConsumePasswordResetToken_ComparesHashNotPlaintext(t *testing.T) { + token := "deadbeef" + sum := sha256.Sum256([]byte(token)) + hash := hex.EncodeToString(sum[:]) + require.Equal(t, hash, hashPasswordResetToken(token)) + require.NotEqual(t, token, hashPasswordResetToken(token)) + + cache := &resetTokenCacheStub{stored: &PasswordResetTokenData{Token: hash}} + svc := NewEmailService(nil, cache) + + require.NoError(t, svc.ConsumePasswordResetToken(context.Background(), "a@b.c", token)) + require.Equal(t, hash, cache.consumedHash) + require.ErrorIs(t, svc.ConsumePasswordResetToken(context.Background(), "a@b.c", token), ErrInvalidResetToken) + + // A legacy plaintext value (issued before upgrade) no longer validates. + cache.stored = &PasswordResetTokenData{Token: token} + require.ErrorIs(t, svc.ConsumePasswordResetToken(context.Background(), "a@b.c", token), ErrInvalidResetToken) +} diff --git a/backend/internal/service/gateway_systemone.go b/backend/internal/service/gateway_systemone.go new file mode 100644 index 000000000000..0add1999df03 --- /dev/null +++ b/backend/internal/service/gateway_systemone.go @@ -0,0 +1,182 @@ +package service + +import ( + "context" + "errors" + "fmt" + "mime" + "net/http" + "strings" + "time" + + "github.com/Wei-Shaw/sub2api/internal/pkg/typesafe" + "github.com/gin-gonic/gin" +) + +type SystemOneForwardResult struct { + ForwardResult + StatusCode int + Body []byte + ContentType string +} + +type SystemOneUpstreamError struct { + StatusCode int +} + +const TypeSafeCredentialRejectedReason GatewayFailureReason = "typesafe_api_key_rejected" + +func (e *SystemOneUpstreamError) Error() string { + return fmt.Sprintf("typesafe upstream rejected request with status %d", e.StatusCode) +} + +func (s *GatewayService) ForwardSystemOne(ctx context.Context, c *gin.Context, account *Account, body []byte) (*SystemOneForwardResult, error) { + started := time.Now() + if account == nil || !account.IsTypeSafe() || account.Type != AccountTypeAPIKey { + return nil, errors.New("invalid typesafe account") + } + key := account.GetTypeSafeAPIKey() + if key == "" { + return nil, errors.New("typesafe api key is missing") + } + baseURL, err := s.validateUpstreamBaseURL(account.GetTypeSafeBaseURL()) + if err != nil { + return nil, err + } + req, err := typesafe.NewSystemOneRequest(ctx, baseURL, key, body) + if err != nil { + return nil, err + } + upstreamURL := req.URL.Scheme + "://" + req.URL.Host + req.URL.Path + proxyURL := "" + if account.ProxyID != nil && account.Proxy != nil { + proxyURL = account.Proxy.URL() + } + resp, err := s.httpUpstream.Do(req, proxyURL, account.ID, account.Concurrency) + if err != nil { + return nil, s.handleUpstreamTransportError(ctx, c, account, err, OpsUpstreamErrorEvent{ + Passthrough: true, + UpstreamURL: upstreamURL, + }) + } + defer func() { _ = resp.Body.Close() }() + + if resp.StatusCode < http.StatusOK || resp.StatusCode >= http.StatusMultipleChoices { + return nil, s.handleSystemOneErrorResponse(ctx, c, account, resp, upstreamURL) + } + + decoded, err := typesafe.DecodeSystemOneResponse(resp.Body) + if err != nil { + // The upstream accepted (and may have charged) this request but the + // gateway cannot relay it; keep an ops trail for reconciliation. + setOpsUpstreamError(c, resp.StatusCode, err.Error(), "") + appendOpsUpstreamError(c, OpsUpstreamErrorEvent{ + Passthrough: true, + ProxyID: opsUpstreamProxyID(account), + ProxyName: opsUpstreamProxyName(account), + Platform: account.Platform, + AccountID: account.ID, + AccountName: account.Name, + UpstreamStatusCode: resp.StatusCode, + UpstreamRequestID: resp.Header.Get("x-request-id"), + UpstreamURL: upstreamURL, + Kind: "response_error", + Message: err.Error(), + }) + return nil, err + } + return &SystemOneForwardResult{ + ForwardResult: ForwardResult{ + RequestID: resp.Header.Get("x-request-id"), + UpstreamHeaders: resp.Header.Clone(), + Usage: ClaudeUsage{InputTokens: decoded.Usage.InputTokens, OutputTokens: decoded.Usage.OutputTokens}, + Model: typesafe.JevLatestModel, + UpstreamResponseModel: decoded.Model, + Duration: time.Since(started), + }, + StatusCode: resp.StatusCode, + Body: decoded.Body, + ContentType: systemOneResponseContentType(resp.Header.Get("Content-Type")), + }, nil +} + +// IsSystemOneRequestErrorStatus reports upstream statuses that describe the +// caller's own payload (malformed, unprocessable, or too large). +func IsSystemOneRequestErrorStatus(status int) bool { + switch status { + case http.StatusBadRequest, http.StatusRequestEntityTooLarge, http.StatusUnprocessableEntity: + return true + default: + return false + } +} + +// handleSystemOneErrorResponse applies the shared account error policy to a +// non-2xx System One response. 400/413/422 describe the caller's own payload, so +// they never touch account state (a tenant must not be able to disable an +// account with bad input) and are not retried elsewhere. Every other status +// goes through the account error policy (custom error codes, temporary +// unschedulable rules, pool mode) and fails over when the status is retryable +// or the policy took the account out of rotation. +func (s *GatewayService) handleSystemOneErrorResponse(ctx context.Context, c *gin.Context, account *Account, resp *http.Response, upstreamURL string) error { + respBody, _ := s.readUpstreamErrorBody(resp) + upstreamMsg := sanitizeUpstreamErrorMessage(strings.TrimSpace(extractUpstreamErrorMessage(respBody))) + setOpsUpstreamError(c, resp.StatusCode, upstreamMsg, "") + event := OpsUpstreamErrorEvent{ + Passthrough: true, + ProxyID: opsUpstreamProxyID(account), + ProxyName: opsUpstreamProxyName(account), + Platform: account.Platform, + AccountID: account.ID, + AccountName: account.Name, + UpstreamStatusCode: resp.StatusCode, + UpstreamRequestID: resp.Header.Get("x-request-id"), + UpstreamURL: upstreamURL, + Kind: "http_error", + Message: upstreamMsg, + } + + if IsSystemOneRequestErrorStatus(resp.StatusCode) { + appendOpsUpstreamError(c, event) + return &SystemOneUpstreamError{StatusCode: resp.StatusCode} + } + + shouldDisable := false + if s.rateLimitService != nil { + shouldDisable = s.rateLimitService.HandleUpstreamError(ctx, account, resp.StatusCode, resp.Header, respBody, typesafe.JevLatestModel) + } + if !shouldDisable && !s.shouldFailoverUpstreamError(resp.StatusCode) { + appendOpsUpstreamError(c, event) + return &SystemOneUpstreamError{StatusCode: resp.StatusCode} + } + + event.Kind = "failover" + appendOpsUpstreamError(c, event) + failoverErr := &UpstreamFailoverError{ + StatusCode: resp.StatusCode, + ResponseBody: respBody, + ResponseHeaders: resp.Header.Clone(), + RetryableOnSameAccount: !shouldDisable && account.IsPoolMode() && account.IsPoolModeRetryableStatus(resp.StatusCode), + } + if resp.StatusCode == http.StatusUnauthorized { + failoverErr.Stage = GatewayFailureStageAccountAuth + failoverErr.Scope = GatewayFailureScopeAccount + failoverErr.Reason = TypeSafeCredentialRejectedReason + failoverErr.NextAccountAction = NextAccountRetry + } + return failoverErr +} + +// systemOneResponseContentType keeps the upstream JSON media type (and its +// charset) but never relays a non-JSON type for a body already validated as +// JSON, so the gateway origin cannot be made to serve it as HTML. +func systemOneResponseContentType(raw string) string { + mediaType, _, err := mime.ParseMediaType(strings.TrimSpace(raw)) + if err != nil { + return "application/json" + } + if mediaType == "application/json" || (strings.HasPrefix(mediaType, "application/") && strings.HasSuffix(mediaType, "+json")) { + return strings.TrimSpace(raw) + } + return "application/json" +} diff --git a/backend/internal/service/gateway_systemone_test.go b/backend/internal/service/gateway_systemone_test.go new file mode 100644 index 000000000000..d3cb1314147d --- /dev/null +++ b/backend/internal/service/gateway_systemone_test.go @@ -0,0 +1,402 @@ +package service + +import ( + "context" + "errors" + "io" + "net/http" + "net/http/httptest" + "strings" + "testing" + "time" + + "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/Wei-Shaw/sub2api/internal/pkg/tlsfingerprint" + "github.com/Wei-Shaw/sub2api/internal/pkg/typesafe" + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" +) + +type systemOneHTTPUpstream struct { + do func(*http.Request) (*http.Response, error) +} + +type systemOnePolicyAccountRepo struct { + AccountRepository + account *Account + rateLimitedCalls int + overloadedCalls int + errorCalls int +} + +func (r *systemOnePolicyAccountRepo) GetByID(context.Context, int64) (*Account, error) { + return r.account, nil +} + +func (r *systemOnePolicyAccountRepo) SetRateLimited(context.Context, int64, time.Time) error { + r.rateLimitedCalls++ + return nil +} + +func (r *systemOnePolicyAccountRepo) SetOverloaded(context.Context, int64, time.Time) error { + r.overloadedCalls++ + return nil +} + +func (r *systemOnePolicyAccountRepo) SetError(context.Context, int64, string) error { + r.errorCalls++ + return nil +} + +func (u *systemOneHTTPUpstream) Do(req *http.Request, _ string, _ int64, _ int) (*http.Response, error) { + return u.do(req) +} + +func (u *systemOneHTTPUpstream) DoWithTLS(req *http.Request, _ string, _ int64, _ int, _ *tlsfingerprint.Profile) (*http.Response, error) { + return u.do(req) +} + +func newSystemOneTestService(upstream HTTPUpstream) *GatewayService { + return &GatewayService{ + cfg: &config.Config{Security: config.SecurityConfig{URLAllowlist: config.URLAllowlistConfig{AllowInsecureHTTP: true}}}, + httpUpstream: upstream, + } +} + +func newSystemOneTestContext() *gin.Context { + gin.SetMode(gin.TestMode) + c, _ := gin.CreateTestContext(httptest.NewRecorder()) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/systemone", nil) + return c +} + +func TestForwardSystemOneForwardsNativeProtocolAndUsage(t *testing.T) { + requestBody := []byte(`{"model":"jev-latest","state":{"text":"sample"},"questions":{"q":{"type":"choice","instructions":"Pick","criteria":{"a":"A","b":"B"}}}}`) + responseBody := []byte(`{"model":"jev-1.13.0","answers":{"q":{"type":"choice","choice":"a"}},"usage":{"input_tokens":123,"output_tokens":7},"provider_extension":{"kept":true}}`) + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + require.Equal(t, typesafeSystemOnePathForTest, r.URL.Path) + require.Equal(t, "Bearer ts-secret", r.Header.Get("Authorization")) + require.Equal(t, "application/json", r.Header.Get("Content-Type")) + got, err := io.ReadAll(r.Body) + require.NoError(t, err) + require.Equal(t, requestBody, got) + w.Header().Set("Content-Type", "application/json; charset=utf-8") + w.Header().Set("x-request-id", "req-jev") + _, err = w.Write(responseBody) + require.NoError(t, err) + })) + defer server.Close() + + svc := newSystemOneTestService(&systemOneHTTPUpstream{do: server.Client().Do}) + account := &Account{ID: 7, Platform: PlatformTypeSafe, Type: AccountTypeAPIKey, Credentials: map[string]any{"base_url": server.URL, "api_key": "ts-secret"}} + result, err := svc.ForwardSystemOne(context.Background(), newSystemOneTestContext(), account, requestBody) + require.NoError(t, err) + require.Equal(t, responseBody, result.Body) + require.Equal(t, http.StatusOK, result.StatusCode) + require.Equal(t, "application/json; charset=utf-8", result.ContentType) + require.Equal(t, "req-jev", result.RequestID) + require.Equal(t, "jev-latest", result.Model) + require.Equal(t, "jev-1.13.0", result.UpstreamResponseModel) + require.Equal(t, 123, result.Usage.InputTokens) + require.Equal(t, 7, result.Usage.OutputTokens) +} + +func TestForwardSystemOneAllowsSuccessfulResponseWithoutModel(t *testing.T) { + responseBody := []byte(`{"answers":{"q":{"type":"noul","answer":"ok"}},"usage":{"input_tokens":12}}`) + upstream := &systemOneHTTPUpstream{do: func(*http.Request) (*http.Response, error) { + return &http.Response{StatusCode: http.StatusOK, Header: make(http.Header), Body: io.NopCloser(strings.NewReader(string(responseBody)))}, nil + }} + svc := newSystemOneTestService(upstream) + account := &Account{ID: 11, Platform: PlatformTypeSafe, Type: AccountTypeAPIKey, Credentials: map[string]any{"base_url": "http://typesafe.test", "api_key": "ts-secret"}} + + result, err := svc.ForwardSystemOne(context.Background(), newSystemOneTestContext(), account, []byte(`{}`)) + require.NoError(t, err) + require.Equal(t, responseBody, result.Body) + require.Empty(t, result.UpstreamResponseModel) + require.Equal(t, 12, result.Usage.InputTokens) +} + +const typesafeSystemOnePathForTest = "/v1/systemone" + +func TestForwardSystemOneSchemaConformancePassthrough(t *testing.T) { + for _, tc := range []struct { + name string + body string + }{ + {"omitted noul instructions", `{"model":"jev-latest","state":"sample","questions":{"q":{"type":"noul","extension":{"kept":true}}}}`}, + {"nullable noul fields", `{"model":"jev-latest","state":"sample","questions":{"q":{"type":"noul","instructions":null,"criteria":null}}}`}, + {"structured noul descriptions", `{"model":"jev-latest","state":"sample","questions":{"q":{"type":"noul","criteria":{"true":{"reason":"Yes"},"false":["No",null]}}}}`}, + {"object choice description", `{"model":"jev-latest","state":"sample","questions":{"q":{"type":"choice","criteria":{"a":{"description":"A","extra":null}}}}}`}, + {"array choice description", `{"model":"jev-latest","state":"sample","questions":{"q":{"type":"choice","instructions":null,"criteria":{"a":["A",null],"b":null}}}}`}, + {"empty choice criteria", `{"model":"jev-latest","state":"sample","questions":{"q":{"type":"choice","criteria":{}}}}`}, + {"one-level score", `{"model":"jev-latest","state":"sample","questions":{"q":{"type":"score","criteria":["only"]}}}`}, + {"object score level", `{"model":"jev-latest","state":"sample","questions":{"q":{"type":"score","instructions":null,"criteria":[{"description":"only","extra":null}]}}}`}, + {"array score level", `{"model":"jev-latest","state":"sample","questions":{"q":{"type":"score","criteria":[["only",null]]}}}`}, + } { + t.Run(tc.name, func(t *testing.T) { + requestBody := []byte(tc.body) + _, err := typesafe.ValidateSystemOneRequest(requestBody) + require.NoError(t, err) + responseBody := []byte(`{"answers":{},"usage":{"input_tokens":12}}`) + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + require.Equal(t, "/v1/systemone", r.URL.Path) + require.Equal(t, "Bearer ts-mock-key", r.Header.Get("Authorization")) + got, err := io.ReadAll(r.Body) + require.NoError(t, err) + require.Equal(t, requestBody, got) + w.Header().Set("Content-Type", "application/json") + _, err = w.Write(responseBody) + require.NoError(t, err) + })) + defer server.Close() + + svc := newSystemOneTestService(&systemOneHTTPUpstream{do: server.Client().Do}) + account := &Account{ID: 12, Platform: PlatformTypeSafe, Type: AccountTypeAPIKey, Credentials: map[string]any{"base_url": server.URL, "api_key": "ts-mock-key"}} + result, err := svc.ForwardSystemOne(context.Background(), newSystemOneTestContext(), account, requestBody) + require.NoError(t, err) + require.Equal(t, responseBody, result.Body) + require.Equal(t, 12, result.Usage.InputTokens) + }) + } +} + +func TestForwardSystemOneErrorPolicy(t *testing.T) { + for _, status := range []int{400, 422, 401, 429, 529, 500, 503} { + t.Run(http.StatusText(status), func(t *testing.T) { + upstream := &systemOneHTTPUpstream{do: func(*http.Request) (*http.Response, error) { + return &http.Response{StatusCode: status, Header: make(http.Header), Body: io.NopCloser(strings.NewReader("private upstream body ts-secret"))}, nil + }} + svc := newSystemOneTestService(upstream) + account := &Account{ID: 8, Platform: PlatformTypeSafe, Type: AccountTypeAPIKey, Credentials: map[string]any{"base_url": "http://typesafe.test", "api_key": "ts-secret"}} + result, err := svc.ForwardSystemOne(context.Background(), newSystemOneTestContext(), account, []byte(`{"private":"request"}`)) + require.Nil(t, result) + require.Error(t, err) + require.NotContains(t, err.Error(), "private") + require.NotContains(t, err.Error(), "ts-secret") + if status == 400 || status == 422 { + var upstreamErr *SystemOneUpstreamError + require.ErrorAs(t, err, &upstreamErr) + require.Equal(t, status, upstreamErr.StatusCode) + return + } + var failoverErr *UpstreamFailoverError + require.ErrorAs(t, err, &failoverErr) + require.True(t, failoverErr.ShouldRetryNextAccount()) + if status == http.StatusUnauthorized { + require.Equal(t, GatewayFailureStageAccountAuth, failoverErr.Stage) + require.Equal(t, GatewayFailureScopeAccount, failoverErr.Scope) + require.Equal(t, TypeSafeCredentialRejectedReason, failoverErr.Reason) + } + }) + } +} + +func TestForwardSystemOneAppliesExistingAccountStatePolicy(t *testing.T) { + for _, tc := range []struct { + name string + status int + wantRateLimited int + wantOverloaded int + wantError int + }{ + {name: "unauthorized", status: http.StatusUnauthorized, wantError: 1}, + {name: "rate limited", status: http.StatusTooManyRequests, wantRateLimited: 1}, + {name: "overloaded", status: 529, wantOverloaded: 1}, + } { + t.Run(tc.name, func(t *testing.T) { + account := &Account{ID: 10, Platform: PlatformTypeSafe, Type: AccountTypeAPIKey, Credentials: map[string]any{"base_url": "http://typesafe.test", "api_key": "ts-secret"}} + repo := &systemOnePolicyAccountRepo{account: account} + upstream := &systemOneHTTPUpstream{do: func(*http.Request) (*http.Response, error) { + return &http.Response{StatusCode: tc.status, Header: make(http.Header), Body: io.NopCloser(strings.NewReader(`{"error":"private"}`))}, nil + }} + svc := newSystemOneTestService(upstream) + svc.rateLimitService = NewRateLimitService(repo, nil, &config.Config{}, nil, nil) + + _, err := svc.ForwardSystemOne(context.Background(), newSystemOneTestContext(), account, []byte(`{}`)) + require.Error(t, err) + require.Equal(t, tc.wantRateLimited, repo.rateLimitedCalls) + require.Equal(t, tc.wantOverloaded, repo.overloadedCalls) + require.Equal(t, tc.wantError, repo.errorCalls) + }) + } +} + +func TestForwardSystemOneTransportAndTimeoutFailOver(t *testing.T) { + for _, transportErr := range []error{errors.New("network unavailable"), context.DeadlineExceeded} { + upstream := &systemOneHTTPUpstream{do: func(*http.Request) (*http.Response, error) { return nil, transportErr }} + svc := newSystemOneTestService(upstream) + account := &Account{ID: 9, Name: "jev", Platform: PlatformTypeSafe, Type: AccountTypeAPIKey, Credentials: map[string]any{"base_url": "http://typesafe.test", "api_key": "ts-secret"}} + _, err := svc.ForwardSystemOne(context.Background(), newSystemOneTestContext(), account, []byte(`{}`)) + var failoverErr *UpstreamFailoverError + require.ErrorAs(t, err, &failoverErr) + } +} + +func newSystemOneStatusUpstream(status int, body string) *systemOneHTTPUpstream { + return &systemOneHTTPUpstream{do: func(*http.Request) (*http.Response, error) { + return &http.Response{StatusCode: status, Header: http.Header{"X-Request-Id": []string{"req-err"}}, Body: io.NopCloser(strings.NewReader(body))}, nil + }} +} + +func systemOneOpsEvents(t *testing.T, c *gin.Context) []*OpsUpstreamErrorEvent { + t.Helper() + raw, ok := c.Get(OpsUpstreamErrorsKey) + require.True(t, ok) + events, ok := raw.([]*OpsUpstreamErrorEvent) + require.True(t, ok) + return events +} + +func TestForwardSystemOneRequestErrorsNeverTouchAccountState(t *testing.T) { + for _, status := range []int{http.StatusBadRequest, http.StatusRequestEntityTooLarge, http.StatusUnprocessableEntity} { + t.Run(http.StatusText(status), func(t *testing.T) { + // Even a custom error-code rule must not let client input disable the account. + account := &Account{ID: 21, Platform: PlatformTypeSafe, Type: AccountTypeAPIKey, Credentials: map[string]any{ + "base_url": "http://typesafe.test", "api_key": "ts-secret", + "custom_error_codes_enabled": true, "custom_error_codes": []any{float64(status)}, + }} + repo := &systemOnePolicyAccountRepo{account: account} + svc := newSystemOneTestService(newSystemOneStatusUpstream(status, `{"detail":"bad question"}`)) + svc.rateLimitService = NewRateLimitService(repo, nil, &config.Config{}, nil, nil) + c := newSystemOneTestContext() + + _, err := svc.ForwardSystemOne(context.Background(), c, account, []byte(`{}`)) + var upstreamErr *SystemOneUpstreamError + require.ErrorAs(t, err, &upstreamErr) + require.Equal(t, status, upstreamErr.StatusCode) + require.Zero(t, repo.errorCalls+repo.rateLimitedCalls+repo.overloadedCalls) + events := systemOneOpsEvents(t, c) + require.Len(t, events, 1) + require.Equal(t, "http_error", events[0].Kind) + require.Equal(t, status, events[0].UpstreamStatusCode) + require.Equal(t, "req-err", events[0].UpstreamRequestID) + }) + } +} + +func TestForwardSystemOneAccountLevelFailuresFailOver(t *testing.T) { + for _, tc := range []struct { + name string + status int + wantError int + }{ + {name: "payment required", status: http.StatusPaymentRequired, wantError: 1}, + {name: "forbidden", status: http.StatusForbidden, wantError: 1}, + } { + t.Run(tc.name, func(t *testing.T) { + account := &Account{ID: 22, Platform: PlatformTypeSafe, Type: AccountTypeAPIKey, Credentials: map[string]any{"base_url": "http://typesafe.test", "api_key": "ts-secret"}} + repo := &systemOnePolicyAccountRepo{account: account} + svc := newSystemOneTestService(newSystemOneStatusUpstream(tc.status, `{"detail":"account problem"}`)) + svc.rateLimitService = NewRateLimitService(repo, nil, &config.Config{}, nil, nil) + c := newSystemOneTestContext() + + _, err := svc.ForwardSystemOne(context.Background(), c, account, []byte(`{}`)) + var failoverErr *UpstreamFailoverError + require.ErrorAs(t, err, &failoverErr) + require.True(t, failoverErr.ShouldRetryNextAccount()) + require.Equal(t, tc.status, failoverErr.StatusCode) + require.Equal(t, tc.wantError, repo.errorCalls) + require.Equal(t, "failover", systemOneOpsEvents(t, c)[0].Kind) + }) + } +} + +func TestForwardSystemOneHonorsCustomErrorCodes(t *testing.T) { + account := &Account{ID: 23, Platform: PlatformTypeSafe, Type: AccountTypeAPIKey, Credentials: map[string]any{ + "base_url": "http://typesafe.test", "api_key": "ts-secret", + "custom_error_codes_enabled": true, "custom_error_codes": []any{float64(http.StatusNotFound)}, + }} + repo := &systemOnePolicyAccountRepo{account: account} + svc := newSystemOneTestService(newSystemOneStatusUpstream(http.StatusNotFound, `{"detail":"missing"}`)) + svc.rateLimitService = NewRateLimitService(repo, nil, &config.Config{}, nil, nil) + + _, err := svc.ForwardSystemOne(context.Background(), newSystemOneTestContext(), account, []byte(`{}`)) + var failoverErr *UpstreamFailoverError + require.ErrorAs(t, err, &failoverErr) + require.Equal(t, 1, repo.errorCalls) +} + +func TestForwardSystemOneUnhandledStatusDoesNotFailOver(t *testing.T) { + account := &Account{ID: 24, Platform: PlatformTypeSafe, Type: AccountTypeAPIKey, Credentials: map[string]any{"base_url": "http://typesafe.test", "api_key": "ts-secret"}} + repo := &systemOnePolicyAccountRepo{account: account} + svc := newSystemOneTestService(newSystemOneStatusUpstream(http.StatusConflict, `{"detail":"conflict"}`)) + svc.rateLimitService = NewRateLimitService(repo, nil, &config.Config{}, nil, nil) + + _, err := svc.ForwardSystemOne(context.Background(), newSystemOneTestContext(), account, []byte(`{}`)) + var upstreamErr *SystemOneUpstreamError + require.ErrorAs(t, err, &upstreamErr) + require.Equal(t, http.StatusConflict, upstreamErr.StatusCode) + require.Zero(t, repo.errorCalls) +} + +func TestForwardSystemOneNormalizesResponseContentType(t *testing.T) { + for _, tc := range []struct{ upstream, want string }{ + {"application/json; charset=utf-8", "application/json; charset=utf-8"}, + {"application/problem+json", "application/problem+json"}, + {"text/html; charset=utf-8", "application/json"}, + {"", "application/json"}, + {"not a media type;;", "application/json"}, + } { + upstream := &systemOneHTTPUpstream{do: func(*http.Request) (*http.Response, error) { + header := make(http.Header) + if tc.upstream != "" { + header.Set("Content-Type", tc.upstream) + } + return &http.Response{StatusCode: http.StatusOK, Header: header, Body: io.NopCloser(strings.NewReader(`{"answers":{}}`))}, nil + }} + account := &Account{ID: 25, Platform: PlatformTypeSafe, Type: AccountTypeAPIKey, Credentials: map[string]any{"base_url": "http://typesafe.test", "api_key": "ts-secret"}} + result, err := newSystemOneTestService(upstream).ForwardSystemOne(context.Background(), newSystemOneTestContext(), account, []byte(`{}`)) + require.NoError(t, err) + require.Equal(t, tc.want, result.ContentType, tc.upstream) + } +} + +func TestForwardSystemOneRejectsOversizedResponse(t *testing.T) { + oversized := `{"pad":"` + strings.Repeat("a", typesafe.MaxSystemOneResponseBytes) + `"}` + upstream := &systemOneHTTPUpstream{do: func(*http.Request) (*http.Response, error) { + return &http.Response{StatusCode: http.StatusOK, Header: make(http.Header), Body: io.NopCloser(strings.NewReader(oversized))}, nil + }} + account := &Account{ID: 26, Platform: PlatformTypeSafe, Type: AccountTypeAPIKey, Credentials: map[string]any{"base_url": "http://typesafe.test", "api_key": "ts-secret"}} + c := newSystemOneTestContext() + _, err := newSystemOneTestService(upstream).ForwardSystemOne(context.Background(), c, account, []byte(`{}`)) + require.ErrorIs(t, err, typesafe.ErrSystemOneResponseTooLarge) + events := systemOneOpsEvents(t, c) + require.Len(t, events, 1) + require.Equal(t, "response_error", events[0].Kind) + require.Equal(t, http.StatusOK, events[0].UpstreamStatusCode) +} + +func TestForwardSystemOneBillsLenientUsageShapes(t *testing.T) { + upstream := &systemOneHTTPUpstream{do: func(*http.Request) (*http.Response, error) { + return &http.Response{StatusCode: http.StatusOK, Header: make(http.Header), Body: io.NopCloser(strings.NewReader(`{"answers":{},"usage":{"input_tokens":"21","output_tokens":1.0}}`))}, nil + }} + account := &Account{ID: 27, Platform: PlatformTypeSafe, Type: AccountTypeAPIKey, Credentials: map[string]any{"base_url": "http://typesafe.test", "api_key": "ts-secret"}} + result, err := newSystemOneTestService(upstream).ForwardSystemOne(context.Background(), newSystemOneTestContext(), account, []byte(`{}`)) + require.NoError(t, err) + require.Equal(t, 21, result.Usage.InputTokens) + require.Equal(t, 1, result.Usage.OutputTokens) +} + +func TestTypeSafeAccountBaseURLNeverFallsBackToAnthropic(t *testing.T) { + account := &Account{Platform: PlatformTypeSafe, Type: AccountTypeAPIKey, Credentials: map[string]any{"api_key": "ts-secret"}} + require.Equal(t, typesafe.DefaultBaseURL, account.GetBaseURL()) + require.Equal(t, typesafe.DefaultBaseURL, account.GetTypeSafeBaseURL()) + for raw, want := range map[string]string{ + "https://api.typesafe.ai/v1": "https://api.typesafe.ai", + "https://api.typesafe.ai/V1/": "https://api.typesafe.ai", + "https://proxy.example/typesafe/": "https://proxy.example/typesafe", + "https://proxy.example/apiv1": "https://proxy.example/apiv1", + " https://proxy.example/x/v1/ ": "https://proxy.example/x", + "/v1": typesafe.DefaultBaseURL, + } { + withBase := &Account{Platform: PlatformTypeSafe, Type: AccountTypeAPIKey, Credentials: map[string]any{"base_url": raw}} + require.Equal(t, want, withBase.GetTypeSafeBaseURL(), raw) + } + anthropic := &Account{Platform: PlatformAnthropic, Type: AccountTypeAPIKey, Credentials: map[string]any{"api_key": "sk"}} + require.Equal(t, "https://api.anthropic.com", anthropic.GetBaseURL()) +} + +func TestTypeSafeModelsListCandidates(t *testing.T) { + require.Equal(t, []string{typesafe.JevLatestModel}, defaultModelsListCandidateIDs(PlatformTypeSafe)) + require.NotContains(t, compositeDefaultModelsListCandidateIDs(), typesafe.JevLatestModel) +} diff --git a/backend/internal/service/grok_upstream_headers.go b/backend/internal/service/grok_upstream_headers.go index 4b24e4c13b3e..75e5be5c71c1 100644 --- a/backend/internal/service/grok_upstream_headers.go +++ b/backend/internal/service/grok_upstream_headers.go @@ -17,7 +17,7 @@ const ( grokClientModeHeader = xai.CLIClientMode ) -// defaultGrokUpstreamUserAgent is the pinned Grok CLI / workspace UA. +// defaultGrokUpstreamUserAgent 使用固定版本的官方交互式 CLI UA。 // Grok upstream must not forward Claude Code / Codex / browser client UAs. func defaultGrokUpstreamUserAgent() string { return xai.CLIUserAgent(xai.ResolveCLIVersion()) diff --git a/backend/internal/service/grok_upstream_headers_test.go b/backend/internal/service/grok_upstream_headers_test.go index 76908a0aa52a..7c2bcafa9927 100644 --- a/backend/internal/service/grok_upstream_headers_test.go +++ b/backend/internal/service/grok_upstream_headers_test.go @@ -28,7 +28,7 @@ func TestApplyDefaultGrokUpstreamHeadersUsesCLIUserAgent(t *testing.T) { } func TestApplyDefaultGrokUpstreamHeadersHonorsCLIVersionOverride(t *testing.T) { - t.Setenv(xai.CLIVersionEnv, "1.0.45-alpha.1") + t.Setenv(xai.CLIVersionEnv, "1.0.47-alpha.1") req, err := http.NewRequest(http.MethodGet, "https://api.x.ai/v1/responses", nil) require.NoError(t, err) @@ -36,9 +36,9 @@ func TestApplyDefaultGrokUpstreamHeadersHonorsCLIVersionOverride(t *testing.T) { applyDefaultGrokUpstreamHeaders(req) - require.Equal(t, "1.0.45-alpha.1", req.Header.Get("x-grok-client-version")) - require.Equal(t, xai.CLIUserAgent("1.0.45-alpha.1"), req.Header.Get("User-Agent")) - require.Equal(t, "grok-shell", req.Header.Get("x-grok-client-identifier")) + require.Equal(t, "1.0.47-alpha.1", req.Header.Get("x-grok-client-version")) + require.Equal(t, xai.CLIUserAgent("1.0.47-alpha.1"), req.Header.Get("User-Agent")) + require.Equal(t, "grok-pager", req.Header.Get("x-grok-client-identifier")) } func TestResolveGrokUpstreamUserAgentNeverPassthrough(t *testing.T) { diff --git a/backend/internal/service/openai_gateway_grok.go b/backend/internal/service/openai_gateway_grok.go index 06a17a961e2e..7aaa649a7425 100644 --- a/backend/internal/service/openai_gateway_grok.go +++ b/backend/internal/service/openai_gateway_grok.go @@ -1613,8 +1613,8 @@ func applyGrokCLIHeaders(headers http.Header) { headers.Set("X-Grok-Client-Version", version) headers.Set("x-grok-client-version", version) headers.Set("x-grok-client-identifier", xai.CLIClientIdentifier) - // Historical mode value expected by some unit tests / older CLI probes. - headers.Set("X-Grok-Client-Mode", "interactive") + // 对齐官方 CLI 交互模式,网关请求与额度探测共用身份。 + headers.Set("X-Grok-Client-Mode", xai.CLIClientMode) } func (s *OpenAIGatewayService) updateGrokUsageSnapshot(ctx context.Context, account *Account, snapshot *xai.QuotaSnapshot) { diff --git a/backend/internal/service/payment_config_service.go b/backend/internal/service/payment_config_service.go index aac1cc8e263b..ad5583da27a6 100644 --- a/backend/internal/service/payment_config_service.go +++ b/backend/internal/service/payment_config_service.go @@ -62,12 +62,18 @@ type PaymentConfig struct { // SubscriptionUSDToCNYRate 为 0 时订阅换算关闭(兼容存量行为)。 SubscriptionUSDToCNYRate float64 `json:"subscription_usd_to_cny_rate"` RechargeFeeRate float64 `json:"recharge_fee_rate"` - LoadBalanceStrategy string `json:"load_balance_strategy"` - ProductNamePrefix string `json:"product_name_prefix"` - ProductNameSuffix string `json:"product_name_suffix"` - HelpImageURL string `json:"help_image_url"` - HelpText string `json:"help_text"` - StripePublishableKey string `json:"stripe_publishable_key,omitempty"` + // RechargeBonusTiers 余额充值优惠阶梯(按 MinAmount 升序);空表示无优惠。 + RechargeBonusTiers []RechargeBonusTier `json:"recharge_bonus_tiers"` + // RechargeBonusMode 阶梯模式:bonus(赠金)/ discount(折扣),已归一化。 + RechargeBonusMode string `json:"recharge_bonus_mode"` + // RechargeBonusNotice 充值页展示的 Markdown 活动文案;空表示不展示。 + RechargeBonusNotice string `json:"recharge_bonus_notice"` + LoadBalanceStrategy string `json:"load_balance_strategy"` + ProductNamePrefix string `json:"product_name_prefix"` + ProductNameSuffix string `json:"product_name_suffix"` + HelpImageURL string `json:"help_image_url"` + HelpText string `json:"help_text"` + StripePublishableKey string `json:"stripe_publishable_key,omitempty"` // Cancel rate limit settings CancelRateLimitEnabled bool `json:"cancel_rate_limit_enabled"` @@ -95,11 +101,15 @@ type UpdatePaymentConfigRequest struct { BalanceRechargeMultiplier *float64 `json:"balance_recharge_multiplier"` SubscriptionUSDToCNYRate *float64 `json:"subscription_usd_to_cny_rate"` RechargeFeeRate *float64 `json:"recharge_fee_rate"` - LoadBalanceStrategy *string `json:"load_balance_strategy"` - ProductNamePrefix *string `json:"product_name_prefix"` - ProductNameSuffix *string `json:"product_name_suffix"` - HelpImageURL *string `json:"help_image_url"` - HelpText *string `json:"help_text"` + // RechargeBonusTiers nil 表示不更新;空切片表示清空阶梯。 + RechargeBonusTiers *[]RechargeBonusTier `json:"recharge_bonus_tiers"` + RechargeBonusMode *string `json:"recharge_bonus_mode"` + RechargeBonusNotice *string `json:"recharge_bonus_notice"` + LoadBalanceStrategy *string `json:"load_balance_strategy"` + ProductNamePrefix *string `json:"product_name_prefix"` + ProductNameSuffix *string `json:"product_name_suffix"` + HelpImageURL *string `json:"help_image_url"` + HelpText *string `json:"help_text"` // Cancel rate limit settings CancelRateLimitEnabled *bool `json:"cancel_rate_limit_enabled"` @@ -220,6 +230,7 @@ func (s *PaymentConfigService) GetPaymentConfig(ctx context.Context) (*PaymentCo SettingPaymentEnabled, SettingMinRechargeAmount, SettingMaxRechargeAmount, SettingDailyRechargeLimit, SettingOrderTimeoutMinutes, SettingMaxPendingOrders, SettingEnabledPaymentTypes, SettingBalancePayDisabled, SettingBalanceRechargeMult, SettingSubscriptionUSDToCNYRate, SettingRechargeFeeRate, SettingLoadBalanceStrategy, + SettingRechargeBonusTiers, SettingRechargeBonusMode, SettingRechargeBonusNotice, SettingProductNamePrefix, SettingProductNameSuffix, SettingHelpImageURL, SettingHelpText, SettingCancelRateLimitOn, SettingCancelRateLimitMax, @@ -250,6 +261,8 @@ func (s *PaymentConfigService) parsePaymentConfig(vals map[string]string) *Payme BalanceRechargeMultiplier: normalizeBalanceRechargeMultiplier(pcParseFloat(vals[SettingBalanceRechargeMult], defaultBalanceRechargeMultiplier)), SubscriptionUSDToCNYRate: normalizeSubscriptionUSDToCNYRate(pcParseFloat(vals[SettingSubscriptionUSDToCNYRate], 0)), RechargeFeeRate: pcParseFloat(vals[SettingRechargeFeeRate], 0), + RechargeBonusTiers: parseRechargeBonusTiers(vals[SettingRechargeBonusTiers]), + RechargeBonusNotice: vals[SettingRechargeBonusNotice], LoadBalanceStrategy: vals[SettingLoadBalanceStrategy], ProductNamePrefix: vals[SettingProductNamePrefix], ProductNameSuffix: vals[SettingProductNameSuffix], @@ -265,6 +278,7 @@ func (s *PaymentConfigService) parsePaymentConfig(vals map[string]string) *Payme AlipayForceQRCode: vals[SettingAlipayForceQRCode] == "true", AlipayMobilePrecreateDeepLink: vals[SettingAlipayMobilePrecreateDeepLink] == "true", } + cfg.RechargeBonusMode, _ = NormalizeRechargeBonusMode(vals[SettingRechargeBonusMode]) cfg.AlipayMobilePrecreateDeepLink = pcEnvBoolOverride( SettingAlipayMobilePrecreateDeepLink, cfg.AlipayMobilePrecreateDeepLink, @@ -343,6 +357,15 @@ func (s *PaymentConfigService) UpdatePaymentConfig(ctx context.Context, req Upda return infraerrors.BadRequest("INVALID_RECHARGE_FEE_RATE", "recharge fee rate allows at most 2 decimal places") } } + rechargeBonusTiersValue, rechargeBonusModeValue, err := s.resolveRechargeBonusUpdate(ctx, req) + if err != nil { + return err + } + if req.RechargeBonusNotice != nil { + if err := validateRechargeBonusNotice(*req.RechargeBonusNotice); err != nil { + return infraerrors.BadRequest("INVALID_RECHARGE_BONUS_NOTICE", err.Error()) + } + } m := make(map[string]string) if req.Enabled != nil { m[SettingPaymentEnabled] = formatBoolOrEmpty(req.Enabled) @@ -377,6 +400,15 @@ func (s *PaymentConfigService) UpdatePaymentConfig(ctx context.Context, req Upda if req.RechargeFeeRate != nil { m[SettingRechargeFeeRate] = formatNonNegativeFloat(req.RechargeFeeRate) } + if req.RechargeBonusTiers != nil { + m[SettingRechargeBonusTiers] = rechargeBonusTiersValue + } + if req.RechargeBonusMode != nil { + m[SettingRechargeBonusMode] = rechargeBonusModeValue + } + if req.RechargeBonusNotice != nil { + m[SettingRechargeBonusNotice] = strings.TrimSpace(*req.RechargeBonusNotice) + } if req.LoadBalanceStrategy != nil { m[SettingLoadBalanceStrategy] = derefStr(req.LoadBalanceStrategy) } diff --git a/backend/internal/service/payment_fulfillment.go b/backend/internal/service/payment_fulfillment.go index 2b70a90ef65c..4c363d3e0373 100644 --- a/backend/internal/service/payment_fulfillment.go +++ b/backend/internal/service/payment_fulfillment.go @@ -741,7 +741,10 @@ func affiliateRebateBaseAmount(o *dbent.PaymentOrder) float64 { return 0 } switch o.OrderType { - case payment.OrderTypeBalance, payment.OrderTypeSubscription: + case payment.OrderTypeBalance: + // 返利只按实充部分计算,赠送额度不参与 + return paymentOrderAmountWithoutBonus(o) + case payment.OrderTypeSubscription: return o.Amount default: return 0 diff --git a/backend/internal/service/payment_order.go b/backend/internal/service/payment_order.go index da7178dd7fee..becefc6a1a02 100644 --- a/backend/internal/service/payment_order.go +++ b/backend/internal/service/payment_order.go @@ -53,14 +53,6 @@ func (s *PaymentService) CreateOrder(ctx context.Context, req CreateOrderRequest if s.notificationEmailService != nil { s.notificationEmailService.RememberRecipientLocale(ctx, req.UserID, user.Email, req.Locale) } - orderAmount := req.Amount - limitAmount := req.Amount - if plan != nil { - orderAmount = plan.Price - limitAmount = plan.Price - } else if req.OrderType == payment.OrderTypeBalance { - orderAmount = calculateCreditedBalance(req.Amount, cfg.BalanceRechargeMultiplier) - } feeRate := cfg.RechargeFeeRate methodCurrency := payment.DefaultPaymentCurrency if s.configService != nil { @@ -69,6 +61,19 @@ func (s *PaymentService) CreateOrder(ctx context.Context, req CreateOrderRequest return nil, err } } + orderAmount := req.Amount + limitAmount := req.Amount + bonusAmount := 0.0 + if plan != nil { + orderAmount = plan.Price + limitAmount = plan.Price + } else if req.OrderType == payment.OrderTypeBalance { + // 阈值按支付金额命中。赠金模式:到账 = 基数 + 赠送;折扣模式:到账 = 基数,实付基数按折扣减少。 + quote := quoteRechargeBonus(cfg, req.Amount, methodCurrency) + limitAmount = quote.PayBase + bonusAmount = quote.Bonus + orderAmount = quote.Credited + } payAmountStr, payAmount, err := calculateCreateOrderPayAmountForOrderType(limitAmount, feeRate, methodCurrency, req.OrderType, cfg.SubscriptionUSDToCNYRate) if err != nil { return nil, err @@ -100,7 +105,7 @@ func (s *PaymentService) CreateOrder(ctx context.Context, req CreateOrderRequest if oauthResp != nil { return oauthResp, nil } - order, err := s.createOrderInTx(ctx, req, user, plan, cfg, orderAmount, limitAmount, feeRate, payAmount, sel) + order, err := s.createOrderInTx(ctx, req, user, plan, cfg, orderAmount, limitAmount, feeRate, payAmount, bonusAmount, sel) if err != nil { return nil, err } @@ -149,7 +154,7 @@ func (s *PaymentService) validateSubOrder(ctx context.Context, req CreateOrderRe return plan, nil } -func (s *PaymentService) createOrderInTx(ctx context.Context, req CreateOrderRequest, user *User, plan *dbent.SubscriptionPlan, cfg *PaymentConfig, orderAmount, limitAmount, feeRate, payAmount float64, sel *payment.InstanceSelection) (*dbent.PaymentOrder, error) { +func (s *PaymentService) createOrderInTx(ctx context.Context, req CreateOrderRequest, user *User, plan *dbent.SubscriptionPlan, cfg *PaymentConfig, orderAmount, limitAmount, feeRate, payAmount, bonusAmount float64, sel *payment.InstanceSelection) (*dbent.PaymentOrder, error) { tx, err := s.entClient.Tx(ctx) if err != nil { return nil, fmt.Errorf("begin transaction: %w", err) @@ -185,6 +190,7 @@ func (s *PaymentService) createOrderInTx(ctx context.Context, req CreateOrderReq SetAmount(orderAmount). SetPayAmount(payAmount). SetFeeRate(feeRate). + SetBonusAmount(bonusAmount). SetRechargeCode(""). SetOutTradeNo(outTradeNo). SetPaymentType(req.PaymentType). @@ -471,6 +477,7 @@ func (s *PaymentService) invokeProvider(ctx context.Context, order *dbent.Paymen s.writeAuditLog(ctx, order.ID, "ORDER_CREATED", fmt.Sprintf("user:%d", req.UserID), map[string]any{ "paymentAmount": req.Amount, "creditedAmount": order.Amount, + "bonusAmount": order.BonusAmount, "payAmount": order.PayAmount, "paymentType": req.PaymentType, "orderType": req.OrderType, @@ -735,6 +742,7 @@ func buildCreateOrderResponse(order *dbent.PaymentOrder, req CreateOrderRequest, Amount: order.Amount, PayAmount: payAmount, FeeRate: order.FeeRate, + BonusAmount: order.BonusAmount, Status: OrderStatusPending, ResultType: resultType, PaymentType: req.PaymentType, diff --git a/backend/internal/service/payment_order_provider_snapshot_test.go b/backend/internal/service/payment_order_provider_snapshot_test.go index 127202bc2394..19047959ed98 100644 --- a/backend/internal/service/payment_order_provider_snapshot_test.go +++ b/backend/internal/service/payment_order_provider_snapshot_test.go @@ -88,6 +88,7 @@ func TestCreateOrderInTx_WritesProviderSnapshot(t *testing.T) { 88, 0, 88, + 0, &payment.InstanceSelection{ InstanceID: strconv.FormatInt(instance.ID, 10), ProviderKey: payment.TypeAlipay, diff --git a/backend/internal/service/payment_recharge_bonus.go b/backend/internal/service/payment_recharge_bonus.go new file mode 100644 index 000000000000..8d6c1525aa10 --- /dev/null +++ b/backend/internal/service/payment_recharge_bonus.go @@ -0,0 +1,337 @@ +package service + +import ( + "context" + "encoding/json" + "fmt" + "log/slog" + "math" + "sort" + "strings" + "unicode/utf8" + + dbent "github.com/Wei-Shaw/sub2api/ent" + "github.com/Wei-Shaw/sub2api/internal/payment" + infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" + "github.com/shopspring/decimal" +) + +// 充值优惠阶梯:余额充值订单按用户输入的支付金额命中阶梯(取不超过该金额的最大 MinAmount), +// 整条阶梯只有一种模式(RECHARGE_BONUS_MODE): +// - bonus(赠金):实付不变,在到账基数(输入 × 充值倍率)之上额外赠送 BonusPercent% 的 USD 余额; +// - discount(折扣):到账不变(输入 × 倍率),实付基数按 BonusPercent% 打折。 +// +// 两种模式落库形态相同:amount 为到账总额,bonus_amount 为其中「免费」的 USD 部分,pay_amount 为实收。 +// 订阅订单不参与。优惠在下单时按当时配置计算并落库,后续改配置不影响已建订单。 +const ( + // SettingRechargeBonusTiers 存 JSON 数组(RechargeBonusTier 列表),空/缺失表示未启用优惠。 + SettingRechargeBonusTiers = "RECHARGE_BONUS_TIERS" + // SettingRechargeBonusNotice 充值页金额卡顶部展示的 Markdown 活动文案,空表示不展示。 + SettingRechargeBonusNotice = "RECHARGE_BONUS_NOTICE" + // SettingRechargeBonusMode 阶梯模式:bonus / discount;空/非法按 bonus 解析(兼容早期配置)。 + SettingRechargeBonusMode = "RECHARGE_BONUS_MODE" +) + +const ( + RechargeBonusModeBonus = "bonus" + RechargeBonusModeDiscount = "discount" +) + +const ( + maxRechargeBonusTiers = 20 + maxRechargeBonusPercent = 1000 + maxRechargeBonusNoticeRunes = 10000 + rechargeBonusAmountEpsilon = 1e-9 +) + +// RechargeBonusTier 一个优惠档位:支付金额 ≥ MinAmount 时按 BonusPercent% 赠送(bonus)或打折(discount)。 +type RechargeBonusTier struct { + MinAmount float64 `json:"min_amount"` + BonusPercent float64 `json:"bonus_percent"` +} + +// NormalizeRechargeBonusMode 归一化模式;空按 bonus。第二个返回值表示输入是否合法。 +func NormalizeRechargeBonusMode(raw string) (string, bool) { + switch strings.ToLower(strings.TrimSpace(raw)) { + case "", RechargeBonusModeBonus: + return RechargeBonusModeBonus, true + case RechargeBonusModeDiscount: + return RechargeBonusModeDiscount, true + default: + return RechargeBonusModeBonus, false + } +} + +// ValidateRechargeBonusTiersForMode 折扣模式下百分比必须 < 100,否则实付为 0 或负数。 +func ValidateRechargeBonusTiersForMode(mode string, tiers []RechargeBonusTier) error { + if mode != RechargeBonusModeDiscount { + return nil + } + for _, tier := range tiers { + if tier.BonusPercent >= 100 { + return fmt.Errorf("discount percent must be less than 100 (tier with min amount %s)", + decimal.NewFromFloat(tier.MinAmount).Round(2).String()) + } + } + return nil +} + +func rechargeBonusValueValid(v float64, max float64) bool { + if math.IsNaN(v) || math.IsInf(v, 0) || v < 0 || v > max { + return false + } + d := decimal.NewFromFloat(v) + return d.Equal(d.Round(2)) +} + +// NormalizeRechargeBonusTiers 严格归一化(写路径):任何非法项直接报错; +// 成功时返回按 MinAmount 升序排序的副本。 +func NormalizeRechargeBonusTiers(raw []RechargeBonusTier) ([]RechargeBonusTier, error) { + if len(raw) == 0 { + return []RechargeBonusTier{}, nil + } + if len(raw) > maxRechargeBonusTiers { + return nil, fmt.Errorf("recharge bonus tiers exceed limit of %d", maxRechargeBonusTiers) + } + out := make([]RechargeBonusTier, 0, len(raw)) + seen := make(map[string]struct{}, len(raw)) + for _, tier := range raw { + if !rechargeBonusValueValid(tier.MinAmount, math.MaxFloat64) { + return nil, fmt.Errorf("recharge bonus tier min amount must be a non-negative number with at most 2 decimals") + } + if !rechargeBonusValueValid(tier.BonusPercent, maxRechargeBonusPercent) { + return nil, fmt.Errorf("recharge bonus tier percent must be between 0 and %d with at most 2 decimals", maxRechargeBonusPercent) + } + key := decimal.NewFromFloat(tier.MinAmount).Round(2).String() + if _, dup := seen[key]; dup { + return nil, fmt.Errorf("duplicate recharge bonus tier min amount: %s", key) + } + seen[key] = struct{}{} + out = append(out, RechargeBonusTier{MinAmount: tier.MinAmount, BonusPercent: tier.BonusPercent}) + } + sortRechargeBonusTiers(out) + return out, nil +} + +func sortRechargeBonusTiers(tiers []RechargeBonusTier) { + sort.SliceStable(tiers, func(i, j int) bool { + return tiers[i].MinAmount < tiers[j].MinAmount + }) +} + +// encodeRechargeBonusTiers 序列化为设置值;空列表存空串,与「未配置」保持同一形态。 +func encodeRechargeBonusTiers(tiers []RechargeBonusTier) (string, error) { + if len(tiers) == 0 { + return "", nil + } + raw, err := json.Marshal(tiers) + if err != nil { + return "", fmt.Errorf("marshal recharge bonus tiers: %w", err) + } + return string(raw), nil +} + +// parseRechargeBonusTiers 宽松解析(读路径):非法条目丢弃而非报错,避免历史错配置阻断下单。 +// 同一 MinAmount 重复时保留先出现的档位。始终返回非 nil 切片,便于 JSON 输出为 []。 +func parseRechargeBonusTiers(raw string) []RechargeBonusTier { + out := make([]RechargeBonusTier, 0) + raw = strings.TrimSpace(raw) + if raw == "" { + return out + } + var items []RechargeBonusTier + if err := json.Unmarshal([]byte(raw), &items); err != nil { + slog.Warn("[Payment] parseRechargeBonusTiers: unmarshal failed", "error", err) + return out + } + seen := make(map[string]struct{}, len(items)) + for _, tier := range items { + if !rechargeBonusValueValid(tier.MinAmount, math.MaxFloat64) || !rechargeBonusValueValid(tier.BonusPercent, maxRechargeBonusPercent) { + continue + } + key := decimal.NewFromFloat(tier.MinAmount).Round(2).String() + if _, dup := seen[key]; dup { + continue + } + seen[key] = struct{}{} + out = append(out, tier) + } + sortRechargeBonusTiers(out) + return out +} + +func validateRechargeBonusNotice(notice string) error { + if utf8.RuneCountInString(notice) > maxRechargeBonusNoticeRunes { + return fmt.Errorf("recharge bonus notice exceeds %d characters", maxRechargeBonusNoticeRunes) + } + return nil +} + +// resolveRechargeBonusUpdate 校验并归一化阶梯/模式更新。任一字段缺省时读取现值做交叉校验 +// (折扣模式下所有档位百分比必须 < 100)。返回值仅在对应请求字段非 nil 时有意义。 +func (s *PaymentConfigService) resolveRechargeBonusUpdate(ctx context.Context, req UpdatePaymentConfigRequest) (tiersValue string, modeValue string, err error) { + if req.RechargeBonusTiers == nil && req.RechargeBonusMode == nil { + return "", "", nil + } + stored := map[string]string{} + if (req.RechargeBonusTiers == nil || req.RechargeBonusMode == nil) && s != nil && s.settingRepo != nil { + stored, err = s.settingRepo.GetMultiple(ctx, []string{SettingRechargeBonusTiers, SettingRechargeBonusMode}) + if err != nil { + return "", "", fmt.Errorf("get recharge bonus settings: %w", err) + } + } + + var tiers []RechargeBonusTier + if req.RechargeBonusTiers != nil { + tiers, err = NormalizeRechargeBonusTiers(*req.RechargeBonusTiers) + if err != nil { + return "", "", infraerrors.BadRequest("INVALID_RECHARGE_BONUS_TIERS", err.Error()) + } + } else { + tiers = parseRechargeBonusTiers(stored[SettingRechargeBonusTiers]) + } + + var mode string + if req.RechargeBonusMode != nil { + normalized, ok := NormalizeRechargeBonusMode(*req.RechargeBonusMode) + if !ok { + return "", "", infraerrors.BadRequest("INVALID_RECHARGE_BONUS_MODE", "recharge bonus mode must be bonus or discount") + } + mode = normalized + } else { + mode, _ = NormalizeRechargeBonusMode(stored[SettingRechargeBonusMode]) + } + + if err := ValidateRechargeBonusTiersForMode(mode, tiers); err != nil { + return "", "", infraerrors.BadRequest("INVALID_RECHARGE_BONUS_TIERS", err.Error()) + } + tiersValue, err = encodeRechargeBonusTiers(tiers) + if err != nil { + return "", "", err + } + return tiersValue, mode, nil +} + +// matchRechargeBonusTier 返回不超过 paymentAmount 的最大档位;tiers 需已按 MinAmount 升序。 +func matchRechargeBonusTier(tiers []RechargeBonusTier, paymentAmount float64) (RechargeBonusTier, bool) { + if math.IsNaN(paymentAmount) || math.IsInf(paymentAmount, 0) || paymentAmount <= 0 { + return RechargeBonusTier{}, false + } + var matched RechargeBonusTier + found := false + for _, tier := range tiers { + if paymentAmount+rechargeBonusAmountEpsilon < tier.MinAmount { + break + } + matched = tier + found = true + } + return matched, found +} + +// calculateRechargeBonus 赠送金额 = 到账基数 × 百分比,保留两位小数(四舍五入)。 +func calculateRechargeBonus(baseCredited, bonusPercent float64) float64 { + if baseCredited <= 0 || bonusPercent <= 0 || math.IsNaN(baseCredited) || math.IsNaN(bonusPercent) { + return 0 + } + return decimal.NewFromFloat(baseCredited). + Mul(decimal.NewFromFloat(bonusPercent)). + Div(decimal.NewFromInt(100)). + Round(2). + InexactFloat64() +} + +// addRechargeBonus 到账总额 = 基数 + 赠送,两位小数。 +func addRechargeBonus(baseCredited, bonus float64) float64 { + return decimal.NewFromFloat(baseCredited). + Add(decimal.NewFromFloat(bonus)). + Round(2). + InexactFloat64() +} + +// calculateDiscountedPayBase 折扣模式实付基数 = 支付金额 × (1 − 百分比),按币种精度四舍五入。 +func calculateDiscountedPayBase(paymentAmount, discountPercent float64, currency string) float64 { + digits := int32(payment.CurrencyMaxFractionDigits(currency)) + return decimal.NewFromFloat(paymentAmount). + Mul(decimal.NewFromInt(100).Sub(decimal.NewFromFloat(discountPercent))). + Div(decimal.NewFromInt(100)). + Round(digits). + InexactFloat64() +} + +// rechargeBonusQuote 一笔余额充值的报价结果。 +type rechargeBonusQuote struct { + // PayBase 网关收款基数(支付币种,不含手续费);赠金模式等于支付金额,折扣模式为折后金额。 + PayBase float64 + // Credited 到账总额(USD),含 Bonus。 + Credited float64 + // Bonus 免费额度(USD):赠金模式为额外赠送,折扣模式为未付费却到账的部分。 + Bonus float64 + // Percent 命中档位的百分比;未命中或未产生优惠时为 0。 + Percent float64 +} + +// quoteRechargeBonus 按配置模式报价。currency 用于折扣模式实付基数的精度。 +// 未配置阶梯、未命中、或折扣百分比 ≥ 100(非法历史数据,fail-safe)时按无优惠处理。 +func quoteRechargeBonus(cfg *PaymentConfig, paymentAmount float64, currency string) rechargeBonusQuote { + multiplier := defaultBalanceRechargeMultiplier + var tiers []RechargeBonusTier + mode := RechargeBonusModeBonus + if cfg != nil { + multiplier = cfg.BalanceRechargeMultiplier + tiers = cfg.RechargeBonusTiers + mode, _ = NormalizeRechargeBonusMode(cfg.RechargeBonusMode) + } + base := calculateCreditedBalance(paymentAmount, multiplier) + quote := rechargeBonusQuote{PayBase: paymentAmount, Credited: base} + + tier, ok := matchRechargeBonusTier(tiers, paymentAmount) + if !ok || tier.BonusPercent <= 0 { + return quote + } + switch mode { + case RechargeBonusModeDiscount: + if tier.BonusPercent >= 100 { + return quote + } + payBase := calculateDiscountedPayBase(paymentAmount, tier.BonusPercent, currency) + if payBase <= 0 || payBase >= paymentAmount { + return quote + } + paidCredit := calculateCreditedBalance(payBase, multiplier) + bonus := decimal.NewFromFloat(base).Sub(decimal.NewFromFloat(paidCredit)).Round(2).InexactFloat64() + if bonus < 0 { + bonus = 0 + } + quote.PayBase = payBase + quote.Bonus = bonus + quote.Percent = tier.BonusPercent + default: + bonus := calculateRechargeBonus(base, tier.BonusPercent) + if bonus <= 0 { + return quote + } + quote.Bonus = bonus + quote.Credited = addRechargeBonus(base, bonus) + quote.Percent = tier.BonusPercent + } + return quote +} + +// paymentOrderAmountWithoutBonus 订单到账金额剔除免费额度后的实付部分(USD),用于推广返利基数。 +func paymentOrderAmountWithoutBonus(o *dbent.PaymentOrder) float64 { + if o == nil { + return 0 + } + if o.OrderType != payment.OrderTypeBalance || o.BonusAmount <= 0 { + return o.Amount + } + base := decimal.NewFromFloat(o.Amount). + Sub(decimal.NewFromFloat(o.BonusAmount)). + Round(2). + InexactFloat64() + if base < 0 { + return 0 + } + return base +} diff --git a/backend/internal/service/payment_recharge_bonus_test.go b/backend/internal/service/payment_recharge_bonus_test.go new file mode 100644 index 000000000000..00af53bfe6d9 --- /dev/null +++ b/backend/internal/service/payment_recharge_bonus_test.go @@ -0,0 +1,349 @@ +//go:build unit + +package service + +import ( + "context" + "testing" + + dbent "github.com/Wei-Shaw/sub2api/ent" + "github.com/Wei-Shaw/sub2api/internal/payment" + "github.com/stretchr/testify/require" +) + +func TestNormalizeRechargeBonusTiers(t *testing.T) { + t.Run("sorts ascending by min amount", func(t *testing.T) { + out, err := NormalizeRechargeBonusTiers([]RechargeBonusTier{ + {MinAmount: 1000, BonusPercent: 35}, + {MinAmount: 100, BonusPercent: 20}, + {MinAmount: 500, BonusPercent: 30}, + }) + require.NoError(t, err) + require.Equal(t, []RechargeBonusTier{ + {MinAmount: 100, BonusPercent: 20}, + {MinAmount: 500, BonusPercent: 30}, + {MinAmount: 1000, BonusPercent: 35}, + }, out) + }) + + t.Run("empty input yields empty non-nil slice", func(t *testing.T) { + out, err := NormalizeRechargeBonusTiers(nil) + require.NoError(t, err) + require.NotNil(t, out) + require.Len(t, out, 0) + }) + + t.Run("allows zero threshold and zero percent", func(t *testing.T) { + out, err := NormalizeRechargeBonusTiers([]RechargeBonusTier{{MinAmount: 0, BonusPercent: 0}}) + require.NoError(t, err) + require.Len(t, out, 1) + }) + + t.Run("rejects invalid values", func(t *testing.T) { + cases := map[string][]RechargeBonusTier{ + "negative min": {{MinAmount: -1, BonusPercent: 10}}, + "min three decimals": {{MinAmount: 100.123, BonusPercent: 10}}, + "negative percent": {{MinAmount: 100, BonusPercent: -5}}, + "percent over limit": {{MinAmount: 100, BonusPercent: 1000.01}}, + "percent 3 decimals": {{MinAmount: 100, BonusPercent: 12.345}}, + "duplicate min": {{MinAmount: 100, BonusPercent: 10}, {MinAmount: 100, BonusPercent: 20}}, + "duplicate min 2 dec": {{MinAmount: 100, BonusPercent: 10}, {MinAmount: 100.00, BonusPercent: 20}}, + } + for name, tiers := range cases { + _, err := NormalizeRechargeBonusTiers(tiers) + require.Error(t, err, name) + } + }) + + t.Run("rejects too many tiers", func(t *testing.T) { + tiers := make([]RechargeBonusTier, 0, maxRechargeBonusTiers+1) + for i := 0; i <= maxRechargeBonusTiers; i++ { + tiers = append(tiers, RechargeBonusTier{MinAmount: float64(i + 1), BonusPercent: 1}) + } + _, err := NormalizeRechargeBonusTiers(tiers) + require.Error(t, err) + }) +} + +func TestParseRechargeBonusTiers(t *testing.T) { + t.Run("empty or invalid json yields empty slice", func(t *testing.T) { + require.NotNil(t, parseRechargeBonusTiers("")) + require.Len(t, parseRechargeBonusTiers(""), 0) + require.Len(t, parseRechargeBonusTiers("not json"), 0) + }) + + t.Run("drops invalid entries keeps first duplicate and sorts", func(t *testing.T) { + raw := `[{"min_amount":500,"bonus_percent":30},{"min_amount":-1,"bonus_percent":5},` + + `{"min_amount":100,"bonus_percent":20},{"min_amount":100,"bonus_percent":99},` + + `{"min_amount":50,"bonus_percent":5000}]` + out := parseRechargeBonusTiers(raw) + require.Equal(t, []RechargeBonusTier{ + {MinAmount: 100, BonusPercent: 20}, + {MinAmount: 500, BonusPercent: 30}, + }, out) + }) + + t.Run("round trips encode", func(t *testing.T) { + tiers := []RechargeBonusTier{{MinAmount: 100, BonusPercent: 20}, {MinAmount: 500, BonusPercent: 30}} + encoded, err := encodeRechargeBonusTiers(tiers) + require.NoError(t, err) + require.Equal(t, tiers, parseRechargeBonusTiers(encoded)) + + empty, err := encodeRechargeBonusTiers(nil) + require.NoError(t, err) + require.Equal(t, "", empty) + }) +} + +func TestMatchRechargeBonusTier(t *testing.T) { + tiers := []RechargeBonusTier{ + {MinAmount: 100, BonusPercent: 20}, + {MinAmount: 500, BonusPercent: 30}, + {MinAmount: 1000, BonusPercent: 35}, + } + cases := []struct { + amount float64 + percent float64 + ok bool + }{ + {amount: 0, ok: false}, + {amount: 99.99, ok: false}, + {amount: 100, percent: 20, ok: true}, + {amount: 499.99, percent: 20, ok: true}, + {amount: 500, percent: 30, ok: true}, + {amount: 1000, percent: 35, ok: true}, + {amount: 1000000, percent: 35, ok: true}, + } + for _, tc := range cases { + tier, ok := matchRechargeBonusTier(tiers, tc.amount) + require.Equal(t, tc.ok, ok, "amount %v", tc.amount) + if ok { + require.Equal(t, tc.percent, tier.BonusPercent, "amount %v", tc.amount) + } + } + + t.Run("float boundary 0.1+0.2 still matches 0.3 threshold", func(t *testing.T) { + _, ok := matchRechargeBonusTier([]RechargeBonusTier{{MinAmount: 0.3, BonusPercent: 1}}, 0.1+0.2) + require.True(t, ok) + }) + + t.Run("no tiers never matches", func(t *testing.T) { + _, ok := matchRechargeBonusTier(nil, 100) + require.False(t, ok) + }) +} + +func TestCalculateRechargeBonusRounding(t *testing.T) { + require.Equal(t, 20.0, calculateRechargeBonus(100, 20)) + require.Equal(t, 120.0, addRechargeBonus(100, 20)) + // 33.33 * 15% = 4.9995 → 5.00 + require.Equal(t, 5.0, calculateRechargeBonus(33.33, 15)) + require.Zero(t, calculateRechargeBonus(100, 0)) + require.Zero(t, calculateRechargeBonus(0, 20)) + + // 阈值按支付金额命中,赠送按到账基数计算:1000 CNY × 0.14 = 140 USD,命中 1000 档 30% → 42 + cfg := &PaymentConfig{ + BalanceRechargeMultiplier: 0.14, + RechargeBonusTiers: []RechargeBonusTier{{MinAmount: 100, BonusPercent: 20}, {MinAmount: 500, BonusPercent: 30}}, + } + require.Equal(t, rechargeBonusQuote{PayBase: 1000, Credited: 182, Bonus: 42, Percent: 30}, quoteRechargeBonus(cfg, 1000, "CNY")) + // 命中 0% 档位视为无优惠 + cfg.RechargeBonusTiers = []RechargeBonusTier{{MinAmount: 10, BonusPercent: 0}} + require.Equal(t, rechargeBonusQuote{PayBase: 50, Credited: 7}, quoteRechargeBonus(cfg, 50, "CNY")) +} + +func TestParsePaymentConfigRechargeBonus(t *testing.T) { + svc := &PaymentConfigService{} + + t.Run("defaults", func(t *testing.T) { + cfg := svc.parsePaymentConfig(map[string]string{}) + require.NotNil(t, cfg.RechargeBonusTiers) + require.Len(t, cfg.RechargeBonusTiers, 0) + require.Equal(t, "", cfg.RechargeBonusNotice) + }) + + t.Run("reads tiers and notice", func(t *testing.T) { + cfg := svc.parsePaymentConfig(map[string]string{ + SettingRechargeBonusTiers: `[{"min_amount":500,"bonus_percent":30},{"min_amount":100,"bonus_percent":20}]`, + SettingRechargeBonusNotice: "**满 100 送 20%**", + }) + require.Equal(t, []RechargeBonusTier{{MinAmount: 100, BonusPercent: 20}, {MinAmount: 500, BonusPercent: 30}}, cfg.RechargeBonusTiers) + require.Equal(t, "**满 100 送 20%**", cfg.RechargeBonusNotice) + }) +} + +func TestUpdatePaymentConfigRechargeBonus(t *testing.T) { + ctx := context.Background() + + t.Run("persists normalized tiers and trimmed notice", func(t *testing.T) { + repo := &paymentConfigSettingRepoStub{values: map[string]string{}} + svc := &PaymentConfigService{settingRepo: repo} + tiers := []RechargeBonusTier{{MinAmount: 500, BonusPercent: 30}, {MinAmount: 100, BonusPercent: 20}} + notice := " 活动文案 " + require.NoError(t, svc.UpdatePaymentConfig(ctx, UpdatePaymentConfigRequest{ + RechargeBonusTiers: &tiers, + RechargeBonusNotice: ¬ice, + })) + require.Equal(t, `[{"min_amount":100,"bonus_percent":20},{"min_amount":500,"bonus_percent":30}]`, repo.updates[SettingRechargeBonusTiers]) + require.Equal(t, "活动文案", repo.updates[SettingRechargeBonusNotice]) + + cfg, err := svc.GetPaymentConfig(ctx) + require.NoError(t, err) + require.Equal(t, []RechargeBonusTier{{MinAmount: 100, BonusPercent: 20}, {MinAmount: 500, BonusPercent: 30}}, cfg.RechargeBonusTiers) + }) + + t.Run("empty tiers clears setting and omitted fields are untouched", func(t *testing.T) { + repo := &paymentConfigSettingRepoStub{values: map[string]string{ + SettingRechargeBonusTiers: `[{"min_amount":100,"bonus_percent":20}]`, + SettingRechargeBonusNotice: "keep me", + }} + svc := &PaymentConfigService{settingRepo: repo} + empty := []RechargeBonusTier{} + require.NoError(t, svc.UpdatePaymentConfig(ctx, UpdatePaymentConfigRequest{RechargeBonusTiers: &empty})) + value, ok := repo.updates[SettingRechargeBonusTiers] + require.True(t, ok) + require.Equal(t, "", value) + _, touched := repo.updates[SettingRechargeBonusNotice] + require.False(t, touched) + require.Equal(t, "keep me", repo.values[SettingRechargeBonusNotice]) + }) + + t.Run("rejects invalid tiers", func(t *testing.T) { + repo := &paymentConfigSettingRepoStub{values: map[string]string{}} + svc := &PaymentConfigService{settingRepo: repo} + bad := []RechargeBonusTier{{MinAmount: 100, BonusPercent: 20}, {MinAmount: 100, BonusPercent: 30}} + err := svc.UpdatePaymentConfig(ctx, UpdatePaymentConfigRequest{RechargeBonusTiers: &bad}) + require.Error(t, err) + require.Nil(t, repo.updates) + }) +} + +func TestAffiliateRebateBaseAmountExcludesRechargeBonus(t *testing.T) { + require.Equal(t, 100.0, affiliateRebateBaseAmount(&dbent.PaymentOrder{ + OrderType: payment.OrderTypeBalance, Amount: 130, BonusAmount: 30, + })) + require.Equal(t, 130.0, affiliateRebateBaseAmount(&dbent.PaymentOrder{ + OrderType: payment.OrderTypeBalance, Amount: 130, + })) + // 订阅订单不受 bonus 字段影响 + require.Equal(t, 50.0, affiliateRebateBaseAmount(&dbent.PaymentOrder{ + OrderType: payment.OrderTypeSubscription, Amount: 50, BonusAmount: 30, + })) + // 异常数据:赠送大于总额时钳到 0 + require.Equal(t, 0.0, affiliateRebateBaseAmount(&dbent.PaymentOrder{ + OrderType: payment.OrderTypeBalance, Amount: 10, BonusAmount: 30, + })) +} + +func TestNormalizeRechargeBonusMode(t *testing.T) { + for raw, want := range map[string]string{"": RechargeBonusModeBonus, "bonus": RechargeBonusModeBonus, " Discount ": RechargeBonusModeDiscount} { + mode, ok := NormalizeRechargeBonusMode(raw) + require.True(t, ok, raw) + require.Equal(t, want, mode, raw) + } + mode, ok := NormalizeRechargeBonusMode("cashback") + require.False(t, ok) + require.Equal(t, RechargeBonusModeBonus, mode) +} + +func TestValidateRechargeBonusTiersForMode(t *testing.T) { + tiers := []RechargeBonusTier{{MinAmount: 100, BonusPercent: 20}, {MinAmount: 500, BonusPercent: 100}} + require.NoError(t, ValidateRechargeBonusTiersForMode(RechargeBonusModeBonus, tiers)) + require.Error(t, ValidateRechargeBonusTiersForMode(RechargeBonusModeDiscount, tiers)) + require.NoError(t, ValidateRechargeBonusTiersForMode(RechargeBonusModeDiscount, tiers[:1])) +} + +func TestQuoteRechargeBonus(t *testing.T) { + tiers := []RechargeBonusTier{{MinAmount: 100, BonusPercent: 20}, {MinAmount: 500, BonusPercent: 50}} + + t.Run("bonus mode keeps pay base and inflates credit", func(t *testing.T) { + cfg := &PaymentConfig{BalanceRechargeMultiplier: 1, RechargeBonusTiers: tiers, RechargeBonusMode: RechargeBonusModeBonus} + q := quoteRechargeBonus(cfg, 100, "USD") + require.Equal(t, rechargeBonusQuote{PayBase: 100, Credited: 120, Bonus: 20, Percent: 20}, q) + + // 未命中:无优惠 + q = quoteRechargeBonus(cfg, 50, "USD") + require.Equal(t, rechargeBonusQuote{PayBase: 50, Credited: 50}, q) + }) + + t.Run("discount mode keeps credit and reduces pay base", func(t *testing.T) { + cfg := &PaymentConfig{BalanceRechargeMultiplier: 1, RechargeBonusTiers: tiers, RechargeBonusMode: RechargeBonusModeDiscount} + q := quoteRechargeBonus(cfg, 500, "USD") + require.Equal(t, rechargeBonusQuote{PayBase: 250, Credited: 500, Bonus: 250, Percent: 50}, q) + + // 倍率 0.14:1000 CNY 到账 140 USD;20% off 实付 800 CNY,免费部分 = 140 − 112 = 28 USD + cfg.BalanceRechargeMultiplier = 0.14 + q = quoteRechargeBonus(cfg, 1000, "CNY") + require.Equal(t, rechargeBonusQuote{PayBase: 500, Credited: 140, Bonus: 70, Percent: 50}, q) + q = quoteRechargeBonus(cfg, 200, "CNY") + require.Equal(t, rechargeBonusQuote{PayBase: 160, Credited: 28, Bonus: 5.6, Percent: 20}, q) + }) + + t.Run("discount rounds pay base to currency precision", func(t *testing.T) { + cfg := &PaymentConfig{BalanceRechargeMultiplier: 1, RechargeBonusTiers: []RechargeBonusTier{{MinAmount: 1, BonusPercent: 15}}, RechargeBonusMode: RechargeBonusModeDiscount} + require.Equal(t, 85.85, quoteRechargeBonus(cfg, 101, "USD").PayBase) + require.Equal(t, 86.0, quoteRechargeBonus(cfg, 101, "JPY").PayBase) + }) + + t.Run("discount percent at or above 100 is ignored fail-safe", func(t *testing.T) { + cfg := &PaymentConfig{BalanceRechargeMultiplier: 1, RechargeBonusTiers: []RechargeBonusTier{{MinAmount: 1, BonusPercent: 100}}, RechargeBonusMode: RechargeBonusModeDiscount} + require.Equal(t, rechargeBonusQuote{PayBase: 100, Credited: 100}, quoteRechargeBonus(cfg, 100, "USD")) + }) + + t.Run("nil config and empty tiers yield plain conversion", func(t *testing.T) { + require.Equal(t, rechargeBonusQuote{PayBase: 100, Credited: 100}, quoteRechargeBonus(nil, 100, "USD")) + require.Equal(t, rechargeBonusQuote{PayBase: 100, Credited: 14}, quoteRechargeBonus(&PaymentConfig{BalanceRechargeMultiplier: 0.14}, 100, "CNY")) + }) +} + +func TestParsePaymentConfigRechargeBonusMode(t *testing.T) { + svc := &PaymentConfigService{} + require.Equal(t, RechargeBonusModeBonus, svc.parsePaymentConfig(map[string]string{}).RechargeBonusMode) + require.Equal(t, RechargeBonusModeDiscount, svc.parsePaymentConfig(map[string]string{SettingRechargeBonusMode: "discount"}).RechargeBonusMode) + require.Equal(t, RechargeBonusModeBonus, svc.parsePaymentConfig(map[string]string{SettingRechargeBonusMode: "junk"}).RechargeBonusMode) +} + +func TestUpdatePaymentConfigRechargeBonusMode(t *testing.T) { + ctx := context.Background() + + t.Run("persists discount mode with valid tiers", func(t *testing.T) { + repo := &paymentConfigSettingRepoStub{values: map[string]string{}} + svc := &PaymentConfigService{settingRepo: repo} + mode := "discount" + tiers := []RechargeBonusTier{{MinAmount: 100, BonusPercent: 20}} + require.NoError(t, svc.UpdatePaymentConfig(ctx, UpdatePaymentConfigRequest{RechargeBonusTiers: &tiers, RechargeBonusMode: &mode})) + require.Equal(t, "discount", repo.updates[SettingRechargeBonusMode]) + cfg, err := svc.GetPaymentConfig(ctx) + require.NoError(t, err) + require.Equal(t, RechargeBonusModeDiscount, cfg.RechargeBonusMode) + }) + + t.Run("rejects unknown mode", func(t *testing.T) { + repo := &paymentConfigSettingRepoStub{values: map[string]string{}} + svc := &PaymentConfigService{settingRepo: repo} + mode := "cashback" + require.Error(t, svc.UpdatePaymentConfig(ctx, UpdatePaymentConfigRequest{RechargeBonusMode: &mode})) + require.Nil(t, repo.updates) + }) + + t.Run("switching to discount with stored tiers at 100 percent is rejected", func(t *testing.T) { + repo := &paymentConfigSettingRepoStub{values: map[string]string{ + SettingRechargeBonusTiers: `[{"min_amount":100,"bonus_percent":100}]`, + }} + svc := &PaymentConfigService{settingRepo: repo} + mode := "discount" + require.Error(t, svc.UpdatePaymentConfig(ctx, UpdatePaymentConfigRequest{RechargeBonusMode: &mode})) + require.Nil(t, repo.updates) + }) + + t.Run("saving tiers at 100 percent while stored mode is discount is rejected", func(t *testing.T) { + repo := &paymentConfigSettingRepoStub{values: map[string]string{SettingRechargeBonusMode: "discount"}} + svc := &PaymentConfigService{settingRepo: repo} + tiers := []RechargeBonusTier{{MinAmount: 100, BonusPercent: 100}} + require.Error(t, svc.UpdatePaymentConfig(ctx, UpdatePaymentConfigRequest{RechargeBonusTiers: &tiers})) + require.Nil(t, repo.updates) + // 同样的档位在赠金模式下合法 + repo.values[SettingRechargeBonusMode] = "bonus" + require.NoError(t, svc.UpdatePaymentConfig(ctx, UpdatePaymentConfigRequest{RechargeBonusTiers: &tiers})) + }) +} diff --git a/backend/internal/service/payment_service.go b/backend/internal/service/payment_service.go index 792a842a2348..dae1a7d9384b 100644 --- a/backend/internal/service/payment_service.go +++ b/backend/internal/service/payment_service.go @@ -92,6 +92,7 @@ type CreateOrderResponse struct { Amount float64 `json:"amount"` PayAmount float64 `json:"pay_amount"` FeeRate float64 `json:"fee_rate"` + BonusAmount float64 `json:"bonus_amount"` Status string `json:"status"` ResultType payment.CreatePaymentResultType `json:"result_type,omitempty"` PaymentType string `json:"payment_type"` diff --git a/backend/internal/service/scheduler_snapshot_service.go b/backend/internal/service/scheduler_snapshot_service.go index 14d43f0c8743..6f999c189c5b 100644 --- a/backend/internal/service/scheduler_snapshot_service.go +++ b/backend/internal/service/scheduler_snapshot_service.go @@ -609,7 +609,7 @@ func (s *SchedulerSnapshotService) handleBulkAccountEvent(ctx context.Context, p } accountGroupIDs := s.normalizeGroupIDs(account.GroupIDs) switch account.Platform { - case PlatformAnthropic, PlatformGemini, PlatformOpenAI, PlatformGrok, PlatformKimi, PlatformZhipu, PlatformDeepseek, PlatformMiniMax, PlatformOpenCodeGo: + case PlatformAnthropic, PlatformGemini, PlatformOpenAI, PlatformGrok, PlatformKimi, PlatformZhipu, PlatformDeepseek, PlatformMiniMax, PlatformOpenCodeGo, PlatformTypeSafe: addPlatformGroups(account.Platform, accountGroupIDs) case PlatformAntigravity: // 批量更新可能刚关闭 mixed_scheduling,仍需清理两个兼容平台的旧快照。 @@ -824,8 +824,8 @@ func (s *SchedulerSnapshotService) rebuildByAccount(ctx context.Context, account return s.rebuildBuckets(ctx, buckets, reason) } -func schedulerSnapshotPlatforms() [10]string { - return [10]string{PlatformAnthropic, PlatformGemini, PlatformOpenAI, PlatformAntigravity, PlatformGrok, PlatformKimi, PlatformZhipu, PlatformDeepseek, PlatformMiniMax, PlatformOpenCodeGo} +func schedulerSnapshotPlatforms() [11]string { + return [11]string{PlatformAnthropic, PlatformGemini, PlatformOpenAI, PlatformAntigravity, PlatformGrok, PlatformKimi, PlatformZhipu, PlatformDeepseek, PlatformMiniMax, PlatformOpenCodeGo, PlatformTypeSafe} } // 生命周期辅助函数有意排除 group0;full rebuild 构造 group0 canonical 集时必须显式调用 canonical helper。 diff --git a/backend/internal/service/upstream_billing_probe.go b/backend/internal/service/upstream_billing_probe.go index fc7cdfd5f53d..ef240a637fc1 100644 --- a/backend/internal/service/upstream_billing_probe.go +++ b/backend/internal/service/upstream_billing_probe.go @@ -1112,7 +1112,8 @@ func IsUpstreamBillingProbeIdentity(platform, accountType string) bool { } switch platform { case PlatformOpenAI, PlatformAnthropic, PlatformGemini, PlatformAntigravity, PlatformGrok, - PlatformKimi, PlatformZhipu, PlatformDeepseek, PlatformMiniMax, PlatformOpenCodeGo: + PlatformKimi, PlatformZhipu, PlatformDeepseek, PlatformMiniMax, PlatformOpenCodeGo, + PlatformTypeSafe: return true default: return false @@ -1151,6 +1152,7 @@ var upstreamBillingProbeOfficialAPIDomains = []string{ "bigmodel.cn", "deepseek.com", "opencode.ai", + "typesafe.ai", } func upstreamBillingProbeTargetIsOfficialAPI(baseURL string) bool { diff --git a/backend/internal/service/upstream_billing_probe_multiplatform_test.go b/backend/internal/service/upstream_billing_probe_multiplatform_test.go index 3153845f47fb..f17f4fe3f602 100644 --- a/backend/internal/service/upstream_billing_probe_multiplatform_test.go +++ b/backend/internal/service/upstream_billing_probe_multiplatform_test.go @@ -17,6 +17,7 @@ func TestUpstreamBillingProbeIdentityCoversAllAPIKeyPlatforms(t *testing.T) { for _, platform := range []string{ PlatformOpenAI, PlatformGrok, PlatformAnthropic, PlatformGemini, PlatformAntigravity, PlatformKimi, PlatformZhipu, PlatformDeepseek, PlatformMiniMax, PlatformOpenCodeGo, + PlatformTypeSafe, } { require.True(t, IsUpstreamBillingProbeIdentity(platform, AccountTypeAPIKey), platform) require.True(t, isUpstreamBillingProbeAccount(&Account{Platform: platform, Type: AccountTypeAPIKey}), platform) @@ -178,6 +179,8 @@ func TestUpstreamBillingProbeOfficialAPIBaseURLIsUnsupportedWithoutRequest(t *te {PlatformDeepseek, "https://api.deepseek.com/anthropic"}, {PlatformOpenCodeGo, "https://opencode.ai/zen/go/v1"}, {PlatformOpenCodeGo, "https://opencode.ai/zen/go"}, + {PlatformTypeSafe, "https://api.typesafe.ai"}, + {PlatformTypeSafe, "https://api.typesafe.ai/v1"}, } for i, tc := range cases { account := &Account{ @@ -221,6 +224,7 @@ func TestUpstreamBillingProbeOfficialAPIHostMatchingIsNormalized(t *testing.T) { require.True(t, upstreamBillingProbeTargetIsOfficialAPI("https://api.deepseek.com/anthropic")) require.True(t, upstreamBillingProbeTargetIsOfficialAPI("https://opencode.ai/zen/go/v1")) require.True(t, upstreamBillingProbeTargetIsOfficialAPI("https://opencode.ai/zen/go")) + require.True(t, upstreamBillingProbeTargetIsOfficialAPI("https://api.typesafe.ai")) // 相似但不同的注册域不拦:中转完全可能叫 *-x.ai 之外的任何名字。 require.False(t, upstreamBillingProbeTargetIsOfficialAPI("https://relay.example/v1")) require.False(t, upstreamBillingProbeTargetIsOfficialAPI("https://notx.ai")) @@ -232,6 +236,7 @@ func TestUpstreamBillingProbeOfficialAPIHostMatchingIsNormalized(t *testing.T) { require.False(t, upstreamBillingProbeTargetIsOfficialAPI("https://notmoonshot.cn/v1")) require.False(t, upstreamBillingProbeTargetIsOfficialAPI("https://moonshot.cn.evil.example/v1")) require.False(t, upstreamBillingProbeTargetIsOfficialAPI("https://kimi.example/v1")) + require.False(t, upstreamBillingProbeTargetIsOfficialAPI("https://nottypesafe.ai")) require.False(t, upstreamBillingProbeTargetIsOfficialAPI("https://notbigmodel.cn")) require.False(t, upstreamBillingProbeTargetIsOfficialAPI("https://deepseek.example.com")) } @@ -280,3 +285,16 @@ func TestUpstreamBillingProbeSetAccountEnabledAcceptsGrokAPIKey(t *testing.T) { err := svc.SetAccountEnabled(context.Background(), grokOAuth.ID, true) require.ErrorIs(t, err, ErrUpstreamBillingProbeAccountInvalid) } + +func TestBuildAccountForCreateAcceptsTypeSafeAPIKeyWithProbeEnabled(t *testing.T) { + enabled := true + account, err := buildAccountForCreate(&CreateAccountInput{ + Name: "typesafe", + Platform: PlatformTypeSafe, + Type: AccountTypeAPIKey, + Credentials: map[string]any{"api_key": "ts-key", "base_url": "https://api.typesafe.ai"}, + ProbeEnabled: &enabled, + }, map[string]any{}) + require.NoError(t, err) + require.Equal(t, true, account.Extra[UpstreamBillingProbeEnabledExtraKey]) +} diff --git a/backend/internal/service/user_service.go b/backend/internal/service/user_service.go index b4eca93b373e..137cc122fc40 100644 --- a/backend/internal/service/user_service.go +++ b/backend/internal/service/user_service.go @@ -4,7 +4,6 @@ import ( "bytes" "context" "crypto/sha256" - "crypto/subtle" "encoding/base64" "encoding/hex" "fmt" @@ -1321,28 +1320,7 @@ func (s *UserService) VerifyAndAddNotifyEmail(ctx context.Context, userID int64, // verifyNotifyCode validates the verification code against the cached data. func verifyNotifyCode(ctx context.Context, cache EmailCache, email, code string) error { - data, err := cache.GetNotifyVerifyCode(ctx, email) - if err != nil || data == nil { - return ErrInvalidVerifyCode - } - if data.Attempts >= maxVerifyCodeAttempts { - return ErrVerifyCodeMaxAttempts - } - if subtle.ConstantTimeCompare([]byte(data.Code), []byte(code)) != 1 { - data.Attempts++ - remaining := time.Until(data.ExpiresAt) - if remaining <= 0 { - return ErrInvalidVerifyCode - } - if err := cache.SetNotifyVerifyCode(ctx, email, data, remaining); err != nil { - slog.Error("failed to update notify verify code attempts", "email", email, "error", err) - } - if data.Attempts >= maxVerifyCodeAttempts { - return ErrVerifyCodeMaxAttempts - } - return ErrInvalidVerifyCode - } - return nil + return verifyCodeWithAttempts(ctx, email, code, cache.GetNotifyVerifyCode, cache.IncrNotifyVerifyCodeAttempts, nil) } // addOrVerifyNotifyEmail adds the email to user's extra notification emails or marks it as verified. diff --git a/backend/migrations/241_add_payment_order_bonus_amount.sql b/backend/migrations/241_add_payment_order_bonus_amount.sql new file mode 100644 index 000000000000..1eeddfadbac4 --- /dev/null +++ b/backend/migrations/241_add_payment_order_bonus_amount.sql @@ -0,0 +1,3 @@ +-- 充值赠送额度:余额充值订单命中赠送阶梯时的赠送 USD 金额。 +-- 已计入 payment_orders.amount(到账总额),单独落列用于订单展示与推广返利基数剔除。 +ALTER TABLE payment_orders ADD COLUMN IF NOT EXISTS bonus_amount DECIMAL(20,2) NOT NULL DEFAULT 0; diff --git a/backend/migrations/241_add_typesafe_platform.sql b/backend/migrations/241_add_typesafe_platform.sql new file mode 100644 index 000000000000..82a81de70d38 --- /dev/null +++ b/backend/migrations/241_add_typesafe_platform.sql @@ -0,0 +1,26 @@ +-- Add TypeSafe (Jev System One) as a first-class platform. +-- +-- 1. user_platform_quotas.platform CHECK +-- 2. composite_model_routes.target_platform CHECK +-- +-- TypeSafe 不是对话模型,不进入渠道监控 provider,因此 channel_monitors / +-- channel_monitor_request_templates 的约束保持不变。 +-- +-- Runs after 238_opencode_go_platform.sql. DROP ... IF EXISTS 保证可重入; +-- 新约束是 238 的超集,存量行瞬时校验通过。 + +ALTER TABLE user_platform_quotas + DROP CONSTRAINT IF EXISTS user_platform_quotas_platform_check; + +ALTER TABLE user_platform_quotas + ADD CONSTRAINT user_platform_quotas_platform_check + CHECK (platform IN ('anthropic', 'openai', 'gemini', 'antigravity', 'grok', + 'kimi', 'zhipu', 'deepseek', 'minimax', 'opencode_go', 'typesafe')); + +ALTER TABLE composite_model_routes + DROP CONSTRAINT IF EXISTS composite_model_routes_target_platform_check; + +ALTER TABLE composite_model_routes + ADD CONSTRAINT composite_model_routes_target_platform_check + CHECK (target_platform IN ('anthropic', 'openai', 'gemini', 'antigravity', 'grok', + 'kimi', 'zhipu', 'deepseek', 'minimax', 'opencode_go', 'typesafe')); diff --git a/backend/migrations/typesafe_platform_migration_test.go b/backend/migrations/typesafe_platform_migration_test.go new file mode 100644 index 000000000000..7eb2380918cc --- /dev/null +++ b/backend/migrations/typesafe_platform_migration_test.go @@ -0,0 +1,21 @@ +package migrations + +import ( + "strings" + "testing" + + "github.com/stretchr/testify/require" +) + +func TestTypeSafePlatformMigration(t *testing.T) { + content, err := FS.ReadFile("241_add_typesafe_platform.sql") + require.NoError(t, err) + + sql := strings.Join(strings.Fields(string(content)), " ") + require.Contains(t, sql, "DROP CONSTRAINT IF EXISTS user_platform_quotas_platform_check") + require.Contains(t, sql, "DROP CONSTRAINT IF EXISTS composite_model_routes_target_platform_check") + require.Contains(t, sql, + "CHECK (platform IN ('anthropic', 'openai', 'gemini', 'antigravity', 'grok', 'kimi', 'zhipu', 'deepseek', 'minimax', 'opencode_go', 'typesafe'))") + require.Contains(t, sql, + "CHECK (target_platform IN ('anthropic', 'openai', 'gemini', 'antigravity', 'grok', 'kimi', 'zhipu', 'deepseek', 'minimax', 'opencode_go', 'typesafe'))") +} diff --git a/docs/typesafe-jev.md b/docs/typesafe-jev.md new file mode 100644 index 000000000000..17cd4b5ab181 --- /dev/null +++ b/docs/typesafe-jev.md @@ -0,0 +1,49 @@ +# TypeSafe / Jev System One + +> 整理自上游 [Wei-Shaw/sub2api v0.2.13](https://github.com/Wei-Shaw/sub2api/tree/v0.2.13) README 中的 TypeSafe / Jev 小节(上游 PR [#7425](https://github.com/Wei-Shaw/sub2api/pull/7425))。本 fork 的 README 结构不同,因此单独成文。 + +## 使用说明 + +Sub2API 支持使用 TypeSafe API Key 账户,通过 Jev 原生、非流式的 System One 协议调用模型。 + +- 平台:`typesafe`;账号类型:API Key +- 默认上游:`https://api.typesafe.ai` +- 对外端点:`POST /v1/systemone` +- 模型:`jev-latest`,TypeSafe 分组的 `/v1/models` 也会返回该模型 +- 问题类型:`noul`、`choice`、`score` + +请求和成功响应保持 System One 原生 JSON 结构。该端点不兼容 Chat Completions、Responses、Anthropic Messages 或流式客户端。 + +问题校验遵循 TypeSafe OpenAPI 的线上协议 schema(SDK v0.5.7 也使用该 schema)。所有问题的 `instructions` 都可以省略或为 `null`。Noul 的 `criteria` 可以省略或为 `null`,其中 `true`/`false` 的描述和 Choice 描述支持字符串、对象、数组或 `null`。Score 的 `criteria` 必须是至少包含一档描述的数组,每档支持字符串、对象或数组;单档也合法。SDK 的整数键 Score 映射会由 SDK 在发送前转换为数组。 + +```bash +curl https://your-sub2api.example.com/v1/systemone \ + -H 'Authorization: Bearer sk-your-sub2api-key' \ + -H 'Content-Type: application/json' \ + --data '{"model":"jev-latest","state":"待评估文本","questions":{"safety":{"type":"noul","instructions":"评估文本是否不安全"}}}' +``` + +`jev-latest` 内置价格为输入 `$0.042/百万 tokens`、输出 `$0`,渠道定价可以覆盖。凭据、欠费、权限、限流、过载、服务端和网络错误(`401`、`402`、`403`、`429`、`529`、`5xx`、传输错误)沿用现有账号错误策略(含自定义错误码与临时不可调度规则)并切换账号;请求错误(`400`、`413`、`422`)不会切换账号重试,也不会改变账号状态。TypeSafe 分组(以及路由到 TypeSafe 的 Composite 请求)调用 Messages、Chat Completions、Responses、count_tokens 时返回 `404`。 + +## English + +Sub2API supports TypeSafe API-key accounts through Jev's native, non-streaming System One protocol. + +- Platform: `typesafe`; account type: API Key +- Default upstream: `https://api.typesafe.ai` +- Public endpoint: `POST /v1/systemone` +- Model: `jev-latest`, also returned by `/v1/models` for TypeSafe groups +- Questions: `noul`, `choice`, and `score` + +Requests and successful responses retain the native System One JSON structure. This endpoint is not compatible with Chat Completions, Responses, Anthropic Messages, or streaming clients. + +Question validation follows the TypeSafe OpenAPI wire schema (also used by SDK v0.5.7). `instructions` may be omitted or `null` for all question types. Noul `criteria` may be omitted or `null`; its `true`/`false` descriptions and Choice descriptions accept strings, objects, arrays, or `null`. Score `criteria` must be a non-empty array of string, object, or array descriptions; a single level is valid. SDK integer-keyed Score maps are normalized to arrays by the SDK before sending. + +```bash +curl https://your-sub2api.example.com/v1/systemone \ + -H 'Authorization: Bearer sk-your-sub2api-key' \ + -H 'Content-Type: application/json' \ + --data '{"model":"jev-latest","state":"Text to evaluate","questions":{"safety":{"type":"noul","instructions":"Evaluate whether the text is unsafe"}}}' +``` + +The built-in `jev-latest` price is `$0.042` per million input tokens and `$0` for output tokens. Channel pricing can override both values. Credential, billing, permission, rate-limit, overload, server, and network failures (`401`, `402`, `403`, `429`, `529`, `5xx`, transport errors) use the existing account error policy (including custom error codes and temporary-unschedulable rules) and fail over to another account; request errors (`400`, `413`, and `422`) are returned without retrying another account and never change account state. TypeSafe groups (and Composite requests routed to TypeSafe) reject Messages, Chat Completions, Responses, and count_tokens requests with `404`. diff --git a/frontend/src/api/__tests__/settings.authSourceDefaults.spec.ts b/frontend/src/api/__tests__/settings.authSourceDefaults.spec.ts index deeb758250e2..3e2b6dda4a2e 100644 --- a/frontend/src/api/__tests__/settings.authSourceDefaults.spec.ts +++ b/frontend/src/api/__tests__/settings.authSourceDefaults.spec.ts @@ -9,13 +9,14 @@ import { type DefaultPlatformQuotasMap, } from "@/api/admin/settings"; -/** 全 null 的 5 平台 map,用于断言归一化默认值 */ +/** 全 null 的 6 平台 map,用于断言归一化默认值 */ const allNullQuotas: DefaultPlatformQuotasMap = { anthropic: { daily: null, weekly: null, monthly: null }, openai: { daily: null, weekly: null, monthly: null }, gemini: { daily: null, weekly: null, monthly: null }, antigravity: { daily: null, weekly: null, monthly: null }, grok: { daily: null, weekly: null, monthly: null }, + typesafe: { daily: null, weekly: null, monthly: null }, } describe("admin settings auth source defaults helpers", () => { @@ -240,9 +241,9 @@ describe("normalizePlatformQuotasMap", () => { expect(result.grok).toEqual({ daily: null, weekly: null, monthly: null }); }); - it("无参数时返回全 5 平台全 null", () => { + it("无参数时返回全 6 平台全 null", () => { const result = normalizePlatformQuotasMap(); - expect(Object.keys(result)).toHaveLength(5); + expect(Object.keys(result)).toHaveLength(6); for (const v of Object.values(result)) { expect(v).toEqual({ daily: null, weekly: null, monthly: null }); } @@ -290,7 +291,7 @@ describe("sanitizePlatformQuotasMap", () => { it("缺失平台填充为全 null", () => { const result = sanitizePlatformQuotasMap({}); - expect(Object.keys(result)).toHaveLength(5); + expect(Object.keys(result)).toHaveLength(6); for (const v of Object.values(result)) { expect(v).toEqual({ daily: null, weekly: null, monthly: null }); } diff --git a/frontend/src/api/admin/settings.ts b/frontend/src/api/admin/settings.ts index e0c381bd335e..6d9d2bce6d44 100644 --- a/frontend/src/api/admin/settings.ts +++ b/frontend/src/api/admin/settings.ts @@ -10,6 +10,7 @@ import type { LoginAgreementDocument, NotifyEmailEntry, } from "@/types"; +import type { RechargeBonusTier } from "@/utils/rechargeBonus"; export interface DefaultSubscriptionSetting { group_id: number; @@ -17,7 +18,7 @@ export interface DefaultSubscriptionSetting { } // ── 平台限额类型 ────────────────────────────────────────────────── -export type PlatformType = "anthropic" | "openai" | "gemini" | "antigravity" | "grok" +export type PlatformType = "anthropic" | "openai" | "gemini" | "antigravity" | "grok" | "typesafe" export type QuotaWindowType = "daily" | "weekly" | "monthly" /** 单平台三档限额;null = 不限制,undefined = 未填(等价 null) */ @@ -30,7 +31,7 @@ export interface PlatformQuotaLimits { /** 全平台默认限额 map(key = PlatformType) */ export type DefaultPlatformQuotasMap = Partial> -const PLATFORMS: PlatformType[] = ["anthropic", "openai", "gemini", "antigravity", "grok"] +const PLATFORMS: PlatformType[] = ["anthropic", "openai", "gemini", "antigravity", "grok", "typesafe"] export type SchedulingThresholdPlatformType = | "openai" @@ -679,6 +680,9 @@ export interface SystemSettings { payment_balance_recharge_multiplier: number; payment_subscription_usd_to_cny_rate: number; payment_recharge_fee_rate: number; + payment_recharge_bonus_tiers?: RechargeBonusTier[]; + payment_recharge_bonus_mode?: string; + payment_recharge_bonus_notice?: string; payment_load_balance_strategy: string; payment_product_name_prefix: string; payment_product_name_suffix: string; @@ -1030,6 +1034,9 @@ export interface UpdateSettingsRequest { payment_balance_recharge_multiplier?: number; payment_subscription_usd_to_cny_rate?: number; payment_recharge_fee_rate?: number; + payment_recharge_bonus_tiers?: RechargeBonusTier[]; + payment_recharge_bonus_mode?: string; + payment_recharge_bonus_notice?: string; payment_load_balance_strategy?: string; payment_product_name_prefix?: string; payment_product_name_suffix?: string; diff --git a/frontend/src/api/admin/users.ts b/frontend/src/api/admin/users.ts index 0b804b1ef7d9..4c79624171e0 100644 --- a/frontend/src/api/admin/users.ts +++ b/frontend/src/api/admin/users.ts @@ -334,7 +334,7 @@ export async function bindUserAuthIdentity( // Keep aligned with backend/internal/service/domain_constants.go AllowedQuotaPlatforms. export const PLATFORM_QUOTA_PLATFORMS = [ 'anthropic', 'openai', 'gemini', 'antigravity', 'grok', - 'kimi', 'zhipu', 'deepseek', 'minimax', 'opencode_go', + 'kimi', 'zhipu', 'deepseek', 'minimax', 'opencode_go', 'typesafe', ] as const export type PlatformQuotaPlatform = typeof PLATFORM_QUOTA_PLATFORMS[number] export type PlatformQuotaWindow = 'daily' | 'weekly' | 'monthly' diff --git a/frontend/src/components/account/AccountPriorityCell.vue b/frontend/src/components/account/AccountPriorityCell.vue new file mode 100644 index 000000000000..b24934149728 --- /dev/null +++ b/frontend/src/components/account/AccountPriorityCell.vue @@ -0,0 +1,182 @@ + + + diff --git a/frontend/src/components/account/CreateAccountModal.vue b/frontend/src/components/account/CreateAccountModal.vue index 9add55b1ab3f..f27a0273a7b6 100644 --- a/frontend/src/components/account/CreateAccountModal.vue +++ b/frontend/src/components/account/CreateAccountModal.vue @@ -228,6 +228,19 @@ OpenCode + @@ -4025,6 +4038,8 @@ const apiKeyBaseUrlPlaceholder = computed(() => { return 'https://generativelanguage.googleapis.com' case 'grok': return 'https://api.x.ai/v1' + case 'typesafe': + return 'https://api.typesafe.ai' default: return 'https://api.anthropic.com' } @@ -4047,6 +4062,8 @@ const apiKeyValuePlaceholder = computed(() => { case 'minimax': case 'opencode_go': return 'sk-...' + case 'typesafe': + return 'ts-...' default: return 'sk-ant-...' } @@ -4260,6 +4277,13 @@ function selectOpenCodeGoPlatform() { resetAdaptiveBaseUrls('opencode_go', openCodeAccountMode.value) openCodeGoProtocolRules.value = cloneOpenCodeGoProtocolRules(defaultOpenCodeProtocolRules(openCodeAccountMode.value)) } +function selectTypeSafePlatform() { + form.platform = 'typesafe' + form.type = 'apikey' + accountCategory.value = 'apikey' + apiKeyBaseUrl.value = 'https://api.typesafe.ai' + allowedModels.value = ['jev-latest'] +} // 账号类型 / 协议变更时同步默认 base url。 watch(openCodeAccountMode, (mode, previousMode) => { if (!isOpenCodeGoPlatform.value) return @@ -4834,12 +4858,20 @@ watch( ? 'https://generativelanguage.googleapis.com' : newPlatform === 'grok' ? 'https://api.x.ai/v1' + : newPlatform === 'typesafe' + ? 'https://api.typesafe.ai' : 'https://api.anthropic.com' } // Clear model-related settings allowedModels.value = [] upstreamModelsPreviewed.value = false modelMappings.value = [] + if (newPlatform === 'typesafe') { + accountCategory.value = 'apikey' + // Grok 等平台会把模式切到映射;TypeSafe 只用白名单写入 jev-latest。 + modelRestrictionMode.value = 'whitelist' + allowedModels.value = ['jev-latest'] + } // Antigravity: 默认使用映射模式并填充默认映射 if (newPlatform === 'antigravity') { antigravityModelRestrictionMode.value = 'mapping' @@ -5784,6 +5816,8 @@ const handleSubmit = async () => { ? 'https://generativelanguage.googleapis.com' : form.platform === 'grok' ? 'https://api.x.ai/v1' + : form.platform === 'typesafe' + ? 'https://api.typesafe.ai' : 'https://api.anthropic.com' // Build credentials with optional model mapping diff --git a/frontend/src/components/account/EditAccountModal.vue b/frontend/src/components/account/EditAccountModal.vue index 103fe865e7c0..e6d4b07dd1ce 100644 --- a/frontend/src/components/account/EditAccountModal.vue +++ b/frontend/src/components/account/EditAccountModal.vue @@ -4256,6 +4256,7 @@ const defaultBaseUrl = computed(() => { if (props.account?.platform === 'openai') return 'https://api.openai.com' if (props.account?.platform === 'gemini') return 'https://generativelanguage.googleapis.com' if (props.account?.platform === 'grok') return 'https://api.x.ai/v1' + if (props.account?.platform === 'typesafe') return 'https://api.typesafe.ai' // CN 供应商:按当前模式/协议回落到官方预设(清空输入框提交时使用), // 不能落到 anthropic 默认值(会被当 CC base 拼出错误端点)。 if ( @@ -4756,6 +4757,8 @@ const syncFormFromAccount = (newAccount: Account | null) => { ? 'https://generativelanguage.googleapis.com' : newAccount.platform === 'grok' ? 'https://api.x.ai/v1' + : newAccount.platform === 'typesafe' + ? 'https://api.typesafe.ai' : newAccount.platform === 'kimi' || newAccount.platform === 'zhipu' || newAccount.platform === 'deepseek' || diff --git a/frontend/src/components/account/__tests__/AccountPriorityCell.spec.ts b/frontend/src/components/account/__tests__/AccountPriorityCell.spec.ts new file mode 100644 index 000000000000..85f7fe37eb91 --- /dev/null +++ b/frontend/src/components/account/__tests__/AccountPriorityCell.spec.ts @@ -0,0 +1,92 @@ +import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest' +import { flushPromises, mount } from '@vue/test-utils' +import AccountPriorityCell from '../AccountPriorityCell.vue' +import type { Account } from '@/types' +import { update } from '@/api/admin/accounts' + +vi.mock('@/api/admin/accounts', () => ({ update: vi.fn() })) +vi.mock('vue-i18n', () => ({ useI18n: () => ({ t: (key: string) => key }) })) + +const account = (overrides: Partial = {}) => ({ + id: 7, name: 'Claude 1', platform: 'anthropic', type: 'oauth', priority: 3, + ...overrides, +}) as Account + +const mountCell = (value = account()) => mount(AccountPriorityCell, { props: { account: value } }) + +beforeEach(() => { + vi.useFakeTimers() + vi.mocked(update).mockReset().mockImplementation(async (id, req) => account({ id, priority: req.priority })) +}) +afterEach(() => { + vi.useRealTimers() +}) + +describe('AccountPriorityCell', () => { + it('batches rapid +/- clicks into a single priority-only update', async () => { + const wrapper = mountCell() + await wrapper.get('[data-testid="account-priority-increment"]').trigger('click') + await wrapper.get('[data-testid="account-priority-increment"]').trigger('click') + await wrapper.get('[data-testid="account-priority-decrement"]').trigger('click') + expect(wrapper.get('[data-testid="account-priority-value"]').text()).toBe('4') + expect(update).not.toHaveBeenCalled() + + await vi.runAllTimersAsync() + await flushPromises() + + expect(update).toHaveBeenCalledTimes(1) + expect(update).toHaveBeenCalledWith(7, { priority: 4 }) + expect(wrapper.emitted('updated')?.[0]?.[0]).toMatchObject({ id: 7, priority: 4 }) + }) + + it('does not go below 1', async () => { + const wrapper = mountCell(account({ priority: 1 })) + const dec = wrapper.get('[data-testid="account-priority-decrement"]') + expect(dec.attributes('disabled')).toBeDefined() + // 到达下限时按钮仍应随悬停显隐,而不是常驻半透明 + expect(dec.classes()).toContain('opacity-0') + expect(dec.classes().some(c => c.startsWith('disabled:opacity'))).toBe(false) + await dec.trigger('click') + await vi.runAllTimersAsync() + expect(update).not.toHaveBeenCalled() + }) + + it('saves a typed value on Enter and ignores unchanged input', async () => { + const wrapper = mountCell() + await wrapper.get('[data-testid="account-priority-value"]').trigger('click') + const input = wrapper.get('[data-testid="account-priority-input"]') + await input.setValue('12') + await input.trigger('keydown', { key: 'Enter' }) + await flushPromises() + expect(update).toHaveBeenCalledWith(7, { priority: 12 }) + + vi.mocked(update).mockClear() + await wrapper.setProps({ account: account({ priority: 12 }) }) + await wrapper.get('[data-testid="account-priority-value"]').trigger('click') + await wrapper.get('[data-testid="account-priority-input"]').trigger('keydown', { key: 'Enter' }) + await flushPromises() + expect(update).not.toHaveBeenCalled() + }) + + it('Escape cancels typing without saving', async () => { + const wrapper = mountCell() + await wrapper.get('[data-testid="account-priority-value"]').trigger('click') + const input = wrapper.get('[data-testid="account-priority-input"]') + await input.setValue('50') + await input.trigger('keydown', { key: 'Escape' }) + await flushPromises() + expect(update).not.toHaveBeenCalled() + expect(wrapper.get('[data-testid="account-priority-value"]').text()).toBe('3') + }) + + it('reverts and emits an error when the update fails', async () => { + vi.mocked(update).mockRejectedValueOnce(new Error('boom')) + const wrapper = mountCell() + await wrapper.get('[data-testid="account-priority-increment"]').trigger('click') + await vi.runAllTimersAsync() + await flushPromises() + expect(wrapper.get('[data-testid="account-priority-value"]').text()).toBe('3') + expect(wrapper.emitted('error')).toHaveLength(1) + expect(wrapper.emitted('updated')).toBeUndefined() + }) +}) diff --git a/frontend/src/components/admin/payment/AdminOrderDetail.vue b/frontend/src/components/admin/payment/AdminOrderDetail.vue index 28bd614ac83d..277cae5e7aeb 100644 --- a/frontend/src/components/admin/payment/AdminOrderDetail.vue +++ b/frontend/src/components/admin/payment/AdminOrderDetail.vue @@ -29,7 +29,11 @@

{{ t('payment.orders.payAmount') }}

{{ paymentAmountSymbol }}{{ order.pay_amount.toFixed(2) }}

-
+
+

{{ t('payment.orders.bonusAmount') }}

+

+{{ creditedAmountSymbol }}{{ (order.bonus_amount ?? 0).toFixed(2) }}

+
+

{{ t('payment.orders.creditedAmount') }}

{{ creditedAmountSymbol }}{{ order.amount.toFixed(2) }}

diff --git a/frontend/src/components/admin/payment/AdminOrderTable.vue b/frontend/src/components/admin/payment/AdminOrderTable.vue index 5290866cae0e..dbbb5e87e0fc 100644 --- a/frontend/src/components/admin/payment/AdminOrderTable.vue +++ b/frontend/src/components/admin/payment/AdminOrderTable.vue @@ -57,8 +57,11 @@ ({{ row.fee_rate }}%) -
+
{{ t('payment.orders.creditedAmount') }}: {{ creditedAmountSymbol }}{{ row.amount.toFixed(2) }} + + ({{ t('payment.orders.bonusIncluded', { amount: creditedAmountSymbol + (row.bonus_amount ?? 0).toFixed(2) }) }}) +
diff --git a/frontend/src/components/admin/settings/RechargeBonusTierEditor.vue b/frontend/src/components/admin/settings/RechargeBonusTierEditor.vue new file mode 100644 index 000000000000..6ef92a9a449c --- /dev/null +++ b/frontend/src/components/admin/settings/RechargeBonusTierEditor.vue @@ -0,0 +1,297 @@ + + + diff --git a/frontend/src/components/admin/user/__tests__/UserPlatformQuotaModal.spec.ts b/frontend/src/components/admin/user/__tests__/UserPlatformQuotaModal.spec.ts index b102ab56c19a..1021d08f6708 100644 --- a/frontend/src/components/admin/user/__tests__/UserPlatformQuotaModal.spec.ts +++ b/frontend/src/components/admin/user/__tests__/UserPlatformQuotaModal.spec.ts @@ -75,7 +75,7 @@ beforeEach(() => { }) describe('UserPlatformQuotaModal', () => { - it.each([0, 4, 14])('does not turn a negative limit in input %s into unlimited', async (index) => { + it.each([0, 4, 14, 17])('does not turn a negative limit in input %s into unlimited', async (index) => { const w = await mountAndOpen() await w.findAll('input[type=number]')[index].setValue('-1') await w.findAll('button').find(b => b.text() === 'admin.users.platformQuota.save')!.trigger('click') @@ -101,12 +101,12 @@ describe('UserPlatformQuotaModal', () => { expect(apiMocks.getPlatformQuotas).toHaveBeenCalledWith(99) }) - it('renders all ten supported platforms with empty limits', async () => { + it('renders all eleven supported platforms with empty limits', async () => { const w = await mountAndOpen() const rows = w.findAll('tbody tr') expect(rows.map(row => row.find('td').text())).toEqual([ 'anthropic', 'openai', 'gemini', 'antigravity', 'grok', - 'kimi', 'zhipu', 'deepseek', 'minimax', 'opencode_go', + 'kimi', 'zhipu', 'deepseek', 'minimax', 'opencode_go', 'typesafe', ]) for (const row of rows) { const inputs = row.findAll('input[type=number]') @@ -137,7 +137,7 @@ describe('UserPlatformQuotaModal', () => { : item) expect(apiMocks.updatePlatformQuotas).toHaveBeenCalledTimes(1) expect(apiMocks.updatePlatformQuotas).toHaveBeenCalledWith(99, expect.arrayContaining(expected)) - expect(apiMocks.updatePlatformQuotas.mock.calls[0][1]).toHaveLength(10) + expect(apiMocks.updatePlatformQuotas.mock.calls[0][1]).toHaveLength(11) expect(w.emitted('success')).toHaveLength(1) w.unmount() }, @@ -152,13 +152,13 @@ describe('UserPlatformQuotaModal', () => { }) const w = await mountAndOpen() const inputs = w.findAll('input[type=number]') - // 10 platforms × 3 windows = 30 inputs - expect(inputs.length).toBe(30) + // 11 platforms × 3 windows = 33 inputs + expect(inputs.length).toBe(33) // 第一个 input 是 anthropic.daily = 10 expect((inputs[0].element as HTMLInputElement).value).toBe('10') }) - it('保存提交完整 10 platform payload', async () => { + it('保存提交完整 11 platform payload', async () => { apiMocks.getPlatformQuotas.mockResolvedValueOnce({ platform_quotas: [ { platform: 'openai', daily_limit_usd: null, weekly_limit_usd: 20, monthly_limit_usd: null, @@ -175,7 +175,7 @@ describe('UserPlatformQuotaModal', () => { expect(apiMocks.updatePlatformQuotas).toHaveBeenCalledTimes(1) const [uid, payload] = apiMocks.updatePlatformQuotas.mock.calls[0] expect(uid).toBe(99) - expect(payload).toHaveLength(10) // 10 platforms always submitted + expect(payload).toHaveLength(11) // 11 platforms always submitted const openai = payload.find((p: any) => p.platform === 'openai') expect(openai.weekly_limit_usd).toBe(20) }) @@ -240,7 +240,7 @@ describe('UserPlatformQuotaModal', () => { it('未配置限额的平台重置按钮禁用并提示不可用', async () => { const w = await mountAndOpen() const resetBtns = w.findAll('button').filter((b) => b.text() === '↻') - expect(resetBtns.length).toBe(30) // 10 平台 × 3 窗口 + expect(resetBtns.length).toBe(33) // 11 平台 × 3 窗口 for (const b of resetBtns) { expect((b.element as HTMLButtonElement).disabled).toBe(true) expect(b.attributes('title')).toBe('admin.users.platformQuota.reset.unavailable') diff --git a/frontend/src/components/keys/UseKeyModal.vue b/frontend/src/components/keys/UseKeyModal.vue index 383181285a6a..8b696800b263 100644 --- a/frontend/src/components/keys/UseKeyModal.vue +++ b/frontend/src/components/keys/UseKeyModal.vue @@ -364,6 +364,8 @@ const defaultClientTab = computed(() => { return 'gemini' case 'antigravity': return 'claude' + case 'typesafe': + return 'systemone' default: return 'claude' } @@ -491,6 +493,10 @@ const clientTabs = computed((): TabConfig[] => { { id: 'codex', label: t('keys.useKeyModal.cliTabs.codexCli'), icon: TerminalIcon }, { id: 'opencode', label: t('keys.useKeyModal.cliTabs.opencode'), icon: TerminalIcon } ] + case 'typesafe': + return [ + { id: 'systemone', label: t('keys.useKeyModal.cliTabs.systemOne'), icon: TerminalIcon } + ] case 'deepseek': case 'minimax': case 'composite': @@ -575,6 +581,8 @@ const platformDescription = computed(() => { return activeClientTab.value === 'codex' ? t('keys.useKeyModal.composite.codexDescription') : t('keys.useKeyModal.composite.description') + case 'typesafe': + return t('keys.useKeyModal.typesafe.description') default: return t('keys.useKeyModal.description') } @@ -632,6 +640,8 @@ const platformNote = computed(() => { return activeClientTab.value === 'codex' ? t('keys.useKeyModal.composite.codexNote') : t('keys.useKeyModal.note') + case 'typesafe': + return t('keys.useKeyModal.typesafe.note') default: return t('keys.useKeyModal.note') } @@ -760,6 +770,8 @@ const currentFiles = computed((): FileConfig[] => { } switch (props.platform) { + case 'typesafe': + return [generateSystemOneCurl(baseRoot, apiKey)] case 'openai': if (activeClientTab.value === 'claude') { // Anthropic clients append /v1/messages themselves. @@ -814,6 +826,46 @@ const currentFiles = computed((): FileConfig[] => { } }) +function generateSystemOneCurl(baseUrl: string, apiKey: string): FileConfig { + const endpoint = `${baseUrl}/v1/systemone` + const payload = `{ + "model": "jev-latest", + "state": "Text to evaluate", + "questions": { + "safety": { + "type": "noul", + "instructions": "Evaluate whether the text is unsafe" + } + } +}` + if (activeTab.value === 'powershell') { + return { + path: 'PowerShell', + content: `$headers = @{ Authorization = "Bearer ${apiKey}" } +$body = @' +${payload} +'@ +Invoke-RestMethod -Method Post -Uri "${endpoint}" -Headers $headers -ContentType "application/json" -Body $body` + } + } + if (activeTab.value === 'cmd') { + return { + path: 'Command Prompt', + content: `curl -X POST "${endpoint}" ^ + -H "Authorization: Bearer ${apiKey}" ^ + -H "Content-Type: application/json" ^ + --data "{\"model\":\"jev-latest\",\"state\":\"Text to evaluate\",\"questions\":{\"safety\":{\"type\":\"noul\",\"instructions\":\"Evaluate whether the text is unsafe\"}}}"` + } + } + return { + path: 'Terminal', + content: `curl -X POST "${endpoint}" \\ + -H "Authorization: Bearer ${apiKey}" \\ + -H "Content-Type: application/json" \\ + --data '${payload}'` + } +} + function generateAnthropicFiles(baseUrl: string, apiKey: string): FileConfig[] { let path: string let content: string @@ -1278,6 +1330,7 @@ function generateRoutedCodexFiles( deepseek: 'DeepSeek', minimax: 'MiniMax', opencode_go: 'OpenCode', + typesafe: 'TypeSafe / Jev', composite: 'Composite' } const label = labels[platform] diff --git a/frontend/src/components/payment/AmountInput.vue b/frontend/src/components/payment/AmountInput.vue index 0e0c9d339ad3..a16029e94bf2 100644 --- a/frontend/src/components/payment/AmountInput.vue +++ b/frontend/src/components/payment/AmountInput.vue @@ -5,20 +5,45 @@ -
+
@@ -48,16 +73,31 @@