diff --git a/README.md b/README.md index 41a5aca11c..99753e4569 100644 --- a/README.md +++ b/README.md @@ -49,9 +49,13 @@ Sub2API is an AI API gateway platform designed to distribute and manage API quot - + + + + +
pinccpincc PinCC is the official relay service built on Sub2API, offering stable access to Claude Code, Codex, Gemini and other popular models — ready to use, no deployment or maintenance required.
PackyCodeThanks to PackyCode for sponsoring this project! PackyCode is a reliable and efficient API relay service provider, offering relay services for Claude Code, Codex, Gemini, and more. PackyCode provides special discounts for our software users: register using this link and enter the "sub2api" promo code during first recharge to get 10% off.
## Ecosystem diff --git a/README_CN.md b/README_CN.md index 3380cce7ff..8b6feaba0d 100644 --- a/README_CN.md +++ b/README_CN.md @@ -48,9 +48,13 @@ Sub2API 是一个 AI API 网关平台,用于分发和管理 AI 产品订阅的 - + + + + +
pinccpincc PinCC 是基于 Sub2API 搭建的官方中转服务,提供 Claude Code、Codex、Gemini 等主流模型的稳定中转,开箱即用,免去自建部署与运维烦恼。
PackyCode感谢 PackyCode 赞助了本项目!PackyCode 是一家稳定、高效的API中转服务商,提供 Claude Code、Codex、Gemini 等多种中转服务。PackyCode 为本软件的用户提供了特别优惠,使用此链接注册并在充值时填写"sub2api"优惠码,首次充值可以享受9折优惠!
## 生态项目 diff --git a/README_JA.md b/README_JA.md index c60b1a8e04..1266bd845c 100644 --- a/README_JA.md +++ b/README_JA.md @@ -49,9 +49,13 @@ Sub2API は、AI 製品のサブスクリプションから API クォータを - + + + + +
pinccpincc PinCC は Sub2API 上に構築された公式リレーサービスで、Claude Code、Codex、Gemini などの人気モデルへの安定したアクセスを提供します。デプロイやメンテナンスは不要で、すぐにご利用いただけます。
PackyCodePackyCode のご支援に感謝します!PackyCode は Claude Code、Codex、Gemini などのリレーサービスを提供する信頼性の高い API 中継プラットフォームです。本ソフト利用者向けに特別割引があります:このリンクで登録し、チャージ時に「sub2api」クーポンを入力すると 10% オフになります。
## エコシステム diff --git a/assets/partners/logos/packycode.png b/assets/partners/logos/packycode.png new file mode 100644 index 0000000000..4fc7eecc75 Binary files /dev/null and b/assets/partners/logos/packycode.png differ diff --git a/backend/cmd/server/VERSION b/backend/cmd/server/VERSION index 15419f54b6..c1d04e7243 100644 --- a/backend/cmd/server/VERSION +++ b/backend/cmd/server/VERSION @@ -1 +1 @@ -0.1.106.1 +0.1.108.1 diff --git a/backend/ent/group.go b/backend/ent/group.go index b901122cd8..3932da2bc7 100644 --- a/backend/ent/group.go +++ b/backend/ent/group.go @@ -70,6 +70,10 @@ type Group struct { SortOrder int `json:"sort_order,omitempty"` // 是否允许 /v1/messages 调度到此 OpenAI 分组 AllowMessagesDispatch bool `json:"allow_messages_dispatch,omitempty"` + // 仅允许非 apikey 类型账号关联到此分组 + RequireOauthOnly bool `json:"require_oauth_only,omitempty"` + // 调度时仅允许 privacy 已成功设置的账号 + RequirePrivacySet bool `json:"require_privacy_set,omitempty"` // 默认映射模型 ID,当账号级映射找不到时使用此值 DefaultMappedModel string `json:"default_mapped_model,omitempty"` // Edges holds the relations/edges for other nodes in the graph. @@ -180,7 +184,7 @@ func (*Group) scanValues(columns []string) ([]any, error) { switch columns[i] { case group.FieldModelRouting, group.FieldSupportedModelScopes: values[i] = new([]byte) - case group.FieldIsExclusive, group.FieldClaudeCodeOnly, group.FieldModelRoutingEnabled, group.FieldMcpXMLInject, group.FieldAllowMessagesDispatch: + case group.FieldIsExclusive, group.FieldClaudeCodeOnly, group.FieldModelRoutingEnabled, group.FieldMcpXMLInject, group.FieldAllowMessagesDispatch, group.FieldRequireOauthOnly, group.FieldRequirePrivacySet: values[i] = new(sql.NullBool) case group.FieldRateMultiplier, group.FieldDailyLimitUsd, group.FieldWeeklyLimitUsd, group.FieldMonthlyLimitUsd, group.FieldImagePrice1k, group.FieldImagePrice2k, group.FieldImagePrice4k: values[i] = new(sql.NullFloat64) @@ -381,6 +385,18 @@ func (_m *Group) assignValues(columns []string, values []any) error { } else if value.Valid { _m.AllowMessagesDispatch = value.Bool } + case group.FieldRequireOauthOnly: + if value, ok := values[i].(*sql.NullBool); !ok { + return fmt.Errorf("unexpected type %T for field require_oauth_only", values[i]) + } else if value.Valid { + _m.RequireOauthOnly = value.Bool + } + case group.FieldRequirePrivacySet: + if value, ok := values[i].(*sql.NullBool); !ok { + return fmt.Errorf("unexpected type %T for field require_privacy_set", values[i]) + } else if value.Valid { + _m.RequirePrivacySet = value.Bool + } case group.FieldDefaultMappedModel: if value, ok := values[i].(*sql.NullString); !ok { return fmt.Errorf("unexpected type %T for field default_mapped_model", values[i]) @@ -561,6 +577,12 @@ func (_m *Group) String() string { builder.WriteString("allow_messages_dispatch=") builder.WriteString(fmt.Sprintf("%v", _m.AllowMessagesDispatch)) builder.WriteString(", ") + builder.WriteString("require_oauth_only=") + builder.WriteString(fmt.Sprintf("%v", _m.RequireOauthOnly)) + builder.WriteString(", ") + builder.WriteString("require_privacy_set=") + builder.WriteString(fmt.Sprintf("%v", _m.RequirePrivacySet)) + builder.WriteString(", ") builder.WriteString("default_mapped_model=") builder.WriteString(_m.DefaultMappedModel) builder.WriteByte(')') diff --git a/backend/ent/group/group.go b/backend/ent/group/group.go index 79549a90aa..21a7c2cb76 100644 --- a/backend/ent/group/group.go +++ b/backend/ent/group/group.go @@ -67,6 +67,10 @@ const ( FieldSortOrder = "sort_order" // FieldAllowMessagesDispatch holds the string denoting the allow_messages_dispatch field in the database. FieldAllowMessagesDispatch = "allow_messages_dispatch" + // FieldRequireOauthOnly holds the string denoting the require_oauth_only field in the database. + FieldRequireOauthOnly = "require_oauth_only" + // FieldRequirePrivacySet holds the string denoting the require_privacy_set field in the database. + FieldRequirePrivacySet = "require_privacy_set" // FieldDefaultMappedModel holds the string denoting the default_mapped_model field in the database. FieldDefaultMappedModel = "default_mapped_model" // EdgeAPIKeys holds the string denoting the api_keys edge name in mutations. @@ -170,6 +174,8 @@ var Columns = []string{ FieldSupportedModelScopes, FieldSortOrder, FieldAllowMessagesDispatch, + FieldRequireOauthOnly, + FieldRequirePrivacySet, FieldDefaultMappedModel, } @@ -238,6 +244,10 @@ var ( DefaultSortOrder int // DefaultAllowMessagesDispatch holds the default value on creation for the "allow_messages_dispatch" field. DefaultAllowMessagesDispatch bool + // DefaultRequireOauthOnly holds the default value on creation for the "require_oauth_only" field. + DefaultRequireOauthOnly bool + // DefaultRequirePrivacySet holds the default value on creation for the "require_privacy_set" field. + DefaultRequirePrivacySet bool // DefaultDefaultMappedModel holds the default value on creation for the "default_mapped_model" field. DefaultDefaultMappedModel string // DefaultMappedModelValidator is a validator for the "default_mapped_model" field. It is called by the builders before save. @@ -372,6 +382,16 @@ func ByAllowMessagesDispatch(opts ...sql.OrderTermOption) OrderOption { return sql.OrderByField(FieldAllowMessagesDispatch, opts...).ToFunc() } +// ByRequireOauthOnly orders the results by the require_oauth_only field. +func ByRequireOauthOnly(opts ...sql.OrderTermOption) OrderOption { + return sql.OrderByField(FieldRequireOauthOnly, opts...).ToFunc() +} + +// ByRequirePrivacySet orders the results by the require_privacy_set field. +func ByRequirePrivacySet(opts ...sql.OrderTermOption) OrderOption { + return sql.OrderByField(FieldRequirePrivacySet, opts...).ToFunc() +} + // ByDefaultMappedModel orders the results by the default_mapped_model field. func ByDefaultMappedModel(opts ...sql.OrderTermOption) OrderOption { return sql.OrderByField(FieldDefaultMappedModel, opts...).ToFunc() diff --git a/backend/ent/group/where.go b/backend/ent/group/where.go index 419917b712..cba2ce5f0e 100644 --- a/backend/ent/group/where.go +++ b/backend/ent/group/where.go @@ -175,6 +175,16 @@ func AllowMessagesDispatch(v bool) predicate.Group { return predicate.Group(sql.FieldEQ(FieldAllowMessagesDispatch, v)) } +// RequireOauthOnly applies equality check predicate on the "require_oauth_only" field. It's identical to RequireOauthOnlyEQ. +func RequireOauthOnly(v bool) predicate.Group { + return predicate.Group(sql.FieldEQ(FieldRequireOauthOnly, v)) +} + +// RequirePrivacySet applies equality check predicate on the "require_privacy_set" field. It's identical to RequirePrivacySetEQ. +func RequirePrivacySet(v bool) predicate.Group { + return predicate.Group(sql.FieldEQ(FieldRequirePrivacySet, v)) +} + // DefaultMappedModel applies equality check predicate on the "default_mapped_model" field. It's identical to DefaultMappedModelEQ. func DefaultMappedModel(v string) predicate.Group { return predicate.Group(sql.FieldEQ(FieldDefaultMappedModel, v)) @@ -1225,6 +1235,26 @@ func AllowMessagesDispatchNEQ(v bool) predicate.Group { return predicate.Group(sql.FieldNEQ(FieldAllowMessagesDispatch, v)) } +// RequireOauthOnlyEQ applies the EQ predicate on the "require_oauth_only" field. +func RequireOauthOnlyEQ(v bool) predicate.Group { + return predicate.Group(sql.FieldEQ(FieldRequireOauthOnly, v)) +} + +// RequireOauthOnlyNEQ applies the NEQ predicate on the "require_oauth_only" field. +func RequireOauthOnlyNEQ(v bool) predicate.Group { + return predicate.Group(sql.FieldNEQ(FieldRequireOauthOnly, v)) +} + +// RequirePrivacySetEQ applies the EQ predicate on the "require_privacy_set" field. +func RequirePrivacySetEQ(v bool) predicate.Group { + return predicate.Group(sql.FieldEQ(FieldRequirePrivacySet, v)) +} + +// RequirePrivacySetNEQ applies the NEQ predicate on the "require_privacy_set" field. +func RequirePrivacySetNEQ(v bool) predicate.Group { + return predicate.Group(sql.FieldNEQ(FieldRequirePrivacySet, v)) +} + // DefaultMappedModelEQ applies the EQ predicate on the "default_mapped_model" field. func DefaultMappedModelEQ(v string) predicate.Group { return predicate.Group(sql.FieldEQ(FieldDefaultMappedModel, v)) diff --git a/backend/ent/group_create.go b/backend/ent/group_create.go index bbc4b4f62b..a8c30b184d 100644 --- a/backend/ent/group_create.go +++ b/backend/ent/group_create.go @@ -368,6 +368,34 @@ func (_c *GroupCreate) SetNillableAllowMessagesDispatch(v *bool) *GroupCreate { return _c } +// SetRequireOauthOnly sets the "require_oauth_only" field. +func (_c *GroupCreate) SetRequireOauthOnly(v bool) *GroupCreate { + _c.mutation.SetRequireOauthOnly(v) + return _c +} + +// SetNillableRequireOauthOnly sets the "require_oauth_only" field if the given value is not nil. +func (_c *GroupCreate) SetNillableRequireOauthOnly(v *bool) *GroupCreate { + if v != nil { + _c.SetRequireOauthOnly(*v) + } + return _c +} + +// SetRequirePrivacySet sets the "require_privacy_set" field. +func (_c *GroupCreate) SetRequirePrivacySet(v bool) *GroupCreate { + _c.mutation.SetRequirePrivacySet(v) + return _c +} + +// SetNillableRequirePrivacySet sets the "require_privacy_set" field if the given value is not nil. +func (_c *GroupCreate) SetNillableRequirePrivacySet(v *bool) *GroupCreate { + if v != nil { + _c.SetRequirePrivacySet(*v) + } + return _c +} + // SetDefaultMappedModel sets the "default_mapped_model" field. func (_c *GroupCreate) SetDefaultMappedModel(v string) *GroupCreate { _c.mutation.SetDefaultMappedModel(v) @@ -571,6 +599,14 @@ func (_c *GroupCreate) defaults() error { v := group.DefaultAllowMessagesDispatch _c.mutation.SetAllowMessagesDispatch(v) } + if _, ok := _c.mutation.RequireOauthOnly(); !ok { + v := group.DefaultRequireOauthOnly + _c.mutation.SetRequireOauthOnly(v) + } + if _, ok := _c.mutation.RequirePrivacySet(); !ok { + v := group.DefaultRequirePrivacySet + _c.mutation.SetRequirePrivacySet(v) + } if _, ok := _c.mutation.DefaultMappedModel(); !ok { v := group.DefaultDefaultMappedModel _c.mutation.SetDefaultMappedModel(v) @@ -645,6 +681,12 @@ func (_c *GroupCreate) check() error { if _, ok := _c.mutation.AllowMessagesDispatch(); !ok { return &ValidationError{Name: "allow_messages_dispatch", err: errors.New(`ent: missing required field "Group.allow_messages_dispatch"`)} } + if _, ok := _c.mutation.RequireOauthOnly(); !ok { + return &ValidationError{Name: "require_oauth_only", err: errors.New(`ent: missing required field "Group.require_oauth_only"`)} + } + if _, ok := _c.mutation.RequirePrivacySet(); !ok { + return &ValidationError{Name: "require_privacy_set", err: errors.New(`ent: missing required field "Group.require_privacy_set"`)} + } if _, ok := _c.mutation.DefaultMappedModel(); !ok { return &ValidationError{Name: "default_mapped_model", err: errors.New(`ent: missing required field "Group.default_mapped_model"`)} } @@ -784,6 +826,14 @@ func (_c *GroupCreate) createSpec() (*Group, *sqlgraph.CreateSpec) { _spec.SetField(group.FieldAllowMessagesDispatch, field.TypeBool, value) _node.AllowMessagesDispatch = value } + if value, ok := _c.mutation.RequireOauthOnly(); ok { + _spec.SetField(group.FieldRequireOauthOnly, field.TypeBool, value) + _node.RequireOauthOnly = value + } + if value, ok := _c.mutation.RequirePrivacySet(); ok { + _spec.SetField(group.FieldRequirePrivacySet, field.TypeBool, value) + _node.RequirePrivacySet = value + } if value, ok := _c.mutation.DefaultMappedModel(); ok { _spec.SetField(group.FieldDefaultMappedModel, field.TypeString, value) _node.DefaultMappedModel = value @@ -1376,6 +1426,30 @@ func (u *GroupUpsert) UpdateAllowMessagesDispatch() *GroupUpsert { return u } +// SetRequireOauthOnly sets the "require_oauth_only" field. +func (u *GroupUpsert) SetRequireOauthOnly(v bool) *GroupUpsert { + u.Set(group.FieldRequireOauthOnly, v) + return u +} + +// UpdateRequireOauthOnly sets the "require_oauth_only" field to the value that was provided on create. +func (u *GroupUpsert) UpdateRequireOauthOnly() *GroupUpsert { + u.SetExcluded(group.FieldRequireOauthOnly) + return u +} + +// SetRequirePrivacySet sets the "require_privacy_set" field. +func (u *GroupUpsert) SetRequirePrivacySet(v bool) *GroupUpsert { + u.Set(group.FieldRequirePrivacySet, v) + return u +} + +// UpdateRequirePrivacySet sets the "require_privacy_set" field to the value that was provided on create. +func (u *GroupUpsert) UpdateRequirePrivacySet() *GroupUpsert { + u.SetExcluded(group.FieldRequirePrivacySet) + return u +} + // SetDefaultMappedModel sets the "default_mapped_model" field. func (u *GroupUpsert) SetDefaultMappedModel(v string) *GroupUpsert { u.Set(group.FieldDefaultMappedModel, v) @@ -1937,6 +2011,34 @@ func (u *GroupUpsertOne) UpdateAllowMessagesDispatch() *GroupUpsertOne { }) } +// SetRequireOauthOnly sets the "require_oauth_only" field. +func (u *GroupUpsertOne) SetRequireOauthOnly(v bool) *GroupUpsertOne { + return u.Update(func(s *GroupUpsert) { + s.SetRequireOauthOnly(v) + }) +} + +// UpdateRequireOauthOnly sets the "require_oauth_only" field to the value that was provided on create. +func (u *GroupUpsertOne) UpdateRequireOauthOnly() *GroupUpsertOne { + return u.Update(func(s *GroupUpsert) { + s.UpdateRequireOauthOnly() + }) +} + +// SetRequirePrivacySet sets the "require_privacy_set" field. +func (u *GroupUpsertOne) SetRequirePrivacySet(v bool) *GroupUpsertOne { + return u.Update(func(s *GroupUpsert) { + s.SetRequirePrivacySet(v) + }) +} + +// UpdateRequirePrivacySet sets the "require_privacy_set" field to the value that was provided on create. +func (u *GroupUpsertOne) UpdateRequirePrivacySet() *GroupUpsertOne { + return u.Update(func(s *GroupUpsert) { + s.UpdateRequirePrivacySet() + }) +} + // SetDefaultMappedModel sets the "default_mapped_model" field. func (u *GroupUpsertOne) SetDefaultMappedModel(v string) *GroupUpsertOne { return u.Update(func(s *GroupUpsert) { @@ -2666,6 +2768,34 @@ func (u *GroupUpsertBulk) UpdateAllowMessagesDispatch() *GroupUpsertBulk { }) } +// SetRequireOauthOnly sets the "require_oauth_only" field. +func (u *GroupUpsertBulk) SetRequireOauthOnly(v bool) *GroupUpsertBulk { + return u.Update(func(s *GroupUpsert) { + s.SetRequireOauthOnly(v) + }) +} + +// UpdateRequireOauthOnly sets the "require_oauth_only" field to the value that was provided on create. +func (u *GroupUpsertBulk) UpdateRequireOauthOnly() *GroupUpsertBulk { + return u.Update(func(s *GroupUpsert) { + s.UpdateRequireOauthOnly() + }) +} + +// SetRequirePrivacySet sets the "require_privacy_set" field. +func (u *GroupUpsertBulk) SetRequirePrivacySet(v bool) *GroupUpsertBulk { + return u.Update(func(s *GroupUpsert) { + s.SetRequirePrivacySet(v) + }) +} + +// UpdateRequirePrivacySet sets the "require_privacy_set" field to the value that was provided on create. +func (u *GroupUpsertBulk) UpdateRequirePrivacySet() *GroupUpsertBulk { + return u.Update(func(s *GroupUpsert) { + s.UpdateRequirePrivacySet() + }) +} + // SetDefaultMappedModel sets the "default_mapped_model" field. func (u *GroupUpsertBulk) SetDefaultMappedModel(v string) *GroupUpsertBulk { return u.Update(func(s *GroupUpsert) { diff --git a/backend/ent/group_update.go b/backend/ent/group_update.go index 1e041a432b..aa1a83d421 100644 --- a/backend/ent/group_update.go +++ b/backend/ent/group_update.go @@ -510,6 +510,34 @@ func (_u *GroupUpdate) SetNillableAllowMessagesDispatch(v *bool) *GroupUpdate { return _u } +// SetRequireOauthOnly sets the "require_oauth_only" field. +func (_u *GroupUpdate) SetRequireOauthOnly(v bool) *GroupUpdate { + _u.mutation.SetRequireOauthOnly(v) + return _u +} + +// SetNillableRequireOauthOnly sets the "require_oauth_only" field if the given value is not nil. +func (_u *GroupUpdate) SetNillableRequireOauthOnly(v *bool) *GroupUpdate { + if v != nil { + _u.SetRequireOauthOnly(*v) + } + return _u +} + +// SetRequirePrivacySet sets the "require_privacy_set" field. +func (_u *GroupUpdate) SetRequirePrivacySet(v bool) *GroupUpdate { + _u.mutation.SetRequirePrivacySet(v) + return _u +} + +// SetNillableRequirePrivacySet sets the "require_privacy_set" field if the given value is not nil. +func (_u *GroupUpdate) SetNillableRequirePrivacySet(v *bool) *GroupUpdate { + if v != nil { + _u.SetRequirePrivacySet(*v) + } + return _u +} + // SetDefaultMappedModel sets the "default_mapped_model" field. func (_u *GroupUpdate) SetDefaultMappedModel(v string) *GroupUpdate { _u.mutation.SetDefaultMappedModel(v) @@ -975,6 +1003,12 @@ func (_u *GroupUpdate) sqlSave(ctx context.Context) (_node int, err error) { if value, ok := _u.mutation.AllowMessagesDispatch(); ok { _spec.SetField(group.FieldAllowMessagesDispatch, field.TypeBool, value) } + if value, ok := _u.mutation.RequireOauthOnly(); ok { + _spec.SetField(group.FieldRequireOauthOnly, field.TypeBool, value) + } + if value, ok := _u.mutation.RequirePrivacySet(); ok { + _spec.SetField(group.FieldRequirePrivacySet, field.TypeBool, value) + } if value, ok := _u.mutation.DefaultMappedModel(); ok { _spec.SetField(group.FieldDefaultMappedModel, field.TypeString, value) } @@ -1767,6 +1801,34 @@ func (_u *GroupUpdateOne) SetNillableAllowMessagesDispatch(v *bool) *GroupUpdate return _u } +// SetRequireOauthOnly sets the "require_oauth_only" field. +func (_u *GroupUpdateOne) SetRequireOauthOnly(v bool) *GroupUpdateOne { + _u.mutation.SetRequireOauthOnly(v) + return _u +} + +// SetNillableRequireOauthOnly sets the "require_oauth_only" field if the given value is not nil. +func (_u *GroupUpdateOne) SetNillableRequireOauthOnly(v *bool) *GroupUpdateOne { + if v != nil { + _u.SetRequireOauthOnly(*v) + } + return _u +} + +// SetRequirePrivacySet sets the "require_privacy_set" field. +func (_u *GroupUpdateOne) SetRequirePrivacySet(v bool) *GroupUpdateOne { + _u.mutation.SetRequirePrivacySet(v) + return _u +} + +// SetNillableRequirePrivacySet sets the "require_privacy_set" field if the given value is not nil. +func (_u *GroupUpdateOne) SetNillableRequirePrivacySet(v *bool) *GroupUpdateOne { + if v != nil { + _u.SetRequirePrivacySet(*v) + } + return _u +} + // SetDefaultMappedModel sets the "default_mapped_model" field. func (_u *GroupUpdateOne) SetDefaultMappedModel(v string) *GroupUpdateOne { _u.mutation.SetDefaultMappedModel(v) @@ -2262,6 +2324,12 @@ func (_u *GroupUpdateOne) sqlSave(ctx context.Context) (_node *Group, err error) if value, ok := _u.mutation.AllowMessagesDispatch(); ok { _spec.SetField(group.FieldAllowMessagesDispatch, field.TypeBool, value) } + if value, ok := _u.mutation.RequireOauthOnly(); ok { + _spec.SetField(group.FieldRequireOauthOnly, field.TypeBool, value) + } + if value, ok := _u.mutation.RequirePrivacySet(); ok { + _spec.SetField(group.FieldRequirePrivacySet, field.TypeBool, value) + } if value, ok := _u.mutation.DefaultMappedModel(); ok { _spec.SetField(group.FieldDefaultMappedModel, field.TypeString, value) } diff --git a/backend/ent/migrate/schema.go b/backend/ent/migrate/schema.go index 765959b55e..5400bf9319 100644 --- a/backend/ent/migrate/schema.go +++ b/backend/ent/migrate/schema.go @@ -404,6 +404,8 @@ var ( {Name: "supported_model_scopes", Type: field.TypeJSON, SchemaType: map[string]string{"postgres": "jsonb"}}, {Name: "sort_order", Type: field.TypeInt, Default: 0}, {Name: "allow_messages_dispatch", Type: field.TypeBool, Default: false}, + {Name: "require_oauth_only", Type: field.TypeBool, Default: false}, + {Name: "require_privacy_set", Type: field.TypeBool, Default: false}, {Name: "default_mapped_model", Type: field.TypeString, Size: 100, Default: ""}, } // GroupsTable holds the schema information for the "groups" table. diff --git a/backend/ent/mutation.go b/backend/ent/mutation.go index ec1b4098c1..d206039af4 100644 --- a/backend/ent/mutation.go +++ b/backend/ent/mutation.go @@ -8243,6 +8243,8 @@ type GroupMutation struct { sort_order *int addsort_order *int allow_messages_dispatch *bool + require_oauth_only *bool + require_privacy_set *bool default_mapped_model *string clearedFields map[string]struct{} api_keys map[int64]struct{} @@ -9688,6 +9690,78 @@ func (m *GroupMutation) ResetAllowMessagesDispatch() { m.allow_messages_dispatch = nil } +// SetRequireOauthOnly sets the "require_oauth_only" field. +func (m *GroupMutation) SetRequireOauthOnly(b bool) { + m.require_oauth_only = &b +} + +// RequireOauthOnly returns the value of the "require_oauth_only" field in the mutation. +func (m *GroupMutation) RequireOauthOnly() (r bool, exists bool) { + v := m.require_oauth_only + if v == nil { + return + } + return *v, true +} + +// OldRequireOauthOnly returns the old "require_oauth_only" field's value of the Group entity. +// If the Group object wasn't provided to the builder, the object is fetched from the database. +// An error is returned if the mutation operation is not UpdateOne, or the database query fails. +func (m *GroupMutation) OldRequireOauthOnly(ctx context.Context) (v bool, err error) { + if !m.op.Is(OpUpdateOne) { + return v, errors.New("OldRequireOauthOnly is only allowed on UpdateOne operations") + } + if m.id == nil || m.oldValue == nil { + return v, errors.New("OldRequireOauthOnly requires an ID field in the mutation") + } + oldValue, err := m.oldValue(ctx) + if err != nil { + return v, fmt.Errorf("querying old value for OldRequireOauthOnly: %w", err) + } + return oldValue.RequireOauthOnly, nil +} + +// ResetRequireOauthOnly resets all changes to the "require_oauth_only" field. +func (m *GroupMutation) ResetRequireOauthOnly() { + m.require_oauth_only = nil +} + +// SetRequirePrivacySet sets the "require_privacy_set" field. +func (m *GroupMutation) SetRequirePrivacySet(b bool) { + m.require_privacy_set = &b +} + +// RequirePrivacySet returns the value of the "require_privacy_set" field in the mutation. +func (m *GroupMutation) RequirePrivacySet() (r bool, exists bool) { + v := m.require_privacy_set + if v == nil { + return + } + return *v, true +} + +// OldRequirePrivacySet returns the old "require_privacy_set" field's value of the Group entity. +// If the Group object wasn't provided to the builder, the object is fetched from the database. +// An error is returned if the mutation operation is not UpdateOne, or the database query fails. +func (m *GroupMutation) OldRequirePrivacySet(ctx context.Context) (v bool, err error) { + if !m.op.Is(OpUpdateOne) { + return v, errors.New("OldRequirePrivacySet is only allowed on UpdateOne operations") + } + if m.id == nil || m.oldValue == nil { + return v, errors.New("OldRequirePrivacySet requires an ID field in the mutation") + } + oldValue, err := m.oldValue(ctx) + if err != nil { + return v, fmt.Errorf("querying old value for OldRequirePrivacySet: %w", err) + } + return oldValue.RequirePrivacySet, nil +} + +// ResetRequirePrivacySet resets all changes to the "require_privacy_set" field. +func (m *GroupMutation) ResetRequirePrivacySet() { + m.require_privacy_set = nil +} + // SetDefaultMappedModel sets the "default_mapped_model" field. func (m *GroupMutation) SetDefaultMappedModel(s string) { m.default_mapped_model = &s @@ -10082,7 +10156,7 @@ func (m *GroupMutation) Type() string { // order to get all numeric fields that were incremented/decremented, call // AddedFields(). func (m *GroupMutation) Fields() []string { - fields := make([]string, 0, 27) + fields := make([]string, 0, 29) if m.created_at != nil { fields = append(fields, group.FieldCreatedAt) } @@ -10161,6 +10235,12 @@ func (m *GroupMutation) Fields() []string { if m.allow_messages_dispatch != nil { fields = append(fields, group.FieldAllowMessagesDispatch) } + if m.require_oauth_only != nil { + fields = append(fields, group.FieldRequireOauthOnly) + } + if m.require_privacy_set != nil { + fields = append(fields, group.FieldRequirePrivacySet) + } if m.default_mapped_model != nil { fields = append(fields, group.FieldDefaultMappedModel) } @@ -10224,6 +10304,10 @@ func (m *GroupMutation) Field(name string) (ent.Value, bool) { return m.SortOrder() case group.FieldAllowMessagesDispatch: return m.AllowMessagesDispatch() + case group.FieldRequireOauthOnly: + return m.RequireOauthOnly() + case group.FieldRequirePrivacySet: + return m.RequirePrivacySet() case group.FieldDefaultMappedModel: return m.DefaultMappedModel() } @@ -10287,6 +10371,10 @@ func (m *GroupMutation) OldField(ctx context.Context, name string) (ent.Value, e return m.OldSortOrder(ctx) case group.FieldAllowMessagesDispatch: return m.OldAllowMessagesDispatch(ctx) + case group.FieldRequireOauthOnly: + return m.OldRequireOauthOnly(ctx) + case group.FieldRequirePrivacySet: + return m.OldRequirePrivacySet(ctx) case group.FieldDefaultMappedModel: return m.OldDefaultMappedModel(ctx) } @@ -10480,6 +10568,20 @@ func (m *GroupMutation) SetField(name string, value ent.Value) error { } m.SetAllowMessagesDispatch(v) return nil + case group.FieldRequireOauthOnly: + v, ok := value.(bool) + if !ok { + return fmt.Errorf("unexpected type %T for field %s", value, name) + } + m.SetRequireOauthOnly(v) + return nil + case group.FieldRequirePrivacySet: + v, ok := value.(bool) + if !ok { + return fmt.Errorf("unexpected type %T for field %s", value, name) + } + m.SetRequirePrivacySet(v) + return nil case group.FieldDefaultMappedModel: v, ok := value.(string) if !ok { @@ -10818,6 +10920,12 @@ func (m *GroupMutation) ResetField(name string) error { case group.FieldAllowMessagesDispatch: m.ResetAllowMessagesDispatch() return nil + case group.FieldRequireOauthOnly: + m.ResetRequireOauthOnly() + return nil + case group.FieldRequirePrivacySet: + m.ResetRequirePrivacySet() + return nil case group.FieldDefaultMappedModel: m.ResetDefaultMappedModel() return nil diff --git a/backend/ent/runtime/runtime.go b/backend/ent/runtime/runtime.go index 12763eddb8..803b7bc24b 100644 --- a/backend/ent/runtime/runtime.go +++ b/backend/ent/runtime/runtime.go @@ -454,8 +454,16 @@ func init() { groupDescAllowMessagesDispatch := groupFields[22].Descriptor() // group.DefaultAllowMessagesDispatch holds the default value on creation for the allow_messages_dispatch field. group.DefaultAllowMessagesDispatch = groupDescAllowMessagesDispatch.Default.(bool) + // groupDescRequireOauthOnly is the schema descriptor for require_oauth_only field. + groupDescRequireOauthOnly := groupFields[23].Descriptor() + // group.DefaultRequireOauthOnly holds the default value on creation for the require_oauth_only field. + group.DefaultRequireOauthOnly = groupDescRequireOauthOnly.Default.(bool) + // groupDescRequirePrivacySet is the schema descriptor for require_privacy_set field. + groupDescRequirePrivacySet := groupFields[24].Descriptor() + // group.DefaultRequirePrivacySet holds the default value on creation for the require_privacy_set field. + group.DefaultRequirePrivacySet = groupDescRequirePrivacySet.Default.(bool) // groupDescDefaultMappedModel is the schema descriptor for default_mapped_model field. - groupDescDefaultMappedModel := groupFields[23].Descriptor() + groupDescDefaultMappedModel := groupFields[25].Descriptor() // group.DefaultDefaultMappedModel holds the default value on creation for the default_mapped_model field. group.DefaultDefaultMappedModel = groupDescDefaultMappedModel.Default.(string) // group.DefaultMappedModelValidator is a validator for the "default_mapped_model" field. It is called by the builders before save. diff --git a/backend/ent/schema/group.go b/backend/ent/schema/group.go index c791a4e222..0a6aeaecad 100644 --- a/backend/ent/schema/group.go +++ b/backend/ent/schema/group.go @@ -83,6 +83,7 @@ func (Group) Fields() []ent.Field { Nillable(). SchemaType(map[string]string{dialect.Postgres: "decimal(20,8)"}), + // Claude Code 客户端限制 (added by migration 029) field.Bool("claude_code_only"). Default(false). Comment("allow Claude Code client only"), @@ -120,6 +121,12 @@ func (Group) Fields() []ent.Field { field.Bool("allow_messages_dispatch"). Default(false). Comment("是否允许 /v1/messages 调度到此 OpenAI 分组"), + field.Bool("require_oauth_only"). + Default(false). + Comment("仅允许非 apikey 类型账号关联到此分组"), + field.Bool("require_privacy_set"). + Default(false). + Comment("调度时仅允许 privacy 已成功设置的账号"), field.String("default_mapped_model"). MaxLen(100). Default(""). diff --git a/backend/internal/handler/admin/account_handler.go b/backend/internal/handler/admin/account_handler.go index 714f39f487..9a16f39433 100644 --- a/backend/internal/handler/admin/account_handler.go +++ b/backend/internal/handler/admin/account_handler.go @@ -839,6 +839,7 @@ func (h *AccountHandler) refreshSingleAccount(ctx context.Context, account *serv if updateErr != nil { return nil, "", fmt.Errorf("failed to update credentials: %w", updateErr) } + h.adminService.EnsureAntigravityPrivacy(ctx, updatedAccount) return updatedAccount, "missing_project_id_temporary", nil } diff --git a/backend/internal/handler/admin/group_handler.go b/backend/internal/handler/admin/group_handler.go index 6aba4ef606..458ed35d47 100644 --- a/backend/internal/handler/admin/group_handler.go +++ b/backend/internal/handler/admin/group_handler.go @@ -106,6 +106,8 @@ type CreateGroupRequest struct { SupportedModelScopes []string `json:"supported_model_scopes"` // OpenAI Messages 调度配置(仅 openai 平台使用) AllowMessagesDispatch bool `json:"allow_messages_dispatch"` + RequireOAuthOnly bool `json:"require_oauth_only"` + RequirePrivacySet bool `json:"require_privacy_set"` DefaultMappedModel string `json:"default_mapped_model"` // 从指定分组复制账号(创建后自动绑定) CopyAccountsFromGroupIDs []int64 `json:"copy_accounts_from_group_ids"` @@ -138,6 +140,8 @@ type UpdateGroupRequest struct { SupportedModelScopes *[]string `json:"supported_model_scopes"` // OpenAI Messages 调度配置(仅 openai 平台使用) AllowMessagesDispatch *bool `json:"allow_messages_dispatch"` + RequireOAuthOnly *bool `json:"require_oauth_only"` + RequirePrivacySet *bool `json:"require_privacy_set"` DefaultMappedModel *string `json:"default_mapped_model"` // 从指定分组复制账号(同步操作:先清空当前分组的账号绑定,再绑定源分组的账号) CopyAccountsFromGroupIDs []int64 `json:"copy_accounts_from_group_ids"` @@ -250,6 +254,8 @@ func (h *GroupHandler) Create(c *gin.Context) { MCPXMLInject: req.MCPXMLInject, SupportedModelScopes: req.SupportedModelScopes, AllowMessagesDispatch: req.AllowMessagesDispatch, + RequireOAuthOnly: req.RequireOAuthOnly, + RequirePrivacySet: req.RequirePrivacySet, DefaultMappedModel: req.DefaultMappedModel, CopyAccountsFromGroupIDs: req.CopyAccountsFromGroupIDs, }) @@ -298,6 +304,8 @@ func (h *GroupHandler) Update(c *gin.Context) { MCPXMLInject: req.MCPXMLInject, SupportedModelScopes: req.SupportedModelScopes, AllowMessagesDispatch: req.AllowMessagesDispatch, + RequireOAuthOnly: req.RequireOAuthOnly, + RequirePrivacySet: req.RequirePrivacySet, DefaultMappedModel: req.DefaultMappedModel, CopyAccountsFromGroupIDs: req.CopyAccountsFromGroupIDs, }) diff --git a/backend/internal/handler/dto/mappers.go b/backend/internal/handler/dto/mappers.go index f606c3ad70..2eab670e75 100644 --- a/backend/internal/handler/dto/mappers.go +++ b/backend/internal/handler/dto/mappers.go @@ -174,6 +174,8 @@ func groupFromServiceBase(g *service.Group) Group { FallbackGroupID: g.FallbackGroupID, FallbackGroupIDOnInvalidRequest: g.FallbackGroupIDOnInvalidRequest, AllowMessagesDispatch: g.AllowMessagesDispatch, + RequireOAuthOnly: g.RequireOAuthOnly, + RequirePrivacySet: g.RequirePrivacySet, CreatedAt: g.CreatedAt, UpdatedAt: g.UpdatedAt, } diff --git a/backend/internal/handler/dto/types.go b/backend/internal/handler/dto/types.go index b7f97e87ec..82065deb72 100644 --- a/backend/internal/handler/dto/types.go +++ b/backend/internal/handler/dto/types.go @@ -91,6 +91,10 @@ type Group struct { // OpenAI Messages 调度开关(用户侧需要此字段判断是否展示 Claude Code 教程) AllowMessagesDispatch bool `json:"allow_messages_dispatch"` + // 账号过滤控制(仅 OpenAI/Antigravity 平台有效) + RequireOAuthOnly bool `json:"require_oauth_only"` + RequirePrivacySet bool `json:"require_privacy_set"` + CreatedAt time.Time `json:"created_at"` UpdatedAt time.Time `json:"updated_at"` } diff --git a/backend/internal/handler/gemini_v1beta_handler.go b/backend/internal/handler/gemini_v1beta_handler.go index 45b5842fa5..d200c17ce5 100644 --- a/backend/internal/handler/gemini_v1beta_handler.go +++ b/backend/internal/handler/gemini_v1beta_handler.go @@ -121,7 +121,7 @@ func (h *GatewayHandler) GeminiV1BetaGetModel(c *gin.Context) { googleError(c, http.StatusBadGateway, err.Error()) return } - if shouldFallbackGeminiModels(res) { + if shouldFallbackGeminiModel(modelName, res) { c.JSON(http.StatusOK, gemini.FallbackModel(modelName)) return } @@ -682,6 +682,16 @@ func shouldFallbackGeminiModels(res *service.UpstreamHTTPResult) bool { return false } +func shouldFallbackGeminiModel(modelName string, res *service.UpstreamHTTPResult) bool { + if shouldFallbackGeminiModels(res) { + return true + } + if res == nil || res.StatusCode != http.StatusNotFound { + return false + } + return gemini.HasFallbackModel(modelName) +} + // extractGeminiCLISessionHash 从 Gemini CLI 请求中提取会话标识。 // 组合 x-gemini-api-privileged-user-id header 和请求体中的 tmp 目录哈希。 // diff --git a/backend/internal/handler/gemini_v1beta_handler_test.go b/backend/internal/handler/gemini_v1beta_handler_test.go index 82b30ee46e..29d7cc4169 100644 --- a/backend/internal/handler/gemini_v1beta_handler_test.go +++ b/backend/internal/handler/gemini_v1beta_handler_test.go @@ -3,6 +3,7 @@ package handler import ( + "net/http" "testing" "github.com/Wei-Shaw/sub2api/internal/service" @@ -141,3 +142,28 @@ func TestGeminiV1BetaHandler_GetModelAntigravityFallback(t *testing.T) { }) } } + +func TestShouldFallbackGeminiModel_KnownFallbackOn404(t *testing.T) { + t.Parallel() + + res := &service.UpstreamHTTPResult{StatusCode: http.StatusNotFound} + require.True(t, shouldFallbackGeminiModel("gemini-3.1-pro-preview-customtools", res)) +} + +func TestShouldFallbackGeminiModel_UnknownModelOn404(t *testing.T) { + t.Parallel() + + res := &service.UpstreamHTTPResult{StatusCode: http.StatusNotFound} + require.False(t, shouldFallbackGeminiModel("gemini-future-model", res)) +} + +func TestShouldFallbackGeminiModel_DelegatesScopeFallback(t *testing.T) { + t.Parallel() + + res := &service.UpstreamHTTPResult{ + StatusCode: http.StatusForbidden, + Headers: http.Header{"Www-Authenticate": []string{"Bearer error=\"insufficient_scope\""}}, + Body: []byte("insufficient authentication scopes"), + } + require.True(t, shouldFallbackGeminiModel("gemini-future-model", res)) +} diff --git a/backend/internal/pkg/antigravity/oauth.go b/backend/internal/pkg/antigravity/oauth.go index 8a8bed92d9..7c963d9e51 100644 --- a/backend/internal/pkg/antigravity/oauth.go +++ b/backend/internal/pkg/antigravity/oauth.go @@ -50,7 +50,7 @@ const ( ) // defaultUserAgentVersion 可通过环境变量 ANTIGRAVITY_USER_AGENT_VERSION 配置,默认 1.20.5 -var defaultUserAgentVersion = "1.20.5" +var defaultUserAgentVersion = "1.21.9" // defaultClientSecret 可通过环境变量 ANTIGRAVITY_OAUTH_CLIENT_SECRET 配置 var defaultClientSecret = "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" diff --git a/backend/internal/pkg/antigravity/oauth_test.go b/backend/internal/pkg/antigravity/oauth_test.go index 3a093fe657..9850af17e1 100644 --- a/backend/internal/pkg/antigravity/oauth_test.go +++ b/backend/internal/pkg/antigravity/oauth_test.go @@ -690,7 +690,7 @@ func TestConstants_值正确(t *testing.T) { if RedirectURI != "http://localhost:8085/callback" { t.Errorf("RedirectURI 不匹配: got %s", RedirectURI) } - if GetUserAgent() != "antigravity/1.20.5 windows/amd64" { + if GetUserAgent() != "antigravity/1.21.9 windows/amd64" { t.Errorf("UserAgent 不匹配: got %s", GetUserAgent()) } if SessionTTL != 30*time.Minute { diff --git a/backend/internal/pkg/gemini/models.go b/backend/internal/pkg/gemini/models.go index 882d2ebdd8..fac79d1873 100644 --- a/backend/internal/pkg/gemini/models.go +++ b/backend/internal/pkg/gemini/models.go @@ -2,6 +2,8 @@ // It is used when upstream model listing is unavailable (e.g. OAuth token missing AI Studio scopes). package gemini +import "strings" + type Model struct { Name string `json:"name"` DisplayName string `json:"displayName,omitempty"` @@ -23,10 +25,27 @@ func DefaultModels() []Model { {Name: "models/gemini-3-flash-preview", SupportedGenerationMethods: methods}, {Name: "models/gemini-3-pro-preview", SupportedGenerationMethods: methods}, {Name: "models/gemini-3.1-pro-preview", SupportedGenerationMethods: methods}, + {Name: "models/gemini-3.1-pro-preview-customtools", SupportedGenerationMethods: methods}, {Name: "models/gemini-3.1-flash-image", SupportedGenerationMethods: methods}, } } +func HasFallbackModel(model string) bool { + trimmed := strings.TrimSpace(model) + if trimmed == "" { + return false + } + if !strings.HasPrefix(trimmed, "models/") { + trimmed = "models/" + trimmed + } + for _, model := range DefaultModels() { + if model.Name == trimmed { + return true + } + } + return false +} + func FallbackModelsList() ModelsListResponse { return ModelsListResponse{Models: DefaultModels()} } diff --git a/backend/internal/pkg/gemini/models_test.go b/backend/internal/pkg/gemini/models_test.go index b80047fb73..1d20c0e62d 100644 --- a/backend/internal/pkg/gemini/models_test.go +++ b/backend/internal/pkg/gemini/models_test.go @@ -2,7 +2,7 @@ package gemini import "testing" -func TestDefaultModels_ContainsImageModels(t *testing.T) { +func TestDefaultModels_ContainsFallbackCatalogModels(t *testing.T) { t.Parallel() models := DefaultModels() @@ -13,6 +13,7 @@ func TestDefaultModels_ContainsImageModels(t *testing.T) { required := []string{ "models/gemini-2.5-flash-image", + "models/gemini-3.1-pro-preview-customtools", "models/gemini-3.1-flash-image", } @@ -26,3 +27,17 @@ func TestDefaultModels_ContainsImageModels(t *testing.T) { } } } + +func TestHasFallbackModel_RecognizesCustomtoolsModel(t *testing.T) { + t.Parallel() + + if !HasFallbackModel("gemini-3.1-pro-preview-customtools") { + t.Fatalf("expected customtools model to exist in fallback catalog") + } + if !HasFallbackModel("models/gemini-3.1-pro-preview-customtools") { + t.Fatalf("expected prefixed customtools model to exist in fallback catalog") + } + if HasFallbackModel("gemini-unknown") { + t.Fatalf("did not expect unknown model to exist in fallback catalog") + } +} diff --git a/backend/internal/repository/api_key_repo.go b/backend/internal/repository/api_key_repo.go index a1f83b8398..b3b12e8113 100644 --- a/backend/internal/repository/api_key_repo.go +++ b/backend/internal/repository/api_key_repo.go @@ -651,6 +651,8 @@ func groupEntityToService(g *dbent.Group) *service.Group { SupportedModelScopes: g.SupportedModelScopes, SortOrder: g.SortOrder, AllowMessagesDispatch: g.AllowMessagesDispatch, + RequireOAuthOnly: g.RequireOauthOnly, + RequirePrivacySet: g.RequirePrivacySet, DefaultMappedModel: g.DefaultMappedModel, CreatedAt: g.CreatedAt, UpdatedAt: g.UpdatedAt, diff --git a/backend/internal/repository/group_repo.go b/backend/internal/repository/group_repo.go index 9cae50e672..a075b586c4 100644 --- a/backend/internal/repository/group_repo.go +++ b/backend/internal/repository/group_repo.go @@ -56,6 +56,8 @@ func (r *groupRepository) Create(ctx context.Context, groupIn *service.Group) er SetModelRoutingEnabled(groupIn.ModelRoutingEnabled). SetMcpXMLInject(groupIn.MCPXMLInject). SetAllowMessagesDispatch(groupIn.AllowMessagesDispatch). + SetRequireOauthOnly(groupIn.RequireOAuthOnly). + SetRequirePrivacySet(groupIn.RequirePrivacySet). SetDefaultMappedModel(groupIn.DefaultMappedModel) // 设置模型路由配置 @@ -120,6 +122,8 @@ func (r *groupRepository) Update(ctx context.Context, groupIn *service.Group) er SetModelRoutingEnabled(groupIn.ModelRoutingEnabled). SetMcpXMLInject(groupIn.MCPXMLInject). SetAllowMessagesDispatch(groupIn.AllowMessagesDispatch). + SetRequireOauthOnly(groupIn.RequireOAuthOnly). + SetRequirePrivacySet(groupIn.RequirePrivacySet). SetDefaultMappedModel(groupIn.DefaultMappedModel) // 显式处理可空字段:nil 需要 clear,非 nil 需要 set。 diff --git a/backend/internal/server/api_contract_test.go b/backend/internal/server/api_contract_test.go index d7aaabfdfd..43cfdc51db 100644 --- a/backend/internal/server/api_contract_test.go +++ b/backend/internal/server/api_contract_test.go @@ -204,10 +204,12 @@ func TestAPIContracts(t *testing.T) { "image_price_1k": null, "image_price_2k": null, "image_price_4k": null, - "claude_code_only": false, - "allow_messages_dispatch": false, - "fallback_group_id": null, - "fallback_group_id_on_invalid_request": null, + "claude_code_only": false, + "allow_messages_dispatch": false, + "fallback_group_id": null, + "fallback_group_id_on_invalid_request": null, + "require_oauth_only": false, + "require_privacy_set": false, "created_at": "2025-01-02T03:04:05Z", "updated_at": "2025-01-02T03:04:05Z" } diff --git a/backend/internal/service/account.go b/backend/internal/service/account.go index a1449ffd89..512195e334 100644 --- a/backend/internal/service/account.go +++ b/backend/internal/service/account.go @@ -141,6 +141,21 @@ func (a *Account) IsOAuth() bool { return a.Type == AccountTypeOAuth || a.Type == AccountTypeSetupToken } +// IsPrivacySet 检查账号的 privacy 是否已成功设置。 +// OpenAI: privacy_mode == "training_off" +// Antigravity: privacy_mode == "privacy_set" +// 其他平台: 无 privacy 概念,始终返回 true +func (a *Account) IsPrivacySet() bool { + switch a.Platform { + case PlatformOpenAI: + return a.getExtraString("privacy_mode") == PrivacyModeTrainingOff + case PlatformAntigravity: + return a.getExtraString("privacy_mode") == AntigravityPrivacySet + default: + return true + } +} + func (a *Account) IsGemini() bool { return a.Platform == PlatformGemini } @@ -500,6 +515,45 @@ func ensureAntigravityDefaultPassthroughs(mapping map[string]string, models []st } } +func normalizeRequestedModelForLookup(platform, requestedModel string) string { + trimmed := strings.TrimSpace(requestedModel) + if trimmed == "" { + return "" + } + if platform != PlatformGemini && platform != PlatformAntigravity { + return trimmed + } + if trimmed == "gemini-3.1-pro-preview-customtools" { + return "gemini-3.1-pro-preview" + } + return trimmed +} + +func mappingSupportsRequestedModel(mapping map[string]string, requestedModel string) bool { + if requestedModel == "" { + return false + } + if _, exists := mapping[requestedModel]; exists { + return true + } + for pattern := range mapping { + if matchWildcard(pattern, requestedModel) { + return true + } + } + return false +} + +func resolveRequestedModelInMapping(mapping map[string]string, requestedModel string) (mappedModel string, matched bool) { + if requestedModel == "" { + return "", false + } + if mappedModel, exists := mapping[requestedModel]; exists { + return mappedModel, true + } + return matchWildcardMappingResult(mapping, requestedModel) +} + // IsModelSupported 检查模型是否在 model_mapping 中(支持通配符) // 如果未配置 mapping,返回 true(允许所有模型) func (a *Account) IsModelSupported(requestedModel string) bool { @@ -507,17 +561,11 @@ func (a *Account) IsModelSupported(requestedModel string) bool { if len(mapping) == 0 { return true // 无映射 = 允许所有 } - // 精确匹配 - if _, exists := mapping[requestedModel]; exists { + if mappingSupportsRequestedModel(mapping, requestedModel) { return true } - // 通配符匹配 - for pattern := range mapping { - if matchWildcard(pattern, requestedModel) { - return true - } - } - return false + normalized := normalizeRequestedModelForLookup(a.Platform, requestedModel) + return normalized != requestedModel && mappingSupportsRequestedModel(mapping, normalized) } // GetMappedModel 获取映射后的模型名(支持通配符,最长优先匹配) @@ -534,12 +582,16 @@ func (a *Account) ResolveMappedModel(requestedModel string) (mappedModel string, if len(mapping) == 0 { return requestedModel, false } - // 精确匹配优先 - if mappedModel, exists := mapping[requestedModel]; exists { + if mappedModel, matched := resolveRequestedModelInMapping(mapping, requestedModel); matched { return mappedModel, true } - // 通配符匹配(最长优先) - return matchWildcardMappingResult(mapping, requestedModel) + normalized := normalizeRequestedModelForLookup(a.Platform, requestedModel) + if normalized != requestedModel { + if mappedModel, matched := resolveRequestedModelInMapping(mapping, normalized); matched { + return mappedModel, true + } + } + return requestedModel, false } func (a *Account) GetBaseURL() string { @@ -1727,22 +1779,47 @@ func (a *Account) GetRPMStrategy() string { } // GetRPMStickyBuffer 获取 RPM 粘性缓冲数量 -// tiered 模式下的黄区大小,默认为 base_rpm 的 20%(至少 1) +// Cache-driven: buffer = concurrency + maxSessions(覆盖幽灵窗口 + 稳态会话需求) +// floor = baseRPM / 5(向后兼容 maxSessions=0 且 concurrency=0 场景) func (a *Account) GetRPMStickyBuffer() int { if a.Extra == nil { return 0 } + + // 手动 override 最高优先级 if v, ok := a.Extra["rpm_sticky_buffer"]; ok { val := parseExtraInt(v) if val > 0 { return val } } + base := a.GetBaseRPM() - buffer := base / 5 - if buffer < 1 && base > 0 { - buffer = 1 + if base <= 0 { + return 0 } + + // Cache-driven buffer = concurrency + maxSessions + conc := a.Concurrency + if conc < 0 { + conc = 0 + } + sess := a.GetMaxSessions() + if sess < 0 { + sess = 0 + } + + buffer := conc + sess + + // floor: 向后兼容 + floor := base / 5 + if floor < 1 { + floor = 1 + } + if buffer < floor { + buffer = floor + } + return buffer } diff --git a/backend/internal/service/account_rpm_test.go b/backend/internal/service/account_rpm_test.go index 9d91f3e0ca..40298263a8 100644 --- a/backend/internal/service/account_rpm_test.go +++ b/backend/internal/service/account_rpm_test.go @@ -90,28 +90,47 @@ func TestCheckRPMSchedulability(t *testing.T) { func TestGetRPMStickyBuffer(t *testing.T) { tests := []struct { - name string - extra map[string]any - expected int + name string + concurrency int + extra map[string]any + expected int }{ - {"nil extra", nil, 0}, - {"no keys", map[string]any{}, 0}, - {"base_rpm=0", map[string]any{"base_rpm": 0}, 0}, - {"base_rpm=1 min buffer 1", map[string]any{"base_rpm": 1}, 1}, - {"base_rpm=4 min buffer 1", map[string]any{"base_rpm": 4}, 1}, - {"base_rpm=5 buffer 1", map[string]any{"base_rpm": 5}, 1}, - {"base_rpm=10 buffer 2", map[string]any{"base_rpm": 10}, 2}, - {"base_rpm=15 buffer 3", map[string]any{"base_rpm": 15}, 3}, - {"base_rpm=100 buffer 20", map[string]any{"base_rpm": 100}, 20}, - {"custom buffer=5", map[string]any{"base_rpm": 10, "rpm_sticky_buffer": 5}, 5}, - {"custom buffer=0 fallback to default", map[string]any{"base_rpm": 10, "rpm_sticky_buffer": 0}, 2}, - {"custom buffer negative fallback", map[string]any{"base_rpm": 10, "rpm_sticky_buffer": -1}, 2}, - {"custom buffer with float", map[string]any{"base_rpm": 10, "rpm_sticky_buffer": float64(7)}, 7}, - {"json.Number base_rpm", map[string]any{"base_rpm": json.Number("10")}, 2}, + // 基础退化 + {"nil extra", 0, nil, 0}, + {"no keys", 0, map[string]any{}, 0}, + {"base_rpm=0", 0, map[string]any{"base_rpm": 0}, 0}, + + // 新公式: concurrency + maxSessions, floor = base/5 + {"conc=3 sess=10 → 13", 3, map[string]any{"base_rpm": 15, "max_sessions": 10}, 13}, + {"conc=2 sess=5 → 7", 2, map[string]any{"base_rpm": 10, "max_sessions": 5}, 7}, + {"conc=3 sess=15 → 18", 3, map[string]any{"base_rpm": 30, "max_sessions": 15}, 18}, + + // floor 生效 (conc+sess < base/5) + {"conc=0 sess=0 base=15 → floor 3", 0, map[string]any{"base_rpm": 15}, 3}, + {"conc=0 sess=0 base=10 → floor 2", 0, map[string]any{"base_rpm": 10}, 2}, + {"conc=0 sess=0 base=1 → floor 1", 0, map[string]any{"base_rpm": 1}, 1}, + {"conc=0 sess=0 base=4 → floor 1", 0, map[string]any{"base_rpm": 4}, 1}, + {"conc=1 sess=0 base=15 → floor 3", 1, map[string]any{"base_rpm": 15}, 3}, + + // 手动 override + {"custom buffer=5", 3, map[string]any{"base_rpm": 10, "rpm_sticky_buffer": 5, "max_sessions": 10}, 5}, + {"custom buffer=0 fallback", 3, map[string]any{"base_rpm": 10, "rpm_sticky_buffer": 0, "max_sessions": 10}, 13}, + {"custom buffer negative fallback", 3, map[string]any{"base_rpm": 10, "rpm_sticky_buffer": -1, "max_sessions": 10}, 13}, + {"custom buffer with float", 3, map[string]any{"base_rpm": 10, "rpm_sticky_buffer": float64(7)}, 7}, + + // 负值 clamp + {"negative concurrency clamped", -5, map[string]any{"base_rpm": 15, "max_sessions": 10}, 10}, + {"negative maxSessions clamped", 3, map[string]any{"base_rpm": 15, "max_sessions": -5}, 3}, + + // 高并发低会话 + {"conc=10 sess=5 → 15", 10, map[string]any{"base_rpm": 10, "max_sessions": 5}, 15}, + + // json.Number + {"json.Number base_rpm", 3, map[string]any{"base_rpm": json.Number("10"), "max_sessions": json.Number("5")}, 8}, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - a := &Account{Extra: tt.extra} + a := &Account{Concurrency: tt.concurrency, Extra: tt.extra} if got := a.GetRPMStickyBuffer(); got != tt.expected { t.Errorf("GetRPMStickyBuffer() = %d, want %d", got, tt.expected) } diff --git a/backend/internal/service/account_service.go b/backend/internal/service/account_service.go index 30b774f3a8..3189a7290f 100644 --- a/backend/internal/service/account_service.go +++ b/backend/internal/service/account_service.go @@ -173,6 +173,19 @@ func (s *AccountService) Create(ctx context.Context, req CreateAccountRequest) ( return nil, fmt.Errorf("create account: %w", err) } + // require_oauth_only 检查:apikey 类型账号不可加入限制分组 + if account.Type == AccountTypeAPIKey && len(req.GroupIDs) > 0 { + for _, gid := range req.GroupIDs { + g, err := s.groupRepo.GetByID(ctx, gid) + if err != nil { + return nil, err + } + if g.RequireOAuthOnly && (g.Platform == PlatformOpenAI || g.Platform == PlatformAntigravity || g.Platform == PlatformAnthropic || g.Platform == PlatformGemini) { + return nil, fmt.Errorf("分组 [%s] 仅允许 OAuth 账号,apikey 类型账号无法加入", g.Name) + } + } + } + // 绑定分组 if len(req.GroupIDs) > 0 { if err := s.accountRepo.BindGroups(ctx, account.ID, req.GroupIDs); err != nil { @@ -276,6 +289,19 @@ func (s *AccountService) Update(ctx context.Context, id int64, req UpdateAccount return nil, fmt.Errorf("update account: %w", err) } + // require_oauth_only 检查 + if account.Type == AccountTypeAPIKey && req.GroupIDs != nil { + for _, gid := range *req.GroupIDs { + g, err := s.groupRepo.GetByID(ctx, gid) + if err != nil { + return nil, err + } + if g.RequireOAuthOnly && (g.Platform == PlatformOpenAI || g.Platform == PlatformAntigravity || g.Platform == PlatformAnthropic || g.Platform == PlatformGemini) { + return nil, fmt.Errorf("分组 [%s] 仅允许 OAuth 账号,apikey 类型账号无法加入", g.Name) + } + } + } + // 绑定分组 if req.GroupIDs != nil { if err := s.accountRepo.BindGroups(ctx, account.ID, *req.GroupIDs); err != nil { diff --git a/backend/internal/service/account_test_service.go b/backend/internal/service/account_test_service.go index 1cbce1ea0f..55865945c6 100644 --- a/backend/internal/service/account_test_service.go +++ b/backend/internal/service/account_test_service.go @@ -531,6 +531,11 @@ func (s *AccountTestService) testOpenAIAccountConnection(c *gin.Context, account account.RateLimitResetAt = resetAt } } + // 401 Unauthorized: 标记账号为永久错误 + if resp.StatusCode == http.StatusUnauthorized && s.accountRepo != nil { + errMsg := fmt.Sprintf("Authentication failed (401): %s", string(body)) + _ = s.accountRepo.SetError(ctx, account.ID, errMsg) + } return s.sendErrorAndEnd(c, fmt.Sprintf("API returned %d: %s", resp.StatusCode, string(body))) } diff --git a/backend/internal/service/account_wildcard_test.go b/backend/internal/service/account_wildcard_test.go index 0d7ffffa8a..d903b940a5 100644 --- a/backend/internal/service/account_wildcard_test.go +++ b/backend/internal/service/account_wildcard_test.go @@ -133,6 +133,7 @@ func TestMatchWildcardMappingResult(t *testing.T) { func TestAccountIsModelSupported(t *testing.T) { tests := []struct { name string + platform string credentials map[string]any requestedModel string expected bool @@ -184,6 +185,17 @@ func TestAccountIsModelSupported(t *testing.T) { requestedModel: "claude-opus-4-5-thinking", expected: true, }, + { + name: "gemini customtools alias matches normalized mapping", + platform: PlatformGemini, + credentials: map[string]any{ + "model_mapping": map[string]any{ + "gemini-3.1-pro-preview": "gemini-3.1-pro-preview", + }, + }, + requestedModel: "gemini-3.1-pro-preview-customtools", + expected: true, + }, { name: "wildcard match not supported", credentials: map[string]any{ @@ -199,6 +211,7 @@ func TestAccountIsModelSupported(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { account := &Account{ + Platform: tt.platform, Credentials: tt.credentials, } result := account.IsModelSupported(tt.requestedModel) @@ -212,6 +225,7 @@ func TestAccountIsModelSupported(t *testing.T) { func TestAccountGetMappedModel(t *testing.T) { tests := []struct { name string + platform string credentials map[string]any requestedModel string expected string @@ -223,6 +237,13 @@ func TestAccountGetMappedModel(t *testing.T) { requestedModel: "claude-sonnet-4-5", expected: "claude-sonnet-4-5", }, + { + name: "no mapping preserves gemini customtools model", + platform: PlatformGemini, + credentials: nil, + requestedModel: "gemini-3.1-pro-preview-customtools", + expected: "gemini-3.1-pro-preview-customtools", + }, // 精确匹配 { @@ -250,6 +271,29 @@ func TestAccountGetMappedModel(t *testing.T) { }, // 无匹配返回原始模型 + { + name: "gemini customtools alias resolves through normalized mapping", + platform: PlatformGemini, + credentials: map[string]any{ + "model_mapping": map[string]any{ + "gemini-3.1-pro-preview": "gemini-3.1-pro-preview", + }, + }, + requestedModel: "gemini-3.1-pro-preview-customtools", + expected: "gemini-3.1-pro-preview", + }, + { + name: "gemini customtools exact mapping wins over normalized fallback", + platform: PlatformGemini, + credentials: map[string]any{ + "model_mapping": map[string]any{ + "gemini-3.1-pro-preview": "gemini-3.1-pro-preview", + "gemini-3.1-pro-preview-customtools": "gemini-3.1-pro-preview-customtools", + }, + }, + requestedModel: "gemini-3.1-pro-preview-customtools", + expected: "gemini-3.1-pro-preview-customtools", + }, { name: "no match returns original", credentials: map[string]any{ @@ -265,6 +309,7 @@ func TestAccountGetMappedModel(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { account := &Account{ + Platform: tt.platform, Credentials: tt.credentials, } result := account.GetMappedModel(tt.requestedModel) @@ -278,6 +323,7 @@ func TestAccountGetMappedModel(t *testing.T) { func TestAccountResolveMappedModel(t *testing.T) { tests := []struct { name string + platform string credentials map[string]any requestedModel string expectedModel string @@ -312,6 +358,31 @@ func TestAccountResolveMappedModel(t *testing.T) { expectedModel: "gpt-5.4", expectedMatch: true, }, + { + name: "gemini customtools alias reports normalized match", + platform: PlatformGemini, + credentials: map[string]any{ + "model_mapping": map[string]any{ + "gemini-3.1-pro-preview": "gemini-3.1-pro-preview", + }, + }, + requestedModel: "gemini-3.1-pro-preview-customtools", + expectedModel: "gemini-3.1-pro-preview", + expectedMatch: true, + }, + { + name: "gemini customtools exact mapping reports exact match", + platform: PlatformGemini, + credentials: map[string]any{ + "model_mapping": map[string]any{ + "gemini-3.1-pro-preview": "gemini-3.1-pro-preview", + "gemini-3.1-pro-preview-customtools": "gemini-3.1-pro-preview-customtools", + }, + }, + requestedModel: "gemini-3.1-pro-preview-customtools", + expectedModel: "gemini-3.1-pro-preview-customtools", + expectedMatch: true, + }, { name: "missing mapping reports unmatched", credentials: map[string]any{ @@ -328,6 +399,7 @@ func TestAccountResolveMappedModel(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { account := &Account{ + Platform: tt.platform, Credentials: tt.credentials, } mappedModel, matched := account.ResolveMappedModel(tt.requestedModel) diff --git a/backend/internal/service/admin_service.go b/backend/internal/service/admin_service.go index 80d8064696..8032f8717a 100644 --- a/backend/internal/service/admin_service.go +++ b/backend/internal/service/admin_service.go @@ -5,6 +5,7 @@ import ( "errors" "fmt" "io" + "log/slog" "net/http" "strings" "time" @@ -153,6 +154,8 @@ type CreateGroupInput struct { // OpenAI Messages 调度配置(仅 openai 平台使用) AllowMessagesDispatch bool DefaultMappedModel string + RequireOAuthOnly bool + RequirePrivacySet bool // 从指定分组复制账号(创建分组后在同一事务内绑定) CopyAccountsFromGroupIDs []int64 } @@ -185,6 +188,8 @@ type UpdateGroupInput struct { // OpenAI Messages 调度配置(仅 openai 平台使用) AllowMessagesDispatch *bool DefaultMappedModel *string + RequireOAuthOnly *bool + RequirePrivacySet *bool // 从指定分组复制账号(同步操作:先清空当前分组的账号绑定,再绑定源分组的账号) CopyAccountsFromGroupIDs []int64 } @@ -900,12 +905,35 @@ func (s *adminServiceImpl) CreateGroup(ctx context.Context, input *CreateGroupIn MCPXMLInject: mcpXMLInject, SupportedModelScopes: input.SupportedModelScopes, AllowMessagesDispatch: input.AllowMessagesDispatch, + RequireOAuthOnly: input.RequireOAuthOnly, + RequirePrivacySet: input.RequirePrivacySet, DefaultMappedModel: input.DefaultMappedModel, } if err := s.groupRepo.Create(ctx, group); err != nil { return nil, err } + // require_oauth_only: 过滤掉 apikey 类型账号 + if group.RequireOAuthOnly && (group.Platform == PlatformOpenAI || group.Platform == PlatformAntigravity || group.Platform == PlatformAnthropic || group.Platform == PlatformGemini) && len(accountIDsToCopy) > 0 { + accounts, err := s.accountRepo.GetByIDs(ctx, accountIDsToCopy) + if err != nil { + return nil, fmt.Errorf("failed to fetch accounts for oauth filter: %w", err) + } + oauthIDs := make(map[int64]struct{}, len(accounts)) + for _, acc := range accounts { + if acc.Type != AccountTypeAPIKey { + oauthIDs[acc.ID] = struct{}{} + } + } + var filtered []int64 + for _, aid := range accountIDsToCopy { + if _, ok := oauthIDs[aid]; ok { + filtered = append(filtered, aid) + } + } + accountIDsToCopy = filtered + } + // 如果有需要复制的账号,绑定到新分组 if len(accountIDsToCopy) > 0 { if err := s.groupRepo.BindAccountsToGroup(ctx, group.ID, accountIDsToCopy); err != nil { @@ -1098,6 +1126,12 @@ func (s *adminServiceImpl) UpdateGroup(ctx context.Context, id int64, input *Upd if input.AllowMessagesDispatch != nil { group.AllowMessagesDispatch = *input.AllowMessagesDispatch } + if input.RequireOAuthOnly != nil { + group.RequireOAuthOnly = *input.RequireOAuthOnly + } + if input.RequirePrivacySet != nil { + group.RequirePrivacySet = *input.RequirePrivacySet + } if input.DefaultMappedModel != nil { group.DefaultMappedModel = *input.DefaultMappedModel } @@ -1145,6 +1179,27 @@ func (s *adminServiceImpl) UpdateGroup(ctx context.Context, id int64, input *Upd return nil, fmt.Errorf("failed to clear existing account bindings: %w", err) } + // require_oauth_only: 过滤掉 apikey 类型账号 + if group.RequireOAuthOnly && (group.Platform == PlatformOpenAI || group.Platform == PlatformAntigravity || group.Platform == PlatformAnthropic || group.Platform == PlatformGemini) && len(accountIDsToCopy) > 0 { + accounts, err := s.accountRepo.GetByIDs(ctx, accountIDsToCopy) + if err != nil { + return nil, fmt.Errorf("failed to fetch accounts for oauth filter: %w", err) + } + oauthIDs := make(map[int64]struct{}, len(accounts)) + for _, acc := range accounts { + if acc.Type != AccountTypeAPIKey { + oauthIDs[acc.ID] = struct{}{} + } + } + var filtered []int64 + for _, aid := range accountIDsToCopy { + if _, ok := oauthIDs[aid]; ok { + filtered = append(filtered, aid) + } + } + accountIDsToCopy = filtered + } + // 再绑定源分组的账号 if len(accountIDsToCopy) > 0 { if err := s.groupRepo.BindAccountsToGroup(ctx, id, accountIDsToCopy); err != nil { @@ -1507,6 +1562,31 @@ func (s *adminServiceImpl) CreateAccount(ctx context.Context, input *CreateAccou } } + // OAuth 账号:创建后异步设置隐私。 + // 使用 Ensure(幂等)而非 Force:新建账号 Extra 为空时效果相同,但更安全。 + if account.Type == AccountTypeOAuth { + switch account.Platform { + case PlatformOpenAI: + go func() { + defer func() { + if r := recover(); r != nil { + slog.Error("create_account_openai_privacy_panic", "account_id", account.ID, "recover", r) + } + }() + s.EnsureOpenAIPrivacy(context.Background(), account) + }() + case PlatformAntigravity: + go func() { + defer func() { + if r := recover(); r != nil { + slog.Error("create_account_antigravity_privacy_panic", "account_id", account.ID, "recover", r) + } + }() + s.EnsureAntigravityPrivacy(context.Background(), account) + }() + } + } + return account, nil } @@ -2625,16 +2705,14 @@ func (s *adminServiceImpl) ForceOpenAIPrivacy(ctx context.Context, account *Acco } // EnsureAntigravityPrivacy 检查 Antigravity OAuth 账号隐私状态。 -// 如果 Extra["privacy_mode"] 已存在(无论成功或失败),直接跳过。 -// 仅对从未设置过隐私的账号执行 setUserSettings + fetchUserInfo 流程。 -// 用户可通过前端 ForceAntigravityPrivacy(SetPrivacy 按钮)强制重新设置。 +// 仅当 privacy_mode 已成功设置("privacy_set")时跳过; +// 未设置或之前失败("privacy_set_failed")均会重试。 func (s *adminServiceImpl) EnsureAntigravityPrivacy(ctx context.Context, account *Account) string { if account.Platform != PlatformAntigravity || account.Type != AccountTypeOAuth { return "" } - // 已设置过则跳过(无论成功或失败),用户可通过 Force 手动重试 if account.Extra != nil { - if existing, ok := account.Extra["privacy_mode"].(string); ok && existing != "" { + if existing, ok := account.Extra["privacy_mode"].(string); ok && existing == AntigravityPrivacySet { return existing } } diff --git a/backend/internal/service/antigravity_model_mapping_test.go b/backend/internal/service/antigravity_model_mapping_test.go index efef2a4b90..a151e1a139 100644 --- a/backend/internal/service/antigravity_model_mapping_test.go +++ b/backend/internal/service/antigravity_model_mapping_test.go @@ -268,6 +268,12 @@ func TestMapAntigravityModel_WildcardTargetEqualsRequest(t *testing.T) { requestedModel: "gemini-2.5-flash", expected: "gemini-2.5-flash", }, + { + name: "customtools alias falls back to normalized preview mapping", + modelMapping: map[string]any{"gemini-3.1-pro-preview": "gemini-3.1-pro-high"}, + requestedModel: "gemini-3.1-pro-preview-customtools", + expected: "gemini-3.1-pro-high", + }, } for _, tt := range tests { diff --git a/backend/internal/service/antigravity_oauth_service.go b/backend/internal/service/antigravity_oauth_service.go index a300d8982b..3a4600db0e 100644 --- a/backend/internal/service/antigravity_oauth_service.go +++ b/backend/internal/service/antigravity_oauth_service.go @@ -91,6 +91,7 @@ type AntigravityTokenInfo struct { ProjectID string `json:"project_id,omitempty"` ProjectIDMissing bool `json:"-"` PlanType string `json:"-"` + PrivacyMode string `json:"-"` } // ExchangeCode 用 authorization code 交换 token @@ -159,6 +160,9 @@ func (s *AntigravityOAuthService) ExchangeCode(ctx context.Context, input *Antig } } + // 令牌刚获取,立即设置隐私(不依赖后续账号创建流程) + result.PrivacyMode = setAntigravityPrivacy(ctx, result.AccessToken, result.ProjectID, proxyURL) + return result, nil } @@ -248,6 +252,9 @@ func (s *AntigravityOAuthService) ValidateRefreshToken(ctx context.Context, refr } } + // 令牌刚获取,立即设置隐私 + tokenInfo.PrivacyMode = setAntigravityPrivacy(ctx, tokenInfo.AccessToken, tokenInfo.ProjectID, proxyURL) + return tokenInfo, nil } diff --git a/backend/internal/service/billing_service.go b/backend/internal/service/billing_service.go index 77ee172c69..763abadbfc 100644 --- a/backend/internal/service/billing_service.go +++ b/backend/internal/service/billing_service.go @@ -422,7 +422,7 @@ type CostInput struct { RateMultiplier float64 ServiceTier string // "priority","flex","" 等 Resolver *ModelPricingResolver // 定价解析器 - Resolved *ResolvedPricing // 可选:已解析的定价,跳过重复 Resolve + Resolved *ResolvedPricing // 可选:预解析的定价结果(避免重复 Resolve 调用) } // CalculateCostUnified 统一计费入口,支持三种计费模式。 @@ -433,6 +433,7 @@ func (s *BillingService) CalculateCostUnified(input CostInput) (*CostBreakdown, return s.calculateCostInternal(input.Model, input.Tokens, input.RateMultiplier, input.ServiceTier, nil) } + // 优先使用预解析结果,避免重复 Resolve 调用 resolved := input.Resolved if resolved == nil { resolved = input.Resolver.Resolve(input.Ctx, PricingInput{ diff --git a/backend/internal/service/channel_service.go b/backend/internal/service/channel_service.go index 9667cb9844..79678c4b80 100644 --- a/backend/internal/service/channel_service.go +++ b/backend/internal/service/channel_service.go @@ -281,6 +281,7 @@ func (s *ChannelService) fetchChannelData(ctx context.Context) ([]Channel, map[i return nil, nil, fmt.Errorf("list all channels: %w", err) } + // 收集所有 groupID,批量查询 platform var allGroupIDs []int64 for i := range channels { allGroupIDs = append(allGroupIDs, channels[i].GroupIDs...) @@ -308,6 +309,7 @@ func populateChannelCache(channels []Channel, groupPlatforms map[int64]string) * for i := range channels { ch := &channels[i] cache.byID[ch.ID] = ch + for _, gid := range ch.GroupIDs { cache.channelByGroupID[gid] = ch platform := groupPlatforms[gid] @@ -315,6 +317,9 @@ func populateChannelCache(channels []Channel, groupPlatforms map[int64]string) * expandMappingToCache(cache, ch, gid, platform) } } + + // 通配符条目保持配置顺序(最先匹配到优先) + return cache } @@ -479,10 +484,7 @@ func (s *ChannelService) ResolveChannelMapping(ctx context.Context, groupID int6 // 返回 true 表示模型被限制(不在允许列表中)。 // 如果渠道未启用模型限制或分组无渠道关联,返回 false。 func (s *ChannelService) IsModelRestricted(ctx context.Context, groupID int64, model string) bool { - lk, err := s.lookupGroupChannel(ctx, groupID) - if err != nil { - slog.Warn("failed to load channel cache for model restriction check", "group_id", groupID, "error", err) - } + lk, _ := s.lookupGroupChannel(ctx, groupID) if lk == nil { return false } @@ -798,6 +800,7 @@ func (s *ChannelService) invalidateAuthCacheForGroups(ctx context.Context, group // Delete 删除渠道 func (s *ChannelService) Delete(ctx context.Context, id int64) error { + // 先获取关联分组用于失效缓存 groupIDs, err := s.repo.GetGroupIDs(ctx, id) if err != nil { slog.Warn("failed to get group IDs before delete", "channel_id", id, "error", err) diff --git a/backend/internal/service/gateway_multiplatform_test.go b/backend/internal/service/gateway_multiplatform_test.go index 5bd0ef16f2..728328373c 100644 --- a/backend/internal/service/gateway_multiplatform_test.go +++ b/backend/internal/service/gateway_multiplatform_test.go @@ -3139,7 +3139,7 @@ func TestGatewayService_GroupResolution_ReusesContextGroup(t *testing.T) { account, err := svc.SelectAccountForModelWithExclusions(ctx, &groupID, "", "claude-3-5-sonnet-20241022", nil) require.NoError(t, err) require.NotNil(t, account) - require.Equal(t, 0, groupRepo.getByIDCalls) + require.Equal(t, 1, groupRepo.getByIDCalls) // +1 for require_privacy_set check require.Equal(t, 0, groupRepo.getByIDLiteCalls) } @@ -3182,7 +3182,7 @@ func TestGatewayService_GroupResolution_IgnoresInvalidContextGroup(t *testing.T) account, err := svc.SelectAccountForModelWithExclusions(ctx, &groupID, "", "claude-3-5-sonnet-20241022", nil) require.NoError(t, err) require.NotNil(t, account) - require.Equal(t, 0, groupRepo.getByIDCalls) + require.Equal(t, 1, groupRepo.getByIDCalls) // +1 for require_privacy_set check require.Equal(t, 1, groupRepo.getByIDLiteCalls) } @@ -3252,7 +3252,7 @@ func TestGatewayService_GroupResolution_FallbackUsesLiteOnce(t *testing.T) { account, err := svc.SelectAccountForModelWithExclusions(ctx, &groupID, "", "claude-3-5-sonnet-20241022", nil) require.NoError(t, err) require.NotNil(t, account) - require.Equal(t, 0, groupRepo.getByIDCalls) + require.Equal(t, 1, groupRepo.getByIDCalls) // +1 for require_privacy_set check require.Equal(t, 1, groupRepo.getByIDLiteCalls) } diff --git a/backend/internal/service/gateway_service.go b/backend/internal/service/gateway_service.go index 5d285fb68f..e8e4343f8a 100644 --- a/backend/internal/service/gateway_service.go +++ b/backend/internal/service/gateway_service.go @@ -1430,19 +1430,24 @@ func (s *GatewayService) SelectAccountWithLoadAwareness(ctx context.Context, gro if containsInt64(routingAccountIDs, stickyAccountID) && !isExcluded(stickyAccountID) { // 粘性账号在路由列表中,优先使用 if stickyAccount, ok := accountByID[stickyAccountID]; ok { - if s.isAccountSchedulableForSelection(stickyAccount) && + var stickyCacheMissReason string + + gatePass := s.isAccountSchedulableForSelection(stickyAccount) && s.isAccountAllowedForPlatform(stickyAccount, platform, useMixed) && (requestedModel == "" || s.isModelSupportedByAccountWithContext(ctx, stickyAccount, requestedModel)) && s.isAccountSchedulableForModelSelection(ctx, stickyAccount, requestedModel) && s.isAccountSchedulableForQuota(stickyAccount) && - s.isAccountSchedulableForWindowCost(ctx, stickyAccount, true) && + s.isAccountSchedulableForWindowCost(ctx, stickyAccount, true) - s.isAccountSchedulableForRPM(ctx, stickyAccount, true) { // 粘性会话窗口费用+RPM 检查 + rpmPass := gatePass && s.isAccountSchedulableForRPM(ctx, stickyAccount, true) + + if rpmPass { // 粘性会话窗口费用+RPM 检查 result, err := s.tryAcquireAccountSlot(ctx, stickyAccountID, stickyAccount.Concurrency) if err == nil && result.Acquired { // 会话数量限制检查 if !s.checkAndRegisterSession(ctx, stickyAccount, sessionHash) { result.ReleaseFunc() // 释放槽位 + stickyCacheMissReason = "session_limit" // 继续到负载感知选择 } else { if s.debugModelRoutingEnabled() { @@ -1456,27 +1461,49 @@ func (s *GatewayService) SelectAccountWithLoadAwareness(ctx context.Context, gro } } - waitingCount, _ := s.concurrencyService.GetAccountWaitingCount(ctx, stickyAccountID) - if waitingCount < cfg.StickySessionMaxWaiting { - // 会话数量限制检查(等待计划也需要占用会话配额) - if !s.checkAndRegisterSession(ctx, stickyAccount, sessionHash) { - // 会话限制已满,继续到负载感知选择 + if stickyCacheMissReason == "" { + waitingCount, _ := s.concurrencyService.GetAccountWaitingCount(ctx, stickyAccountID) + if waitingCount < cfg.StickySessionMaxWaiting { + // 会话数量限制检查(等待计划也需要占用会话配额) + if !s.checkAndRegisterSession(ctx, stickyAccount, sessionHash) { + stickyCacheMissReason = "session_limit" + // 会话限制已满,继续到负载感知选择 + } else { + return &AccountSelectionResult{ + Account: stickyAccount, + WaitPlan: &AccountWaitPlan{ + AccountID: stickyAccountID, + MaxConcurrency: stickyAccount.Concurrency, + Timeout: cfg.StickySessionWaitTimeout, + MaxWaiting: cfg.StickySessionMaxWaiting, + }, + }, nil + } } else { - return &AccountSelectionResult{ - Account: stickyAccount, - WaitPlan: &AccountWaitPlan{ - AccountID: stickyAccountID, - MaxConcurrency: stickyAccount.Concurrency, - Timeout: cfg.StickySessionWaitTimeout, - MaxWaiting: cfg.StickySessionMaxWaiting, - }, - }, nil + stickyCacheMissReason = "wait_queue_full" } } // 粘性账号槽位满且等待队列已满,继续使用负载感知选择 + } else if !gatePass { + stickyCacheMissReason = "gate_check" + } else { + stickyCacheMissReason = "rpm_red" + } + + // 记录粘性缓存未命中的结构化日志 + if stickyCacheMissReason != "" { + baseRPM := stickyAccount.GetBaseRPM() + var currentRPM int + if count, ok := rpmFromPrefetchContext(ctx, stickyAccount.ID); ok { + currentRPM = count + } + logger.LegacyPrintf("service.gateway", "[StickyCacheMiss] reason=%s account_id=%d session=%s current_rpm=%d base_rpm=%d", + stickyCacheMissReason, stickyAccountID, shortSessionHash(sessionHash), currentRPM, baseRPM) } } else { _ = s.cache.DeleteSessionAccountID(ctx, derefGroupID(groupID), sessionHash) + logger.LegacyPrintf("service.gateway", "[StickyCacheMiss] reason=account_cleared account_id=%d session=%s current_rpm=0 base_rpm=0", + stickyAccountID, shortSessionHash(sessionHash)) } } } @@ -2673,6 +2700,12 @@ func (s *GatewayService) selectAccountForModelWithPlatform(ctx context.Context, preferOAuth := platform == PlatformGemini routingAccountIDs := s.routingAccountIDsForRequest(ctx, groupID, requestedModel, platform) + // require_privacy_set: 获取分组信息 + var schedGroup *Group + if groupID != nil && s.groupRepo != nil { + schedGroup, _ = s.groupRepo.GetByID(ctx, *groupID) + } + var accounts []Account accountsLoaded := false @@ -2744,6 +2777,12 @@ func (s *GatewayService) selectAccountForModelWithPlatform(ctx context.Context, if !s.isAccountSchedulableForSelection(acc) { continue } + // require_privacy_set: 跳过 privacy 未设置的账号并标记异常 + if schedGroup != nil && schedGroup.RequirePrivacySet && !acc.IsPrivacySet() { + _ = s.accountRepo.SetError(ctx, acc.ID, + fmt.Sprintf("Privacy not set, required by group [%s]", schedGroup.Name)) + continue + } if requestedModel != "" && !s.isModelSupportedByAccountWithContext(ctx, acc, requestedModel) { continue } @@ -2849,6 +2888,12 @@ func (s *GatewayService) selectAccountForModelWithPlatform(ctx context.Context, if !s.isAccountSchedulableForSelection(acc) { continue } + // require_privacy_set: 跳过 privacy 未设置的账号并标记异常 + if schedGroup != nil && schedGroup.RequirePrivacySet && !acc.IsPrivacySet() { + _ = s.accountRepo.SetError(ctx, acc.ID, + fmt.Sprintf("Privacy not set, required by group [%s]", schedGroup.Name)) + continue + } if requestedModel != "" && !s.isModelSupportedByAccountWithContext(ctx, acc, requestedModel) { continue } @@ -2915,6 +2960,12 @@ func (s *GatewayService) selectAccountWithMixedScheduling(ctx context.Context, g preferOAuth := nativePlatform == PlatformGemini routingAccountIDs := s.routingAccountIDsForRequest(ctx, groupID, requestedModel, nativePlatform) + // require_privacy_set: 获取分组信息 + var schedGroup *Group + if groupID != nil && s.groupRepo != nil { + schedGroup, _ = s.groupRepo.GetByID(ctx, *groupID) + } + var accounts []Account accountsLoaded := false @@ -2982,6 +3033,12 @@ func (s *GatewayService) selectAccountWithMixedScheduling(ctx context.Context, g if !s.isAccountSchedulableForSelection(acc) { continue } + // require_privacy_set: 跳过 privacy 未设置的账号并标记异常 + if schedGroup != nil && schedGroup.RequirePrivacySet && !acc.IsPrivacySet() { + _ = s.accountRepo.SetError(ctx, acc.ID, + fmt.Sprintf("Privacy not set, required by group [%s]", schedGroup.Name)) + continue + } // 过滤:原生平台直接通过,antigravity 需要启用混合调度 if acc.Platform == PlatformAntigravity && !acc.IsMixedSchedulingEnabled() { continue @@ -3075,6 +3132,7 @@ func (s *GatewayService) selectAccountWithMixedScheduling(ctx context.Context, g ctx = s.withRPMPrefetch(ctx, accounts) // 3. 按优先级+最久未用选择(考虑模型支持和混合调度) + // needsUpstreamCheck 仅在主选择循环中使用;粘性会话命中时跳过此检查。 needsUpstreamCheck := s.needsUpstreamChannelRestrictionCheck(ctx, groupID) var selected *Account for i := range accounts { @@ -3087,6 +3145,12 @@ func (s *GatewayService) selectAccountWithMixedScheduling(ctx context.Context, g if !s.isAccountSchedulableForSelection(acc) { continue } + // require_privacy_set: 跳过 privacy 未设置的账号并标记异常 + if schedGroup != nil && schedGroup.RequirePrivacySet && !acc.IsPrivacySet() { + _ = s.accountRepo.SetError(ctx, acc.ID, + fmt.Sprintf("Privacy not set, required by group [%s]", schedGroup.Name)) + continue + } // 过滤:原生平台直接通过,antigravity 需要启用混合调度 if acc.Platform == PlatformAntigravity && !acc.IsMixedSchedulingEnabled() { continue @@ -3254,8 +3318,7 @@ func (s *GatewayService) diagnoseSelectionFailure( return selectionFailureDiagnosis{Category: "excluded"} } if !s.isAccountSchedulableForSelection(acc) { - detail := "generic_unschedulable" - return selectionFailureDiagnosis{Category: "unschedulable", Detail: detail} + return selectionFailureDiagnosis{Category: "unschedulable", Detail: "generic_unschedulable"} } if isPlatformFilteredForSelection(acc, platform, allowMixedScheduling) { return selectionFailureDiagnosis{ @@ -3279,7 +3342,6 @@ func (s *GatewayService) diagnoseSelectionFailure( return selectionFailureDiagnosis{Category: "eligible"} } -// GetAccessToken 获取账号凭证 func isPlatformFilteredForSelection(acc *Account, platform string, allowMixedScheduling bool) bool { if acc == nil { return true @@ -7242,9 +7304,6 @@ func buildUsageBillingCommand(requestID string, usageLog *UsageLog, p *postUsage cmd.CacheCreationTokens = usageLog.CacheCreationTokens cmd.CacheReadTokens = usageLog.CacheReadTokens cmd.ImageCount = usageLog.ImageCount - if usageLog.MediaType != nil { - cmd.MediaType = *usageLog.MediaType - } if usageLog.ServiceTier != nil { cmd.ServiceTier = *usageLog.ServiceTier } @@ -7395,11 +7454,11 @@ func writeUsageLogBestEffort(ctx context.Context, repo UsageLogRepository, usage // recordUsageOpts 内部选项,参数化 RecordUsage 与 RecordUsageWithLongContext 的差异点。 type recordUsageOpts struct { - // ParsedRequest(可选,仅 Claude 路径传入) + // Claude Max 策略所需的 ParsedRequest(可选,仅 Claude 路径传入) ParsedRequest *ParsedRequest // EnableClaudePath 启用 Claude 路径特有逻辑: - // - MediaType 字段写入使用日志 + // - Claude Max 缓存计费策略 EnableClaudePath bool // 长上下文计费(仅 Gemini 路径需要) @@ -7424,7 +7483,6 @@ func (s *GatewayService) RecordUsage(ctx context.Context, input *RecordUsageInpu APIKeyService: input.APIKeyService, ChannelUsageFields: input.ChannelUsageFields, }, &recordUsageOpts{ - ParsedRequest: input.ParsedRequest, EnableClaudePath: true, }) } @@ -7490,6 +7548,7 @@ type recordUsageCoreInput struct { // recordUsageCore 是 RecordUsage 和 RecordUsageWithLongContext 的统一实现。 // opts 中的字段控制两者之间的差异行为: +// - ParsedRequest != nil → 启用 Claude Max 缓存计费策略 // - LongContextThreshold > 0 → Token 计费回退走 CalculateCostWithLongContext func (s *GatewayService) recordUsageCore(ctx context.Context, input *recordUsageCoreInput, opts *recordUsageOpts) error { result := input.Result @@ -7748,13 +7807,12 @@ func (s *GatewayService) buildRecordUsageLog( RateMultiplier: multiplier, AccountRateMultiplier: &accountRateMultiplier, BillingType: billingType, - BillingMode: resolveBillingMode(opts, result, cost), + BillingMode: resolveBillingMode(result, cost), Stream: result.Stream, DurationMs: &durationMs, FirstTokenMs: result.FirstTokenMs, ImageCount: result.ImageCount, ImageSize: optionalTrimmedStringPtr(result.ImageSize), - MediaType: resolveMediaType(opts, result), CacheTTLOverridden: cacheTTLOverridden, ChannelID: optionalInt64Ptr(input.ChannelID), ModelMappingChain: optionalTrimmedStringPtr(input.ModelMappingChain), @@ -7778,7 +7836,7 @@ func (s *GatewayService) buildRecordUsageLog( } // resolveBillingMode 根据计费结果和请求类型确定计费模式。 -func resolveBillingMode(opts *recordUsageOpts, result *ForwardResult, cost *CostBreakdown) *string { +func resolveBillingMode(result *ForwardResult, cost *CostBreakdown) *string { var mode string switch { case cost != nil && cost.BillingMode != "": @@ -7791,10 +7849,6 @@ func resolveBillingMode(opts *recordUsageOpts, result *ForwardResult, cost *Cost return &mode } -func resolveMediaType(opts *recordUsageOpts, result *ForwardResult) *string { - return nil -} - func optionalSubscriptionID(subscription *UserSubscription) *int64 { if subscription != nil { return &subscription.ID diff --git a/backend/internal/service/group.go b/backend/internal/service/group.go index 6445abea52..d59af9e1c0 100644 --- a/backend/internal/service/group.go +++ b/backend/internal/service/group.go @@ -50,6 +50,8 @@ type Group struct { // OpenAI Messages 调度配置(仅 openai 平台使用) AllowMessagesDispatch bool + RequireOAuthOnly bool // 仅允许非 apikey 类型账号关联(OpenAI/Antigravity/Anthropic/Gemini) + RequirePrivacySet bool // 调度时仅允许 privacy 已成功设置的账号(OpenAI/Antigravity/Anthropic/Gemini) DefaultMappedModel string CreatedAt time.Time diff --git a/backend/internal/service/openai_account_scheduler.go b/backend/internal/service/openai_account_scheduler.go index 707a9c8b85..cc85c46ef1 100644 --- a/backend/internal/service/openai_account_scheduler.go +++ b/backend/internal/service/openai_account_scheduler.go @@ -4,6 +4,7 @@ import ( "container/heap" "context" "errors" + "fmt" "hash/fnv" "math" "sort" @@ -575,6 +576,12 @@ func (s *defaultOpenAIAccountScheduler) selectByLoadBalance( return nil, 0, 0, 0, errors.New("no available OpenAI accounts") } + // require_privacy_set: 获取分组信息 + var schedGroup *Group + if req.GroupID != nil && s.service.schedulerSnapshot != nil { + schedGroup, _ = s.service.schedulerSnapshot.GetGroupByID(ctx, *req.GroupID) + } + filtered := make([]*Account, 0, len(accounts)) loadReq := make([]AccountWithConcurrency, 0, len(accounts)) for i := range accounts { @@ -587,6 +594,12 @@ func (s *defaultOpenAIAccountScheduler) selectByLoadBalance( if !account.IsSchedulable() || !account.IsOpenAI() { continue } + // require_privacy_set: 跳过 privacy 未设置的账号并标记异常 + if schedGroup != nil && schedGroup.RequirePrivacySet && !account.IsPrivacySet() { + _ = s.service.accountRepo.SetError(ctx, account.ID, + fmt.Sprintf("Privacy not set, required by group [%s]", schedGroup.Name)) + continue + } if req.RequestedModel != "" && !account.IsOpenAIPassthroughEnabled() && !account.IsModelSupported(req.RequestedModel) { continue } diff --git a/backend/internal/service/openai_codex_transform.go b/backend/internal/service/openai_codex_transform.go index d0534d8cd3..21b4874eb3 100644 --- a/backend/internal/service/openai_codex_transform.go +++ b/backend/internal/service/openai_codex_transform.go @@ -85,7 +85,7 @@ func applyCodexOAuthTransform(reqBody map[string]any, isCodexCLI bool, isCompact if v, ok := reqBody["model"].(string); ok { model = v } - normalizedModel := normalizeCodexModel(model) + normalizedModel := strings.TrimSpace(model) if normalizedModel != "" { if model != normalizedModel { reqBody["model"] = normalizedModel diff --git a/backend/internal/service/openai_codex_transform_test.go b/backend/internal/service/openai_codex_transform_test.go index eab88c0960..889ac61598 100644 --- a/backend/internal/service/openai_codex_transform_test.go +++ b/backend/internal/service/openai_codex_transform_test.go @@ -246,6 +246,7 @@ func TestNormalizeCodexModel_Gpt53(t *testing.T) { "gpt-5.3-codex": "gpt-5.3-codex", "gpt-5.3-codex-xhigh": "gpt-5.3-codex", "gpt-5.3-codex-spark": "gpt-5.3-codex", + "gpt 5.3 codex spark": "gpt-5.3-codex", "gpt-5.3-codex-spark-high": "gpt-5.3-codex", "gpt-5.3-codex-spark-xhigh": "gpt-5.3-codex", "gpt 5.3 codex": "gpt-5.3-codex", @@ -256,6 +257,34 @@ func TestNormalizeCodexModel_Gpt53(t *testing.T) { } } +func TestApplyCodexOAuthTransform_PreservesBareSparkModel(t *testing.T) { + reqBody := map[string]any{ + "model": "gpt-5.3-codex-spark", + "input": []any{}, + } + + result := applyCodexOAuthTransform(reqBody, false, false) + + require.Equal(t, "gpt-5.3-codex-spark", reqBody["model"]) + require.Equal(t, "gpt-5.3-codex-spark", result.NormalizedModel) + store, ok := reqBody["store"].(bool) + require.True(t, ok) + require.False(t, store) +} + +func TestApplyCodexOAuthTransform_TrimmedModelWithoutPolicyRewrite(t *testing.T) { + reqBody := map[string]any{ + "model": " gpt-5.3-codex-spark ", + "input": []any{}, + } + + result := applyCodexOAuthTransform(reqBody, false, false) + + require.Equal(t, "gpt-5.3-codex-spark", reqBody["model"]) + require.Equal(t, "gpt-5.3-codex-spark", result.NormalizedModel) + require.True(t, result.Modified) +} + func TestApplyCodexOAuthTransform_CodexCLI_PreservesExistingInstructions(t *testing.T) { // Codex CLI 场景:已有 instructions 时不修改 diff --git a/backend/internal/service/openai_compat_prompt_cache_key_test.go b/backend/internal/service/openai_compat_prompt_cache_key_test.go index eb9148de2d..6ca3e85cd3 100644 --- a/backend/internal/service/openai_compat_prompt_cache_key_test.go +++ b/backend/internal/service/openai_compat_prompt_cache_key_test.go @@ -17,6 +17,7 @@ func TestShouldAutoInjectPromptCacheKeyForCompat(t *testing.T) { require.True(t, shouldAutoInjectPromptCacheKeyForCompat("gpt-5.4")) require.True(t, shouldAutoInjectPromptCacheKeyForCompat("gpt-5.3")) require.True(t, shouldAutoInjectPromptCacheKeyForCompat("gpt-5.3-codex")) + require.True(t, shouldAutoInjectPromptCacheKeyForCompat("gpt-5.3-codex-spark")) require.False(t, shouldAutoInjectPromptCacheKeyForCompat("gpt-4o")) } @@ -62,3 +63,17 @@ func TestDeriveCompatPromptCacheKey_DiffersAcrossSessions(t *testing.T) { k2 := deriveCompatPromptCacheKey(req2, "gpt-5.4") require.NotEqual(t, k1, k2, "different first user messages should yield different keys") } + +func TestDeriveCompatPromptCacheKey_UsesResolvedSparkFamily(t *testing.T) { + req := &apicompat.ChatCompletionsRequest{ + Model: "gpt-5.3-codex-spark", + Messages: []apicompat.ChatMessage{ + {Role: "user", Content: mustRawJSON(t, `"Question A"`)}, + }, + } + + k1 := deriveCompatPromptCacheKey(req, "gpt-5.3-codex-spark") + k2 := deriveCompatPromptCacheKey(req, " openai/gpt-5.3-codex-spark ") + require.NotEmpty(t, k1) + require.Equal(t, k1, k2, "resolved spark family should derive a stable compat cache key") +} diff --git a/backend/internal/service/openai_gateway_chat_completions.go b/backend/internal/service/openai_gateway_chat_completions.go index 026e69d38f..3cada2ebdd 100644 --- a/backend/internal/service/openai_gateway_chat_completions.go +++ b/backend/internal/service/openai_gateway_chat_completions.go @@ -45,12 +45,13 @@ func (s *OpenAIGatewayService) ForwardAsChatCompletions( // 2. Resolve model mapping early so compat prompt_cache_key injection can // derive a stable seed from the final upstream model family. - mappedModel := resolveOpenAIForwardModel(account, originalModel, defaultMappedModel) + billingModel := resolveOpenAIForwardModel(account, originalModel, defaultMappedModel) + upstreamModel := normalizeCodexModel(billingModel) promptCacheKey = strings.TrimSpace(promptCacheKey) compatPromptCacheInjected := false - if promptCacheKey == "" && account.Type == AccountTypeOAuth && shouldAutoInjectPromptCacheKeyForCompat(mappedModel) { - promptCacheKey = deriveCompatPromptCacheKey(&chatReq, mappedModel) + if promptCacheKey == "" && account.Type == AccountTypeOAuth && shouldAutoInjectPromptCacheKeyForCompat(upstreamModel) { + promptCacheKey = deriveCompatPromptCacheKey(&chatReq, upstreamModel) compatPromptCacheInjected = promptCacheKey != "" } @@ -60,12 +61,13 @@ func (s *OpenAIGatewayService) ForwardAsChatCompletions( if err != nil { return nil, fmt.Errorf("convert chat completions to responses: %w", err) } - responsesReq.Model = mappedModel + responsesReq.Model = upstreamModel logFields := []zap.Field{ zap.Int64("account_id", account.ID), zap.String("original_model", originalModel), - zap.String("mapped_model", mappedModel), + zap.String("billing_model", billingModel), + zap.String("upstream_model", upstreamModel), zap.Bool("stream", clientStream), } if compatPromptCacheInjected { @@ -91,6 +93,9 @@ func (s *OpenAIGatewayService) ForwardAsChatCompletions( if account.Type == AccountTypeOAuth { codexResult := applyCodexOAuthTransform(reqBody, false, false) modified = codexResult.Modified + if codexResult.NormalizedModel != "" { + upstreamModel = codexResult.NormalizedModel + } if codexResult.PromptCacheKey != "" { promptCacheKey = codexResult.PromptCacheKey } else if promptCacheKey != "" { @@ -195,9 +200,9 @@ func (s *OpenAIGatewayService) ForwardAsChatCompletions( var result *OpenAIForwardResult var handleErr error if clientStream { - result, handleErr = s.handleChatStreamingResponse(resp, c, originalModel, mappedModel, includeUsage, startTime) + result, handleErr = s.handleChatStreamingResponse(resp, c, originalModel, billingModel, upstreamModel, includeUsage, startTime) } else { - result, handleErr = s.handleChatBufferedStreamingResponse(resp, c, originalModel, mappedModel, startTime) + result, handleErr = s.handleChatBufferedStreamingResponse(resp, c, originalModel, billingModel, upstreamModel, startTime) } // Propagate ServiceTier and ReasoningEffort to result for billing @@ -239,7 +244,8 @@ func (s *OpenAIGatewayService) handleChatBufferedStreamingResponse( resp *http.Response, c *gin.Context, originalModel string, - mappedModel string, + billingModel string, + upstreamModel string, startTime time.Time, ) (*OpenAIForwardResult, error) { requestID := resp.Header.Get("x-request-id") @@ -310,8 +316,8 @@ func (s *OpenAIGatewayService) handleChatBufferedStreamingResponse( RequestID: requestID, Usage: usage, Model: originalModel, - BillingModel: mappedModel, - UpstreamModel: mappedModel, + BillingModel: billingModel, + UpstreamModel: upstreamModel, Stream: false, Duration: time.Since(startTime), }, nil @@ -323,7 +329,8 @@ func (s *OpenAIGatewayService) handleChatStreamingResponse( resp *http.Response, c *gin.Context, originalModel string, - mappedModel string, + billingModel string, + upstreamModel string, includeUsage bool, startTime time.Time, ) (*OpenAIForwardResult, error) { @@ -358,8 +365,8 @@ func (s *OpenAIGatewayService) handleChatStreamingResponse( RequestID: requestID, Usage: usage, Model: originalModel, - BillingModel: mappedModel, - UpstreamModel: mappedModel, + BillingModel: billingModel, + UpstreamModel: upstreamModel, Stream: true, Duration: time.Since(startTime), FirstTokenMs: firstTokenMs, diff --git a/backend/internal/service/openai_gateway_messages.go b/backend/internal/service/openai_gateway_messages.go index 3df91b56e1..dd416269f4 100644 --- a/backend/internal/service/openai_gateway_messages.go +++ b/backend/internal/service/openai_gateway_messages.go @@ -41,6 +41,7 @@ func (s *OpenAIGatewayService) ForwardAsAnthropic( } originalModel := anthropicReq.Model applyOpenAICompatModelNormalization(&anthropicReq) + normalizedModel := anthropicReq.Model clientStream := anthropicReq.Stream // client's original stream preference // 2. Convert Anthropic → Responses @@ -60,13 +61,16 @@ func (s *OpenAIGatewayService) ForwardAsAnthropic( } // 3. Model mapping - mappedModel := resolveOpenAIForwardModel(account, anthropicReq.Model, defaultMappedModel) - responsesReq.Model = mappedModel + billingModel := resolveOpenAIForwardModel(account, normalizedModel, defaultMappedModel) + upstreamModel := normalizeCodexModel(billingModel) + responsesReq.Model = upstreamModel logger.L().Debug("openai messages: model mapping applied", zap.Int64("account_id", account.ID), zap.String("original_model", originalModel), - zap.String("mapped_model", mappedModel), + zap.String("normalized_model", normalizedModel), + zap.String("billing_model", billingModel), + zap.String("upstream_model", upstreamModel), zap.Bool("stream", isStream), ) @@ -82,6 +86,9 @@ func (s *OpenAIGatewayService) ForwardAsAnthropic( return nil, fmt.Errorf("unmarshal for codex transform: %w", err) } codexResult := applyCodexOAuthTransform(reqBody, false, false) + if codexResult.NormalizedModel != "" { + upstreamModel = codexResult.NormalizedModel + } if codexResult.PromptCacheKey != "" { promptCacheKey = codexResult.PromptCacheKey } else if promptCacheKey != "" { @@ -182,10 +189,10 @@ func (s *OpenAIGatewayService) ForwardAsAnthropic( var result *OpenAIForwardResult var handleErr error if clientStream { - result, handleErr = s.handleAnthropicStreamingResponse(resp, c, originalModel, mappedModel, startTime) + result, handleErr = s.handleAnthropicStreamingResponse(resp, c, originalModel, billingModel, upstreamModel, startTime) } else { // Client wants JSON: buffer the streaming response and assemble a JSON reply. - result, handleErr = s.handleAnthropicBufferedStreamingResponse(resp, c, originalModel, mappedModel, startTime) + result, handleErr = s.handleAnthropicBufferedStreamingResponse(resp, c, originalModel, billingModel, upstreamModel, startTime) } // Propagate ServiceTier and ReasoningEffort to result for billing @@ -230,7 +237,8 @@ func (s *OpenAIGatewayService) handleAnthropicBufferedStreamingResponse( resp *http.Response, c *gin.Context, originalModel string, - mappedModel string, + billingModel string, + upstreamModel string, startTime time.Time, ) (*OpenAIForwardResult, error) { requestID := resp.Header.Get("x-request-id") @@ -303,8 +311,8 @@ func (s *OpenAIGatewayService) handleAnthropicBufferedStreamingResponse( RequestID: requestID, Usage: usage, Model: originalModel, - BillingModel: mappedModel, - UpstreamModel: mappedModel, + BillingModel: billingModel, + UpstreamModel: upstreamModel, Stream: false, Duration: time.Since(startTime), }, nil @@ -319,7 +327,8 @@ func (s *OpenAIGatewayService) handleAnthropicStreamingResponse( resp *http.Response, c *gin.Context, originalModel string, - mappedModel string, + billingModel string, + upstreamModel string, startTime time.Time, ) (*OpenAIForwardResult, error) { requestID := resp.Header.Get("x-request-id") @@ -352,8 +361,8 @@ func (s *OpenAIGatewayService) handleAnthropicStreamingResponse( RequestID: requestID, Usage: usage, Model: originalModel, - BillingModel: mappedModel, - UpstreamModel: mappedModel, + BillingModel: billingModel, + UpstreamModel: upstreamModel, Stream: true, Duration: time.Since(startTime), FirstTokenMs: firstTokenMs, diff --git a/backend/internal/service/openai_gateway_record_usage_test.go b/backend/internal/service/openai_gateway_record_usage_test.go index e2b164c09d..38b97b1190 100644 --- a/backend/internal/service/openai_gateway_record_usage_test.go +++ b/backend/internal/service/openai_gateway_record_usage_test.go @@ -933,6 +933,89 @@ func TestOpenAIGatewayServiceRecordUsage_BillsMappedRequestsUsingRequestedModel( require.Equal(t, expectedCost.ActualCost, userRepo.lastAmount) } +func TestOpenAIGatewayServiceRecordUsage_ChannelMappedDoesNotOverrideBillingModelWhenUnmapped(t *testing.T) { + usageRepo := &openAIRecordUsageLogRepoStub{inserted: true} + userRepo := &openAIRecordUsageUserRepoStub{} + subRepo := &openAIRecordUsageSubRepoStub{} + svc := newOpenAIRecordUsageServiceForTest(usageRepo, userRepo, subRepo, nil) + usage := OpenAIUsage{InputTokens: 20, OutputTokens: 10} + + // When channel did NOT map the model (ChannelMappedModel == OriginalModel), + // billing should use result.BillingModel (the actual model used after group + // DefaultMappedModel resolution), not the unmapped original model. + expectedCost, err := svc.billingService.CalculateCost("gpt-5.1", UsageTokens{ + InputTokens: 20, + OutputTokens: 10, + }, 1.1) + require.NoError(t, err) + + err = svc.RecordUsage(context.Background(), &OpenAIRecordUsageInput{ + Result: &OpenAIForwardResult{ + RequestID: "resp_channel_unmapped_billing", + Model: "glm", + BillingModel: "gpt-5.1", + UpstreamModel: "gpt-5.1", + Usage: usage, + Duration: time.Second, + }, + APIKey: &APIKey{ID: 10}, + User: &User{ID: 20}, + Account: &Account{ID: 30}, + ChannelUsageFields: ChannelUsageFields{ + ChannelID: 1, + OriginalModel: "glm", + ChannelMappedModel: "glm", // channel did NOT map + BillingModelSource: BillingModelSourceChannelMapped, + }, + }) + + require.NoError(t, err) + require.NotNil(t, usageRepo.lastLog) + require.Equal(t, expectedCost.ActualCost, usageRepo.lastLog.ActualCost) + require.True(t, usageRepo.lastLog.ActualCost > 0, "cost must not be zero") +} + +func TestOpenAIGatewayServiceRecordUsage_ChannelMappedOverridesBillingModelWhenMapped(t *testing.T) { + usageRepo := &openAIRecordUsageLogRepoStub{inserted: true} + userRepo := &openAIRecordUsageUserRepoStub{} + subRepo := &openAIRecordUsageSubRepoStub{} + svc := newOpenAIRecordUsageServiceForTest(usageRepo, userRepo, subRepo, nil) + usage := OpenAIUsage{InputTokens: 20, OutputTokens: 10} + + // When channel DID map the model (ChannelMappedModel != OriginalModel), + // billing should use the channel-mapped model, honoring admin intent. + expectedCost, err := svc.billingService.CalculateCost("gpt-5.1", UsageTokens{ + InputTokens: 20, + OutputTokens: 10, + }, 1.1) + require.NoError(t, err) + + err = svc.RecordUsage(context.Background(), &OpenAIRecordUsageInput{ + Result: &OpenAIForwardResult{ + RequestID: "resp_channel_mapped_billing", + Model: "glm", + BillingModel: "gpt-5.1-codex", + UpstreamModel: "gpt-5.1-codex", + Usage: usage, + Duration: time.Second, + }, + APIKey: &APIKey{ID: 10}, + User: &User{ID: 20}, + Account: &Account{ID: 30}, + ChannelUsageFields: ChannelUsageFields{ + ChannelID: 1, + OriginalModel: "glm", + ChannelMappedModel: "gpt-5.1", // channel mapped glm → gpt-5.1 + BillingModelSource: BillingModelSourceChannelMapped, + }, + }) + + require.NoError(t, err) + require.NotNil(t, usageRepo.lastLog) + require.Equal(t, expectedCost.ActualCost, usageRepo.lastLog.ActualCost) + require.True(t, usageRepo.lastLog.ActualCost > 0, "cost must not be zero") +} + func TestOpenAIGatewayServiceRecordUsage_SubscriptionBillingSetsSubscriptionFields(t *testing.T) { usageRepo := &openAIRecordUsageLogRepoStub{inserted: true} userRepo := &openAIRecordUsageUserRepoStub{} diff --git a/backend/internal/service/openai_gateway_service.go b/backend/internal/service/openai_gateway_service.go index 4b095ee4c1..fcf3d028e4 100644 --- a/backend/internal/service/openai_gateway_service.go +++ b/backend/internal/service/openai_gateway_service.go @@ -1515,6 +1515,9 @@ func (s *OpenAIGatewayService) SelectAccountWithLoadAwareness(ctx context.Contex if needsUpstreamCheck && s.isUpstreamModelRestrictedByChannel(ctx, *groupID, acc, requestedModel) { continue } + if needsUpstreamCheck && s.isUpstreamModelRestrictedByChannel(ctx, *groupID, acc, requestedModel) { + continue + } candidates = append(candidates, acc) } @@ -1928,29 +1931,29 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco } // 对所有请求执行模型映射(包含 Codex CLI)。 - mappedModel := account.GetMappedModel(reqModel) - if mappedModel != reqModel { - logger.LegacyPrintf("service.openai_gateway", "[OpenAI] Model mapping applied: %s -> %s (account: %s, isCodexCLI: %v)", reqModel, mappedModel, account.Name, isCodexCLI) - reqBody["model"] = mappedModel + billingModel := account.GetMappedModel(reqModel) + if billingModel != reqModel { + logger.LegacyPrintf("service.openai_gateway", "[OpenAI] Model mapping applied: %s -> %s (account: %s, isCodexCLI: %v)", reqModel, billingModel, account.Name, isCodexCLI) + reqBody["model"] = billingModel bodyModified = true - markPatchSet("model", mappedModel) + markPatchSet("model", billingModel) } + upstreamModel := billingModel // 针对所有 OpenAI 账号执行 Codex 模型名规范化,确保上游识别一致。 if model, ok := reqBody["model"].(string); ok { - normalizedModel := normalizeCodexModel(model) - if normalizedModel != "" && normalizedModel != model { - logger.LegacyPrintf("service.openai_gateway", "[OpenAI] Codex model normalization: %s -> %s (account: %s, type: %s, isCodexCLI: %v)", - model, normalizedModel, account.Name, account.Type, isCodexCLI) - reqBody["model"] = normalizedModel - mappedModel = normalizedModel + upstreamModel = normalizeCodexModel(model) + if upstreamModel != "" && upstreamModel != model { + logger.LegacyPrintf("service.openai_gateway", "[OpenAI] Upstream model resolved: %s -> %s (account: %s, type: %s, isCodexCLI: %v)", + model, upstreamModel, account.Name, account.Type, isCodexCLI) + reqBody["model"] = upstreamModel bodyModified = true - markPatchSet("model", normalizedModel) + markPatchSet("model", upstreamModel) } // 移除 gpt-5.2-codex 以下的版本 verbosity 参数 // 确保高版本模型向低版本模型映射不报错 - if !SupportsVerbosity(normalizedModel) { + if !SupportsVerbosity(upstreamModel) { if text, ok := reqBody["text"].(map[string]any); ok { delete(text, "verbosity") } @@ -1974,7 +1977,7 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco disablePatch() } if codexResult.NormalizedModel != "" { - mappedModel = codexResult.NormalizedModel + upstreamModel = codexResult.NormalizedModel } if codexResult.PromptCacheKey != "" { promptCacheKey = codexResult.PromptCacheKey @@ -2091,7 +2094,7 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco "forward_start account_id=%d account_type=%s model=%s stream=%v has_previous_response_id=%v", account.ID, account.Type, - mappedModel, + upstreamModel, reqStream, hasPreviousResponseID, ) @@ -2180,7 +2183,7 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco isCodexCLI, reqStream, originalModel, - mappedModel, + upstreamModel, startTime, attempt, wsLastFailureReason, @@ -2281,7 +2284,7 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco firstTokenMs, wsAttempts, ) - wsResult.UpstreamModel = mappedModel + wsResult.UpstreamModel = upstreamModel return wsResult, nil } s.writeOpenAIWSFallbackErrorResponse(c, account, wsErr) @@ -2386,14 +2389,14 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco var usage *OpenAIUsage var firstTokenMs *int if reqStream { - streamResult, err := s.handleStreamingResponse(ctx, resp, c, account, startTime, originalModel, mappedModel) + streamResult, err := s.handleStreamingResponse(ctx, resp, c, account, startTime, originalModel, upstreamModel) if err != nil { return nil, err } usage = streamResult.usage firstTokenMs = streamResult.firstTokenMs } else { - usage, err = s.handleNonStreamingResponse(ctx, resp, c, account, originalModel, mappedModel) + usage, err = s.handleNonStreamingResponse(ctx, resp, c, account, originalModel, upstreamModel) if err != nil { return nil, err } @@ -2417,7 +2420,7 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco RequestID: resp.Header.Get("x-request-id"), Usage: *usage, Model: originalModel, - UpstreamModel: mappedModel, + UpstreamModel: upstreamModel, ServiceTier: serviceTier, ReasoningEffort: reasoningEffort, Stream: reqStream, @@ -4277,7 +4280,7 @@ func (s *OpenAIGatewayService) RecordUsage(ctx context.Context, input *OpenAIRec if result.BillingModel != "" { billingModel = strings.TrimSpace(result.BillingModel) } - if input.BillingModelSource == BillingModelSourceChannelMapped && input.ChannelMappedModel != "" { + if input.BillingModelSource == BillingModelSourceChannelMapped && input.ChannelMappedModel != "" && input.ChannelMappedModel != input.OriginalModel { billingModel = input.ChannelMappedModel } if input.BillingModelSource == BillingModelSourceRequested && input.OriginalModel != "" { diff --git a/backend/internal/service/openai_model_mapping_test.go b/backend/internal/service/openai_model_mapping_test.go index edbb968bd1..5ce2602c1c 100644 --- a/backend/internal/service/openai_model_mapping_test.go +++ b/backend/internal/service/openai_model_mapping_test.go @@ -74,13 +74,28 @@ func TestResolveOpenAIForwardModel_PreventsClaudeModelFromFallingBackToGpt51(t * Credentials: map[string]any{}, } - withoutDefault := resolveOpenAIForwardModel(account, "claude-opus-4-6", "") - if got := normalizeCodexModel(withoutDefault); got != "gpt-5.1" { - t.Fatalf("normalizeCodexModel(%q) = %q, want %q", withoutDefault, got, "gpt-5.1") + withoutDefault := normalizeCodexModel(resolveOpenAIForwardModel(account, "claude-opus-4-6", "")) + if withoutDefault != "gpt-5.1" { + t.Fatalf("normalizeCodexModel(...) = %q, want %q", withoutDefault, "gpt-5.1") } - withDefault := resolveOpenAIForwardModel(account, "claude-opus-4-6", "gpt-5.4") - if got := normalizeCodexModel(withDefault); got != "gpt-5.4" { - t.Fatalf("normalizeCodexModel(%q) = %q, want %q", withDefault, got, "gpt-5.4") + withDefault := normalizeCodexModel(resolveOpenAIForwardModel(account, "claude-opus-4-6", "gpt-5.4")) + if withDefault != "gpt-5.4" { + t.Fatalf("normalizeCodexModel(...) = %q, want %q", withDefault, "gpt-5.4") + } +} + +func TestNormalizeCodexModel(t *testing.T) { + cases := map[string]string{ + "gpt-5.3-codex-spark": "gpt-5.3-codex", + "gpt-5.3-codex-spark-high": "gpt-5.3-codex", + "gpt-5.3-codex-spark-xhigh": "gpt-5.3-codex", + "gpt-5.3": "gpt-5.3-codex", + } + + for input, expected := range cases { + if got := normalizeCodexModel(input); got != expected { + t.Fatalf("normalizeCodexModel(%q) = %q, want %q", input, got, expected) + } } } diff --git a/backend/internal/service/openai_oauth_service.go b/backend/internal/service/openai_oauth_service.go index e7b1f2339c..dc094d43ce 100644 --- a/backend/internal/service/openai_oauth_service.go +++ b/backend/internal/service/openai_oauth_service.go @@ -112,18 +112,19 @@ type OpenAIExchangeCodeInput struct { // OpenAITokenInfo represents the token information for OpenAI type OpenAITokenInfo struct { - AccessToken string `json:"access_token"` - RefreshToken string `json:"refresh_token"` - IDToken string `json:"id_token,omitempty"` - ExpiresIn int64 `json:"expires_in"` - ExpiresAt int64 `json:"expires_at"` - ClientID string `json:"client_id,omitempty"` - Email string `json:"email,omitempty"` - ChatGPTAccountID string `json:"chatgpt_account_id,omitempty"` - ChatGPTUserID string `json:"chatgpt_user_id,omitempty"` - OrganizationID string `json:"organization_id,omitempty"` - PlanType string `json:"plan_type,omitempty"` - PrivacyMode string `json:"privacy_mode,omitempty"` + AccessToken string `json:"access_token"` + RefreshToken string `json:"refresh_token"` + IDToken string `json:"id_token,omitempty"` + ExpiresIn int64 `json:"expires_in"` + ExpiresAt int64 `json:"expires_at"` + ClientID string `json:"client_id,omitempty"` + Email string `json:"email,omitempty"` + ChatGPTAccountID string `json:"chatgpt_account_id,omitempty"` + ChatGPTUserID string `json:"chatgpt_user_id,omitempty"` + OrganizationID string `json:"organization_id,omitempty"` + PlanType string `json:"plan_type,omitempty"` + SubscriptionExpiresAt string `json:"subscription_expires_at,omitempty"` + PrivacyMode string `json:"privacy_mode,omitempty"` } // ExchangeCode exchanges authorization code for tokens @@ -199,6 +200,8 @@ func (s *OpenAIOAuthService) ExchangeCode(ctx context.Context, input *OpenAIExch tokenInfo.PlanType = userInfo.PlanType } + s.enrichTokenInfo(ctx, tokenInfo, proxyURL) + return tokenInfo, nil } @@ -244,31 +247,40 @@ func (s *OpenAIOAuthService) RefreshTokenWithClientID(ctx context.Context, refre tokenInfo.PlanType = userInfo.PlanType } - // id_token 中缺少 plan_type 时(如 Mobile RT),尝试通过 ChatGPT backend-api 补全 - if tokenInfo.PlanType == "" && tokenInfo.AccessToken != "" && s.privacyClientFactory != nil { - // 从 access_token JWT 中提取 orgID(poid),用于匹配正确的账号 - orgID := tokenInfo.OrganizationID - if orgID == "" { - if atClaims, err := openai.DecodeIDToken(tokenInfo.AccessToken); err == nil && atClaims.OpenAIAuth != nil { - orgID = atClaims.OpenAIAuth.POID - } + s.enrichTokenInfo(ctx, tokenInfo, proxyURL) + + return tokenInfo, nil +} + +// enrichTokenInfo 通过 ChatGPT backend-api 补全 tokenInfo 并设置隐私(best-effort)。 +// 从 accounts/check 获取最新 plan_type、subscription_expires_at、email, +// 然后尝试关闭训练数据共享。适用于所有获取/刷新 token 的路径。 +func (s *OpenAIOAuthService) enrichTokenInfo(ctx context.Context, tokenInfo *OpenAITokenInfo, proxyURL string) { + if tokenInfo.AccessToken == "" || s.privacyClientFactory == nil { + return + } + + // 从 access_token JWT 中提取 orgID(poid),用于匹配正确的账号 + orgID := tokenInfo.OrganizationID + if orgID == "" { + if atClaims, err := openai.DecodeIDToken(tokenInfo.AccessToken); err == nil && atClaims.OpenAIAuth != nil { + orgID = atClaims.OpenAIAuth.POID } - if info := fetchChatGPTAccountInfo(ctx, s.privacyClientFactory, tokenInfo.AccessToken, proxyURL, orgID); info != nil { - if tokenInfo.PlanType == "" && info.PlanType != "" { - tokenInfo.PlanType = info.PlanType - } - if tokenInfo.Email == "" && info.Email != "" { - tokenInfo.Email = info.Email - } + } + if info := fetchChatGPTAccountInfo(ctx, s.privacyClientFactory, tokenInfo.AccessToken, proxyURL, orgID); info != nil { + if info.PlanType != "" { + tokenInfo.PlanType = info.PlanType + } + if info.SubscriptionExpiresAt != "" { + tokenInfo.SubscriptionExpiresAt = info.SubscriptionExpiresAt + } + if tokenInfo.Email == "" && info.Email != "" { + tokenInfo.Email = info.Email } } // 尝试设置隐私(关闭训练数据共享),best-effort - if tokenInfo.AccessToken != "" && s.privacyClientFactory != nil { - tokenInfo.PrivacyMode = disableOpenAITraining(ctx, s.privacyClientFactory, tokenInfo.AccessToken, proxyURL) - } - - return tokenInfo, nil + tokenInfo.PrivacyMode = disableOpenAITraining(ctx, s.privacyClientFactory, tokenInfo.AccessToken, proxyURL) } // RefreshAccountToken refreshes token for an OpenAI OAuth account @@ -347,6 +359,9 @@ func (s *OpenAIOAuthService) BuildAccountCredentials(tokenInfo *OpenAITokenInfo) if tokenInfo.PlanType != "" { creds["plan_type"] = tokenInfo.PlanType } + if tokenInfo.SubscriptionExpiresAt != "" { + creds["subscription_expires_at"] = tokenInfo.SubscriptionExpiresAt + } if strings.TrimSpace(tokenInfo.ClientID) != "" { creds["client_id"] = strings.TrimSpace(tokenInfo.ClientID) } diff --git a/backend/internal/service/openai_privacy_service.go b/backend/internal/service/openai_privacy_service.go index 6bc71ab9d4..da6dbefc93 100644 --- a/backend/internal/service/openai_privacy_service.go +++ b/backend/internal/service/openai_privacy_service.go @@ -56,6 +56,10 @@ func disableOpenAITraining(ctx context.Context, clientFactory PrivacyClientFacto SetHeader("Authorization", "Bearer "+accessToken). SetHeader("Origin", "https://chatgpt.com"). SetHeader("Referer", "https://chatgpt.com/"). + SetHeader("Accept", "application/json"). + SetHeader("sec-fetch-mode", "cors"). + SetHeader("sec-fetch-site", "same-origin"). + SetHeader("sec-fetch-dest", "empty"). SetQueryParam("feature", "training_allowed"). SetQueryParam("value", "false"). Patch(openAISettingsURL) @@ -84,8 +88,9 @@ func disableOpenAITraining(ctx context.Context, clientFactory PrivacyClientFacto // ChatGPTAccountInfo 从 chatgpt.com/backend-api/accounts/check 获取的账号信息 type ChatGPTAccountInfo struct { - PlanType string - Email string + PlanType string + Email string + SubscriptionExpiresAt string // entitlement.expires_at (RFC3339) } const chatGPTAccountsCheckURL = "https://chatgpt.com/backend-api/accounts/check/v4-2023-04-27" @@ -138,14 +143,20 @@ func fetchChatGPTAccountInfo(ctx context.Context, clientFactory PrivacyClientFac // 优先匹配 orgID 对应的账号(access_token JWT 中的 poid) if orgID != "" { - if matched := extractPlanFromAccount(accounts, orgID); matched != "" { - info.PlanType = matched + if acctRaw, ok := accounts[orgID]; ok { + if acct, ok := acctRaw.(map[string]any); ok { + fillAccountInfo(info, acct) + } } } // 未匹配到时,遍历所有账号:优先 is_default,次选非 free if info.PlanType == "" { - var defaultPlan, paidPlan, anyPlan string + type candidate struct { + planType string + expiresAt string + } + var defaultC, paidC, anyC candidate for _, acctRaw := range accounts { acct, ok := acctRaw.(map[string]any) if !ok { @@ -155,26 +166,27 @@ func fetchChatGPTAccountInfo(ctx context.Context, clientFactory PrivacyClientFac if planType == "" { continue } - if anyPlan == "" { - anyPlan = planType + ea := extractEntitlementExpiresAt(acct) + if anyC.planType == "" { + anyC = candidate{planType, ea} } if account, ok := acct["account"].(map[string]any); ok { if isDefault, _ := account["is_default"].(bool); isDefault { - defaultPlan = planType + defaultC = candidate{planType, ea} } } - if !strings.EqualFold(planType, "free") && paidPlan == "" { - paidPlan = planType + if !strings.EqualFold(planType, "free") && paidC.planType == "" { + paidC = candidate{planType, ea} } } // 优先级:default > 非 free > 任意 switch { - case defaultPlan != "": - info.PlanType = defaultPlan - case paidPlan != "": - info.PlanType = paidPlan + case defaultC.planType != "": + info.PlanType, info.SubscriptionExpiresAt = defaultC.planType, defaultC.expiresAt + case paidC.planType != "": + info.PlanType, info.SubscriptionExpiresAt = paidC.planType, paidC.expiresAt default: - info.PlanType = anyPlan + info.PlanType, info.SubscriptionExpiresAt = anyC.planType, anyC.expiresAt } } @@ -183,21 +195,14 @@ func fetchChatGPTAccountInfo(ctx context.Context, clientFactory PrivacyClientFac return nil } - slog.Info("chatgpt_account_check_success", "plan_type", info.PlanType, "org_id", orgID) + slog.Info("chatgpt_account_check_success", "plan_type", info.PlanType, "subscription_expires_at", info.SubscriptionExpiresAt, "org_id", orgID) return info } -// extractPlanFromAccount 从 accounts map 中按 key(account_id)精确匹配并提取 plan_type -func extractPlanFromAccount(accounts map[string]any, accountKey string) string { - acctRaw, ok := accounts[accountKey] - if !ok { - return "" - } - acct, ok := acctRaw.(map[string]any) - if !ok { - return "" - } - return extractPlanType(acct) +// fillAccountInfo 从单个 account 对象中提取 plan_type 和 subscription_expires_at +func fillAccountInfo(info *ChatGPTAccountInfo, acct map[string]any) { + info.PlanType = extractPlanType(acct) + info.SubscriptionExpiresAt = extractEntitlementExpiresAt(acct) } // extractPlanType 从单个 account 对象中提取 plan_type @@ -215,6 +220,17 @@ func extractPlanType(acct map[string]any) string { return "" } +// extractEntitlementExpiresAt 从 entitlement 中提取 expires_at。 +// 预期为 RFC3339 字符串格式,如 "2026-05-02T20:32:12+00:00"。 +func extractEntitlementExpiresAt(acct map[string]any) string { + entitlement, ok := acct["entitlement"].(map[string]any) + if !ok { + return "" + } + ea, _ := entitlement["expires_at"].(string) + return ea +} + func truncate(s string, n int) string { if len(s) <= n { return s diff --git a/backend/internal/service/openai_ws_forwarder.go b/backend/internal/service/openai_ws_forwarder.go index 4f1837c444..6d45baab36 100644 --- a/backend/internal/service/openai_ws_forwarder.go +++ b/backend/internal/service/openai_ws_forwarder.go @@ -2515,12 +2515,9 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient( } normalized = next } - mappedModel := account.GetMappedModel(originalModel) - if normalizedModel := normalizeCodexModel(mappedModel); normalizedModel != "" { - mappedModel = normalizedModel - } - if mappedModel != originalModel { - next, setErr := applyPayloadMutation(normalized, "model", mappedModel) + upstreamModel := normalizeCodexModel(account.GetMappedModel(originalModel)) + if upstreamModel != originalModel { + next, setErr := applyPayloadMutation(normalized, "model", upstreamModel) if setErr != nil { return openAIWSClientPayload{}, NewOpenAIWSClientCloseError(coderws.StatusPolicyViolation, "invalid websocket request payload", setErr) } @@ -2776,10 +2773,7 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient( mappedModel := "" var mappedModelBytes []byte if originalModel != "" { - mappedModel = account.GetMappedModel(originalModel) - if normalizedModel := normalizeCodexModel(mappedModel); normalizedModel != "" { - mappedModel = normalizedModel - } + mappedModel = normalizeCodexModel(account.GetMappedModel(originalModel)) needModelReplace = mappedModel != "" && mappedModel != originalModel if needModelReplace { mappedModelBytes = []byte(mappedModel) diff --git a/backend/internal/service/ratelimit_service.go b/backend/internal/service/ratelimit_service.go index aa0ae200c2..4f5b57cc97 100644 --- a/backend/internal/service/ratelimit_service.go +++ b/backend/internal/service/ratelimit_service.go @@ -161,6 +161,16 @@ func (s *RateLimitService) HandleUpstreamError(ctx context.Context, account *Acc shouldDisable = true break } + // OpenAI: {"detail":"Unauthorized"} 表示 token 完全无效(非标准 OpenAI 错误格式),直接标记 error + if account.Platform == PlatformOpenAI && gjson.GetBytes(responseBody, "detail").String() == "Unauthorized" { + msg := "Unauthorized (401): account authentication failed permanently" + if upstreamMsg != "" { + msg = "Unauthorized (401): " + upstreamMsg + } + s.handleAuthError(ctx, account, msg) + shouldDisable = true + break + } // OAuth 账号在 401 错误时临时不可调度(给 token 刷新窗口);非 OAuth 账号保持原有 SetError 行为。 // Antigravity 除外:其 401 由 applyErrorPolicy 的 temp_unschedulable_rules 自行控制。 if account.Type == AccountTypeOAuth && account.Platform != PlatformAntigravity { diff --git a/backend/internal/service/scheduler_snapshot_service.go b/backend/internal/service/scheduler_snapshot_service.go index 4c9540f115..d1330abb87 100644 --- a/backend/internal/service/scheduler_snapshot_service.go +++ b/backend/internal/service/scheduler_snapshot_service.go @@ -152,6 +152,14 @@ func (s *SchedulerSnapshotService) GetAccount(ctx context.Context, accountID int return s.accountRepo.GetByID(fallbackCtx, accountID) } +// GetGroupByID 获取分组信息(供调度器使用) +func (s *SchedulerSnapshotService) GetGroupByID(ctx context.Context, groupID int64) (*Group, error) { + if s.groupRepo == nil { + return nil, nil + } + return s.groupRepo.GetByID(ctx, groupID) +} + // UpdateAccountInCache 立即更新 Redis 中单个账号的数据(用于模型限流后立即生效) func (s *SchedulerSnapshotService) UpdateAccountInCache(ctx context.Context, account *Account) error { if s.cache == nil || account == nil { diff --git a/backend/internal/service/token_refresh_service.go b/backend/internal/service/token_refresh_service.go index d95f9a9f62..22f4aa29fd 100644 --- a/backend/internal/service/token_refresh_service.go +++ b/backend/internal/service/token_refresh_service.go @@ -292,6 +292,7 @@ func (s *TokenRefreshService) refreshWithRetry(ctx context.Context, account *Acc } // 刷新失败但 access_token 可能仍有效,尝试设置隐私 s.ensureOpenAIPrivacy(ctx, account) + s.ensureAntigravityPrivacy(ctx, account) return err } @@ -321,6 +322,7 @@ func (s *TokenRefreshService) refreshWithRetry(ctx context.Context, account *Acc // 刷新失败但 access_token 可能仍有效,尝试设置隐私 s.ensureOpenAIPrivacy(ctx, account) + s.ensureAntigravityPrivacy(ctx, account) // 设置临时不可调度 10 分钟(不标记 error,保持 status=active 让下个刷新周期能继续尝试) until := time.Now().Add(tokenRefreshTempUnschedDuration) @@ -474,15 +476,14 @@ func (s *TokenRefreshService) ensureOpenAIPrivacy(ctx context.Context, account * } // ensureAntigravityPrivacy 后台刷新中检查 Antigravity OAuth 账号隐私状态。 -// 仅做 Extra["privacy_mode"] 存在性检查,不发起 HTTP 请求,避免每轮循环产生额外网络开销。 -// 用户可通过前端 SetPrivacy 按钮强制重新设置。 +// 仅当 privacy_mode 已成功设置("privacy_set")时跳过; +// 未设置或之前失败("privacy_set_failed")均会重试。 func (s *TokenRefreshService) ensureAntigravityPrivacy(ctx context.Context, account *Account) { if account.Platform != PlatformAntigravity || account.Type != AccountTypeOAuth { return } - // 已设置过(无论成功或失败)则跳过,不发 HTTP if account.Extra != nil { - if _, ok := account.Extra["privacy_mode"]; ok { + if mode, ok := account.Extra["privacy_mode"].(string); ok && mode == AntigravityPrivacySet { return } } diff --git a/backend/internal/service/token_refresher.go b/backend/internal/service/token_refresher.go index 5248457a0d..916c226715 100644 --- a/backend/internal/service/token_refresher.go +++ b/backend/internal/service/token_refresher.go @@ -93,11 +93,11 @@ func (r *OpenAITokenRefresher) CanRefresh(account *Account) bool { } // NeedsRefresh 检查token是否需要刷新 -// 基于 expires_at 字段判断是否在刷新窗口内 +// expires_at 缺失且处于限流状态时需要刷新,防止限流期间 token 静默过期 func (r *OpenAITokenRefresher) NeedsRefresh(account *Account, refreshWindow time.Duration) bool { expiresAt := account.GetCredentialAsTime("expires_at") if expiresAt == nil { - return false + return account.IsRateLimited() } return time.Until(*expiresAt) < refreshWindow diff --git a/backend/migrations/081_add_group_account_filter.sql b/backend/migrations/081_add_group_account_filter.sql new file mode 100644 index 0000000000..0afb21d945 --- /dev/null +++ b/backend/migrations/081_add_group_account_filter.sql @@ -0,0 +1,2 @@ +ALTER TABLE groups ADD COLUMN IF NOT EXISTS require_oauth_only BOOLEAN NOT NULL DEFAULT false; +ALTER TABLE groups ADD COLUMN IF NOT EXISTS require_privacy_set BOOLEAN NOT NULL DEFAULT false; diff --git a/frontend/src/components/charts/TokenUsageTrend.vue b/frontend/src/components/charts/TokenUsageTrend.vue index a255fb03e6..4cd126b931 100644 --- a/frontend/src/components/charts/TokenUsageTrend.vue +++ b/frontend/src/components/charts/TokenUsageTrend.vue @@ -64,7 +64,8 @@ const chartColors = computed(() => ({ input: '#3b82f6', output: '#10b981', cacheCreation: '#f59e0b', - cacheRead: '#06b6d4' + cacheRead: '#06b6d4', + cacheHitRate: '#8b5cf6' })) const chartData = computed(() => { @@ -104,6 +105,19 @@ const chartData = computed(() => { backgroundColor: `${chartColors.value.cacheRead}20`, fill: true, tension: 0.3 + }, + { + label: 'Cache Hit Rate', + data: props.trendData.map((d) => { + const total = d.cache_read_tokens + d.cache_creation_tokens + return total > 0 ? (d.cache_read_tokens / total) * 100 : 0 + }), + borderColor: chartColors.value.cacheHitRate, + backgroundColor: `${chartColors.value.cacheHitRate}20`, + borderDash: [5, 5], + fill: false, + tension: 0.3, + yAxisID: 'yPercent' } ] } @@ -132,6 +146,9 @@ const lineOptions = computed(() => ({ tooltip: { callbacks: { label: (context: any) => { + if (context.dataset.yAxisID === 'yPercent') { + return `${context.dataset.label}: ${context.raw.toFixed(1)}%` + } return `${context.dataset.label}: ${formatTokens(context.raw)}` }, footer: (tooltipItems: any) => { @@ -168,6 +185,21 @@ const lineOptions = computed(() => ({ }, callback: (value: string | number) => formatTokens(Number(value)) } + }, + yPercent: { + position: 'right' as const, + min: 0, + max: 100, + grid: { + drawOnChartArea: false + }, + ticks: { + color: chartColors.value.cacheHitRate, + font: { + size: 10 + }, + callback: (value: string | number) => `${value}%` + } } } })) diff --git a/frontend/src/components/common/PlatformTypeBadge.vue b/frontend/src/components/common/PlatformTypeBadge.vue index e01b383a9e..1ebc889276 100644 --- a/frontend/src/components/common/PlatformTypeBadge.vue +++ b/frontend/src/components/common/PlatformTypeBadge.vue @@ -45,6 +45,10 @@ {{ privacyBadge.label }} + +
+ {{ expiresLabel }} +
@@ -62,6 +66,7 @@ interface Props { type: AccountType planType?: string privacyMode?: string + subscriptionExpiresAt?: string } const props = defineProps() @@ -141,6 +146,22 @@ const planBadgeClass = computed(() => { return typeClass.value }) +// Subscription expiration label (non-free only) +const expiresLabel = computed(() => { + if (!props.subscriptionExpiresAt || !props.planType) return '' + if (props.planType.toLowerCase() === 'free') return '' + try { + const d = new Date(props.subscriptionExpiresAt) + if (isNaN(d.getTime())) return '' + const yyyy = d.getFullYear() + const mm = String(d.getMonth() + 1).padStart(2, '0') + const dd = String(d.getDate()).padStart(2, '0') + return `${t('admin.accounts.subscriptionExpires')} ${yyyy}-${mm}-${dd}` + } catch { + return '' + } +}) + // Privacy badge — shows different states for OpenAI/Antigravity OAuth privacy setting const privacyBadge = computed(() => { if (props.type !== 'oauth' || !props.privacyMode) return null diff --git a/frontend/src/components/user/dashboard/UserDashboardCharts.vue b/frontend/src/components/user/dashboard/UserDashboardCharts.vue index 2214859294..73e88c3b28 100644 --- a/frontend/src/components/user/dashboard/UserDashboardCharts.vue +++ b/frontend/src/components/user/dashboard/UserDashboardCharts.vue @@ -7,6 +7,9 @@ {{ t('dashboard.timeRange') }}: +
{{ t('dashboard.granularity') }}:
@@ -74,7 +77,7 @@ import { Chart as ChartJS, CategoryScale, LinearScale, PointElement, LineElement ChartJS.register(CategoryScale, LinearScale, PointElement, LineElement, ArcElement, Title, Tooltip, Legend, Filler) const props = defineProps<{ loading: boolean, startDate: string, endDate: string, granularity: string, trend: TrendDataPoint[], models: ModelStat[] }>() -defineEmits(['update:startDate', 'update:endDate', 'update:granularity', 'dateRangeChange', 'granularityChange']) +defineEmits(['update:startDate', 'update:endDate', 'update:granularity', 'dateRangeChange', 'granularityChange', 'refresh']) const { t } = useI18n() const modelData = computed(() => !props.models?.length ? null : { diff --git a/frontend/src/i18n/locales/en.ts b/frontend/src/i18n/locales/en.ts index 41deee28a5..fca0b20a60 100644 --- a/frontend/src/i18n/locales/en.ts +++ b/frontend/src/i18n/locales/en.ts @@ -2063,6 +2063,7 @@ export default { privacyAntigravityFailed: 'Privacy setting failed', setPrivacy: 'Set Privacy', subscriptionAbnormal: 'Abnormal', + subscriptionExpires: 'Expires', // Capacity status tooltips capacity: { windowCost: { diff --git a/frontend/src/i18n/locales/zh.ts b/frontend/src/i18n/locales/zh.ts index d04f9e7ae2..c2d7546ac6 100644 --- a/frontend/src/i18n/locales/zh.ts +++ b/frontend/src/i18n/locales/zh.ts @@ -2114,6 +2114,7 @@ export default { privacyAntigravityFailed: '隐私设置失败', setPrivacy: '设置隐私', subscriptionAbnormal: '异常', + subscriptionExpires: '到期', // 容量状态提示 capacity: { windowCost: { diff --git a/frontend/src/types/index.ts b/frontend/src/types/index.ts index f4bc29e059..d7150f2a87 100644 --- a/frontend/src/types/index.ts +++ b/frontend/src/types/index.ts @@ -388,6 +388,8 @@ export interface Group { fallback_group_id_on_invalid_request: number | null // OpenAI Messages 调度开关(用户侧需要此字段判断是否展示 Claude Code 教程) allow_messages_dispatch?: boolean + require_oauth_only: boolean + require_privacy_set: boolean created_at: string updated_at: string } @@ -491,6 +493,8 @@ export interface CreateGroupRequest { fallback_group_id_on_invalid_request?: number | null mcp_xml_inject?: boolean supported_model_scopes?: string[] + require_oauth_only?: boolean + require_privacy_set?: boolean // 从指定分组复制账号 copy_accounts_from_group_ids?: number[] } @@ -514,6 +518,8 @@ export interface UpdateGroupRequest { fallback_group_id_on_invalid_request?: number | null mcp_xml_inject?: boolean supported_model_scopes?: string[] + require_oauth_only?: boolean + require_privacy_set?: boolean copy_accounts_from_group_ids?: number[] } diff --git a/frontend/src/views/admin/AccountsView.vue b/frontend/src/views/admin/AccountsView.vue index a39e07e3d8..e45c7ee6d7 100644 --- a/frontend/src/views/admin/AccountsView.vue +++ b/frontend/src/views/admin/AccountsView.vue @@ -185,7 +185,7 @@