diff --git a/.github/audit-exceptions.yml b/.github/audit-exceptions.yml index 61baa6e380..b71422a78a 100644 --- a/.github/audit-exceptions.yml +++ b/.github/audit-exceptions.yml @@ -5,14 +5,14 @@ exceptions: severity: high reason: "Admin export only; switched to dynamic import to reduce exposure (CVE-2023-30533)" mitigation: "Load only on export; restrict export permissions and data scope" - expires_on: "2026-07-07" + expires_on: "2026-07-06" owner: "security@your-domain" - package: xlsx advisory: "GHSA-5pgg-2g8v-p4x9" severity: high reason: "Admin export only; switched to dynamic import to reduce exposure (CVE-2024-22363)" mitigation: "Load only on export; restrict export permissions and data scope" - expires_on: "2026-07-07" + expires_on: "2026-07-06" owner: "security@your-domain" - package: lodash advisory: "GHSA-r5fr-rjxr-66jc" diff --git a/.github/workflows/release.yml b/.github/workflows/release.yml index c51b3c075e..b729c575ea 100644 --- a/.github/workflows/release.yml +++ b/.github/workflows/release.yml @@ -115,7 +115,7 @@ jobs: - name: Verify Go version run: | - go version | grep -q 'go1.26.1' + go version | grep -q 'go1.26.2' # Docker setup for GoReleaser - name: Set up QEMU diff --git a/README.md b/README.md index 99753e4569..25bef4730b 100644 --- a/README.md +++ b/README.md @@ -45,17 +45,41 @@ Sub2API is an AI API gateway platform designed to distribute and manage API quot - **Admin Dashboard** - Web interface for monitoring and management - **External System Integration** - Embed external systems (e.g. payment, ticketing) via iframe to extend the admin dashboard -## Don't Want to Self-Host? +## ❤️ Sponsors + +> [Want to appear here?](mailto:support@pincc.ai) + + + + + + + + + + + + + + + + + + + + + +
pincc 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.
PackyCode Thanks 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.
PoixeAiThanks to Poixe Ai for sponsoring this project! Poixe AI provides reliable LLM API services. You can leverage the platform's API endpoints to seamlessly build AI-powered products. Additionally, you can become a vendor by providing AI API resources to the platform and earn revenue. Register through the exclusive sub2api referral link and receive a bonus of $5 USD on your first top-up.
CTokThanks to CTok.ai for sponsoring this project! CTok.ai is dedicated to building a one-stop AI programming tool service platform. We offer professional Claude Code packages and technical community services, with support for Google Gemini and OpenAI Codex. Through carefully designed plans and a professional tech community, we provide developers with reliable service guarantees and continuous technical support, making AI-assisted programming a true productivity tool. Click here to register!
silkapiThanks to SilkAPI for sponsoring this project! SilkAPI is a relay service built on Sub2API, specializing in providing high-speed and stable Codex API relay.
ylscodeThanks to YLS Code for sponsoring this project! YLS Code is dedicated to building secure enterprise-grade Coding Agent productivity services, offering stable and fast Codex / Claude / Gemini subscription services along with pay-as-you-go API options for flexible choices. Register now for a limited-time 3-day Codex trial bonus!
## Ecosystem diff --git a/README_CN.md b/README_CN.md index 8b6feaba0d..003a9530cb 100644 --- a/README_CN.md +++ b/README_CN.md @@ -44,17 +44,41 @@ Sub2API 是一个 AI API 网关平台,用于分发和管理 AI 产品订阅的 - **管理后台** - Web 界面进行监控和管理 - **外部系统集成** - 支持通过 iframe 嵌入外部系统(如支付、工单等),扩展管理后台功能 -## 不想自建?试试官方中转 +## ❤️ 赞助商 + +> [想出现在这里?](mailto:support@pincc.ai) + + + + + + + + + + + + + + + + + + + + + +
pincc PinCC 是基于 Sub2API 搭建的官方中转服务,提供 Claude Code、Codex、Gemini 等主流模型的稳定中转,开箱即用,免去自建部署与运维烦恼。
PackyCode 感谢 PackyCode 赞助了本项目!PackyCode 是一家稳定、高效的API中转服务商,提供 Claude Code、Codex、Gemini 等多种中转服务。PackyCode 为本软件的用户提供了特别优惠,使用此链接注册并在充值时填写"sub2api"优惠码,首次充值可以享受9折优惠!
PoixeAI感谢 Poixe AI 赞助了本项目!Poixe AI 提供可靠的 AI 模型接口服务,您可以使用平台提供的 LLM API 接口轻松构建 AI 产品,同时也可以成为供应商,为平台提供大模型资源以赚取收益。通过 此链接 专属链接注册,充值额外赠送 $5 美金
CTok感谢 CTok.ai 赞助了本项目!CTok.ai 致力于打造一站式 AI 编程工具服务平台。我们提供 Claude Code 专业套餐及技术社群服务,同时支持 Google Gemini 和 OpenAI Codex。通过精心设计的套餐方案和专业的技术社群,为开发者提供稳定的服务保障和持续的技术支持,让 AI 辅助编程真正成为开发者的生产力工具。点击这里注册!
silkapi感谢 丝绸API 赞助了本项目! 丝绸API 是基于 Sub2API 搭建的中转服务,专注于提供 Codex 高速稳定API中转。
silkapi感谢 伊莉思Code 赞助了本项目! 伊莉思Code 致力于构建安全的企业级Coding Agent生产力服务,提供稳定快速的 Codex / Claude / Gemini 订阅服务与即用即付API多种方案灵活选择,限时注册赠送 3 天 Codex 试用福利!
## 生态项目 diff --git a/README_JA.md b/README_JA.md index 1266bd845c..818e944b68 100644 --- a/README_JA.md +++ b/README_JA.md @@ -45,7 +45,9 @@ Sub2API は、AI 製品のサブスクリプションから API クォータを - **管理ダッシュボード** - 監視・管理のための Web インターフェース - **外部システム連携** - 外部システム(決済、チケット管理など)を iframe 経由で管理ダッシュボードに埋め込み可能 -## セルフホストが不要な方へ +## ❤️ スポンサー + +> [こちらに掲載しませんか?](mailto:support@pincc.ai) @@ -56,6 +58,27 @@ Sub2API は、AI 製品のサブスクリプションから API クォータを + + + + + + + + + + + + + + + + + + + + +
PackyCode PackyCode のご支援に感謝します!PackyCode は Claude Code、Codex、Gemini などのリレーサービスを提供する信頼性の高い API 中継プラットフォームです。本ソフト利用者向けに特別割引があります:このリンクで登録し、チャージ時に「sub2api」クーポンを入力すると 10% オフになります。
PoixeAiPoixe AI のご支援に感謝します!Poixe AI は信頼性の高い LLM API サービスを提供しています。プラットフォームの API エンドポイントを活用して、AI 搭載プロダクトをシームレスに構築できます。また、ベンダーとして AI API リソースをプラットフォームに提供し、収益を得ることも可能です。専用の sub2api 紹介リンクから登録すると、初回チャージ時に $5 USD のボーナスがもらえます。
CTokCTok.ai のご支援に感謝します!CTok.ai はワンストップ AI プログラミングツールサービスプラットフォームの構築に取り組んでいます。Claude Code の専用プランと技術コミュニティサービスを提供し、Google Gemini や OpenAI Codex もサポートしています。丁寧に設計されたプランと専門的な技術コミュニティを通じて、開発者に安定したサービス保証と継続的な技術サポートを提供し、AI アシスト プログラミングを真の生産性向上ツールにします。こちらから登録!
silkapiSilkAPI のご支援に感謝します!SilkAPI は Sub2API をベースに構築された中継サービスで、高速かつ安定した Codex API 中継の提供に特化しています。
ylscodeYLS Code のご支援に感謝します!YLS Code は安全なエンタープライズグレードの Coding Agent 生産性サービスの構築に取り組んでおり、安定かつ高速な Codex / Claude / Gemini サブスクリプションサービスと従量課金 API の柔軟なプランを提供しています。期間限定で新規登録者に 3 日間の Codex 試用特典をプレゼント中!
## エコシステム diff --git a/assets/partners/logos/ctok.png b/assets/partners/logos/ctok.png new file mode 100644 index 0000000000..cf6fcf1706 Binary files /dev/null and b/assets/partners/logos/ctok.png differ diff --git a/assets/partners/logos/poixe.png b/assets/partners/logos/poixe.png new file mode 100644 index 0000000000..aa89cb06c9 Binary files /dev/null and b/assets/partners/logos/poixe.png differ diff --git a/assets/partners/logos/silkapi.png b/assets/partners/logos/silkapi.png new file mode 100644 index 0000000000..97afbda926 Binary files /dev/null and b/assets/partners/logos/silkapi.png differ diff --git a/assets/partners/logos/ylscode.png b/assets/partners/logos/ylscode.png new file mode 100644 index 0000000000..4d374f04c2 Binary files /dev/null and b/assets/partners/logos/ylscode.png differ diff --git a/backend/cmd/server/wire_gen.go b/backend/cmd/server/wire_gen.go index 7766580e7d..ffae9b39af 100644 --- a/backend/cmd/server/wire_gen.go +++ b/backend/cmd/server/wire_gen.go @@ -98,7 +98,7 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) { } dashboardAggregationService := service.ProvideDashboardAggregationService(dashboardAggregationRepository, timingWheelService, configConfig) dashboardHandler := admin.NewDashboardHandler(dashboardService, dashboardAggregationService) - schedulerCache := repository.NewSchedulerCache(redisClient) + schedulerCache := repository.ProvideSchedulerCache(redisClient, configConfig) accountRepository := repository.NewAccountRepository(client, db, schedulerCache) proxyRepository := repository.NewProxyRepository(client, db) proxyExitInfoProber := repository.NewProxyExitInfoProber(configConfig) diff --git a/backend/ent/group.go b/backend/ent/group.go index 3932da2bc7..2dce468ed4 100644 --- a/backend/ent/group.go +++ b/backend/ent/group.go @@ -11,6 +11,7 @@ import ( "entgo.io/ent" "entgo.io/ent/dialect/sql" "github.com/Wei-Shaw/sub2api/ent/group" + "github.com/Wei-Shaw/sub2api/internal/domain" ) // Group is the model entity for the Group schema. @@ -76,6 +77,8 @@ type Group struct { RequirePrivacySet bool `json:"require_privacy_set,omitempty"` // 默认映射模型 ID,当账号级映射找不到时使用此值 DefaultMappedModel string `json:"default_mapped_model,omitempty"` + // OpenAI Messages 调度模型配置:按 Claude 系列/精确模型映射到目标 GPT 模型 + MessagesDispatchModelConfig domain.OpenAIMessagesDispatchModelConfig `json:"messages_dispatch_model_config,omitempty"` // Edges holds the relations/edges for other nodes in the graph. // The values are being populated by the GroupQuery when eager-loading is set. Edges GroupEdges `json:"edges"` @@ -182,7 +185,7 @@ func (*Group) scanValues(columns []string) ([]any, error) { values := make([]any, len(columns)) for i := range columns { switch columns[i] { - case group.FieldModelRouting, group.FieldSupportedModelScopes: + case group.FieldModelRouting, group.FieldSupportedModelScopes, group.FieldMessagesDispatchModelConfig: values[i] = new([]byte) case group.FieldIsExclusive, group.FieldClaudeCodeOnly, group.FieldModelRoutingEnabled, group.FieldMcpXMLInject, group.FieldAllowMessagesDispatch, group.FieldRequireOauthOnly, group.FieldRequirePrivacySet: values[i] = new(sql.NullBool) @@ -403,6 +406,14 @@ func (_m *Group) assignValues(columns []string, values []any) error { } else if value.Valid { _m.DefaultMappedModel = value.String } + case group.FieldMessagesDispatchModelConfig: + if value, ok := values[i].(*[]byte); !ok { + return fmt.Errorf("unexpected type %T for field messages_dispatch_model_config", values[i]) + } else if value != nil && len(*value) > 0 { + if err := json.Unmarshal(*value, &_m.MessagesDispatchModelConfig); err != nil { + return fmt.Errorf("unmarshal field messages_dispatch_model_config: %w", err) + } + } default: _m.selectValues.Set(columns[i], values[i]) } @@ -585,6 +596,9 @@ func (_m *Group) String() string { builder.WriteString(", ") builder.WriteString("default_mapped_model=") builder.WriteString(_m.DefaultMappedModel) + builder.WriteString(", ") + builder.WriteString("messages_dispatch_model_config=") + builder.WriteString(fmt.Sprintf("%v", _m.MessagesDispatchModelConfig)) builder.WriteByte(')') return builder.String() } diff --git a/backend/ent/group/group.go b/backend/ent/group/group.go index 21a7c2cb76..b1371630dd 100644 --- a/backend/ent/group/group.go +++ b/backend/ent/group/group.go @@ -8,6 +8,7 @@ import ( "entgo.io/ent" "entgo.io/ent/dialect/sql" "entgo.io/ent/dialect/sql/sqlgraph" + "github.com/Wei-Shaw/sub2api/internal/domain" ) const ( @@ -73,6 +74,8 @@ const ( FieldRequirePrivacySet = "require_privacy_set" // FieldDefaultMappedModel holds the string denoting the default_mapped_model field in the database. FieldDefaultMappedModel = "default_mapped_model" + // FieldMessagesDispatchModelConfig holds the string denoting the messages_dispatch_model_config field in the database. + FieldMessagesDispatchModelConfig = "messages_dispatch_model_config" // EdgeAPIKeys holds the string denoting the api_keys edge name in mutations. EdgeAPIKeys = "api_keys" // EdgeRedeemCodes holds the string denoting the redeem_codes edge name in mutations. @@ -177,6 +180,7 @@ var Columns = []string{ FieldRequireOauthOnly, FieldRequirePrivacySet, FieldDefaultMappedModel, + FieldMessagesDispatchModelConfig, } var ( @@ -252,6 +256,8 @@ var ( DefaultDefaultMappedModel string // DefaultMappedModelValidator is a validator for the "default_mapped_model" field. It is called by the builders before save. DefaultMappedModelValidator func(string) error + // DefaultMessagesDispatchModelConfig holds the default value on creation for the "messages_dispatch_model_config" field. + DefaultMessagesDispatchModelConfig domain.OpenAIMessagesDispatchModelConfig ) // OrderOption defines the ordering options for the Group queries. diff --git a/backend/ent/group_create.go b/backend/ent/group_create.go index a8c30b184d..f412fa4070 100644 --- a/backend/ent/group_create.go +++ b/backend/ent/group_create.go @@ -18,6 +18,7 @@ import ( "github.com/Wei-Shaw/sub2api/ent/usagelog" "github.com/Wei-Shaw/sub2api/ent/user" "github.com/Wei-Shaw/sub2api/ent/usersubscription" + "github.com/Wei-Shaw/sub2api/internal/domain" ) // GroupCreate is the builder for creating a Group entity. @@ -410,6 +411,20 @@ func (_c *GroupCreate) SetNillableDefaultMappedModel(v *string) *GroupCreate { return _c } +// SetMessagesDispatchModelConfig sets the "messages_dispatch_model_config" field. +func (_c *GroupCreate) SetMessagesDispatchModelConfig(v domain.OpenAIMessagesDispatchModelConfig) *GroupCreate { + _c.mutation.SetMessagesDispatchModelConfig(v) + return _c +} + +// SetNillableMessagesDispatchModelConfig sets the "messages_dispatch_model_config" field if the given value is not nil. +func (_c *GroupCreate) SetNillableMessagesDispatchModelConfig(v *domain.OpenAIMessagesDispatchModelConfig) *GroupCreate { + if v != nil { + _c.SetMessagesDispatchModelConfig(*v) + } + return _c +} + // AddAPIKeyIDs adds the "api_keys" edge to the APIKey entity by IDs. func (_c *GroupCreate) AddAPIKeyIDs(ids ...int64) *GroupCreate { _c.mutation.AddAPIKeyIDs(ids...) @@ -611,6 +626,10 @@ func (_c *GroupCreate) defaults() error { v := group.DefaultDefaultMappedModel _c.mutation.SetDefaultMappedModel(v) } + if _, ok := _c.mutation.MessagesDispatchModelConfig(); !ok { + v := group.DefaultMessagesDispatchModelConfig + _c.mutation.SetMessagesDispatchModelConfig(v) + } return nil } @@ -695,6 +714,9 @@ func (_c *GroupCreate) check() error { return &ValidationError{Name: "default_mapped_model", err: fmt.Errorf(`ent: validator failed for field "Group.default_mapped_model": %w`, err)} } } + if _, ok := _c.mutation.MessagesDispatchModelConfig(); !ok { + return &ValidationError{Name: "messages_dispatch_model_config", err: errors.New(`ent: missing required field "Group.messages_dispatch_model_config"`)} + } return nil } @@ -838,6 +860,10 @@ func (_c *GroupCreate) createSpec() (*Group, *sqlgraph.CreateSpec) { _spec.SetField(group.FieldDefaultMappedModel, field.TypeString, value) _node.DefaultMappedModel = value } + if value, ok := _c.mutation.MessagesDispatchModelConfig(); ok { + _spec.SetField(group.FieldMessagesDispatchModelConfig, field.TypeJSON, value) + _node.MessagesDispatchModelConfig = value + } if nodes := _c.mutation.APIKeysIDs(); len(nodes) > 0 { edge := &sqlgraph.EdgeSpec{ Rel: sqlgraph.O2M, @@ -1462,6 +1488,18 @@ func (u *GroupUpsert) UpdateDefaultMappedModel() *GroupUpsert { return u } +// SetMessagesDispatchModelConfig sets the "messages_dispatch_model_config" field. +func (u *GroupUpsert) SetMessagesDispatchModelConfig(v domain.OpenAIMessagesDispatchModelConfig) *GroupUpsert { + u.Set(group.FieldMessagesDispatchModelConfig, v) + return u +} + +// UpdateMessagesDispatchModelConfig sets the "messages_dispatch_model_config" field to the value that was provided on create. +func (u *GroupUpsert) UpdateMessagesDispatchModelConfig() *GroupUpsert { + u.SetExcluded(group.FieldMessagesDispatchModelConfig) + return u +} + // UpdateNewValues updates the mutable fields using the new values that were set on create. // Using this option is equivalent to using: // @@ -2053,6 +2091,20 @@ func (u *GroupUpsertOne) UpdateDefaultMappedModel() *GroupUpsertOne { }) } +// SetMessagesDispatchModelConfig sets the "messages_dispatch_model_config" field. +func (u *GroupUpsertOne) SetMessagesDispatchModelConfig(v domain.OpenAIMessagesDispatchModelConfig) *GroupUpsertOne { + return u.Update(func(s *GroupUpsert) { + s.SetMessagesDispatchModelConfig(v) + }) +} + +// UpdateMessagesDispatchModelConfig sets the "messages_dispatch_model_config" field to the value that was provided on create. +func (u *GroupUpsertOne) UpdateMessagesDispatchModelConfig() *GroupUpsertOne { + return u.Update(func(s *GroupUpsert) { + s.UpdateMessagesDispatchModelConfig() + }) +} + // Exec executes the query. func (u *GroupUpsertOne) Exec(ctx context.Context) error { if len(u.create.conflict) == 0 { @@ -2810,6 +2862,20 @@ func (u *GroupUpsertBulk) UpdateDefaultMappedModel() *GroupUpsertBulk { }) } +// SetMessagesDispatchModelConfig sets the "messages_dispatch_model_config" field. +func (u *GroupUpsertBulk) SetMessagesDispatchModelConfig(v domain.OpenAIMessagesDispatchModelConfig) *GroupUpsertBulk { + return u.Update(func(s *GroupUpsert) { + s.SetMessagesDispatchModelConfig(v) + }) +} + +// UpdateMessagesDispatchModelConfig sets the "messages_dispatch_model_config" field to the value that was provided on create. +func (u *GroupUpsertBulk) UpdateMessagesDispatchModelConfig() *GroupUpsertBulk { + return u.Update(func(s *GroupUpsert) { + s.UpdateMessagesDispatchModelConfig() + }) +} + // Exec executes the query. func (u *GroupUpsertBulk) Exec(ctx context.Context) error { if u.create.err != nil { diff --git a/backend/ent/group_update.go b/backend/ent/group_update.go index aa1a83d421..7b6d619325 100644 --- a/backend/ent/group_update.go +++ b/backend/ent/group_update.go @@ -20,6 +20,7 @@ import ( "github.com/Wei-Shaw/sub2api/ent/usagelog" "github.com/Wei-Shaw/sub2api/ent/user" "github.com/Wei-Shaw/sub2api/ent/usersubscription" + "github.com/Wei-Shaw/sub2api/internal/domain" ) // GroupUpdate is the builder for updating Group entities. @@ -552,6 +553,20 @@ func (_u *GroupUpdate) SetNillableDefaultMappedModel(v *string) *GroupUpdate { return _u } +// SetMessagesDispatchModelConfig sets the "messages_dispatch_model_config" field. +func (_u *GroupUpdate) SetMessagesDispatchModelConfig(v domain.OpenAIMessagesDispatchModelConfig) *GroupUpdate { + _u.mutation.SetMessagesDispatchModelConfig(v) + return _u +} + +// SetNillableMessagesDispatchModelConfig sets the "messages_dispatch_model_config" field if the given value is not nil. +func (_u *GroupUpdate) SetNillableMessagesDispatchModelConfig(v *domain.OpenAIMessagesDispatchModelConfig) *GroupUpdate { + if v != nil { + _u.SetMessagesDispatchModelConfig(*v) + } + return _u +} + // AddAPIKeyIDs adds the "api_keys" edge to the APIKey entity by IDs. func (_u *GroupUpdate) AddAPIKeyIDs(ids ...int64) *GroupUpdate { _u.mutation.AddAPIKeyIDs(ids...) @@ -1012,6 +1027,9 @@ func (_u *GroupUpdate) sqlSave(ctx context.Context) (_node int, err error) { if value, ok := _u.mutation.DefaultMappedModel(); ok { _spec.SetField(group.FieldDefaultMappedModel, field.TypeString, value) } + if value, ok := _u.mutation.MessagesDispatchModelConfig(); ok { + _spec.SetField(group.FieldMessagesDispatchModelConfig, field.TypeJSON, value) + } if _u.mutation.APIKeysCleared() { edge := &sqlgraph.EdgeSpec{ Rel: sqlgraph.O2M, @@ -1843,6 +1861,20 @@ func (_u *GroupUpdateOne) SetNillableDefaultMappedModel(v *string) *GroupUpdateO return _u } +// SetMessagesDispatchModelConfig sets the "messages_dispatch_model_config" field. +func (_u *GroupUpdateOne) SetMessagesDispatchModelConfig(v domain.OpenAIMessagesDispatchModelConfig) *GroupUpdateOne { + _u.mutation.SetMessagesDispatchModelConfig(v) + return _u +} + +// SetNillableMessagesDispatchModelConfig sets the "messages_dispatch_model_config" field if the given value is not nil. +func (_u *GroupUpdateOne) SetNillableMessagesDispatchModelConfig(v *domain.OpenAIMessagesDispatchModelConfig) *GroupUpdateOne { + if v != nil { + _u.SetMessagesDispatchModelConfig(*v) + } + return _u +} + // AddAPIKeyIDs adds the "api_keys" edge to the APIKey entity by IDs. func (_u *GroupUpdateOne) AddAPIKeyIDs(ids ...int64) *GroupUpdateOne { _u.mutation.AddAPIKeyIDs(ids...) @@ -2333,6 +2365,9 @@ func (_u *GroupUpdateOne) sqlSave(ctx context.Context) (_node *Group, err error) if value, ok := _u.mutation.DefaultMappedModel(); ok { _spec.SetField(group.FieldDefaultMappedModel, field.TypeString, value) } + if value, ok := _u.mutation.MessagesDispatchModelConfig(); ok { + _spec.SetField(group.FieldMessagesDispatchModelConfig, field.TypeJSON, value) + } if _u.mutation.APIKeysCleared() { edge := &sqlgraph.EdgeSpec{ Rel: sqlgraph.O2M, diff --git a/backend/ent/migrate/schema.go b/backend/ent/migrate/schema.go index db76f3f507..e947b2e882 100644 --- a/backend/ent/migrate/schema.go +++ b/backend/ent/migrate/schema.go @@ -407,6 +407,7 @@ var ( {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: ""}, + {Name: "messages_dispatch_model_config", Type: field.TypeJSON, SchemaType: map[string]string{"postgres": "jsonb"}}, } // GroupsTable holds the schema information for the "groups" table. GroupsTable = &schema.Table{ diff --git a/backend/ent/mutation.go b/backend/ent/mutation.go index f873eb788b..6b2fa83869 100644 --- a/backend/ent/mutation.go +++ b/backend/ent/mutation.go @@ -8254,6 +8254,7 @@ type GroupMutation struct { require_oauth_only *bool require_privacy_set *bool default_mapped_model *string + messages_dispatch_model_config *domain.OpenAIMessagesDispatchModelConfig clearedFields map[string]struct{} api_keys map[int64]struct{} removedapi_keys map[int64]struct{} @@ -9806,6 +9807,42 @@ func (m *GroupMutation) ResetDefaultMappedModel() { m.default_mapped_model = nil } +// SetMessagesDispatchModelConfig sets the "messages_dispatch_model_config" field. +func (m *GroupMutation) SetMessagesDispatchModelConfig(damdmc domain.OpenAIMessagesDispatchModelConfig) { + m.messages_dispatch_model_config = &damdmc +} + +// MessagesDispatchModelConfig returns the value of the "messages_dispatch_model_config" field in the mutation. +func (m *GroupMutation) MessagesDispatchModelConfig() (r domain.OpenAIMessagesDispatchModelConfig, exists bool) { + v := m.messages_dispatch_model_config + if v == nil { + return + } + return *v, true +} + +// OldMessagesDispatchModelConfig returns the old "messages_dispatch_model_config" 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) OldMessagesDispatchModelConfig(ctx context.Context) (v domain.OpenAIMessagesDispatchModelConfig, err error) { + if !m.op.Is(OpUpdateOne) { + return v, errors.New("OldMessagesDispatchModelConfig is only allowed on UpdateOne operations") + } + if m.id == nil || m.oldValue == nil { + return v, errors.New("OldMessagesDispatchModelConfig requires an ID field in the mutation") + } + oldValue, err := m.oldValue(ctx) + if err != nil { + return v, fmt.Errorf("querying old value for OldMessagesDispatchModelConfig: %w", err) + } + return oldValue.MessagesDispatchModelConfig, nil +} + +// ResetMessagesDispatchModelConfig resets all changes to the "messages_dispatch_model_config" field. +func (m *GroupMutation) ResetMessagesDispatchModelConfig() { + m.messages_dispatch_model_config = nil +} + // AddAPIKeyIDs adds the "api_keys" edge to the APIKey entity by ids. func (m *GroupMutation) AddAPIKeyIDs(ids ...int64) { if m.api_keys == nil { @@ -10164,7 +10201,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, 29) + fields := make([]string, 0, 30) if m.created_at != nil { fields = append(fields, group.FieldCreatedAt) } @@ -10252,6 +10289,9 @@ func (m *GroupMutation) Fields() []string { if m.default_mapped_model != nil { fields = append(fields, group.FieldDefaultMappedModel) } + if m.messages_dispatch_model_config != nil { + fields = append(fields, group.FieldMessagesDispatchModelConfig) + } return fields } @@ -10318,6 +10358,8 @@ func (m *GroupMutation) Field(name string) (ent.Value, bool) { return m.RequirePrivacySet() case group.FieldDefaultMappedModel: return m.DefaultMappedModel() + case group.FieldMessagesDispatchModelConfig: + return m.MessagesDispatchModelConfig() } return nil, false } @@ -10385,6 +10427,8 @@ func (m *GroupMutation) OldField(ctx context.Context, name string) (ent.Value, e return m.OldRequirePrivacySet(ctx) case group.FieldDefaultMappedModel: return m.OldDefaultMappedModel(ctx) + case group.FieldMessagesDispatchModelConfig: + return m.OldMessagesDispatchModelConfig(ctx) } return nil, fmt.Errorf("unknown Group field %s", name) } @@ -10597,6 +10641,13 @@ func (m *GroupMutation) SetField(name string, value ent.Value) error { } m.SetDefaultMappedModel(v) return nil + case group.FieldMessagesDispatchModelConfig: + v, ok := value.(domain.OpenAIMessagesDispatchModelConfig) + if !ok { + return fmt.Errorf("unexpected type %T for field %s", value, name) + } + m.SetMessagesDispatchModelConfig(v) + return nil } return fmt.Errorf("unknown Group field %s", name) } @@ -10937,6 +10988,9 @@ func (m *GroupMutation) ResetField(name string) error { case group.FieldDefaultMappedModel: m.ResetDefaultMappedModel() return nil + case group.FieldMessagesDispatchModelConfig: + m.ResetMessagesDispatchModelConfig() + return nil } return fmt.Errorf("unknown Group field %s", name) } diff --git a/backend/ent/runtime/runtime.go b/backend/ent/runtime/runtime.go index 385d040dbb..821b7d66e0 100644 --- a/backend/ent/runtime/runtime.go +++ b/backend/ent/runtime/runtime.go @@ -32,6 +32,7 @@ import ( "github.com/Wei-Shaw/sub2api/ent/userattributedefinition" "github.com/Wei-Shaw/sub2api/ent/userattributevalue" "github.com/Wei-Shaw/sub2api/ent/usersubscription" + "github.com/Wei-Shaw/sub2api/internal/domain" ) // The init function reads all schema descriptors with runtime code @@ -472,6 +473,10 @@ func init() { group.DefaultDefaultMappedModel = groupDescDefaultMappedModel.Default.(string) // group.DefaultMappedModelValidator is a validator for the "default_mapped_model" field. It is called by the builders before save. group.DefaultMappedModelValidator = groupDescDefaultMappedModel.Validators[0].(func(string) error) + // groupDescMessagesDispatchModelConfig is the schema descriptor for messages_dispatch_model_config field. + groupDescMessagesDispatchModelConfig := groupFields[26].Descriptor() + // group.DefaultMessagesDispatchModelConfig holds the default value on creation for the messages_dispatch_model_config field. + group.DefaultMessagesDispatchModelConfig = groupDescMessagesDispatchModelConfig.Default.(domain.OpenAIMessagesDispatchModelConfig) idempotencyrecordMixin := schema.IdempotencyRecord{}.Mixin() idempotencyrecordMixinFields0 := idempotencyrecordMixin[0].Fields() _ = idempotencyrecordMixinFields0 diff --git a/backend/ent/schema/group.go b/backend/ent/schema/group.go index 0a6aeaecad..00420f9f88 100644 --- a/backend/ent/schema/group.go +++ b/backend/ent/schema/group.go @@ -131,6 +131,10 @@ func (Group) Fields() []ent.Field { MaxLen(100). Default(""). Comment("默认映射模型 ID,当账号级映射找不到时使用此值"), + field.JSON("messages_dispatch_model_config", domain.OpenAIMessagesDispatchModelConfig{}). + Default(domain.OpenAIMessagesDispatchModelConfig{}). + SchemaType(map[string]string{dialect.Postgres: "jsonb"}). + Comment("OpenAI Messages 调度模型配置:按 Claude 系列/精确模型映射到目标 GPT 模型"), } } diff --git a/backend/go.sum b/backend/go.sum index 0c407c395c..d26d2e2c9f 100644 --- a/backend/go.sum +++ b/backend/go.sum @@ -181,8 +181,6 @@ github.com/icholy/digest v1.1.0 h1:HfGg9Irj7i+IX1o1QAmPfIBNu/Q5A5Tu3n/MED9k9H4= github.com/icholy/digest v1.1.0/go.mod h1:QNrsSGQ5v7v9cReDI0+eyjsXGUoRSUZQHeQ5C4XLa0Y= github.com/imroc/req/v3 v3.57.0 h1:LMTUjNRUybUkTPn8oJDq8Kg3JRBOBTcnDhKu7mzupKI= github.com/imroc/req/v3 v3.57.0/go.mod h1:JL62ey1nvSLq81HORNcosvlf7SxZStONNqOprg0Pz00= -github.com/inconshreveable/mousetrap v1.1.0 h1:wN+x4NVGpMsO7ErUn/mUI3vEoE6Jt13X2s0bqwp9tc8= -github.com/inconshreveable/mousetrap v1.1.0/go.mod h1:vpF70FUmC8bwa3OWnCshd2FqLfsEA9PFc4w1p2J65bw= github.com/jackc/pgpassfile v1.0.0 h1:/6Hmqy13Ss2zCq62VdNG8tM1wchn8zjSGOBJ6icpsIM= github.com/jackc/pgpassfile v1.0.0/go.mod h1:CEx0iS5ambNFdcRtxPj5JhEz+xB6uRky5eyVu/W2HEg= github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 h1:iCEnooe7UlwOQYpKFhBabPMi4aNAfoODPEFNiAnClxo= @@ -218,8 +216,6 @@ github.com/mattn/go-colorable v0.1.13/go.mod h1:7S9/ev0klgBDR4GtXTXX8a3vIGJpMovk github.com/mattn/go-isatty v0.0.16/go.mod h1:kYGgaQfpe5nmfYZH+SKPsOc2e4SrIfOl2e/yFXSvRLM= github.com/mattn/go-isatty v0.0.20 h1:xfD0iDuEKnDkl03q4limB+vH+GxLEtL/jb4xVJSWWEY= github.com/mattn/go-isatty v0.0.20/go.mod h1:W+V8PltTTMOvKvAeJH7IuucS94S2C6jfK/D7dTCTo3Y= -github.com/mattn/go-runewidth v0.0.15 h1:UNAjwbU9l54TA3KzvqLGxwWjHmMgBUVhBiTjelZgg3U= -github.com/mattn/go-runewidth v0.0.15/go.mod h1:Jdepj2loyihRzMpdS35Xk/zdY8IAYHsh153qUoGf23w= github.com/mattn/go-sqlite3 v1.14.17 h1:mCRHCLDUBXgpKAqIKsaAaAsrAlbkeomtRFKXh2L6YIM= github.com/mattn/go-sqlite3 v1.14.17/go.mod h1:2eHXhiwb8IkHr+BDWZGa96P6+rkvnG63S2DGjv9HUNg= github.com/mdelapenya/tlscert v0.2.0 h1:7H81W6Z/4weDvZBNOfQte5GpIMo0lGYEeWbkGp5LJHI= @@ -253,8 +249,6 @@ github.com/morikuni/aec v1.0.0 h1:nP9CBfwrvYnBRgY6qfDQkygYDmYwOilePFkwzv4dU8A= github.com/morikuni/aec v1.0.0/go.mod h1:BbKIizmSmc5MMPqRYbxO4ZU0S0+P200+tUnFx7PXmsc= github.com/ncruces/go-strftime v1.0.0 h1:HMFp8mLCTPp341M/ZnA4qaf7ZlsbTc+miZjCLOFAw7w= github.com/ncruces/go-strftime v1.0.0/go.mod h1:Fwc5htZGVVkseilnfgOVb9mKy6w1naJmn9CehxcKcls= -github.com/olekukonko/tablewriter v0.0.5 h1:P2Ga83D34wi1o9J6Wh1mRuqd4mF/x/lgBS7N7AbDhec= -github.com/olekukonko/tablewriter v0.0.5/go.mod h1:hPp6KlRPjbx+hW8ykQs1w3UBbZlj6HuIJcUGPhkA7kY= github.com/opencontainers/go-digest v1.0.0 h1:apOUWs51W5PlhuyGyz9FCeeBIOUDA/6nW8Oi/yOhh5U= github.com/opencontainers/go-digest v1.0.0/go.mod h1:0JzlMkj0TRzQZfJkVvzbP0HBR3IKzErnv2BNG4W4MAM= github.com/opencontainers/image-spec v1.1.1 h1:y0fUlFfIZhPF1W537XOLg0/fcx6zcHCJwooC2xJA040= @@ -284,8 +278,6 @@ github.com/refraction-networking/utls v1.8.2 h1:j4Q1gJj0xngdeH+Ox/qND11aEfhpgoEv github.com/refraction-networking/utls v1.8.2/go.mod h1:jkSOEkLqn+S/jtpEHPOsVv/4V4EVnelwbMQl4vCWXAM= github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec h1:W09IVJc94icq4NjY3clb7Lk8O1qJ8BdBEF8z0ibU0rE= github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec/go.mod h1:qqbHyh8v60DhA7CoWK5oRCqLrMHRGoxYCSS9EjAz6Eo= -github.com/rivo/uniseg v0.2.0 h1:S1pD9weZBuJdFmowNwbpi7BJ8TNftyUImj/0WQi72jY= -github.com/rivo/uniseg v0.2.0/go.mod h1:J6wj4VEh+S6ZtnVlnTBMWIodfgj8LQOQFoIToxlJtxc= github.com/robfig/cron/v3 v3.0.1 h1:WdRxkvbJztn8LMz/QEvLN5sBU+xKpSqwwUO1Pjr4qDs= github.com/robfig/cron/v3 v3.0.1/go.mod h1:eQICP3HwyT7UooqI/z+Ov+PtYAWygg1TEWWzGIFLtro= github.com/rogpeppe/go-internal v1.13.1 h1:KvO1DLK/DRN07sQ1LQKScxyZJuNnedQ5/wKSR38lUII= @@ -318,8 +310,6 @@ github.com/spf13/afero v1.11.0 h1:WJQKhtpdm3v2IzqG8VMqrr6Rf3UYpEF239Jy9wNepM8= github.com/spf13/afero v1.11.0/go.mod h1:GH9Y3pIexgf1MTIWtNGyogA5MwRIDXGUr+hbWNoBjkY= github.com/spf13/cast v1.6.0 h1:GEiTHELF+vaR5dhz3VqZfFSzZjYbgeKDpBxQVS4GYJ0= github.com/spf13/cast v1.6.0/go.mod h1:ancEpBxwJDODSW/UG4rDrAqiKolqNNh2DX3mk86cAdo= -github.com/spf13/cobra v1.7.0 h1:hyqWnYt1ZQShIddO5kBpj3vu05/++x6tJ6dg8EC572I= -github.com/spf13/cobra v1.7.0/go.mod h1:uLxZILRyS/50WlhOIKD7W6V5bgeIt+4sICxh6uRMrb0= github.com/spf13/pflag v1.0.5 h1:iy+VFUOCP1a+8yFto/drg2CJ5u0yRoB7fZw3DKv/JXA= github.com/spf13/pflag v1.0.5/go.mod h1:McXfInJRrz4CZXVZOBLb0bTZqETkiAhM9Iw0y3An2Bg= github.com/spf13/viper v1.18.2 h1:LUXCnvUvSM6FXAsj6nnfc8Q2tp1dIgUfY9Kc8GsSOiQ= diff --git a/backend/internal/config/config.go b/backend/internal/config/config.go index 981c52e3e4..dd9a4e588f 100644 --- a/backend/internal/config/config.go +++ b/backend/internal/config/config.go @@ -65,6 +65,7 @@ type Config struct { JWT JWTConfig `mapstructure:"jwt"` Totp TotpConfig `mapstructure:"totp"` LinuxDo LinuxDoConnectConfig `mapstructure:"linuxdo_connect"` + OIDC OIDCConnectConfig `mapstructure:"oidc_connect"` Default DefaultConfig `mapstructure:"default"` RateLimit RateLimitConfig `mapstructure:"rate_limit"` Pricing PricingConfig `mapstructure:"pricing"` @@ -184,6 +185,34 @@ type LinuxDoConnectConfig struct { UserInfoUsernamePath string `mapstructure:"userinfo_username_path"` } +type OIDCConnectConfig struct { + Enabled bool `mapstructure:"enabled"` + ProviderName string `mapstructure:"provider_name"` // 显示名: "Keycloak" 等 + ClientID string `mapstructure:"client_id"` + ClientSecret string `mapstructure:"client_secret"` + IssuerURL string `mapstructure:"issuer_url"` + DiscoveryURL string `mapstructure:"discovery_url"` + AuthorizeURL string `mapstructure:"authorize_url"` + TokenURL string `mapstructure:"token_url"` + UserInfoURL string `mapstructure:"userinfo_url"` + JWKSURL string `mapstructure:"jwks_url"` + Scopes string `mapstructure:"scopes"` // 默认 "openid email profile" + RedirectURL string `mapstructure:"redirect_url"` // 后端回调地址(需在提供方后台登记) + FrontendRedirectURL string `mapstructure:"frontend_redirect_url"` // 前端接收 token 的路由(默认:/auth/oidc/callback) + TokenAuthMethod string `mapstructure:"token_auth_method"` // client_secret_post / client_secret_basic / none + UsePKCE bool `mapstructure:"use_pkce"` + ValidateIDToken bool `mapstructure:"validate_id_token"` + AllowedSigningAlgs string `mapstructure:"allowed_signing_algs"` // 默认 "RS256,ES256,PS256" + ClockSkewSeconds int `mapstructure:"clock_skew_seconds"` // 默认 120 + RequireEmailVerified bool `mapstructure:"require_email_verified"` // 默认 false + + // 可选:用于从 userinfo JSON 中提取字段的 gjson 路径。 + // 为空时,服务端会尝试一组常见字段名。 + UserInfoEmailPath string `mapstructure:"userinfo_email_path"` + UserInfoIDPath string `mapstructure:"userinfo_id_path"` + UserInfoUsernamePath string `mapstructure:"userinfo_username_path"` +} + // TokenRefreshConfig OAuth token自动刷新配置 type TokenRefreshConfig struct { // 是否启用自动刷新 @@ -318,6 +347,12 @@ type GatewayConfig struct { // ForceCodexCLI: 强制将 OpenAI `/v1/responses` 请求按 Codex CLI 处理。 // 用于网关未透传/改写 User-Agent 时的兼容兜底(默认关闭,避免影响其他客户端)。 ForceCodexCLI bool `mapstructure:"force_codex_cli"` + // ForcedCodexInstructionsTemplateFile: 服务端强制附加到 Codex 顶层 instructions 的模板文件路径。 + // 模板渲染后会直接覆盖最终 instructions;若需要保留客户端 system 转换结果,请在模板中显式引用 {{ .ExistingInstructions }}。 + ForcedCodexInstructionsTemplateFile string `mapstructure:"forced_codex_instructions_template_file"` + // ForcedCodexInstructionsTemplate: 启动时从模板文件读取并缓存的模板内容。 + // 该字段不直接参与配置反序列化,仅用于请求热路径避免重复读盘。 + ForcedCodexInstructionsTemplate string `mapstructure:"-"` // OpenAIPassthroughAllowTimeoutHeaders: OpenAI 透传模式是否放行客户端超时头 // 关闭(默认)可避免 x-stainless-timeout 等头导致上游提前断流。 OpenAIPassthroughAllowTimeoutHeaders bool `mapstructure:"openai_passthrough_allow_timeout_headers"` @@ -620,6 +655,10 @@ type GatewaySchedulingConfig struct { // 负载计算 LoadBatchEnabled bool `mapstructure:"load_batch_enabled"` + // 快照桶读取时的 MGET 分块大小 + SnapshotMGetChunkSize int `mapstructure:"snapshot_mget_chunk_size"` + // 快照重建时的缓存写入分块大小 + SnapshotWriteChunkSize int `mapstructure:"snapshot_write_chunk_size"` // 过期槽位清理周期(0 表示禁用) SlotCleanupInterval time.Duration `mapstructure:"slot_cleanup_interval"` @@ -968,6 +1007,23 @@ func load(allowMissingJWTSecret bool) (*Config, error) { cfg.LinuxDo.UserInfoEmailPath = strings.TrimSpace(cfg.LinuxDo.UserInfoEmailPath) cfg.LinuxDo.UserInfoIDPath = strings.TrimSpace(cfg.LinuxDo.UserInfoIDPath) cfg.LinuxDo.UserInfoUsernamePath = strings.TrimSpace(cfg.LinuxDo.UserInfoUsernamePath) + cfg.OIDC.ProviderName = strings.TrimSpace(cfg.OIDC.ProviderName) + cfg.OIDC.ClientID = strings.TrimSpace(cfg.OIDC.ClientID) + cfg.OIDC.ClientSecret = strings.TrimSpace(cfg.OIDC.ClientSecret) + cfg.OIDC.IssuerURL = strings.TrimSpace(cfg.OIDC.IssuerURL) + cfg.OIDC.DiscoveryURL = strings.TrimSpace(cfg.OIDC.DiscoveryURL) + cfg.OIDC.AuthorizeURL = strings.TrimSpace(cfg.OIDC.AuthorizeURL) + cfg.OIDC.TokenURL = strings.TrimSpace(cfg.OIDC.TokenURL) + cfg.OIDC.UserInfoURL = strings.TrimSpace(cfg.OIDC.UserInfoURL) + cfg.OIDC.JWKSURL = strings.TrimSpace(cfg.OIDC.JWKSURL) + cfg.OIDC.Scopes = strings.TrimSpace(cfg.OIDC.Scopes) + cfg.OIDC.RedirectURL = strings.TrimSpace(cfg.OIDC.RedirectURL) + cfg.OIDC.FrontendRedirectURL = strings.TrimSpace(cfg.OIDC.FrontendRedirectURL) + cfg.OIDC.TokenAuthMethod = strings.ToLower(strings.TrimSpace(cfg.OIDC.TokenAuthMethod)) + cfg.OIDC.AllowedSigningAlgs = strings.TrimSpace(cfg.OIDC.AllowedSigningAlgs) + cfg.OIDC.UserInfoEmailPath = strings.TrimSpace(cfg.OIDC.UserInfoEmailPath) + cfg.OIDC.UserInfoIDPath = strings.TrimSpace(cfg.OIDC.UserInfoIDPath) + cfg.OIDC.UserInfoUsernamePath = strings.TrimSpace(cfg.OIDC.UserInfoUsernamePath) cfg.Dashboard.KeyPrefix = strings.TrimSpace(cfg.Dashboard.KeyPrefix) cfg.CORS.AllowedOrigins = normalizeStringSlice(cfg.CORS.AllowedOrigins) cfg.Security.ResponseHeaders.AdditionalAllowed = normalizeStringSlice(cfg.Security.ResponseHeaders.AdditionalAllowed) @@ -979,6 +1035,14 @@ func load(allowMissingJWTSecret bool) (*Config, error) { cfg.Log.Environment = strings.TrimSpace(cfg.Log.Environment) cfg.Log.StacktraceLevel = strings.ToLower(strings.TrimSpace(cfg.Log.StacktraceLevel)) cfg.Log.Output.FilePath = strings.TrimSpace(cfg.Log.Output.FilePath) + cfg.Gateway.ForcedCodexInstructionsTemplateFile = strings.TrimSpace(cfg.Gateway.ForcedCodexInstructionsTemplateFile) + if cfg.Gateway.ForcedCodexInstructionsTemplateFile != "" { + content, err := os.ReadFile(cfg.Gateway.ForcedCodexInstructionsTemplateFile) + if err != nil { + return nil, fmt.Errorf("read forced codex instructions template %q: %w", cfg.Gateway.ForcedCodexInstructionsTemplateFile, err) + } + cfg.Gateway.ForcedCodexInstructionsTemplate = string(content) + } // 兼容旧键 gateway.openai_ws.sticky_previous_response_ttl_seconds。 // 新键未配置(<=0)时回退旧键;新键优先。 @@ -1138,6 +1202,30 @@ func setDefaults() { viper.SetDefault("linuxdo_connect.userinfo_id_path", "") viper.SetDefault("linuxdo_connect.userinfo_username_path", "") + // Generic OIDC OAuth 登录 + viper.SetDefault("oidc_connect.enabled", false) + viper.SetDefault("oidc_connect.provider_name", "OIDC") + viper.SetDefault("oidc_connect.client_id", "") + viper.SetDefault("oidc_connect.client_secret", "") + viper.SetDefault("oidc_connect.issuer_url", "") + viper.SetDefault("oidc_connect.discovery_url", "") + viper.SetDefault("oidc_connect.authorize_url", "") + viper.SetDefault("oidc_connect.token_url", "") + viper.SetDefault("oidc_connect.userinfo_url", "") + viper.SetDefault("oidc_connect.jwks_url", "") + viper.SetDefault("oidc_connect.scopes", "openid email profile") + viper.SetDefault("oidc_connect.redirect_url", "") + viper.SetDefault("oidc_connect.frontend_redirect_url", "/auth/oidc/callback") + viper.SetDefault("oidc_connect.token_auth_method", "client_secret_post") + viper.SetDefault("oidc_connect.use_pkce", false) + viper.SetDefault("oidc_connect.validate_id_token", true) + viper.SetDefault("oidc_connect.allowed_signing_algs", "RS256,ES256,PS256") + viper.SetDefault("oidc_connect.clock_skew_seconds", 120) + viper.SetDefault("oidc_connect.require_email_verified", false) + viper.SetDefault("oidc_connect.userinfo_email_path", "") + viper.SetDefault("oidc_connect.userinfo_id_path", "") + viper.SetDefault("oidc_connect.userinfo_username_path", "") + // Database viper.SetDefault("database.host", "localhost") viper.SetDefault("database.port", 5432) @@ -1340,6 +1428,8 @@ func setDefaults() { viper.SetDefault("gateway.scheduling.fallback_max_waiting", 100) viper.SetDefault("gateway.scheduling.fallback_selection_mode", "last_used") viper.SetDefault("gateway.scheduling.load_batch_enabled", true) + viper.SetDefault("gateway.scheduling.snapshot_mget_chunk_size", 128) + viper.SetDefault("gateway.scheduling.snapshot_write_chunk_size", 256) viper.SetDefault("gateway.scheduling.slot_cleanup_interval", 30*time.Second) viper.SetDefault("gateway.scheduling.db_fallback_enabled", true) viper.SetDefault("gateway.scheduling.db_fallback_timeout_seconds", 0) @@ -1572,6 +1662,87 @@ func (c *Config) Validate() error { warnIfInsecureURL("linuxdo_connect.redirect_url", c.LinuxDo.RedirectURL) warnIfInsecureURL("linuxdo_connect.frontend_redirect_url", c.LinuxDo.FrontendRedirectURL) } + if c.OIDC.Enabled { + if strings.TrimSpace(c.OIDC.ClientID) == "" { + return fmt.Errorf("oidc_connect.client_id is required when oidc_connect.enabled=true") + } + if strings.TrimSpace(c.OIDC.IssuerURL) == "" { + return fmt.Errorf("oidc_connect.issuer_url is required when oidc_connect.enabled=true") + } + if strings.TrimSpace(c.OIDC.RedirectURL) == "" { + return fmt.Errorf("oidc_connect.redirect_url is required when oidc_connect.enabled=true") + } + if strings.TrimSpace(c.OIDC.FrontendRedirectURL) == "" { + return fmt.Errorf("oidc_connect.frontend_redirect_url is required when oidc_connect.enabled=true") + } + if !scopeContainsOpenID(c.OIDC.Scopes) { + return fmt.Errorf("oidc_connect.scopes must contain openid") + } + + method := strings.ToLower(strings.TrimSpace(c.OIDC.TokenAuthMethod)) + switch method { + case "", "client_secret_post", "client_secret_basic", "none": + default: + return fmt.Errorf("oidc_connect.token_auth_method must be one of: client_secret_post/client_secret_basic/none") + } + if method == "none" && !c.OIDC.UsePKCE { + return fmt.Errorf("oidc_connect.use_pkce must be true when oidc_connect.token_auth_method=none") + } + if (method == "" || method == "client_secret_post" || method == "client_secret_basic") && + strings.TrimSpace(c.OIDC.ClientSecret) == "" { + return fmt.Errorf("oidc_connect.client_secret is required when oidc_connect.enabled=true and token_auth_method is client_secret_post/client_secret_basic") + } + if c.OIDC.ClockSkewSeconds < 0 || c.OIDC.ClockSkewSeconds > 600 { + return fmt.Errorf("oidc_connect.clock_skew_seconds must be between 0 and 600") + } + if c.OIDC.ValidateIDToken && strings.TrimSpace(c.OIDC.AllowedSigningAlgs) == "" { + return fmt.Errorf("oidc_connect.allowed_signing_algs is required when oidc_connect.validate_id_token=true") + } + + if err := ValidateAbsoluteHTTPURL(c.OIDC.IssuerURL); err != nil { + return fmt.Errorf("oidc_connect.issuer_url invalid: %w", err) + } + if v := strings.TrimSpace(c.OIDC.DiscoveryURL); v != "" { + if err := ValidateAbsoluteHTTPURL(v); err != nil { + return fmt.Errorf("oidc_connect.discovery_url invalid: %w", err) + } + } + if v := strings.TrimSpace(c.OIDC.AuthorizeURL); v != "" { + if err := ValidateAbsoluteHTTPURL(v); err != nil { + return fmt.Errorf("oidc_connect.authorize_url invalid: %w", err) + } + } + if v := strings.TrimSpace(c.OIDC.TokenURL); v != "" { + if err := ValidateAbsoluteHTTPURL(v); err != nil { + return fmt.Errorf("oidc_connect.token_url invalid: %w", err) + } + } + if v := strings.TrimSpace(c.OIDC.UserInfoURL); v != "" { + if err := ValidateAbsoluteHTTPURL(v); err != nil { + return fmt.Errorf("oidc_connect.userinfo_url invalid: %w", err) + } + } + if v := strings.TrimSpace(c.OIDC.JWKSURL); v != "" { + if err := ValidateAbsoluteHTTPURL(v); err != nil { + return fmt.Errorf("oidc_connect.jwks_url invalid: %w", err) + } + } + if err := ValidateAbsoluteHTTPURL(c.OIDC.RedirectURL); err != nil { + return fmt.Errorf("oidc_connect.redirect_url invalid: %w", err) + } + if err := ValidateFrontendRedirectURL(c.OIDC.FrontendRedirectURL); err != nil { + return fmt.Errorf("oidc_connect.frontend_redirect_url invalid: %w", err) + } + + warnIfInsecureURL("oidc_connect.issuer_url", c.OIDC.IssuerURL) + warnIfInsecureURL("oidc_connect.discovery_url", c.OIDC.DiscoveryURL) + warnIfInsecureURL("oidc_connect.authorize_url", c.OIDC.AuthorizeURL) + warnIfInsecureURL("oidc_connect.token_url", c.OIDC.TokenURL) + warnIfInsecureURL("oidc_connect.userinfo_url", c.OIDC.UserInfoURL) + warnIfInsecureURL("oidc_connect.jwks_url", c.OIDC.JWKSURL) + warnIfInsecureURL("oidc_connect.redirect_url", c.OIDC.RedirectURL) + warnIfInsecureURL("oidc_connect.frontend_redirect_url", c.OIDC.FrontendRedirectURL) + } if c.Billing.CircuitBreaker.Enabled { if c.Billing.CircuitBreaker.FailureThreshold <= 0 { return fmt.Errorf("billing.circuit_breaker.failure_threshold must be positive") @@ -2001,6 +2172,12 @@ func (c *Config) Validate() error { if c.Gateway.Scheduling.FallbackMaxWaiting <= 0 { return fmt.Errorf("gateway.scheduling.fallback_max_waiting must be positive") } + if c.Gateway.Scheduling.SnapshotMGetChunkSize <= 0 { + return fmt.Errorf("gateway.scheduling.snapshot_mget_chunk_size must be positive") + } + if c.Gateway.Scheduling.SnapshotWriteChunkSize <= 0 { + return fmt.Errorf("gateway.scheduling.snapshot_write_chunk_size must be positive") + } if c.Gateway.Scheduling.SlotCleanupInterval < 0 { return fmt.Errorf("gateway.scheduling.slot_cleanup_interval must be non-negative") } @@ -2184,6 +2361,15 @@ func ValidateFrontendRedirectURL(raw string) error { return nil } +func scopeContainsOpenID(scopes string) bool { + for _, scope := range strings.Fields(strings.ToLower(strings.TrimSpace(scopes))) { + if scope == "openid" { + return true + } + } + return false +} + // isHTTPScheme 检查是否为 HTTP 或 HTTPS 协议 func isHTTPScheme(scheme string) bool { return strings.EqualFold(scheme, "http") || strings.EqualFold(scheme, "https") diff --git a/backend/internal/config/config_test.go b/backend/internal/config/config_test.go index 2de5451ee0..fe181a2f30 100644 --- a/backend/internal/config/config_test.go +++ b/backend/internal/config/config_test.go @@ -1,6 +1,8 @@ package config import ( + "os" + "path/filepath" "strings" "testing" "time" @@ -223,6 +225,23 @@ func TestLoadSchedulingConfigFromEnv(t *testing.T) { } } +func TestLoadForcedCodexInstructionsTemplate(t *testing.T) { + resetViperWithJWTSecret(t) + + tempDir := t.TempDir() + templatePath := filepath.Join(tempDir, "codex-instructions.md.tmpl") + configPath := filepath.Join(tempDir, "config.yaml") + + require.NoError(t, os.WriteFile(templatePath, []byte("server-prefix\n\n{{ .ExistingInstructions }}"), 0o644)) + require.NoError(t, os.WriteFile(configPath, []byte("gateway:\n forced_codex_instructions_template_file: \""+templatePath+"\"\n"), 0o644)) + t.Setenv("DATA_DIR", tempDir) + + cfg, err := Load() + require.NoError(t, err) + require.Equal(t, templatePath, cfg.Gateway.ForcedCodexInstructionsTemplateFile) + require.Equal(t, "server-prefix\n\n{{ .ExistingInstructions }}", cfg.Gateway.ForcedCodexInstructionsTemplate) +} + func TestLoadDefaultSecurityToggles(t *testing.T) { resetViperWithJWTSecret(t) @@ -351,6 +370,60 @@ func TestValidateLinuxDoPKCERequiredForPublicClient(t *testing.T) { } } +func TestValidateOIDCScopesMustContainOpenID(t *testing.T) { + resetViperWithJWTSecret(t) + + cfg, err := Load() + if err != nil { + t.Fatalf("Load() error: %v", err) + } + + cfg.OIDC.Enabled = true + cfg.OIDC.ClientID = "oidc-client" + cfg.OIDC.ClientSecret = "oidc-secret" + cfg.OIDC.IssuerURL = "https://issuer.example.com" + cfg.OIDC.AuthorizeURL = "https://issuer.example.com/auth" + cfg.OIDC.TokenURL = "https://issuer.example.com/token" + cfg.OIDC.JWKSURL = "https://issuer.example.com/jwks" + cfg.OIDC.RedirectURL = "https://example.com/api/v1/auth/oauth/oidc/callback" + cfg.OIDC.FrontendRedirectURL = "/auth/oidc/callback" + cfg.OIDC.Scopes = "profile email" + + err = cfg.Validate() + if err == nil { + t.Fatalf("Validate() expected error when scopes do not include openid, got nil") + } + if !strings.Contains(err.Error(), "oidc_connect.scopes") { + t.Fatalf("Validate() expected oidc_connect.scopes error, got: %v", err) + } +} + +func TestValidateOIDCAllowsIssuerOnlyEndpointsWithDiscoveryFallback(t *testing.T) { + resetViperWithJWTSecret(t) + + cfg, err := Load() + if err != nil { + t.Fatalf("Load() error: %v", err) + } + + cfg.OIDC.Enabled = true + cfg.OIDC.ClientID = "oidc-client" + cfg.OIDC.ClientSecret = "oidc-secret" + cfg.OIDC.IssuerURL = "https://issuer.example.com" + cfg.OIDC.AuthorizeURL = "" + cfg.OIDC.TokenURL = "" + cfg.OIDC.JWKSURL = "" + cfg.OIDC.RedirectURL = "https://example.com/api/v1/auth/oauth/oidc/callback" + cfg.OIDC.FrontendRedirectURL = "/auth/oidc/callback" + cfg.OIDC.Scopes = "openid email profile" + cfg.OIDC.ValidateIDToken = true + + err = cfg.Validate() + if err != nil { + t.Fatalf("Validate() expected issuer-only OIDC config to pass with discovery fallback, got: %v", err) + } +} + func TestLoadDefaultDashboardCacheConfig(t *testing.T) { resetViperWithJWTSecret(t) diff --git a/backend/internal/domain/openai_messages_dispatch.go b/backend/internal/domain/openai_messages_dispatch.go new file mode 100644 index 0000000000..6b018f1c30 --- /dev/null +++ b/backend/internal/domain/openai_messages_dispatch.go @@ -0,0 +1,10 @@ +package domain + +// OpenAIMessagesDispatchModelConfig controls how Anthropic /v1/messages +// requests are mapped onto OpenAI/Codex models. +type OpenAIMessagesDispatchModelConfig struct { + OpusMappedModel string `json:"opus_mapped_model,omitempty"` + SonnetMappedModel string `json:"sonnet_mapped_model,omitempty"` + HaikuMappedModel string `json:"haiku_mapped_model,omitempty"` + ExactModelMappings map[string]string `json:"exact_model_mappings,omitempty"` +} diff --git a/backend/internal/handler/admin/account_data.go b/backend/internal/handler/admin/account_data.go index 20cc09eebc..00da48212a 100644 --- a/backend/internal/handler/admin/account_data.go +++ b/backend/internal/handler/admin/account_data.go @@ -10,6 +10,7 @@ import ( "log/slog" + infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" "github.com/Wei-Shaw/sub2api/internal/pkg/openai" "github.com/Wei-Shaw/sub2api/internal/pkg/response" "github.com/Wei-Shaw/sub2api/internal/service" @@ -359,7 +360,7 @@ func (h *AccountHandler) listAllProxies(ctx context.Context) ([]service.Proxy, e pageSize := dataPageCap var out []service.Proxy for { - items, total, err := h.adminService.ListProxies(ctx, page, pageSize, "", "", "") + items, total, err := h.adminService.ListProxies(ctx, page, pageSize, "", "", "", "created_at", "desc") if err != nil { return nil, err } @@ -372,12 +373,12 @@ func (h *AccountHandler) listAllProxies(ctx context.Context) ([]service.Proxy, e return out, nil } -func (h *AccountHandler) listAccountsFiltered(ctx context.Context, platform, accountType, status, search string) ([]service.Account, error) { +func (h *AccountHandler) listAccountsFiltered(ctx context.Context, platform, accountType, status, search string, groupID int64, privacyMode, sortBy, sortOrder string) ([]service.Account, error) { page := 1 pageSize := dataPageCap var out []service.Account for { - items, total, err := h.adminService.ListAccounts(ctx, page, pageSize, platform, accountType, status, search, 0, "") + items, total, err := h.adminService.ListAccounts(ctx, page, pageSize, platform, accountType, status, search, groupID, privacyMode, sortBy, sortOrder) if err != nil { return nil, err } @@ -409,11 +410,28 @@ func (h *AccountHandler) resolveExportAccounts(ctx context.Context, ids []int64, platform := c.Query("platform") accountType := c.Query("type") status := c.Query("status") + privacyMode := strings.TrimSpace(c.Query("privacy_mode")) search := strings.TrimSpace(c.Query("search")) + sortBy := c.DefaultQuery("sort_by", "name") + sortOrder := c.DefaultQuery("sort_order", "asc") if len(search) > 100 { search = search[:100] } - return h.listAccountsFiltered(ctx, platform, accountType, status, search) + + groupID := int64(0) + if groupIDStr := c.Query("group"); groupIDStr != "" { + if groupIDStr == accountListGroupUngroupedQueryValue { + groupID = service.AccountListGroupUngrouped + } else { + parsedGroupID, parseErr := strconv.ParseInt(groupIDStr, 10, 64) + if parseErr != nil || parsedGroupID <= 0 { + return nil, infraerrors.BadRequest("INVALID_GROUP_FILTER", "invalid group filter") + } + groupID = parsedGroupID + } + } + + return h.listAccountsFiltered(ctx, platform, accountType, status, search, groupID, privacyMode, sortBy, sortOrder) } func (h *AccountHandler) resolveExportProxies(ctx context.Context, accounts []service.Account) ([]service.Proxy, error) { diff --git a/backend/internal/handler/admin/account_data_handler_test.go b/backend/internal/handler/admin/account_data_handler_test.go index 285033a17d..5793983cba 100644 --- a/backend/internal/handler/admin/account_data_handler_test.go +++ b/backend/internal/handler/admin/account_data_handler_test.go @@ -172,6 +172,51 @@ func TestExportDataWithoutProxies(t *testing.T) { require.Nil(t, resp.Data.Accounts[0].ProxyKey) } +func TestExportDataPassesAccountFiltersAndSort(t *testing.T) { + router, adminSvc := setupAccountDataRouter() + adminSvc.accounts = []service.Account{ + {ID: 1, Name: "acc-1", Status: service.StatusActive}, + } + + rec := httptest.NewRecorder() + req := httptest.NewRequest( + http.MethodGet, + "/api/v1/admin/accounts/data?platform=openai&type=oauth&status=active&group=12&privacy_mode=blocked&search=keyword&sort_by=priority&sort_order=desc", + nil, + ) + router.ServeHTTP(rec, req) + require.Equal(t, http.StatusOK, rec.Code) + + require.Equal(t, 1, adminSvc.lastListAccounts.calls) + require.Equal(t, "openai", adminSvc.lastListAccounts.platform) + require.Equal(t, "oauth", adminSvc.lastListAccounts.accountType) + require.Equal(t, "active", adminSvc.lastListAccounts.status) + require.Equal(t, int64(12), adminSvc.lastListAccounts.groupID) + require.Equal(t, "blocked", adminSvc.lastListAccounts.privacyMode) + require.Equal(t, "keyword", adminSvc.lastListAccounts.search) + require.Equal(t, "priority", adminSvc.lastListAccounts.sortBy) + require.Equal(t, "desc", adminSvc.lastListAccounts.sortOrder) +} + +func TestExportDataSelectedIDsOverrideFilters(t *testing.T) { + router, adminSvc := setupAccountDataRouter() + + rec := httptest.NewRecorder() + req := httptest.NewRequest( + http.MethodGet, + "/api/v1/admin/accounts/data?ids=1,2&platform=openai&search=keyword&sort_by=priority&sort_order=desc", + nil, + ) + router.ServeHTTP(rec, req) + require.Equal(t, http.StatusOK, rec.Code) + + var resp dataResponse + require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &resp)) + require.Equal(t, 0, resp.Code) + require.Len(t, resp.Data.Accounts, 2) + require.Equal(t, 0, adminSvc.lastListAccounts.calls) +} + func TestImportDataReusesProxyAndSkipsDefaultGroup(t *testing.T) { router, adminSvc := setupAccountDataRouter() diff --git a/backend/internal/handler/admin/account_handler.go b/backend/internal/handler/admin/account_handler.go index 9a16f39433..9883d007bb 100644 --- a/backend/internal/handler/admin/account_handler.go +++ b/backend/internal/handler/admin/account_handler.go @@ -221,6 +221,8 @@ func (h *AccountHandler) List(c *gin.Context) { status := c.Query("status") search := c.Query("search") privacyMode := strings.TrimSpace(c.Query("privacy_mode")) + sortBy := c.DefaultQuery("sort_by", "name") + sortOrder := c.DefaultQuery("sort_order", "asc") // 标准化和验证 search 参数 search = strings.TrimSpace(search) if len(search) > 100 { @@ -246,7 +248,7 @@ func (h *AccountHandler) List(c *gin.Context) { } } - accounts, total, err := h.adminService.ListAccounts(c.Request.Context(), page, pageSize, platform, accountType, status, search, groupID, privacyMode) + accounts, total, err := h.adminService.ListAccounts(c.Request.Context(), page, pageSize, platform, accountType, status, search, groupID, privacyMode, sortBy, sortOrder) if err != nil { response.ErrorFrom(c, err) return @@ -2035,7 +2037,7 @@ func (h *AccountHandler) BatchRefreshTier(c *gin.Context) { accounts := make([]*service.Account, 0) if len(req.AccountIDs) == 0 { - allAccounts, _, err := h.adminService.ListAccounts(ctx, 1, 10000, "gemini", "oauth", "", "", 0, "") + allAccounts, _, err := h.adminService.ListAccounts(ctx, 1, 10000, "gemini", "oauth", "", "", 0, "", "name", "asc") if err != nil { response.ErrorFrom(c, err) return diff --git a/backend/internal/handler/admin/admin_service_stub_test.go b/backend/internal/handler/admin/admin_service_stub_test.go index 60d68913e8..6d1ef1b6b8 100644 --- a/backend/internal/handler/admin/admin_service_stub_test.go +++ b/backend/internal/handler/admin/admin_service_stub_test.go @@ -31,6 +31,33 @@ type stubAdminService struct { platform string groupIDs []int64 } + lastListAccounts struct { + platform string + accountType string + status string + search string + groupID int64 + privacyMode string + sortBy string + sortOrder string + calls int + } + lastListProxies struct { + protocol string + status string + search string + sortBy string + sortOrder string + calls int + } + lastListRedeemCodes struct { + codeType string + status string + search string + sortBy string + sortOrder string + calls int + } mu sync.Mutex } @@ -99,7 +126,7 @@ func newStubAdminService() *stubAdminService { } } -func (s *stubAdminService) ListUsers(ctx context.Context, page, pageSize int, filters service.UserListFilters) ([]service.User, int64, error) { +func (s *stubAdminService) ListUsers(ctx context.Context, page, pageSize int, filters service.UserListFilters, sortBy, sortOrder string) ([]service.User, int64, error) { return s.users, int64(len(s.users)), nil } @@ -132,7 +159,7 @@ func (s *stubAdminService) UpdateUserBalance(ctx context.Context, userID int64, return &user, nil } -func (s *stubAdminService) GetUserAPIKeys(ctx context.Context, userID int64, page, pageSize int) ([]service.APIKey, int64, error) { +func (s *stubAdminService) GetUserAPIKeys(ctx context.Context, userID int64, page, pageSize int, sortBy, sortOrder string) ([]service.APIKey, int64, error) { return s.apiKeys, int64(len(s.apiKeys)), nil } @@ -140,7 +167,7 @@ func (s *stubAdminService) GetUserUsageStats(ctx context.Context, userID int64, return map[string]any{"user_id": userID}, nil } -func (s *stubAdminService) ListGroups(ctx context.Context, page, pageSize int, platform, status, search string, isExclusive *bool) ([]service.Group, int64, error) { +func (s *stubAdminService) ListGroups(ctx context.Context, page, pageSize int, platform, status, search string, isExclusive *bool, sortBy, sortOrder string) ([]service.Group, int64, error) { return s.groups, int64(len(s.groups)), nil } @@ -187,7 +214,16 @@ func (s *stubAdminService) BatchSetGroupRateMultipliers(_ context.Context, _ int return nil } -func (s *stubAdminService) ListAccounts(ctx context.Context, page, pageSize int, platform, accountType, status, search string, groupID int64, privacyMode string) ([]service.Account, int64, error) { +func (s *stubAdminService) ListAccounts(ctx context.Context, page, pageSize int, platform, accountType, status, search string, groupID int64, privacyMode string, sortBy, sortOrder string) ([]service.Account, int64, error) { + s.lastListAccounts.platform = platform + s.lastListAccounts.accountType = accountType + s.lastListAccounts.status = status + s.lastListAccounts.search = search + s.lastListAccounts.groupID = groupID + s.lastListAccounts.privacyMode = privacyMode + s.lastListAccounts.sortBy = sortBy + s.lastListAccounts.sortOrder = sortOrder + s.lastListAccounts.calls++ return s.accounts, int64(len(s.accounts)), nil } @@ -261,7 +297,13 @@ func (s *stubAdminService) CheckMixedChannelRisk(ctx context.Context, currentAcc return s.checkMixedErr } -func (s *stubAdminService) ListProxies(ctx context.Context, page, pageSize int, protocol, status, search string) ([]service.Proxy, int64, error) { +func (s *stubAdminService) ListProxies(ctx context.Context, page, pageSize int, protocol, status, search string, sortBy, sortOrder string) ([]service.Proxy, int64, error) { + s.lastListProxies.protocol = protocol + s.lastListProxies.status = status + s.lastListProxies.search = search + s.lastListProxies.sortBy = sortBy + s.lastListProxies.sortOrder = sortOrder + s.lastListProxies.calls++ search = strings.TrimSpace(strings.ToLower(search)) filtered := make([]service.Proxy, 0, len(s.proxies)) for _, proxy := range s.proxies { @@ -283,7 +325,7 @@ func (s *stubAdminService) ListProxies(ctx context.Context, page, pageSize int, return filtered, int64(len(filtered)), nil } -func (s *stubAdminService) ListProxiesWithAccountCount(ctx context.Context, page, pageSize int, protocol, status, search string) ([]service.ProxyWithAccountCount, int64, error) { +func (s *stubAdminService) ListProxiesWithAccountCount(ctx context.Context, page, pageSize int, protocol, status, search string, sortBy, sortOrder string) ([]service.ProxyWithAccountCount, int64, error) { return s.proxyCounts, int64(len(s.proxyCounts)), nil } @@ -384,7 +426,13 @@ func (s *stubAdminService) CheckProxyQuality(ctx context.Context, id int64) (*se }, nil } -func (s *stubAdminService) ListRedeemCodes(ctx context.Context, page, pageSize int, codeType, status, search string) ([]service.RedeemCode, int64, error) { +func (s *stubAdminService) ListRedeemCodes(ctx context.Context, page, pageSize int, codeType, status, search string, sortBy, sortOrder string) ([]service.RedeemCode, int64, error) { + s.lastListRedeemCodes.codeType = codeType + s.lastListRedeemCodes.status = status + s.lastListRedeemCodes.search = search + s.lastListRedeemCodes.sortBy = sortBy + s.lastListRedeemCodes.sortOrder = sortOrder + s.lastListRedeemCodes.calls++ return s.redeems, int64(len(s.redeems)), nil } diff --git a/backend/internal/handler/admin/announcement_handler.go b/backend/internal/handler/admin/announcement_handler.go index d1312bc0c7..d3b9d17373 100644 --- a/backend/internal/handler/admin/announcement_handler.go +++ b/backend/internal/handler/admin/announcement_handler.go @@ -52,13 +52,17 @@ func (h *AnnouncementHandler) List(c *gin.Context) { page, pageSize := response.ParsePagination(c) status := strings.TrimSpace(c.Query("status")) search := strings.TrimSpace(c.Query("search")) + sortBy := c.DefaultQuery("sort_by", "created_at") + sortOrder := c.DefaultQuery("sort_order", "desc") if len(search) > 200 { search = search[:200] } params := pagination.PaginationParams{ - Page: page, - PageSize: pageSize, + Page: page, + PageSize: pageSize, + SortBy: sortBy, + SortOrder: sortOrder, } items, paginationResult, err := h.announcementService.List( @@ -227,8 +231,10 @@ func (h *AnnouncementHandler) ListReadStatus(c *gin.Context) { page, pageSize := response.ParsePagination(c) params := pagination.PaginationParams{ - Page: page, - PageSize: pageSize, + Page: page, + PageSize: pageSize, + SortBy: c.DefaultQuery("sort_by", "email"), + SortOrder: c.DefaultQuery("sort_order", "asc"), } search := strings.TrimSpace(c.Query("search")) if len(search) > 200 { diff --git a/backend/internal/handler/admin/announcement_handler_sort_test.go b/backend/internal/handler/admin/announcement_handler_sort_test.go new file mode 100644 index 0000000000..545e619e43 --- /dev/null +++ b/backend/internal/handler/admin/announcement_handler_sort_test.go @@ -0,0 +1,138 @@ +package admin + +import ( + "context" + "net/http" + "net/http/httptest" + "testing" + "time" + + "github.com/Wei-Shaw/sub2api/internal/pkg/pagination" + "github.com/Wei-Shaw/sub2api/internal/service" + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" +) + +type announcementRepoCapture struct { + service.AnnouncementRepository + listParams pagination.PaginationParams +} + +func (r *announcementRepoCapture) List(ctx context.Context, params pagination.PaginationParams, filters service.AnnouncementListFilters) ([]service.Announcement, *pagination.PaginationResult, error) { + r.listParams = params + return []service.Announcement{}, &pagination.PaginationResult{ + Total: 0, + Page: params.Page, + PageSize: params.PageSize, + Pages: 0, + }, nil +} + +func (r *announcementRepoCapture) GetByID(ctx context.Context, id int64) (*service.Announcement, error) { + return &service.Announcement{ + ID: id, + Title: "announcement", + Content: "content", + Status: service.AnnouncementStatusActive, + CreatedAt: time.Now(), + UpdatedAt: time.Now(), + }, nil +} + +type announcementUserRepoCapture struct { + service.UserRepository + listParams pagination.PaginationParams +} + +func (r *announcementUserRepoCapture) ListWithFilters(ctx context.Context, params pagination.PaginationParams, filters service.UserListFilters) ([]service.User, *pagination.PaginationResult, error) { + r.listParams = params + return []service.User{}, &pagination.PaginationResult{ + Total: 0, + Page: params.Page, + PageSize: params.PageSize, + Pages: 0, + }, nil +} + +type announcementReadRepoCapture struct { + service.AnnouncementReadRepository +} + +func (r *announcementReadRepoCapture) GetReadMapByUsers(ctx context.Context, announcementID int64, userIDs []int64) (map[int64]time.Time, error) { + return map[int64]time.Time{}, nil +} + +type announcementUserSubRepoCapture struct { + service.UserSubscriptionRepository +} + +func newAnnouncementSortTestRouter(announcementRepo *announcementRepoCapture, userRepo *announcementUserRepoCapture) *gin.Engine { + gin.SetMode(gin.TestMode) + svc := service.NewAnnouncementService( + announcementRepo, + &announcementReadRepoCapture{}, + userRepo, + &announcementUserSubRepoCapture{}, + ) + handler := NewAnnouncementHandler(svc) + router := gin.New() + router.GET("/admin/announcements", handler.List) + router.GET("/admin/announcements/:id/read-status", handler.ListReadStatus) + return router +} + +func TestAdminAnnouncementListSortParams(t *testing.T) { + announcementRepo := &announcementRepoCapture{} + userRepo := &announcementUserRepoCapture{} + router := newAnnouncementSortTestRouter(announcementRepo, userRepo) + + req := httptest.NewRequest(http.MethodGet, "/admin/announcements?sort_by=title&sort_order=ASC", nil) + rec := httptest.NewRecorder() + router.ServeHTTP(rec, req) + + require.Equal(t, http.StatusOK, rec.Code) + require.Equal(t, "title", announcementRepo.listParams.SortBy) + require.Equal(t, "ASC", announcementRepo.listParams.SortOrder) +} + +func TestAdminAnnouncementListSortDefaults(t *testing.T) { + announcementRepo := &announcementRepoCapture{} + userRepo := &announcementUserRepoCapture{} + router := newAnnouncementSortTestRouter(announcementRepo, userRepo) + + req := httptest.NewRequest(http.MethodGet, "/admin/announcements", nil) + rec := httptest.NewRecorder() + router.ServeHTTP(rec, req) + + require.Equal(t, http.StatusOK, rec.Code) + require.Equal(t, "created_at", announcementRepo.listParams.SortBy) + require.Equal(t, "desc", announcementRepo.listParams.SortOrder) +} + +func TestAdminAnnouncementReadStatusSortParams(t *testing.T) { + announcementRepo := &announcementRepoCapture{} + userRepo := &announcementUserRepoCapture{} + router := newAnnouncementSortTestRouter(announcementRepo, userRepo) + + req := httptest.NewRequest(http.MethodGet, "/admin/announcements/1/read-status?sort_by=balance&sort_order=DESC", nil) + rec := httptest.NewRecorder() + router.ServeHTTP(rec, req) + + require.Equal(t, http.StatusOK, rec.Code) + require.Equal(t, "balance", userRepo.listParams.SortBy) + require.Equal(t, "DESC", userRepo.listParams.SortOrder) +} + +func TestAdminAnnouncementReadStatusSortDefaults(t *testing.T) { + announcementRepo := &announcementRepoCapture{} + userRepo := &announcementUserRepoCapture{} + router := newAnnouncementSortTestRouter(announcementRepo, userRepo) + + req := httptest.NewRequest(http.MethodGet, "/admin/announcements/1/read-status", nil) + rec := httptest.NewRecorder() + router.ServeHTTP(rec, req) + + require.Equal(t, http.StatusOK, rec.Code) + require.Equal(t, "email", userRepo.listParams.SortBy) + require.Equal(t, "asc", userRepo.listParams.SortOrder) +} diff --git a/backend/internal/handler/admin/channel_handler.go b/backend/internal/handler/admin/channel_handler.go index f08f45577f..d6022283dc 100644 --- a/backend/internal/handler/admin/channel_handler.go +++ b/backend/internal/handler/admin/channel_handler.go @@ -249,7 +249,12 @@ func (h *ChannelHandler) List(c *gin.Context) { search = search[:100] } - channels, pag, err := h.channelService.List(c.Request.Context(), pagination.PaginationParams{Page: page, PageSize: pageSize}, status, search) + channels, pag, err := h.channelService.List(c.Request.Context(), pagination.PaginationParams{ + Page: page, + PageSize: pageSize, + SortBy: c.DefaultQuery("sort_by", "created_at"), + SortOrder: c.DefaultQuery("sort_order", "desc"), + }, status, search) if err != nil { response.ErrorFrom(c, err) return diff --git a/backend/internal/handler/admin/group_handler.go b/backend/internal/handler/admin/group_handler.go index 458ed35d47..cb2bd2018e 100644 --- a/backend/internal/handler/admin/group_handler.go +++ b/backend/internal/handler/admin/group_handler.go @@ -105,10 +105,11 @@ type CreateGroupRequest struct { // 支持的模型系列(仅 antigravity 平台使用) 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"` + AllowMessagesDispatch bool `json:"allow_messages_dispatch"` + RequireOAuthOnly bool `json:"require_oauth_only"` + RequirePrivacySet bool `json:"require_privacy_set"` + DefaultMappedModel string `json:"default_mapped_model"` + MessagesDispatchModelConfig service.OpenAIMessagesDispatchModelConfig `json:"messages_dispatch_model_config"` // 从指定分组复制账号(创建后自动绑定) CopyAccountsFromGroupIDs []int64 `json:"copy_accounts_from_group_ids"` } @@ -139,10 +140,11 @@ type UpdateGroupRequest struct { // 支持的模型系列(仅 antigravity 平台使用) 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"` + AllowMessagesDispatch *bool `json:"allow_messages_dispatch"` + RequireOAuthOnly *bool `json:"require_oauth_only"` + RequirePrivacySet *bool `json:"require_privacy_set"` + DefaultMappedModel *string `json:"default_mapped_model"` + MessagesDispatchModelConfig *service.OpenAIMessagesDispatchModelConfig `json:"messages_dispatch_model_config"` // 从指定分组复制账号(同步操作:先清空当前分组的账号绑定,再绑定源分组的账号) CopyAccountsFromGroupIDs []int64 `json:"copy_accounts_from_group_ids"` } @@ -160,6 +162,8 @@ func (h *GroupHandler) List(c *gin.Context) { search = search[:100] } isExclusiveStr := c.Query("is_exclusive") + sortBy := c.DefaultQuery("sort_by", "sort_order") + sortOrder := c.DefaultQuery("sort_order", "asc") var isExclusive *bool if isExclusiveStr != "" { @@ -167,7 +171,7 @@ func (h *GroupHandler) List(c *gin.Context) { isExclusive = &val } - groups, total, err := h.adminService.ListGroups(c.Request.Context(), page, pageSize, platform, status, search, isExclusive) + groups, total, err := h.adminService.ListGroups(c.Request.Context(), page, pageSize, platform, status, search, isExclusive, sortBy, sortOrder) if err != nil { response.ErrorFrom(c, err) return @@ -257,6 +261,7 @@ func (h *GroupHandler) Create(c *gin.Context) { RequireOAuthOnly: req.RequireOAuthOnly, RequirePrivacySet: req.RequirePrivacySet, DefaultMappedModel: req.DefaultMappedModel, + MessagesDispatchModelConfig: req.MessagesDispatchModelConfig, CopyAccountsFromGroupIDs: req.CopyAccountsFromGroupIDs, }) if err != nil { @@ -307,6 +312,7 @@ func (h *GroupHandler) Update(c *gin.Context) { RequireOAuthOnly: req.RequireOAuthOnly, RequirePrivacySet: req.RequirePrivacySet, DefaultMappedModel: req.DefaultMappedModel, + MessagesDispatchModelConfig: req.MessagesDispatchModelConfig, CopyAccountsFromGroupIDs: req.CopyAccountsFromGroupIDs, }) if err != nil { diff --git a/backend/internal/handler/admin/promo_handler.go b/backend/internal/handler/admin/promo_handler.go index 3eafa3801a..77d5f17165 100644 --- a/backend/internal/handler/admin/promo_handler.go +++ b/backend/internal/handler/admin/promo_handler.go @@ -55,8 +55,10 @@ func (h *PromoHandler) List(c *gin.Context) { } params := pagination.PaginationParams{ - Page: page, - PageSize: pageSize, + Page: page, + PageSize: pageSize, + SortBy: c.DefaultQuery("sort_by", "created_at"), + SortOrder: c.DefaultQuery("sort_order", "desc"), } codes, paginationResult, err := h.promoService.List(c.Request.Context(), params, status, search) diff --git a/backend/internal/handler/admin/proxy_data.go b/backend/internal/handler/admin/proxy_data.go index 72ecd6c131..8149ce3b3c 100644 --- a/backend/internal/handler/admin/proxy_data.go +++ b/backend/internal/handler/admin/proxy_data.go @@ -33,11 +33,13 @@ func (h *ProxyHandler) ExportData(c *gin.Context) { protocol := c.Query("protocol") status := c.Query("status") search := strings.TrimSpace(c.Query("search")) + sortBy := c.DefaultQuery("sort_by", "id") + sortOrder := c.DefaultQuery("sort_order", "desc") if len(search) > 100 { search = search[:100] } - proxies, err = h.listProxiesFiltered(ctx, protocol, status, search) + proxies, err = h.listProxiesFiltered(ctx, protocol, status, search, sortBy, sortOrder) if err != nil { response.ErrorFrom(c, err) return @@ -89,7 +91,7 @@ func (h *ProxyHandler) ImportData(c *gin.Context) { ctx := c.Request.Context() result := DataImportResult{} - existingProxies, err := h.listProxiesFiltered(ctx, "", "", "") + existingProxies, err := h.listProxiesFiltered(ctx, "", "", "", "id", "desc") if err != nil { response.ErrorFrom(c, err) return @@ -220,18 +222,33 @@ func parseProxyIDs(c *gin.Context) ([]int64, error) { return ids, nil } -func (h *ProxyHandler) listProxiesFiltered(ctx context.Context, protocol, status, search string) ([]service.Proxy, error) { +func (h *ProxyHandler) listProxiesFiltered(ctx context.Context, protocol, status, search, sortBy, sortOrder string) ([]service.Proxy, error) { page := 1 pageSize := dataPageCap var out []service.Proxy + sortBy = strings.TrimSpace(sortBy) + useAccountCountSort := strings.EqualFold(sortBy, "account_count") for { - items, total, err := h.adminService.ListProxies(ctx, page, pageSize, protocol, status, search) - if err != nil { - return nil, err - } - out = append(out, items...) - if len(out) >= int(total) || len(items) == 0 { - break + if useAccountCountSort { + items, total, err := h.adminService.ListProxiesWithAccountCount(ctx, page, pageSize, protocol, status, search, sortBy, sortOrder) + if err != nil { + return nil, err + } + for i := range items { + out = append(out, items[i].Proxy) + } + if len(out) >= int(total) || len(items) == 0 { + break + } + } else { + items, total, err := h.adminService.ListProxies(ctx, page, pageSize, protocol, status, search, sortBy, sortOrder) + if err != nil { + return nil, err + } + out = append(out, items...) + if len(out) >= int(total) || len(items) == 0 { + break + } } page++ } diff --git a/backend/internal/handler/admin/proxy_data_handler_test.go b/backend/internal/handler/admin/proxy_data_handler_test.go index 803f9b6135..8cd035ed3a 100644 --- a/backend/internal/handler/admin/proxy_data_handler_test.go +++ b/backend/internal/handler/admin/proxy_data_handler_test.go @@ -74,6 +74,10 @@ func TestProxyExportDataRespectsFilters(t *testing.T) { require.Len(t, resp.Data.Proxies, 1) require.Len(t, resp.Data.Accounts, 0) require.Equal(t, "https", resp.Data.Proxies[0].Protocol) + require.Equal(t, 1, adminSvc.lastListProxies.calls) + require.Equal(t, "https", adminSvc.lastListProxies.protocol) + require.Equal(t, "id", adminSvc.lastListProxies.sortBy) + require.Equal(t, "desc", adminSvc.lastListProxies.sortOrder) } func TestProxyExportDataWithSelectedIDs(t *testing.T) { @@ -113,6 +117,96 @@ func TestProxyExportDataWithSelectedIDs(t *testing.T) { require.Len(t, resp.Data.Proxies, 1) require.Equal(t, "https", resp.Data.Proxies[0].Protocol) require.Equal(t, "10.0.0.2", resp.Data.Proxies[0].Host) + require.Equal(t, 0, adminSvc.lastListProxies.calls) +} + +func TestProxyExportDataPassesSortParams(t *testing.T) { + router, adminSvc := setupProxyDataRouter() + + adminSvc.proxies = []service.Proxy{ + { + ID: 1, + Name: "proxy-a", + Protocol: "http", + Host: "127.0.0.1", + Port: 8080, + Username: "user", + Password: "pass", + Status: service.StatusActive, + }, + } + + rec := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodGet, "/api/v1/admin/proxies/data?protocol=http&status=active&search=proxy&sort_by=name&sort_order=asc", nil) + router.ServeHTTP(rec, req) + require.Equal(t, http.StatusOK, rec.Code) + + require.Equal(t, 1, adminSvc.lastListProxies.calls) + require.Equal(t, "http", adminSvc.lastListProxies.protocol) + require.Equal(t, "active", adminSvc.lastListProxies.status) + require.Equal(t, "proxy", adminSvc.lastListProxies.search) + require.Equal(t, "name", adminSvc.lastListProxies.sortBy) + require.Equal(t, "asc", adminSvc.lastListProxies.sortOrder) +} + +func TestProxyExportDataSortByAccountCountUsesAccountCountListing(t *testing.T) { + router, adminSvc := setupProxyDataRouter() + + adminSvc.proxies = []service.Proxy{ + { + ID: 1, + Name: "proxy-id-1", + Protocol: "http", + Host: "127.0.0.1", + Port: 8080, + Status: service.StatusActive, + }, + { + ID: 2, + Name: "proxy-id-2", + Protocol: "http", + Host: "127.0.0.2", + Port: 8081, + Status: service.StatusActive, + }, + } + adminSvc.proxyCounts = []service.ProxyWithAccountCount{ + { + Proxy: service.Proxy{ + ID: 2, + Name: "proxy-count-high", + Protocol: "http", + Host: "127.0.0.2", + Port: 8081, + Status: service.StatusActive, + }, + AccountCount: 9, + }, + { + Proxy: service.Proxy{ + ID: 1, + Name: "proxy-count-low", + Protocol: "http", + Host: "127.0.0.1", + Port: 8080, + Status: service.StatusActive, + }, + AccountCount: 1, + }, + } + + rec := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodGet, "/api/v1/admin/proxies/data?sort_by=account_count&sort_order=desc", nil) + router.ServeHTTP(rec, req) + require.Equal(t, http.StatusOK, rec.Code) + + var resp proxyDataResponse + require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &resp)) + require.Equal(t, 0, resp.Code) + require.Len(t, resp.Data.Proxies, 2) + require.Equal(t, "proxy-count-high", resp.Data.Proxies[0].Name) + require.Equal(t, "proxy-count-low", resp.Data.Proxies[1].Name) + require.Equal(t, 0, adminSvc.lastListProxies.calls) } func TestProxyImportDataReusesAndTriggersLatencyProbe(t *testing.T) { diff --git a/backend/internal/handler/admin/proxy_handler.go b/backend/internal/handler/admin/proxy_handler.go index e8ae0ce2d0..f97fcb0a71 100644 --- a/backend/internal/handler/admin/proxy_handler.go +++ b/backend/internal/handler/admin/proxy_handler.go @@ -52,13 +52,15 @@ func (h *ProxyHandler) List(c *gin.Context) { protocol := c.Query("protocol") status := c.Query("status") search := c.Query("search") + sortBy := c.DefaultQuery("sort_by", "id") + sortOrder := c.DefaultQuery("sort_order", "desc") // 标准化和验证 search 参数 search = strings.TrimSpace(search) if len(search) > 100 { search = search[:100] } - proxies, total, err := h.adminService.ListProxiesWithAccountCount(c.Request.Context(), page, pageSize, protocol, status, search) + proxies, total, err := h.adminService.ListProxiesWithAccountCount(c.Request.Context(), page, pageSize, protocol, status, search, sortBy, sortOrder) if err != nil { response.ErrorFrom(c, err) return diff --git a/backend/internal/handler/admin/redeem_export_handler_test.go b/backend/internal/handler/admin/redeem_export_handler_test.go new file mode 100644 index 0000000000..9983fe319c --- /dev/null +++ b/backend/internal/handler/admin/redeem_export_handler_test.go @@ -0,0 +1,49 @@ +package admin + +import ( + "net/http" + "net/http/httptest" + "testing" + + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" +) + +func setupRedeemExportRouter() (*gin.Engine, *stubAdminService) { + gin.SetMode(gin.TestMode) + router := gin.New() + adminSvc := newStubAdminService() + + h := NewRedeemHandler(adminSvc, nil) + router.GET("/api/v1/admin/redeem-codes/export", h.Export) + return router, adminSvc +} + +func TestRedeemExportPassesSearchAndSort(t *testing.T) { + router, adminSvc := setupRedeemExportRouter() + + rec := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodGet, "/api/v1/admin/redeem-codes/export?type=balance&status=unused&search=ABC&sort_by=value&sort_order=asc", nil) + router.ServeHTTP(rec, req) + require.Equal(t, http.StatusOK, rec.Code) + + require.Equal(t, 1, adminSvc.lastListRedeemCodes.calls) + require.Equal(t, "balance", adminSvc.lastListRedeemCodes.codeType) + require.Equal(t, "unused", adminSvc.lastListRedeemCodes.status) + require.Equal(t, "ABC", adminSvc.lastListRedeemCodes.search) + require.Equal(t, "value", adminSvc.lastListRedeemCodes.sortBy) + require.Equal(t, "asc", adminSvc.lastListRedeemCodes.sortOrder) +} + +func TestRedeemExportSortDefaults(t *testing.T) { + router, adminSvc := setupRedeemExportRouter() + + rec := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodGet, "/api/v1/admin/redeem-codes/export", nil) + router.ServeHTTP(rec, req) + require.Equal(t, http.StatusOK, rec.Code) + + require.Equal(t, 1, adminSvc.lastListRedeemCodes.calls) + require.Equal(t, "id", adminSvc.lastListRedeemCodes.sortBy) + require.Equal(t, "desc", adminSvc.lastListRedeemCodes.sortOrder) +} diff --git a/backend/internal/handler/admin/redeem_handler.go b/backend/internal/handler/admin/redeem_handler.go index c494e5fb8b..24365f3da4 100644 --- a/backend/internal/handler/admin/redeem_handler.go +++ b/backend/internal/handler/admin/redeem_handler.go @@ -59,13 +59,15 @@ func (h *RedeemHandler) List(c *gin.Context) { codeType := c.Query("type") status := c.Query("status") search := c.Query("search") + sortBy := c.DefaultQuery("sort_by", "id") + sortOrder := c.DefaultQuery("sort_order", "desc") // 标准化和验证 search 参数 search = strings.TrimSpace(search) if len(search) > 100 { search = search[:100] } - codes, total, err := h.adminService.ListRedeemCodes(c.Request.Context(), page, pageSize, codeType, status, search) + codes, total, err := h.adminService.ListRedeemCodes(c.Request.Context(), page, pageSize, codeType, status, search, sortBy, sortOrder) if err != nil { response.ErrorFrom(c, err) return @@ -300,9 +302,15 @@ func (h *RedeemHandler) GetStats(c *gin.Context) { func (h *RedeemHandler) Export(c *gin.Context) { codeType := c.Query("type") status := c.Query("status") + search := strings.TrimSpace(c.Query("search")) + sortBy := c.DefaultQuery("sort_by", "id") + sortOrder := c.DefaultQuery("sort_order", "desc") + if len(search) > 100 { + search = search[:100] + } // Get all codes without pagination (use large page size) - codes, _, err := h.adminService.ListRedeemCodes(c.Request.Context(), 1, 10000, codeType, status, "") + codes, _, err := h.adminService.ListRedeemCodes(c.Request.Context(), 1, 10000, codeType, status, search, sortBy, sortOrder) if err != nil { response.ErrorFrom(c, err) return diff --git a/backend/internal/handler/admin/setting_handler.go b/backend/internal/handler/admin/setting_handler.go index 41df5f3ef1..ba7511315c 100644 --- a/backend/internal/handler/admin/setting_handler.go +++ b/backend/internal/handler/admin/setting_handler.go @@ -35,6 +35,15 @@ func generateMenuItemID() (string, error) { return hex.EncodeToString(b), nil } +func scopesContainOpenID(scopes string) bool { + for _, scope := range strings.Fields(strings.ToLower(strings.TrimSpace(scopes))) { + if scope == "openid" { + return true + } + } + return false +} + // SettingHandler 系统设置处理器 type SettingHandler struct { settingService *service.SettingService @@ -109,6 +118,28 @@ func (h *SettingHandler) GetSettings(c *gin.Context) { LinuxDoConnectClientID: settings.LinuxDoConnectClientID, LinuxDoConnectClientSecretConfigured: settings.LinuxDoConnectClientSecretConfigured, LinuxDoConnectRedirectURL: settings.LinuxDoConnectRedirectURL, + OIDCConnectEnabled: settings.OIDCConnectEnabled, + OIDCConnectProviderName: settings.OIDCConnectProviderName, + OIDCConnectClientID: settings.OIDCConnectClientID, + OIDCConnectClientSecretConfigured: settings.OIDCConnectClientSecretConfigured, + OIDCConnectIssuerURL: settings.OIDCConnectIssuerURL, + OIDCConnectDiscoveryURL: settings.OIDCConnectDiscoveryURL, + OIDCConnectAuthorizeURL: settings.OIDCConnectAuthorizeURL, + OIDCConnectTokenURL: settings.OIDCConnectTokenURL, + OIDCConnectUserInfoURL: settings.OIDCConnectUserInfoURL, + OIDCConnectJWKSURL: settings.OIDCConnectJWKSURL, + OIDCConnectScopes: settings.OIDCConnectScopes, + OIDCConnectRedirectURL: settings.OIDCConnectRedirectURL, + OIDCConnectFrontendRedirectURL: settings.OIDCConnectFrontendRedirectURL, + OIDCConnectTokenAuthMethod: settings.OIDCConnectTokenAuthMethod, + OIDCConnectUsePKCE: settings.OIDCConnectUsePKCE, + OIDCConnectValidateIDToken: settings.OIDCConnectValidateIDToken, + OIDCConnectAllowedSigningAlgs: settings.OIDCConnectAllowedSigningAlgs, + OIDCConnectClockSkewSeconds: settings.OIDCConnectClockSkewSeconds, + OIDCConnectRequireEmailVerified: settings.OIDCConnectRequireEmailVerified, + OIDCConnectUserInfoEmailPath: settings.OIDCConnectUserInfoEmailPath, + OIDCConnectUserInfoIDPath: settings.OIDCConnectUserInfoIDPath, + OIDCConnectUserInfoUsernamePath: settings.OIDCConnectUserInfoUsernamePath, SiteName: settings.SiteName, SiteLogo: settings.SiteLogo, SiteSubtitle: settings.SiteSubtitle, @@ -119,6 +150,8 @@ func (h *SettingHandler) GetSettings(c *gin.Context) { HideCcsImportButton: settings.HideCcsImportButton, PurchaseSubscriptionEnabled: settings.PurchaseSubscriptionEnabled, PurchaseSubscriptionURL: settings.PurchaseSubscriptionURL, + TableDefaultPageSize: settings.TableDefaultPageSize, + TablePageSizeOptions: settings.TablePageSizeOptions, CustomMenuItems: dto.ParseCustomMenuItems(settings.CustomMenuItems), CustomEndpoints: dto.ParseCustomEndpoints(settings.CustomEndpoints), DefaultConcurrency: settings.DefaultConcurrency, @@ -141,6 +174,7 @@ func (h *SettingHandler) GetSettings(c *gin.Context) { BackendModeEnabled: settings.BackendModeEnabled, EnableFingerprintUnification: settings.EnableFingerprintUnification, EnableMetadataPassthrough: settings.EnableMetadataPassthrough, + EnableCCHSigning: settings.EnableCCHSigning, PaymentEnabled: paymentCfg.Enabled, PaymentMinAmount: paymentCfg.MinAmount, PaymentMaxAmount: paymentCfg.MaxAmount, @@ -194,6 +228,30 @@ type UpdateSettingsRequest struct { LinuxDoConnectClientSecret string `json:"linuxdo_connect_client_secret"` LinuxDoConnectRedirectURL string `json:"linuxdo_connect_redirect_url"` + // Generic OIDC OAuth 登录 + OIDCConnectEnabled bool `json:"oidc_connect_enabled"` + OIDCConnectProviderName string `json:"oidc_connect_provider_name"` + OIDCConnectClientID string `json:"oidc_connect_client_id"` + OIDCConnectClientSecret string `json:"oidc_connect_client_secret"` + OIDCConnectIssuerURL string `json:"oidc_connect_issuer_url"` + OIDCConnectDiscoveryURL string `json:"oidc_connect_discovery_url"` + OIDCConnectAuthorizeURL string `json:"oidc_connect_authorize_url"` + OIDCConnectTokenURL string `json:"oidc_connect_token_url"` + OIDCConnectUserInfoURL string `json:"oidc_connect_userinfo_url"` + OIDCConnectJWKSURL string `json:"oidc_connect_jwks_url"` + OIDCConnectScopes string `json:"oidc_connect_scopes"` + OIDCConnectRedirectURL string `json:"oidc_connect_redirect_url"` + OIDCConnectFrontendRedirectURL string `json:"oidc_connect_frontend_redirect_url"` + OIDCConnectTokenAuthMethod string `json:"oidc_connect_token_auth_method"` + OIDCConnectUsePKCE bool `json:"oidc_connect_use_pkce"` + OIDCConnectValidateIDToken bool `json:"oidc_connect_validate_id_token"` + OIDCConnectAllowedSigningAlgs string `json:"oidc_connect_allowed_signing_algs"` + OIDCConnectClockSkewSeconds int `json:"oidc_connect_clock_skew_seconds"` + OIDCConnectRequireEmailVerified bool `json:"oidc_connect_require_email_verified"` + OIDCConnectUserInfoEmailPath string `json:"oidc_connect_userinfo_email_path"` + OIDCConnectUserInfoIDPath string `json:"oidc_connect_userinfo_id_path"` + OIDCConnectUserInfoUsernamePath string `json:"oidc_connect_userinfo_username_path"` + // OEM设置 SiteName string `json:"site_name"` SiteLogo string `json:"site_logo"` @@ -205,6 +263,8 @@ type UpdateSettingsRequest struct { HideCcsImportButton bool `json:"hide_ccs_import_button"` PurchaseSubscriptionEnabled *bool `json:"purchase_subscription_enabled"` PurchaseSubscriptionURL *string `json:"purchase_subscription_url"` + TableDefaultPageSize int `json:"table_default_page_size"` + TablePageSizeOptions []int `json:"table_page_size_options"` CustomMenuItems *[]dto.CustomMenuItem `json:"custom_menu_items"` CustomEndpoints *[]dto.CustomEndpoint `json:"custom_endpoints"` @@ -242,6 +302,7 @@ type UpdateSettingsRequest struct { // Gateway forwarding behavior EnableFingerprintUnification *bool `json:"enable_fingerprint_unification"` EnableMetadataPassthrough *bool `json:"enable_metadata_passthrough"` + EnableCCHSigning *bool `json:"enable_cch_signing"` // Payment configuration (integrated into settings, full replace) PaymentEnabled *bool `json:"payment_enabled"` @@ -288,6 +349,13 @@ func (h *SettingHandler) UpdateSettings(c *gin.Context) { if req.DefaultBalance < 0 { req.DefaultBalance = 0 } + // 通用表格配置:兼容旧客户端未传字段时保留当前值。 + if req.TableDefaultPageSize <= 0 { + req.TableDefaultPageSize = previousSettings.TableDefaultPageSize + } + if req.TablePageSizeOptions == nil { + req.TablePageSizeOptions = previousSettings.TablePageSizeOptions + } req.SMTPHost = strings.TrimSpace(req.SMTPHost) req.SMTPUsername = strings.TrimSpace(req.SMTPUsername) req.SMTPPassword = strings.TrimSpace(req.SMTPPassword) @@ -375,6 +443,122 @@ func (h *SettingHandler) UpdateSettings(c *gin.Context) { } } + // Generic OIDC 参数验证 + if req.OIDCConnectEnabled { + req.OIDCConnectProviderName = strings.TrimSpace(req.OIDCConnectProviderName) + req.OIDCConnectClientID = strings.TrimSpace(req.OIDCConnectClientID) + req.OIDCConnectClientSecret = strings.TrimSpace(req.OIDCConnectClientSecret) + req.OIDCConnectIssuerURL = strings.TrimSpace(req.OIDCConnectIssuerURL) + req.OIDCConnectDiscoveryURL = strings.TrimSpace(req.OIDCConnectDiscoveryURL) + req.OIDCConnectAuthorizeURL = strings.TrimSpace(req.OIDCConnectAuthorizeURL) + req.OIDCConnectTokenURL = strings.TrimSpace(req.OIDCConnectTokenURL) + req.OIDCConnectUserInfoURL = strings.TrimSpace(req.OIDCConnectUserInfoURL) + req.OIDCConnectJWKSURL = strings.TrimSpace(req.OIDCConnectJWKSURL) + req.OIDCConnectScopes = strings.TrimSpace(req.OIDCConnectScopes) + req.OIDCConnectRedirectURL = strings.TrimSpace(req.OIDCConnectRedirectURL) + req.OIDCConnectFrontendRedirectURL = strings.TrimSpace(req.OIDCConnectFrontendRedirectURL) + req.OIDCConnectTokenAuthMethod = strings.ToLower(strings.TrimSpace(req.OIDCConnectTokenAuthMethod)) + req.OIDCConnectAllowedSigningAlgs = strings.TrimSpace(req.OIDCConnectAllowedSigningAlgs) + req.OIDCConnectUserInfoEmailPath = strings.TrimSpace(req.OIDCConnectUserInfoEmailPath) + req.OIDCConnectUserInfoIDPath = strings.TrimSpace(req.OIDCConnectUserInfoIDPath) + req.OIDCConnectUserInfoUsernamePath = strings.TrimSpace(req.OIDCConnectUserInfoUsernamePath) + + if req.OIDCConnectProviderName == "" { + req.OIDCConnectProviderName = "OIDC" + } + if req.OIDCConnectClientID == "" { + response.BadRequest(c, "OIDC Client ID is required when enabled") + return + } + if req.OIDCConnectIssuerURL == "" { + response.BadRequest(c, "OIDC Issuer URL is required when enabled") + return + } + if err := config.ValidateAbsoluteHTTPURL(req.OIDCConnectIssuerURL); err != nil { + response.BadRequest(c, "OIDC Issuer URL must be an absolute http(s) URL") + return + } + if req.OIDCConnectDiscoveryURL != "" { + if err := config.ValidateAbsoluteHTTPURL(req.OIDCConnectDiscoveryURL); err != nil { + response.BadRequest(c, "OIDC Discovery URL must be an absolute http(s) URL") + return + } + } + if req.OIDCConnectAuthorizeURL != "" { + if err := config.ValidateAbsoluteHTTPURL(req.OIDCConnectAuthorizeURL); err != nil { + response.BadRequest(c, "OIDC Authorize URL must be an absolute http(s) URL") + return + } + } + if req.OIDCConnectTokenURL != "" { + if err := config.ValidateAbsoluteHTTPURL(req.OIDCConnectTokenURL); err != nil { + response.BadRequest(c, "OIDC Token URL must be an absolute http(s) URL") + return + } + } + if req.OIDCConnectUserInfoURL != "" { + if err := config.ValidateAbsoluteHTTPURL(req.OIDCConnectUserInfoURL); err != nil { + response.BadRequest(c, "OIDC UserInfo URL must be an absolute http(s) URL") + return + } + } + if req.OIDCConnectRedirectURL == "" { + response.BadRequest(c, "OIDC Redirect URL is required when enabled") + return + } + if err := config.ValidateAbsoluteHTTPURL(req.OIDCConnectRedirectURL); err != nil { + response.BadRequest(c, "OIDC Redirect URL must be an absolute http(s) URL") + return + } + if req.OIDCConnectFrontendRedirectURL == "" { + response.BadRequest(c, "OIDC Frontend Redirect URL is required when enabled") + return + } + if err := config.ValidateFrontendRedirectURL(req.OIDCConnectFrontendRedirectURL); err != nil { + response.BadRequest(c, "OIDC Frontend Redirect URL is invalid") + return + } + if !scopesContainOpenID(req.OIDCConnectScopes) { + response.BadRequest(c, "OIDC scopes must contain openid") + return + } + switch req.OIDCConnectTokenAuthMethod { + case "", "client_secret_post", "client_secret_basic", "none": + default: + response.BadRequest(c, "OIDC Token Auth Method must be one of client_secret_post/client_secret_basic/none") + return + } + if req.OIDCConnectTokenAuthMethod == "none" && !req.OIDCConnectUsePKCE { + response.BadRequest(c, "OIDC PKCE must be enabled when token_auth_method=none") + return + } + if req.OIDCConnectClockSkewSeconds < 0 || req.OIDCConnectClockSkewSeconds > 600 { + response.BadRequest(c, "OIDC clock skew seconds must be between 0 and 600") + return + } + if req.OIDCConnectValidateIDToken { + if req.OIDCConnectAllowedSigningAlgs == "" { + response.BadRequest(c, "OIDC Allowed Signing Algs is required when validate_id_token=true") + return + } + } + if req.OIDCConnectJWKSURL != "" { + if err := config.ValidateAbsoluteHTTPURL(req.OIDCConnectJWKSURL); err != nil { + response.BadRequest(c, "OIDC JWKS URL must be an absolute http(s) URL") + return + } + } + if req.OIDCConnectTokenAuthMethod == "" || req.OIDCConnectTokenAuthMethod == "client_secret_post" || req.OIDCConnectTokenAuthMethod == "client_secret_basic" { + if req.OIDCConnectClientSecret == "" { + if previousSettings.OIDCConnectClientSecret == "" { + response.BadRequest(c, "OIDC Client Secret is required when enabled") + return + } + req.OIDCConnectClientSecret = previousSettings.OIDCConnectClientSecret + } + } + } + // “购买订阅”页面配置验证 purchaseEnabled := previousSettings.PurchaseSubscriptionEnabled if req.PurchaseSubscriptionEnabled != nil { @@ -605,6 +789,28 @@ func (h *SettingHandler) UpdateSettings(c *gin.Context) { LinuxDoConnectClientID: req.LinuxDoConnectClientID, LinuxDoConnectClientSecret: req.LinuxDoConnectClientSecret, LinuxDoConnectRedirectURL: req.LinuxDoConnectRedirectURL, + OIDCConnectEnabled: req.OIDCConnectEnabled, + OIDCConnectProviderName: req.OIDCConnectProviderName, + OIDCConnectClientID: req.OIDCConnectClientID, + OIDCConnectClientSecret: req.OIDCConnectClientSecret, + OIDCConnectIssuerURL: req.OIDCConnectIssuerURL, + OIDCConnectDiscoveryURL: req.OIDCConnectDiscoveryURL, + OIDCConnectAuthorizeURL: req.OIDCConnectAuthorizeURL, + OIDCConnectTokenURL: req.OIDCConnectTokenURL, + OIDCConnectUserInfoURL: req.OIDCConnectUserInfoURL, + OIDCConnectJWKSURL: req.OIDCConnectJWKSURL, + OIDCConnectScopes: req.OIDCConnectScopes, + OIDCConnectRedirectURL: req.OIDCConnectRedirectURL, + OIDCConnectFrontendRedirectURL: req.OIDCConnectFrontendRedirectURL, + OIDCConnectTokenAuthMethod: req.OIDCConnectTokenAuthMethod, + OIDCConnectUsePKCE: req.OIDCConnectUsePKCE, + OIDCConnectValidateIDToken: req.OIDCConnectValidateIDToken, + OIDCConnectAllowedSigningAlgs: req.OIDCConnectAllowedSigningAlgs, + OIDCConnectClockSkewSeconds: req.OIDCConnectClockSkewSeconds, + OIDCConnectRequireEmailVerified: req.OIDCConnectRequireEmailVerified, + OIDCConnectUserInfoEmailPath: req.OIDCConnectUserInfoEmailPath, + OIDCConnectUserInfoIDPath: req.OIDCConnectUserInfoIDPath, + OIDCConnectUserInfoUsernamePath: req.OIDCConnectUserInfoUsernamePath, SiteName: req.SiteName, SiteLogo: req.SiteLogo, SiteSubtitle: req.SiteSubtitle, @@ -615,6 +821,8 @@ func (h *SettingHandler) UpdateSettings(c *gin.Context) { HideCcsImportButton: req.HideCcsImportButton, PurchaseSubscriptionEnabled: purchaseEnabled, PurchaseSubscriptionURL: purchaseURL, + TableDefaultPageSize: req.TableDefaultPageSize, + TablePageSizeOptions: req.TablePageSizeOptions, CustomMenuItems: customMenuJSON, CustomEndpoints: customEndpointsJSON, DefaultConcurrency: req.DefaultConcurrency, @@ -667,6 +875,12 @@ func (h *SettingHandler) UpdateSettings(c *gin.Context) { } return previousSettings.EnableMetadataPassthrough }(), + EnableCCHSigning: func() bool { + if req.EnableCCHSigning != nil { + return *req.EnableCCHSigning + } + return previousSettings.EnableCCHSigning + }(), } if err := h.settingService.UpdateSettings(c.Request.Context(), settings); err != nil { @@ -756,6 +970,28 @@ func (h *SettingHandler) UpdateSettings(c *gin.Context) { LinuxDoConnectClientID: updatedSettings.LinuxDoConnectClientID, LinuxDoConnectClientSecretConfigured: updatedSettings.LinuxDoConnectClientSecretConfigured, LinuxDoConnectRedirectURL: updatedSettings.LinuxDoConnectRedirectURL, + OIDCConnectEnabled: updatedSettings.OIDCConnectEnabled, + OIDCConnectProviderName: updatedSettings.OIDCConnectProviderName, + OIDCConnectClientID: updatedSettings.OIDCConnectClientID, + OIDCConnectClientSecretConfigured: updatedSettings.OIDCConnectClientSecretConfigured, + OIDCConnectIssuerURL: updatedSettings.OIDCConnectIssuerURL, + OIDCConnectDiscoveryURL: updatedSettings.OIDCConnectDiscoveryURL, + OIDCConnectAuthorizeURL: updatedSettings.OIDCConnectAuthorizeURL, + OIDCConnectTokenURL: updatedSettings.OIDCConnectTokenURL, + OIDCConnectUserInfoURL: updatedSettings.OIDCConnectUserInfoURL, + OIDCConnectJWKSURL: updatedSettings.OIDCConnectJWKSURL, + OIDCConnectScopes: updatedSettings.OIDCConnectScopes, + OIDCConnectRedirectURL: updatedSettings.OIDCConnectRedirectURL, + OIDCConnectFrontendRedirectURL: updatedSettings.OIDCConnectFrontendRedirectURL, + OIDCConnectTokenAuthMethod: updatedSettings.OIDCConnectTokenAuthMethod, + OIDCConnectUsePKCE: updatedSettings.OIDCConnectUsePKCE, + OIDCConnectValidateIDToken: updatedSettings.OIDCConnectValidateIDToken, + OIDCConnectAllowedSigningAlgs: updatedSettings.OIDCConnectAllowedSigningAlgs, + OIDCConnectClockSkewSeconds: updatedSettings.OIDCConnectClockSkewSeconds, + OIDCConnectRequireEmailVerified: updatedSettings.OIDCConnectRequireEmailVerified, + OIDCConnectUserInfoEmailPath: updatedSettings.OIDCConnectUserInfoEmailPath, + OIDCConnectUserInfoIDPath: updatedSettings.OIDCConnectUserInfoIDPath, + OIDCConnectUserInfoUsernamePath: updatedSettings.OIDCConnectUserInfoUsernamePath, SiteName: updatedSettings.SiteName, SiteLogo: updatedSettings.SiteLogo, SiteSubtitle: updatedSettings.SiteSubtitle, @@ -766,6 +1002,8 @@ func (h *SettingHandler) UpdateSettings(c *gin.Context) { HideCcsImportButton: updatedSettings.HideCcsImportButton, PurchaseSubscriptionEnabled: updatedSettings.PurchaseSubscriptionEnabled, PurchaseSubscriptionURL: updatedSettings.PurchaseSubscriptionURL, + TableDefaultPageSize: updatedSettings.TableDefaultPageSize, + TablePageSizeOptions: updatedSettings.TablePageSizeOptions, CustomMenuItems: dto.ParseCustomMenuItems(updatedSettings.CustomMenuItems), CustomEndpoints: dto.ParseCustomEndpoints(updatedSettings.CustomEndpoints), DefaultConcurrency: updatedSettings.DefaultConcurrency, @@ -788,6 +1026,7 @@ func (h *SettingHandler) UpdateSettings(c *gin.Context) { BackendModeEnabled: updatedSettings.BackendModeEnabled, EnableFingerprintUnification: updatedSettings.EnableFingerprintUnification, EnableMetadataPassthrough: updatedSettings.EnableMetadataPassthrough, + EnableCCHSigning: updatedSettings.EnableCCHSigning, PaymentEnabled: updatedPaymentCfg.Enabled, PaymentMinAmount: updatedPaymentCfg.MinAmount, PaymentMaxAmount: updatedPaymentCfg.MaxAmount, @@ -904,6 +1143,72 @@ func diffSettings(before *service.SystemSettings, after *service.SystemSettings, if before.LinuxDoConnectRedirectURL != after.LinuxDoConnectRedirectURL { changed = append(changed, "linuxdo_connect_redirect_url") } + if before.OIDCConnectEnabled != after.OIDCConnectEnabled { + changed = append(changed, "oidc_connect_enabled") + } + if before.OIDCConnectProviderName != after.OIDCConnectProviderName { + changed = append(changed, "oidc_connect_provider_name") + } + if before.OIDCConnectClientID != after.OIDCConnectClientID { + changed = append(changed, "oidc_connect_client_id") + } + if req.OIDCConnectClientSecret != "" { + changed = append(changed, "oidc_connect_client_secret") + } + if before.OIDCConnectIssuerURL != after.OIDCConnectIssuerURL { + changed = append(changed, "oidc_connect_issuer_url") + } + if before.OIDCConnectDiscoveryURL != after.OIDCConnectDiscoveryURL { + changed = append(changed, "oidc_connect_discovery_url") + } + if before.OIDCConnectAuthorizeURL != after.OIDCConnectAuthorizeURL { + changed = append(changed, "oidc_connect_authorize_url") + } + if before.OIDCConnectTokenURL != after.OIDCConnectTokenURL { + changed = append(changed, "oidc_connect_token_url") + } + if before.OIDCConnectUserInfoURL != after.OIDCConnectUserInfoURL { + changed = append(changed, "oidc_connect_userinfo_url") + } + if before.OIDCConnectJWKSURL != after.OIDCConnectJWKSURL { + changed = append(changed, "oidc_connect_jwks_url") + } + if before.OIDCConnectScopes != after.OIDCConnectScopes { + changed = append(changed, "oidc_connect_scopes") + } + if before.OIDCConnectRedirectURL != after.OIDCConnectRedirectURL { + changed = append(changed, "oidc_connect_redirect_url") + } + if before.OIDCConnectFrontendRedirectURL != after.OIDCConnectFrontendRedirectURL { + changed = append(changed, "oidc_connect_frontend_redirect_url") + } + if before.OIDCConnectTokenAuthMethod != after.OIDCConnectTokenAuthMethod { + changed = append(changed, "oidc_connect_token_auth_method") + } + if before.OIDCConnectUsePKCE != after.OIDCConnectUsePKCE { + changed = append(changed, "oidc_connect_use_pkce") + } + if before.OIDCConnectValidateIDToken != after.OIDCConnectValidateIDToken { + changed = append(changed, "oidc_connect_validate_id_token") + } + if before.OIDCConnectAllowedSigningAlgs != after.OIDCConnectAllowedSigningAlgs { + changed = append(changed, "oidc_connect_allowed_signing_algs") + } + if before.OIDCConnectClockSkewSeconds != after.OIDCConnectClockSkewSeconds { + changed = append(changed, "oidc_connect_clock_skew_seconds") + } + if before.OIDCConnectRequireEmailVerified != after.OIDCConnectRequireEmailVerified { + changed = append(changed, "oidc_connect_require_email_verified") + } + if before.OIDCConnectUserInfoEmailPath != after.OIDCConnectUserInfoEmailPath { + changed = append(changed, "oidc_connect_userinfo_email_path") + } + if before.OIDCConnectUserInfoIDPath != after.OIDCConnectUserInfoIDPath { + changed = append(changed, "oidc_connect_userinfo_id_path") + } + if before.OIDCConnectUserInfoUsernamePath != after.OIDCConnectUserInfoUsernamePath { + changed = append(changed, "oidc_connect_userinfo_username_path") + } if before.SiteName != after.SiteName { changed = append(changed, "site_name") } @@ -988,6 +1293,12 @@ func diffSettings(before *service.SystemSettings, after *service.SystemSettings, if before.PurchaseSubscriptionURL != after.PurchaseSubscriptionURL { changed = append(changed, "purchase_subscription_url") } + if before.TableDefaultPageSize != after.TableDefaultPageSize { + changed = append(changed, "table_default_page_size") + } + if !equalIntSlice(before.TablePageSizeOptions, after.TablePageSizeOptions) { + changed = append(changed, "table_page_size_options") + } if before.CustomMenuItems != after.CustomMenuItems { changed = append(changed, "custom_menu_items") } @@ -997,6 +1308,9 @@ func diffSettings(before *service.SystemSettings, after *service.SystemSettings, if before.EnableMetadataPassthrough != after.EnableMetadataPassthrough { changed = append(changed, "enable_metadata_passthrough") } + if before.EnableCCHSigning != after.EnableCCHSigning { + changed = append(changed, "enable_cch_signing") + } return changed } @@ -1041,6 +1355,18 @@ func equalDefaultSubscriptions(a, b []service.DefaultSubscriptionSetting) bool { return true } +func equalIntSlice(a, b []int) bool { + if len(a) != len(b) { + return false + } + for i := range a { + if a[i] != b[i] { + return false + } + } + return true +} + // TestSMTPRequest 测试SMTP连接请求 type TestSMTPRequest struct { SMTPHost string `json:"smtp_host"` diff --git a/backend/internal/handler/admin/usage_handler.go b/backend/internal/handler/admin/usage_handler.go index 2967b3840e..0857a13880 100644 --- a/backend/internal/handler/admin/usage_handler.go +++ b/backend/internal/handler/admin/usage_handler.go @@ -165,7 +165,12 @@ func (h *UsageHandler) List(c *gin.Context) { endTime = &t } - params := pagination.PaginationParams{Page: page, PageSize: pageSize} + params := pagination.PaginationParams{ + Page: page, + PageSize: pageSize, + SortBy: c.DefaultQuery("sort_by", "created_at"), + SortOrder: c.DefaultQuery("sort_order", "desc"), + } filters := usagestats.UsageLogFilters{ UserID: userID, APIKeyID: apiKeyID, @@ -339,7 +344,7 @@ func (h *UsageHandler) SearchUsers(c *gin.Context) { } // Limit to 30 results - users, _, err := h.adminService.ListUsers(c.Request.Context(), 1, 30, service.UserListFilters{Search: keyword}) + users, _, err := h.adminService.ListUsers(c.Request.Context(), 1, 30, service.UserListFilters{Search: keyword}, "email", "asc") if err != nil { response.ErrorFrom(c, err) return diff --git a/backend/internal/handler/admin/usage_handler_request_type_test.go b/backend/internal/handler/admin/usage_handler_request_type_test.go index 3f158316fe..882cbe9362 100644 --- a/backend/internal/handler/admin/usage_handler_request_type_test.go +++ b/backend/internal/handler/admin/usage_handler_request_type_test.go @@ -15,11 +15,13 @@ import ( type adminUsageRepoCapture struct { service.UsageLogRepository + listParams pagination.PaginationParams listFilters usagestats.UsageLogFilters statsFilters usagestats.UsageLogFilters } func (s *adminUsageRepoCapture) ListWithFilters(ctx context.Context, params pagination.PaginationParams, filters usagestats.UsageLogFilters) ([]service.UsageLog, *pagination.PaginationResult, error) { + s.listParams = params s.listFilters = filters return []service.UsageLog{}, &pagination.PaginationResult{ Total: 0, diff --git a/backend/internal/handler/admin/usage_handler_sort_test.go b/backend/internal/handler/admin/usage_handler_sort_test.go new file mode 100644 index 0000000000..dac826762a --- /dev/null +++ b/backend/internal/handler/admin/usage_handler_sort_test.go @@ -0,0 +1,35 @@ +package admin + +import ( + "net/http" + "net/http/httptest" + "testing" + + "github.com/stretchr/testify/require" +) + +func TestAdminUsageListSortParams(t *testing.T) { + repo := &adminUsageRepoCapture{} + router := newAdminUsageRequestTypeTestRouter(repo) + + req := httptest.NewRequest(http.MethodGet, "/admin/usage?sort_by=model&sort_order=ASC", nil) + rec := httptest.NewRecorder() + router.ServeHTTP(rec, req) + + require.Equal(t, http.StatusOK, rec.Code) + require.Equal(t, "model", repo.listParams.SortBy) + require.Equal(t, "ASC", repo.listParams.SortOrder) +} + +func TestAdminUsageListSortDefaults(t *testing.T) { + repo := &adminUsageRepoCapture{} + router := newAdminUsageRequestTypeTestRouter(repo) + + req := httptest.NewRequest(http.MethodGet, "/admin/usage", nil) + rec := httptest.NewRecorder() + router.ServeHTTP(rec, req) + + require.Equal(t, http.StatusOK, rec.Code) + require.Equal(t, "created_at", repo.listParams.SortBy) + require.Equal(t, "desc", repo.listParams.SortOrder) +} diff --git a/backend/internal/handler/admin/user_handler.go b/backend/internal/handler/admin/user_handler.go index a357657e20..1453bd0739 100644 --- a/backend/internal/handler/admin/user_handler.go +++ b/backend/internal/handler/admin/user_handler.go @@ -91,12 +91,14 @@ func (h *UserHandler) List(c *gin.Context) { GroupName: strings.TrimSpace(c.Query("group_name")), Attributes: parseAttributeFilters(c), } + sortBy := c.DefaultQuery("sort_by", "created_at") + sortOrder := c.DefaultQuery("sort_order", "desc") if raw, ok := c.GetQuery("include_subscriptions"); ok { includeSubscriptions := parseBoolQueryWithDefault(raw, true) filters.IncludeSubscriptions = &includeSubscriptions } - users, total, err := h.adminService.ListUsers(c.Request.Context(), page, pageSize, filters) + users, total, err := h.adminService.ListUsers(c.Request.Context(), page, pageSize, filters, sortBy, sortOrder) if err != nil { response.ErrorFrom(c, err) return @@ -290,8 +292,10 @@ func (h *UserHandler) GetUserAPIKeys(c *gin.Context) { } page, pageSize := response.ParsePagination(c) + sortBy := c.DefaultQuery("sort_by", "created_at") + sortOrder := c.DefaultQuery("sort_order", "desc") - keys, total, err := h.adminService.GetUserAPIKeys(c.Request.Context(), userID, page, pageSize) + keys, total, err := h.adminService.GetUserAPIKeys(c.Request.Context(), userID, page, pageSize, sortBy, sortOrder) if err != nil { response.ErrorFrom(c, err) return diff --git a/backend/internal/handler/api_key_handler.go b/backend/internal/handler/api_key_handler.go index 951aed08db..9d6c6c1523 100644 --- a/backend/internal/handler/api_key_handler.go +++ b/backend/internal/handler/api_key_handler.go @@ -72,7 +72,12 @@ func (h *APIKeyHandler) List(c *gin.Context) { } page, pageSize := response.ParsePagination(c) - params := pagination.PaginationParams{Page: page, PageSize: pageSize} + params := pagination.PaginationParams{ + Page: page, + PageSize: pageSize, + SortBy: c.DefaultQuery("sort_by", "created_at"), + SortOrder: c.DefaultQuery("sort_order", "desc"), + } // Parse filter parameters var filters service.APIKeyListFilters diff --git a/backend/internal/handler/auth_oidc_oauth.go b/backend/internal/handler/auth_oidc_oauth.go new file mode 100644 index 0000000000..9d24df88ab --- /dev/null +++ b/backend/internal/handler/auth_oidc_oauth.go @@ -0,0 +1,873 @@ +package handler + +import ( + "context" + "crypto/ecdsa" + "crypto/elliptic" + "crypto/rsa" + "crypto/sha256" + "encoding/base64" + "encoding/hex" + "encoding/json" + "errors" + "fmt" + "log" + "math/big" + "net/http" + "net/url" + "strconv" + "strings" + "time" + + "github.com/Wei-Shaw/sub2api/internal/config" + infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" + "github.com/Wei-Shaw/sub2api/internal/pkg/oauth" + "github.com/Wei-Shaw/sub2api/internal/pkg/response" + "github.com/Wei-Shaw/sub2api/internal/service" + + "github.com/gin-gonic/gin" + "github.com/golang-jwt/jwt/v5" + "github.com/imroc/req/v3" + "github.com/tidwall/gjson" +) + +const ( + oidcOAuthCookiePath = "/api/v1/auth/oauth/oidc" + oidcOAuthStateCookieName = "oidc_oauth_state" + oidcOAuthVerifierCookie = "oidc_oauth_verifier" + oidcOAuthRedirectCookie = "oidc_oauth_redirect" + oidcOAuthNonceCookie = "oidc_oauth_nonce" + oidcOAuthCookieMaxAgeSec = 10 * 60 // 10 minutes + oidcOAuthDefaultRedirectTo = "/dashboard" + oidcOAuthDefaultFrontendCB = "/auth/oidc/callback" +) + +type oidcTokenResponse struct { + AccessToken string `json:"access_token"` + TokenType string `json:"token_type"` + ExpiresIn int64 `json:"expires_in"` + RefreshToken string `json:"refresh_token,omitempty"` + Scope string `json:"scope,omitempty"` + IDToken string `json:"id_token,omitempty"` +} + +type oidcTokenExchangeError struct { + StatusCode int + ProviderError string + ProviderDescription string + Body string +} + +func (e *oidcTokenExchangeError) Error() string { + if e == nil { + return "" + } + parts := []string{fmt.Sprintf("token exchange status=%d", e.StatusCode)} + if strings.TrimSpace(e.ProviderError) != "" { + parts = append(parts, "error="+strings.TrimSpace(e.ProviderError)) + } + if strings.TrimSpace(e.ProviderDescription) != "" { + parts = append(parts, "error_description="+strings.TrimSpace(e.ProviderDescription)) + } + return strings.Join(parts, " ") +} + +type oidcIDTokenClaims struct { + Email string `json:"email,omitempty"` + EmailVerified *bool `json:"email_verified,omitempty"` + PreferredUsername string `json:"preferred_username,omitempty"` + Name string `json:"name,omitempty"` + Nonce string `json:"nonce,omitempty"` + Azp string `json:"azp,omitempty"` + jwt.RegisteredClaims +} + +type oidcUserInfoClaims struct { + Email string + Username string + Subject string + EmailVerified *bool +} + +type oidcJWKSet struct { + Keys []oidcJWK `json:"keys"` +} + +type oidcJWK struct { + Kty string `json:"kty"` + Kid string `json:"kid"` + Use string `json:"use"` + Alg string `json:"alg"` + + N string `json:"n"` + E string `json:"e"` + + Crv string `json:"crv"` + X string `json:"x"` + Y string `json:"y"` +} + +// OIDCOAuthStart 启动通用 OIDC OAuth 登录流程。 +// GET /api/v1/auth/oauth/oidc/start?redirect=/dashboard +func (h *AuthHandler) OIDCOAuthStart(c *gin.Context) { + cfg, err := h.getOIDCOAuthConfig(c.Request.Context()) + if err != nil { + response.ErrorFrom(c, err) + return + } + + state, err := oauth.GenerateState() + if err != nil { + response.ErrorFrom(c, infraerrors.InternalServer("OAUTH_STATE_GEN_FAILED", "failed to generate oauth state").WithCause(err)) + return + } + + redirectTo := sanitizeFrontendRedirectPath(c.Query("redirect")) + if redirectTo == "" { + redirectTo = oidcOAuthDefaultRedirectTo + } + + secureCookie := isRequestHTTPS(c) + oidcSetCookie(c, oidcOAuthStateCookieName, encodeCookieValue(state), oidcOAuthCookieMaxAgeSec, secureCookie) + oidcSetCookie(c, oidcOAuthRedirectCookie, encodeCookieValue(redirectTo), oidcOAuthCookieMaxAgeSec, secureCookie) + + codeChallenge := "" + if cfg.UsePKCE { + verifier, genErr := oauth.GenerateCodeVerifier() + if genErr != nil { + response.ErrorFrom(c, infraerrors.InternalServer("OAUTH_PKCE_GEN_FAILED", "failed to generate pkce verifier").WithCause(genErr)) + return + } + codeChallenge = oauth.GenerateCodeChallenge(verifier) + oidcSetCookie(c, oidcOAuthVerifierCookie, encodeCookieValue(verifier), oidcOAuthCookieMaxAgeSec, secureCookie) + } + + nonce := "" + if cfg.ValidateIDToken { + nonce, err = oauth.GenerateState() + if err != nil { + response.ErrorFrom(c, infraerrors.InternalServer("OAUTH_NONCE_GEN_FAILED", "failed to generate oauth nonce").WithCause(err)) + return + } + oidcSetCookie(c, oidcOAuthNonceCookie, encodeCookieValue(nonce), oidcOAuthCookieMaxAgeSec, secureCookie) + } + + redirectURI := strings.TrimSpace(cfg.RedirectURL) + if redirectURI == "" { + response.ErrorFrom(c, infraerrors.InternalServer("OAUTH_CONFIG_INVALID", "oauth redirect url not configured")) + return + } + + authURL, err := buildOIDCAuthorizeURL(cfg, state, nonce, codeChallenge, redirectURI) + if err != nil { + response.ErrorFrom(c, infraerrors.InternalServer("OAUTH_BUILD_URL_FAILED", "failed to build oauth authorization url").WithCause(err)) + return + } + + c.Redirect(http.StatusFound, authURL) +} + +// OIDCOAuthCallback 处理 OIDC 回调:校验 id_token、创建/登录用户并重定向到前端。 +// GET /api/v1/auth/oauth/oidc/callback?code=...&state=... +func (h *AuthHandler) OIDCOAuthCallback(c *gin.Context) { + cfg, cfgErr := h.getOIDCOAuthConfig(c.Request.Context()) + if cfgErr != nil { + response.ErrorFrom(c, cfgErr) + return + } + + frontendCallback := strings.TrimSpace(cfg.FrontendRedirectURL) + if frontendCallback == "" { + frontendCallback = oidcOAuthDefaultFrontendCB + } + + if providerErr := strings.TrimSpace(c.Query("error")); providerErr != "" { + redirectOAuthError(c, frontendCallback, "provider_error", providerErr, c.Query("error_description")) + return + } + + code := strings.TrimSpace(c.Query("code")) + state := strings.TrimSpace(c.Query("state")) + if code == "" || state == "" { + redirectOAuthError(c, frontendCallback, "missing_params", "missing code/state", "") + return + } + + secureCookie := isRequestHTTPS(c) + defer func() { + oidcClearCookie(c, oidcOAuthStateCookieName, secureCookie) + oidcClearCookie(c, oidcOAuthVerifierCookie, secureCookie) + oidcClearCookie(c, oidcOAuthRedirectCookie, secureCookie) + oidcClearCookie(c, oidcOAuthNonceCookie, secureCookie) + }() + + expectedState, err := readCookieDecoded(c, oidcOAuthStateCookieName) + if err != nil || expectedState == "" || state != expectedState { + redirectOAuthError(c, frontendCallback, "invalid_state", "invalid oauth state", "") + return + } + + redirectTo, _ := readCookieDecoded(c, oidcOAuthRedirectCookie) + redirectTo = sanitizeFrontendRedirectPath(redirectTo) + if redirectTo == "" { + redirectTo = oidcOAuthDefaultRedirectTo + } + + codeVerifier := "" + if cfg.UsePKCE { + codeVerifier, _ = readCookieDecoded(c, oidcOAuthVerifierCookie) + if codeVerifier == "" { + redirectOAuthError(c, frontendCallback, "missing_verifier", "missing pkce verifier", "") + return + } + } + + expectedNonce := "" + if cfg.ValidateIDToken { + expectedNonce, _ = readCookieDecoded(c, oidcOAuthNonceCookie) + if expectedNonce == "" { + redirectOAuthError(c, frontendCallback, "missing_nonce", "missing oauth nonce", "") + return + } + } + + redirectURI := strings.TrimSpace(cfg.RedirectURL) + if redirectURI == "" { + redirectOAuthError(c, frontendCallback, "config_error", "oauth redirect url not configured", "") + return + } + + tokenResp, err := oidcExchangeCode(c.Request.Context(), cfg, code, redirectURI, codeVerifier) + if err != nil { + description := "" + var exchangeErr *oidcTokenExchangeError + if errors.As(err, &exchangeErr) && exchangeErr != nil { + log.Printf( + "[OIDC OAuth] token exchange failed: status=%d provider_error=%q provider_description=%q body=%s", + exchangeErr.StatusCode, + exchangeErr.ProviderError, + exchangeErr.ProviderDescription, + truncateLogValue(exchangeErr.Body, 2048), + ) + description = exchangeErr.Error() + } else { + log.Printf("[OIDC OAuth] token exchange failed: %v", err) + description = err.Error() + } + redirectOAuthError(c, frontendCallback, "token_exchange_failed", "failed to exchange oauth code", singleLine(description)) + return + } + + if cfg.ValidateIDToken && strings.TrimSpace(tokenResp.IDToken) == "" { + redirectOAuthError(c, frontendCallback, "missing_id_token", "missing id_token", "") + return + } + + idClaims, err := oidcParseAndValidateIDToken(c.Request.Context(), cfg, tokenResp.IDToken, expectedNonce) + if err != nil { + log.Printf("[OIDC OAuth] id_token validation failed: %v", err) + redirectOAuthError(c, frontendCallback, "invalid_id_token", "failed to validate id_token", "") + return + } + + userInfoClaims, err := oidcFetchUserInfo(c.Request.Context(), cfg, tokenResp) + if err != nil { + log.Printf("[OIDC OAuth] userinfo fetch failed: %v", err) + redirectOAuthError(c, frontendCallback, "userinfo_failed", "failed to fetch user info", "") + return + } + + subject := strings.TrimSpace(idClaims.Subject) + if subject == "" { + subject = strings.TrimSpace(userInfoClaims.Subject) + } + if subject == "" { + redirectOAuthError(c, frontendCallback, "missing_subject", "missing subject claim", "") + return + } + issuer := strings.TrimSpace(idClaims.Issuer) + if issuer == "" { + issuer = strings.TrimSpace(cfg.IssuerURL) + } + if issuer == "" { + redirectOAuthError(c, frontendCallback, "missing_issuer", "missing issuer claim", "") + return + } + + emailVerified := userInfoClaims.EmailVerified + if emailVerified == nil { + emailVerified = idClaims.EmailVerified + } + if cfg.RequireEmailVerified { + if emailVerified == nil || !*emailVerified { + redirectOAuthError(c, frontendCallback, "email_not_verified", "email is not verified", "") + return + } + } + + identityKey := oidcIdentityKey(issuer, subject) + email := oidcSelectLoginEmail(userInfoClaims.Email, idClaims.Email, identityKey) + username := firstNonEmpty( + userInfoClaims.Username, + idClaims.PreferredUsername, + idClaims.Name, + oidcFallbackUsername(subject), + ) + + // 传入空邀请码;如果需要邀请码,服务层返回 ErrOAuthInvitationRequired + tokenPair, _, err := h.authService.LoginOrRegisterOAuthWithTokenPair(c.Request.Context(), email, username, "") + if err != nil { + if errors.Is(err, service.ErrOAuthInvitationRequired) { + pendingToken, tokenErr := h.authService.CreatePendingOAuthToken(email, username) + if tokenErr != nil { + redirectOAuthError(c, frontendCallback, "login_failed", "service_error", "") + return + } + fragment := url.Values{} + fragment.Set("error", "invitation_required") + fragment.Set("pending_oauth_token", pendingToken) + fragment.Set("redirect", redirectTo) + redirectWithFragment(c, frontendCallback, fragment) + return + } + redirectOAuthError(c, frontendCallback, "login_failed", infraerrors.Reason(err), infraerrors.Message(err)) + return + } + + fragment := url.Values{} + fragment.Set("access_token", tokenPair.AccessToken) + fragment.Set("refresh_token", tokenPair.RefreshToken) + fragment.Set("expires_in", fmt.Sprintf("%d", tokenPair.ExpiresIn)) + fragment.Set("token_type", "Bearer") + fragment.Set("redirect", redirectTo) + redirectWithFragment(c, frontendCallback, fragment) +} + +type completeOIDCOAuthRequest struct { + PendingOAuthToken string `json:"pending_oauth_token" binding:"required"` + InvitationCode string `json:"invitation_code" binding:"required"` +} + +// CompleteOIDCOAuthRegistration completes a pending OAuth registration by validating +// the invitation code and creating the user account. +// POST /api/v1/auth/oauth/oidc/complete-registration +func (h *AuthHandler) CompleteOIDCOAuthRegistration(c *gin.Context) { + var req completeOIDCOAuthRequest + if err := c.ShouldBindJSON(&req); err != nil { + c.JSON(http.StatusBadRequest, gin.H{"error": "INVALID_REQUEST", "message": err.Error()}) + return + } + + email, username, err := h.authService.VerifyPendingOAuthToken(req.PendingOAuthToken) + if err != nil { + c.JSON(http.StatusUnauthorized, gin.H{"error": "INVALID_TOKEN", "message": "invalid or expired registration token"}) + return + } + + tokenPair, _, err := h.authService.LoginOrRegisterOAuthWithTokenPair(c.Request.Context(), email, username, req.InvitationCode) + if err != nil { + response.ErrorFrom(c, err) + return + } + + c.JSON(http.StatusOK, gin.H{ + "access_token": tokenPair.AccessToken, + "refresh_token": tokenPair.RefreshToken, + "expires_in": tokenPair.ExpiresIn, + "token_type": "Bearer", + }) +} + +func (h *AuthHandler) getOIDCOAuthConfig(ctx context.Context) (config.OIDCConnectConfig, error) { + if h != nil && h.settingSvc != nil { + return h.settingSvc.GetOIDCConnectOAuthConfig(ctx) + } + if h == nil || h.cfg == nil { + return config.OIDCConnectConfig{}, infraerrors.ServiceUnavailable("CONFIG_NOT_READY", "config not loaded") + } + if !h.cfg.OIDC.Enabled { + return config.OIDCConnectConfig{}, infraerrors.NotFound("OAUTH_DISABLED", "oauth login is disabled") + } + return h.cfg.OIDC, nil +} + +func oidcExchangeCode( + ctx context.Context, + cfg config.OIDCConnectConfig, + code string, + redirectURI string, + codeVerifier string, +) (*oidcTokenResponse, error) { + client := req.C().SetTimeout(30 * time.Second) + + form := url.Values{} + form.Set("grant_type", "authorization_code") + form.Set("client_id", cfg.ClientID) + form.Set("code", code) + form.Set("redirect_uri", redirectURI) + if cfg.UsePKCE { + form.Set("code_verifier", codeVerifier) + } + + r := client.R(). + SetContext(ctx). + SetHeader("Accept", "application/json") + + switch strings.ToLower(strings.TrimSpace(cfg.TokenAuthMethod)) { + case "", "client_secret_post": + form.Set("client_secret", cfg.ClientSecret) + case "client_secret_basic": + r.SetBasicAuth(cfg.ClientID, cfg.ClientSecret) + case "none": + default: + return nil, fmt.Errorf("unsupported token_auth_method: %s", cfg.TokenAuthMethod) + } + + resp, err := r.SetFormDataFromValues(form).Post(cfg.TokenURL) + if err != nil { + return nil, fmt.Errorf("request token: %w", err) + } + body := strings.TrimSpace(resp.String()) + if !resp.IsSuccessState() { + providerErr, providerDesc := parseOAuthProviderError(body) + return nil, &oidcTokenExchangeError{ + StatusCode: resp.StatusCode, + ProviderError: providerErr, + ProviderDescription: providerDesc, + Body: body, + } + } + + tokenResp, ok := oidcParseTokenResponse(body) + if !ok { + return nil, &oidcTokenExchangeError{StatusCode: resp.StatusCode, Body: body} + } + if strings.TrimSpace(tokenResp.TokenType) == "" { + tokenResp.TokenType = "Bearer" + } + if strings.TrimSpace(tokenResp.AccessToken) == "" && strings.TrimSpace(tokenResp.IDToken) == "" { + return nil, &oidcTokenExchangeError{StatusCode: resp.StatusCode, Body: body} + } + return tokenResp, nil +} + +func oidcParseTokenResponse(body string) (*oidcTokenResponse, bool) { + body = strings.TrimSpace(body) + if body == "" { + return nil, false + } + + accessToken := strings.TrimSpace(getGJSON(body, "access_token")) + idToken := strings.TrimSpace(getGJSON(body, "id_token")) + if accessToken != "" || idToken != "" { + tokenType := strings.TrimSpace(getGJSON(body, "token_type")) + refreshToken := strings.TrimSpace(getGJSON(body, "refresh_token")) + scope := strings.TrimSpace(getGJSON(body, "scope")) + expiresIn := gjson.Get(body, "expires_in").Int() + return &oidcTokenResponse{ + AccessToken: accessToken, + TokenType: tokenType, + ExpiresIn: expiresIn, + RefreshToken: refreshToken, + Scope: scope, + IDToken: idToken, + }, true + } + + values, err := url.ParseQuery(body) + if err != nil { + return nil, false + } + accessToken = strings.TrimSpace(values.Get("access_token")) + idToken = strings.TrimSpace(values.Get("id_token")) + if accessToken == "" && idToken == "" { + return nil, false + } + expiresIn := int64(0) + if raw := strings.TrimSpace(values.Get("expires_in")); raw != "" { + if v, parseErr := strconv.ParseInt(raw, 10, 64); parseErr == nil { + expiresIn = v + } + } + return &oidcTokenResponse{ + AccessToken: accessToken, + TokenType: strings.TrimSpace(values.Get("token_type")), + ExpiresIn: expiresIn, + RefreshToken: strings.TrimSpace(values.Get("refresh_token")), + Scope: strings.TrimSpace(values.Get("scope")), + IDToken: idToken, + }, true +} + +func oidcFetchUserInfo( + ctx context.Context, + cfg config.OIDCConnectConfig, + token *oidcTokenResponse, +) (*oidcUserInfoClaims, error) { + if strings.TrimSpace(cfg.UserInfoURL) == "" { + return &oidcUserInfoClaims{}, nil + } + if token == nil || strings.TrimSpace(token.AccessToken) == "" { + return nil, errors.New("missing access_token for userinfo request") + } + + client := req.C().SetTimeout(30 * time.Second) + authorization, err := buildBearerAuthorization(token.TokenType, token.AccessToken) + if err != nil { + return nil, fmt.Errorf("invalid token for userinfo request: %w", err) + } + + resp, err := client.R(). + SetContext(ctx). + SetHeader("Accept", "application/json"). + SetHeader("Authorization", authorization). + Get(cfg.UserInfoURL) + if err != nil { + return nil, fmt.Errorf("request userinfo: %w", err) + } + if !resp.IsSuccessState() { + return nil, fmt.Errorf("userinfo status=%d", resp.StatusCode) + } + + return oidcParseUserInfo(resp.String(), cfg), nil +} + +func oidcParseUserInfo(body string, cfg config.OIDCConnectConfig) *oidcUserInfoClaims { + claims := &oidcUserInfoClaims{} + claims.Email = firstNonEmpty( + getGJSON(body, cfg.UserInfoEmailPath), + getGJSON(body, "email"), + getGJSON(body, "user.email"), + getGJSON(body, "data.email"), + getGJSON(body, "attributes.email"), + ) + claims.Username = firstNonEmpty( + getGJSON(body, cfg.UserInfoUsernamePath), + getGJSON(body, "preferred_username"), + getGJSON(body, "username"), + getGJSON(body, "name"), + getGJSON(body, "user.username"), + getGJSON(body, "user.name"), + ) + claims.Subject = firstNonEmpty( + getGJSON(body, cfg.UserInfoIDPath), + getGJSON(body, "sub"), + getGJSON(body, "id"), + getGJSON(body, "user_id"), + getGJSON(body, "uid"), + getGJSON(body, "user.id"), + ) + if verified, ok := getGJSONBool(body, "email_verified"); ok { + claims.EmailVerified = &verified + } + claims.Email = strings.TrimSpace(claims.Email) + claims.Username = strings.TrimSpace(claims.Username) + claims.Subject = strings.TrimSpace(claims.Subject) + return claims +} + +func getGJSONBool(body string, path string) (bool, bool) { + path = strings.TrimSpace(path) + if path == "" { + return false, false + } + res := gjson.Get(body, path) + if !res.Exists() { + return false, false + } + return res.Bool(), true +} + +func buildOIDCAuthorizeURL(cfg config.OIDCConnectConfig, state, nonce, codeChallenge, redirectURI string) (string, error) { + u, err := url.Parse(cfg.AuthorizeURL) + if err != nil { + return "", fmt.Errorf("parse authorize_url: %w", err) + } + + q := u.Query() + q.Set("response_type", "code") + q.Set("client_id", cfg.ClientID) + q.Set("redirect_uri", redirectURI) + if strings.TrimSpace(cfg.Scopes) != "" { + q.Set("scope", cfg.Scopes) + } + q.Set("state", state) + if strings.TrimSpace(nonce) != "" { + q.Set("nonce", nonce) + } + if cfg.UsePKCE { + q.Set("code_challenge", codeChallenge) + q.Set("code_challenge_method", "S256") + } + + u.RawQuery = q.Encode() + return u.String(), nil +} + +func oidcParseAndValidateIDToken(ctx context.Context, cfg config.OIDCConnectConfig, idToken string, expectedNonce string) (*oidcIDTokenClaims, error) { + idToken = strings.TrimSpace(idToken) + if idToken == "" { + return nil, errors.New("missing id_token") + } + allowed := oidcAllowedSigningAlgs(cfg.AllowedSigningAlgs) + if len(allowed) == 0 { + return nil, errors.New("empty allowed signing algorithms") + } + + jwks, err := oidcFetchJWKSet(ctx, cfg.JWKSURL) + if err != nil { + return nil, err + } + leeway := time.Duration(cfg.ClockSkewSeconds) * time.Second + claims := &oidcIDTokenClaims{} + + parsed, err := jwt.ParseWithClaims( + idToken, + claims, + func(token *jwt.Token) (any, error) { + alg := strings.TrimSpace(token.Method.Alg()) + if !containsString(allowed, alg) { + return nil, fmt.Errorf("unexpected signing algorithm: %s", alg) + } + kid, _ := token.Header["kid"].(string) + return oidcFindPublicKey(jwks, strings.TrimSpace(kid), alg) + }, + jwt.WithValidMethods(allowed), + jwt.WithAudience(cfg.ClientID), + jwt.WithIssuer(cfg.IssuerURL), + jwt.WithLeeway(leeway), + ) + if err != nil { + return nil, err + } + if !parsed.Valid { + return nil, errors.New("id_token invalid") + } + if strings.TrimSpace(claims.Subject) == "" { + return nil, errors.New("id_token missing sub") + } + if expectedNonce != "" && strings.TrimSpace(claims.Nonce) != strings.TrimSpace(expectedNonce) { + return nil, errors.New("id_token nonce mismatch") + } + if len(claims.Audience) > 1 { + if strings.TrimSpace(claims.Azp) == "" || strings.TrimSpace(claims.Azp) != strings.TrimSpace(cfg.ClientID) { + return nil, errors.New("id_token azp mismatch") + } + } + return claims, nil +} + +func oidcAllowedSigningAlgs(raw string) []string { + if strings.TrimSpace(raw) == "" { + return []string{"RS256", "ES256", "PS256"} + } + seen := make(map[string]struct{}) + out := make([]string, 0, 4) + for _, part := range strings.Split(raw, ",") { + alg := strings.ToUpper(strings.TrimSpace(part)) + if alg == "" { + continue + } + if _, ok := seen[alg]; ok { + continue + } + seen[alg] = struct{}{} + out = append(out, alg) + } + return out +} + +func oidcFetchJWKSet(ctx context.Context, jwksURL string) (*oidcJWKSet, error) { + jwksURL = strings.TrimSpace(jwksURL) + if jwksURL == "" { + return nil, errors.New("missing jwks_url") + } + resp, err := req.C(). + SetTimeout(30*time.Second). + R(). + SetContext(ctx). + SetHeader("Accept", "application/json"). + Get(jwksURL) + if err != nil { + return nil, fmt.Errorf("request jwks: %w", err) + } + if !resp.IsSuccessState() { + return nil, fmt.Errorf("jwks status=%d", resp.StatusCode) + } + set := &oidcJWKSet{} + if err := json.Unmarshal(resp.Bytes(), set); err != nil { + return nil, fmt.Errorf("parse jwks: %w", err) + } + if len(set.Keys) == 0 { + return nil, errors.New("jwks empty keys") + } + return set, nil +} + +func oidcFindPublicKey(set *oidcJWKSet, kid, alg string) (any, error) { + if set == nil { + return nil, errors.New("jwks not loaded") + } + alg = strings.ToUpper(strings.TrimSpace(alg)) + kid = strings.TrimSpace(kid) + + var lastErr error + for i := range set.Keys { + k := set.Keys[i] + if strings.TrimSpace(k.Use) != "" && !strings.EqualFold(strings.TrimSpace(k.Use), "sig") { + continue + } + if kid != "" && strings.TrimSpace(k.Kid) != kid { + continue + } + if strings.TrimSpace(k.Alg) != "" && !strings.EqualFold(strings.TrimSpace(k.Alg), alg) { + continue + } + pk, err := k.publicKey() + if err != nil { + lastErr = err + continue + } + if pk != nil { + return pk, nil + } + } + if lastErr != nil { + return nil, lastErr + } + if kid != "" { + return nil, fmt.Errorf("jwk not found for kid=%s", kid) + } + return nil, errors.New("jwk not found") +} + +func (k oidcJWK) publicKey() (any, error) { + switch strings.ToUpper(strings.TrimSpace(k.Kty)) { + case "RSA": + n, err := decodeBase64URLBigInt(k.N) + if err != nil { + return nil, fmt.Errorf("decode rsa n: %w", err) + } + eBytes, err := base64.RawURLEncoding.DecodeString(strings.TrimSpace(k.E)) + if err != nil { + return nil, fmt.Errorf("decode rsa e: %w", err) + } + if len(eBytes) == 0 { + return nil, errors.New("empty rsa e") + } + e := 0 + for _, b := range eBytes { + e = (e << 8) | int(b) + } + if e <= 0 { + return nil, errors.New("invalid rsa exponent") + } + if n.Sign() <= 0 { + return nil, errors.New("invalid rsa modulus") + } + return &rsa.PublicKey{N: n, E: e}, nil + case "EC": + var curve elliptic.Curve + switch strings.TrimSpace(k.Crv) { + case "P-256": + curve = elliptic.P256() + case "P-384": + curve = elliptic.P384() + case "P-521": + curve = elliptic.P521() + default: + return nil, fmt.Errorf("unsupported ec curve: %s", k.Crv) + } + x, err := decodeBase64URLBigInt(k.X) + if err != nil { + return nil, fmt.Errorf("decode ec x: %w", err) + } + y, err := decodeBase64URLBigInt(k.Y) + if err != nil { + return nil, fmt.Errorf("decode ec y: %w", err) + } + if !curve.IsOnCurve(x, y) { + return nil, errors.New("ec point is not on curve") + } + return &ecdsa.PublicKey{Curve: curve, X: x, Y: y}, nil + default: + return nil, fmt.Errorf("unsupported jwk kty: %s", k.Kty) + } +} + +func decodeBase64URLBigInt(raw string) (*big.Int, error) { + buf, err := base64.RawURLEncoding.DecodeString(strings.TrimSpace(raw)) + if err != nil { + return nil, err + } + if len(buf) == 0 { + return nil, errors.New("empty value") + } + return new(big.Int).SetBytes(buf), nil +} + +func containsString(values []string, target string) bool { + target = strings.TrimSpace(target) + for _, v := range values { + if strings.EqualFold(strings.TrimSpace(v), target) { + return true + } + } + return false +} + +func oidcIdentityKey(issuer, subject string) string { + issuer = strings.TrimSpace(strings.ToLower(issuer)) + subject = strings.TrimSpace(subject) + return issuer + "\x1f" + subject +} + +func oidcSyntheticEmailFromIdentityKey(identityKey string) string { + identityKey = strings.TrimSpace(identityKey) + if identityKey == "" { + return "" + } + sum := sha256.Sum256([]byte(identityKey)) + return "oidc-" + hex.EncodeToString(sum[:16]) + service.OIDCConnectSyntheticEmailDomain +} + +func oidcSelectLoginEmail(userInfoEmail, idTokenEmail, identityKey string) string { + email := strings.TrimSpace(firstNonEmpty(userInfoEmail, idTokenEmail)) + if email != "" { + return email + } + return oidcSyntheticEmailFromIdentityKey(identityKey) +} + +func oidcFallbackUsername(subject string) string { + subject = strings.TrimSpace(subject) + if subject == "" { + return "oidc_user" + } + sum := sha256.Sum256([]byte(subject)) + return "oidc_" + hex.EncodeToString(sum[:])[:12] +} + +func oidcSetCookie(c *gin.Context, name, value string, maxAgeSec int, secure bool) { + http.SetCookie(c.Writer, &http.Cookie{ + Name: name, + Value: value, + Path: oidcOAuthCookiePath, + MaxAge: maxAgeSec, + HttpOnly: true, + Secure: secure, + SameSite: http.SameSiteLaxMode, + }) +} + +func oidcClearCookie(c *gin.Context, name string, secure bool) { + http.SetCookie(c.Writer, &http.Cookie{ + Name: name, + Value: "", + Path: oidcOAuthCookiePath, + MaxAge: -1, + HttpOnly: true, + Secure: secure, + SameSite: http.SameSiteLaxMode, + }) +} diff --git a/backend/internal/handler/auth_oidc_oauth_test.go b/backend/internal/handler/auth_oidc_oauth_test.go new file mode 100644 index 0000000000..a161aa77cf --- /dev/null +++ b/backend/internal/handler/auth_oidc_oauth_test.go @@ -0,0 +1,120 @@ +package handler + +import ( + "context" + "crypto/rand" + "crypto/rsa" + "encoding/base64" + "encoding/json" + "math/big" + "net/http" + "net/http/httptest" + "testing" + "time" + + "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/golang-jwt/jwt/v5" + "github.com/stretchr/testify/require" +) + +func TestOIDCSyntheticEmailStableAndDistinct(t *testing.T) { + k1 := oidcIdentityKey("https://issuer.example.com", "subject-a") + k2 := oidcIdentityKey("https://issuer.example.com", "subject-b") + + e1 := oidcSyntheticEmailFromIdentityKey(k1) + e1Again := oidcSyntheticEmailFromIdentityKey(k1) + e2 := oidcSyntheticEmailFromIdentityKey(k2) + + require.Equal(t, e1, e1Again) + require.NotEqual(t, e1, e2) + require.Contains(t, e1, "@oidc-connect.invalid") +} + +func TestOIDCSelectLoginEmailPrefersRealEmail(t *testing.T) { + identityKey := oidcIdentityKey("https://issuer.example.com", "subject-a") + + email := oidcSelectLoginEmail("user@example.com", "idtoken@example.com", identityKey) + require.Equal(t, "user@example.com", email) + + email = oidcSelectLoginEmail("", "idtoken@example.com", identityKey) + require.Equal(t, "idtoken@example.com", email) + + email = oidcSelectLoginEmail("", "", identityKey) + require.Contains(t, email, "@oidc-connect.invalid") + require.Equal(t, oidcSyntheticEmailFromIdentityKey(identityKey), email) +} + +func TestBuildOIDCAuthorizeURLIncludesNonceAndPKCE(t *testing.T) { + cfg := config.OIDCConnectConfig{ + AuthorizeURL: "https://issuer.example.com/auth", + ClientID: "cid", + Scopes: "openid email profile", + UsePKCE: true, + } + + u, err := buildOIDCAuthorizeURL(cfg, "state123", "nonce123", "challenge123", "https://app.example.com/callback") + require.NoError(t, err) + require.Contains(t, u, "nonce=nonce123") + require.Contains(t, u, "code_challenge=challenge123") + require.Contains(t, u, "code_challenge_method=S256") + require.Contains(t, u, "scope=openid+email+profile") +} + +func TestOIDCParseAndValidateIDToken(t *testing.T) { + priv, err := rsa.GenerateKey(rand.Reader, 2048) + require.NoError(t, err) + + kid := "kid-1" + jwks := oidcJWKSet{Keys: []oidcJWK{buildRSAJWK(kid, &priv.PublicKey)}} + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + require.NoError(t, json.NewEncoder(w).Encode(jwks)) + })) + defer srv.Close() + + now := time.Now() + claims := oidcIDTokenClaims{ + Nonce: "nonce-ok", + Azp: "client-1", + RegisteredClaims: jwt.RegisteredClaims{ + Issuer: "https://issuer.example.com", + Subject: "subject-1", + Audience: jwt.ClaimStrings{"client-1", "another-aud"}, + IssuedAt: jwt.NewNumericDate(now), + NotBefore: jwt.NewNumericDate(now.Add(-30 * time.Second)), + ExpiresAt: jwt.NewNumericDate(now.Add(5 * time.Minute)), + }, + } + tok := jwt.NewWithClaims(jwt.SigningMethodRS256, claims) + tok.Header["kid"] = kid + signed, err := tok.SignedString(priv) + require.NoError(t, err) + + cfg := config.OIDCConnectConfig{ + ClientID: "client-1", + IssuerURL: "https://issuer.example.com", + JWKSURL: srv.URL, + AllowedSigningAlgs: "RS256", + ClockSkewSeconds: 120, + } + + parsed, err := oidcParseAndValidateIDToken(context.Background(), cfg, signed, "nonce-ok") + require.NoError(t, err) + require.Equal(t, "subject-1", parsed.Subject) + require.Equal(t, "https://issuer.example.com", parsed.Issuer) + + _, err = oidcParseAndValidateIDToken(context.Background(), cfg, signed, "bad-nonce") + require.Error(t, err) +} + +func buildRSAJWK(kid string, pub *rsa.PublicKey) oidcJWK { + n := base64.RawURLEncoding.EncodeToString(pub.N.Bytes()) + e := base64.RawURLEncoding.EncodeToString(big.NewInt(int64(pub.E)).Bytes()) + return oidcJWK{ + Kty: "RSA", + Kid: kid, + Use: "sig", + Alg: "RS256", + N: n, + E: e, + } +} diff --git a/backend/internal/handler/dto/mappers.go b/backend/internal/handler/dto/mappers.go index 2eab670e75..478600eb8c 100644 --- a/backend/internal/handler/dto/mappers.go +++ b/backend/internal/handler/dto/mappers.go @@ -133,16 +133,17 @@ func GroupFromServiceAdmin(g *service.Group) *AdminGroup { return nil } out := &AdminGroup{ - Group: groupFromServiceBase(g), - ModelRouting: g.ModelRouting, - ModelRoutingEnabled: g.ModelRoutingEnabled, - MCPXMLInject: g.MCPXMLInject, - DefaultMappedModel: g.DefaultMappedModel, - SupportedModelScopes: g.SupportedModelScopes, - AccountCount: g.AccountCount, - ActiveAccountCount: g.ActiveAccountCount, - RateLimitedAccountCount: g.RateLimitedAccountCount, - SortOrder: g.SortOrder, + Group: groupFromServiceBase(g), + ModelRouting: g.ModelRouting, + ModelRoutingEnabled: g.ModelRoutingEnabled, + MCPXMLInject: g.MCPXMLInject, + DefaultMappedModel: g.DefaultMappedModel, + MessagesDispatchModelConfig: g.MessagesDispatchModelConfig, + SupportedModelScopes: g.SupportedModelScopes, + AccountCount: g.AccountCount, + ActiveAccountCount: g.ActiveAccountCount, + RateLimitedAccountCount: g.RateLimitedAccountCount, + SortOrder: g.SortOrder, } if len(g.AccountGroups) > 0 { out.AccountGroups = make([]AccountGroup, 0, len(g.AccountGroups)) diff --git a/backend/internal/handler/dto/settings.go b/backend/internal/handler/dto/settings.go index b419b9703d..cbbe92160d 100644 --- a/backend/internal/handler/dto/settings.go +++ b/backend/internal/handler/dto/settings.go @@ -51,6 +51,29 @@ type SystemSettings struct { LinuxDoConnectClientSecretConfigured bool `json:"linuxdo_connect_client_secret_configured"` LinuxDoConnectRedirectURL string `json:"linuxdo_connect_redirect_url"` + OIDCConnectEnabled bool `json:"oidc_connect_enabled"` + OIDCConnectProviderName string `json:"oidc_connect_provider_name"` + OIDCConnectClientID string `json:"oidc_connect_client_id"` + OIDCConnectClientSecretConfigured bool `json:"oidc_connect_client_secret_configured"` + OIDCConnectIssuerURL string `json:"oidc_connect_issuer_url"` + OIDCConnectDiscoveryURL string `json:"oidc_connect_discovery_url"` + OIDCConnectAuthorizeURL string `json:"oidc_connect_authorize_url"` + OIDCConnectTokenURL string `json:"oidc_connect_token_url"` + OIDCConnectUserInfoURL string `json:"oidc_connect_userinfo_url"` + OIDCConnectJWKSURL string `json:"oidc_connect_jwks_url"` + OIDCConnectScopes string `json:"oidc_connect_scopes"` + OIDCConnectRedirectURL string `json:"oidc_connect_redirect_url"` + OIDCConnectFrontendRedirectURL string `json:"oidc_connect_frontend_redirect_url"` + OIDCConnectTokenAuthMethod string `json:"oidc_connect_token_auth_method"` + OIDCConnectUsePKCE bool `json:"oidc_connect_use_pkce"` + OIDCConnectValidateIDToken bool `json:"oidc_connect_validate_id_token"` + OIDCConnectAllowedSigningAlgs string `json:"oidc_connect_allowed_signing_algs"` + OIDCConnectClockSkewSeconds int `json:"oidc_connect_clock_skew_seconds"` + OIDCConnectRequireEmailVerified bool `json:"oidc_connect_require_email_verified"` + OIDCConnectUserInfoEmailPath string `json:"oidc_connect_userinfo_email_path"` + OIDCConnectUserInfoIDPath string `json:"oidc_connect_userinfo_id_path"` + OIDCConnectUserInfoUsernamePath string `json:"oidc_connect_userinfo_username_path"` + SiteName string `json:"site_name"` SiteLogo string `json:"site_logo"` SiteSubtitle string `json:"site_subtitle"` @@ -61,6 +84,8 @@ type SystemSettings struct { HideCcsImportButton bool `json:"hide_ccs_import_button"` PurchaseSubscriptionEnabled bool `json:"purchase_subscription_enabled"` PurchaseSubscriptionURL string `json:"purchase_subscription_url"` + TableDefaultPageSize int `json:"table_default_page_size"` + TablePageSizeOptions []int `json:"table_page_size_options"` CustomMenuItems []CustomMenuItem `json:"custom_menu_items"` CustomEndpoints []CustomEndpoint `json:"custom_endpoints"` @@ -97,6 +122,7 @@ type SystemSettings struct { // Gateway forwarding behavior EnableFingerprintUnification bool `json:"enable_fingerprint_unification"` EnableMetadataPassthrough bool `json:"enable_metadata_passthrough"` + EnableCCHSigning bool `json:"enable_cch_signing"` // Payment configuration PaymentEnabled bool `json:"payment_enabled"` @@ -146,9 +172,14 @@ type PublicSettings struct { HideCcsImportButton bool `json:"hide_ccs_import_button"` PurchaseSubscriptionEnabled bool `json:"purchase_subscription_enabled"` PurchaseSubscriptionURL string `json:"purchase_subscription_url"` + TableDefaultPageSize int `json:"table_default_page_size"` + TablePageSizeOptions []int `json:"table_page_size_options"` CustomMenuItems []CustomMenuItem `json:"custom_menu_items"` CustomEndpoints []CustomEndpoint `json:"custom_endpoints"` LinuxDoOAuthEnabled bool `json:"linuxdo_oauth_enabled"` + OIDCOAuthEnabled bool `json:"oidc_oauth_enabled"` + OIDCOAuthProviderName string `json:"oidc_oauth_provider_name"` + SoraClientEnabled bool `json:"sora_client_enabled"` BackendModeEnabled bool `json:"backend_mode_enabled"` PaymentEnabled bool `json:"payment_enabled"` Version string `json:"version"` @@ -180,10 +211,13 @@ type RectifierSettings struct { // BetaPolicyRule Beta 策略规则 DTO type BetaPolicyRule struct { - BetaToken string `json:"beta_token"` - Action string `json:"action"` - Scope string `json:"scope"` - ErrorMessage string `json:"error_message,omitempty"` + BetaToken string `json:"beta_token"` + Action string `json:"action"` + Scope string `json:"scope"` + ErrorMessage string `json:"error_message,omitempty"` + ModelWhitelist []string `json:"model_whitelist,omitempty"` + FallbackAction string `json:"fallback_action,omitempty"` + FallbackErrorMessage string `json:"fallback_error_message,omitempty"` } // BetaPolicySettings Beta 策略配置 DTO diff --git a/backend/internal/handler/dto/types.go b/backend/internal/handler/dto/types.go index 82065deb72..e026ca6551 100644 --- a/backend/internal/handler/dto/types.go +++ b/backend/internal/handler/dto/types.go @@ -1,6 +1,10 @@ package dto -import "time" +import ( + "time" + + "github.com/Wei-Shaw/sub2api/internal/domain" +) type User struct { ID int64 `json:"id"` @@ -112,7 +116,8 @@ type AdminGroup struct { MCPXMLInject bool `json:"mcp_xml_inject"` // OpenAI Messages 调度配置(仅 openai 平台使用) - DefaultMappedModel string `json:"default_mapped_model"` + DefaultMappedModel string `json:"default_mapped_model"` + MessagesDispatchModelConfig domain.OpenAIMessagesDispatchModelConfig `json:"messages_dispatch_model_config"` // 支持的模型系列(仅 antigravity 平台使用) SupportedModelScopes []string `json:"supported_model_scopes"` diff --git a/backend/internal/handler/gateway_handler_warmup_intercept_unit_test.go b/backend/internal/handler/gateway_handler_warmup_intercept_unit_test.go index 4caef9551b..acea37804f 100644 --- a/backend/internal/handler/gateway_handler_warmup_intercept_unit_test.go +++ b/backend/internal/handler/gateway_handler_warmup_intercept_unit_test.go @@ -34,7 +34,12 @@ func (f *fakeSchedulerCache) GetSnapshot(_ context.Context, _ service.SchedulerB func (f *fakeSchedulerCache) SetSnapshot(_ context.Context, _ service.SchedulerBucket, _ []service.Account) error { return nil } -func (f *fakeSchedulerCache) GetAccount(_ context.Context, _ int64) (*service.Account, error) { +func (f *fakeSchedulerCache) GetAccount(_ context.Context, id int64) (*service.Account, error) { + for _, account := range f.accounts { + if account != nil && account.ID == id { + return account, nil + } + } return nil, nil } func (f *fakeSchedulerCache) SetAccount(_ context.Context, _ *service.Account) error { return nil } diff --git a/backend/internal/handler/openai_gateway_handler.go b/backend/internal/handler/openai_gateway_handler.go index 4747ccfe1e..5319b55d9d 100644 --- a/backend/internal/handler/openai_gateway_handler.go +++ b/backend/internal/handler/openai_gateway_handler.go @@ -47,6 +47,13 @@ func resolveOpenAIForwardDefaultMappedModel(apiKey *service.APIKey, fallbackMode return strings.TrimSpace(apiKey.Group.DefaultMappedModel) } +func resolveOpenAIMessagesDispatchMappedModel(apiKey *service.APIKey, requestedModel string) string { + if apiKey == nil || apiKey.Group == nil { + return "" + } + return strings.TrimSpace(apiKey.Group.ResolveMessagesDispatchModel(requestedModel)) +} + // NewOpenAIGatewayHandler creates a new OpenAIGatewayHandler func NewOpenAIGatewayHandler( gatewayService *service.OpenAIGatewayService, @@ -551,6 +558,7 @@ func (h *OpenAIGatewayHandler) Messages(c *gin.Context) { } reqModel := modelResult.String() routingModel := service.NormalizeOpenAICompatRequestedModel(reqModel) + preferredMappedModel := resolveOpenAIMessagesDispatchMappedModel(apiKey, reqModel) reqStream := gjson.GetBytes(body, "stream").Bool() reqLog = reqLog.With(zap.String("model", reqModel), zap.Bool("stream", reqStream)) @@ -609,17 +617,20 @@ func (h *OpenAIGatewayHandler) Messages(c *gin.Context) { failedAccountIDs := make(map[int64]struct{}) sameAccountRetryCount := make(map[int64]int) var lastFailoverErr *service.UpstreamFailoverError + effectiveMappedModel := preferredMappedModel for { - // 清除上一次迭代的降级模型标记,避免残留影响本次迭代 - c.Set("openai_messages_fallback_model", "") + currentRoutingModel := routingModel + if effectiveMappedModel != "" { + currentRoutingModel = effectiveMappedModel + } reqLog.Debug("openai_messages.account_selecting", zap.Int("excluded_account_count", len(failedAccountIDs))) selection, scheduleDecision, err := h.gatewayService.SelectAccountWithScheduler( c.Request.Context(), apiKey.GroupID, "", // no previous_response_id sessionHash, - routingModel, + currentRoutingModel, failedAccountIDs, service.OpenAIUpstreamTransportAny, ) @@ -628,29 +639,7 @@ func (h *OpenAIGatewayHandler) Messages(c *gin.Context) { zap.Error(err), zap.Int("excluded_account_count", len(failedAccountIDs)), ) - // 首次调度失败 + 有默认映射模型 → 用默认模型重试 if len(failedAccountIDs) == 0 { - defaultModel := "" - if apiKey.Group != nil { - defaultModel = apiKey.Group.DefaultMappedModel - } - if defaultModel != "" && defaultModel != routingModel { - reqLog.Info("openai_messages.fallback_to_default_model", - zap.String("default_mapped_model", defaultModel), - ) - selection, scheduleDecision, err = h.gatewayService.SelectAccountWithScheduler( - c.Request.Context(), - apiKey.GroupID, - "", - sessionHash, - defaultModel, - failedAccountIDs, - service.OpenAIUpstreamTransportAny, - ) - if err == nil && selection != nil { - c.Set("openai_messages_fallback_model", defaultModel) - } - } if err != nil { h.anthropicStreamingAwareError(c, http.StatusServiceUnavailable, "api_error", "Service temporarily unavailable", streamStarted) return @@ -682,9 +671,7 @@ func (h *OpenAIGatewayHandler) Messages(c *gin.Context) { service.SetOpsLatencyMs(c, service.OpsRoutingLatencyMsKey, time.Since(routingStart).Milliseconds()) forwardStart := time.Now() - // Forward 层需要始终拿到 group 默认映射模型,这样未命中账号级映射的 - // Claude 兼容模型才不会在后续 Codex 规范化中意外退化到 gpt-5.1。 - defaultMappedModel := resolveOpenAIForwardDefaultMappedModel(apiKey, c.GetString("openai_messages_fallback_model")) + defaultMappedModel := strings.TrimSpace(effectiveMappedModel) // 应用渠道模型映射到请求体 forwardBody := body if channelMappingMsg.Mapped { diff --git a/backend/internal/handler/openai_gateway_handler_test.go b/backend/internal/handler/openai_gateway_handler_test.go index 7bbf94ecb3..d299fb81e3 100644 --- a/backend/internal/handler/openai_gateway_handler_test.go +++ b/backend/internal/handler/openai_gateway_handler_test.go @@ -360,7 +360,7 @@ func TestResolveOpenAIForwardDefaultMappedModel(t *testing.T) { require.Equal(t, "gpt-5.2", resolveOpenAIForwardDefaultMappedModel(apiKey, " gpt-5.2 ")) }) - t.Run("uses_group_default_on_normal_path", func(t *testing.T) { + t.Run("uses_group_default_when_explicit_fallback_absent", func(t *testing.T) { apiKey := &service.APIKey{ Group: &service.Group{DefaultMappedModel: "gpt-5.4"}, } @@ -376,6 +376,45 @@ func TestResolveOpenAIForwardDefaultMappedModel(t *testing.T) { }) } +func TestResolveOpenAIMessagesDispatchMappedModel(t *testing.T) { + t.Run("exact_claude_model_override_wins", func(t *testing.T) { + apiKey := &service.APIKey{ + Group: &service.Group{ + MessagesDispatchModelConfig: service.OpenAIMessagesDispatchModelConfig{ + SonnetMappedModel: "gpt-5.2", + ExactModelMappings: map[string]string{ + "claude-sonnet-4-5-20250929": "gpt-5.4-mini-high", + }, + }, + }, + } + require.Equal(t, "gpt-5.4-mini", resolveOpenAIMessagesDispatchMappedModel(apiKey, "claude-sonnet-4-5-20250929")) + }) + + t.Run("uses_family_default_when_no_override", func(t *testing.T) { + apiKey := &service.APIKey{Group: &service.Group{}} + require.Equal(t, "gpt-5.4", resolveOpenAIMessagesDispatchMappedModel(apiKey, "claude-opus-4-6")) + require.Equal(t, "gpt-5.3-codex", resolveOpenAIMessagesDispatchMappedModel(apiKey, "claude-sonnet-4-5-20250929")) + require.Equal(t, "gpt-5.4-mini", resolveOpenAIMessagesDispatchMappedModel(apiKey, "claude-haiku-4-5-20251001")) + }) + + t.Run("returns_empty_for_non_claude_or_missing_group", func(t *testing.T) { + require.Empty(t, resolveOpenAIMessagesDispatchMappedModel(nil, "claude-sonnet-4-5-20250929")) + require.Empty(t, resolveOpenAIMessagesDispatchMappedModel(&service.APIKey{}, "claude-sonnet-4-5-20250929")) + require.Empty(t, resolveOpenAIMessagesDispatchMappedModel(&service.APIKey{Group: &service.Group{}}, "gpt-5.4")) + }) + + t.Run("does_not_fall_back_to_group_default_mapped_model", func(t *testing.T) { + apiKey := &service.APIKey{ + Group: &service.Group{ + DefaultMappedModel: "gpt-5.4", + }, + } + require.Empty(t, resolveOpenAIMessagesDispatchMappedModel(apiKey, "gpt-5.4")) + require.Equal(t, "gpt-5.3-codex", resolveOpenAIMessagesDispatchMappedModel(apiKey, "claude-sonnet-4-5-20250929")) + }) +} + func TestOpenAIResponses_MissingDependencies_ReturnsServiceUnavailable(t *testing.T) { gin.SetMode(gin.TestMode) diff --git a/backend/internal/handler/setting_handler.go b/backend/internal/handler/setting_handler.go index 5917fed010..54a92a8c78 100644 --- a/backend/internal/handler/setting_handler.go +++ b/backend/internal/handler/setting_handler.go @@ -51,9 +51,13 @@ func (h *SettingHandler) GetPublicSettings(c *gin.Context) { HideCcsImportButton: settings.HideCcsImportButton, PurchaseSubscriptionEnabled: settings.PurchaseSubscriptionEnabled, PurchaseSubscriptionURL: settings.PurchaseSubscriptionURL, + TableDefaultPageSize: settings.TableDefaultPageSize, + TablePageSizeOptions: settings.TablePageSizeOptions, CustomMenuItems: dto.ParseUserVisibleMenuItems(settings.CustomMenuItems), CustomEndpoints: dto.ParseCustomEndpoints(settings.CustomEndpoints), LinuxDoOAuthEnabled: settings.LinuxDoOAuthEnabled, + OIDCOAuthEnabled: settings.OIDCOAuthEnabled, + OIDCOAuthProviderName: settings.OIDCOAuthProviderName, BackendModeEnabled: settings.BackendModeEnabled, PaymentEnabled: settings.PaymentEnabled, Version: h.version, diff --git a/backend/internal/handler/usage_handler.go b/backend/internal/handler/usage_handler.go index 483f51059b..b850615436 100644 --- a/backend/internal/handler/usage_handler.go +++ b/backend/internal/handler/usage_handler.go @@ -119,7 +119,12 @@ func (h *UsageHandler) List(c *gin.Context) { endTime = &t } - params := pagination.PaginationParams{Page: page, PageSize: pageSize} + params := pagination.PaginationParams{ + Page: page, + PageSize: pageSize, + SortBy: c.DefaultQuery("sort_by", "created_at"), + SortOrder: c.DefaultQuery("sort_order", "desc"), + } filters := usagestats.UsageLogFilters{ UserID: subject.UserID, // Always filter by current user for security APIKeyID: apiKeyID, diff --git a/backend/internal/handler/usage_handler_request_type_test.go b/backend/internal/handler/usage_handler_request_type_test.go index 7c4c79135a..b49ed59ba3 100644 --- a/backend/internal/handler/usage_handler_request_type_test.go +++ b/backend/internal/handler/usage_handler_request_type_test.go @@ -16,10 +16,12 @@ import ( type userUsageRepoCapture struct { service.UsageLogRepository + listParams pagination.PaginationParams listFilters usagestats.UsageLogFilters } func (s *userUsageRepoCapture) ListWithFilters(ctx context.Context, params pagination.PaginationParams, filters usagestats.UsageLogFilters) ([]service.UsageLog, *pagination.PaginationResult, error) { + s.listParams = params s.listFilters = filters return []service.UsageLog{}, &pagination.PaginationResult{ Total: 0, diff --git a/backend/internal/handler/usage_handler_sort_test.go b/backend/internal/handler/usage_handler_sort_test.go new file mode 100644 index 0000000000..1af313b09d --- /dev/null +++ b/backend/internal/handler/usage_handler_sort_test.go @@ -0,0 +1,35 @@ +package handler + +import ( + "net/http" + "net/http/httptest" + "testing" + + "github.com/stretchr/testify/require" +) + +func TestUserUsageListSortParams(t *testing.T) { + repo := &userUsageRepoCapture{} + router := newUserUsageRequestTypeTestRouter(repo) + + req := httptest.NewRequest(http.MethodGet, "/usage?sort_by=model&sort_order=ASC", nil) + rec := httptest.NewRecorder() + router.ServeHTTP(rec, req) + + require.Equal(t, http.StatusOK, rec.Code) + require.Equal(t, "model", repo.listParams.SortBy) + require.Equal(t, "ASC", repo.listParams.SortOrder) +} + +func TestUserUsageListSortDefaults(t *testing.T) { + repo := &userUsageRepoCapture{} + router := newUserUsageRequestTypeTestRouter(repo) + + req := httptest.NewRequest(http.MethodGet, "/usage", nil) + rec := httptest.NewRecorder() + router.ServeHTTP(rec, req) + + require.Equal(t, http.StatusOK, rec.Code) + require.Equal(t, "created_at", repo.listParams.SortBy) + require.Equal(t, "desc", repo.listParams.SortOrder) +} diff --git a/backend/internal/pkg/antigravity/request_transformer.go b/backend/internal/pkg/antigravity/request_transformer.go index 1b45e507fe..d13a84983f 100644 --- a/backend/internal/pkg/antigravity/request_transformer.go +++ b/backend/internal/pkg/antigravity/request_transformer.go @@ -730,13 +730,14 @@ func buildTools(tools []ClaudeTool) []GeminiToolDeclaration { }) } - if len(funcDecls) == 0 { - if !hasWebSearch { - return nil - } - - // Web Search 工具映射 - return []GeminiToolDeclaration{{ + var declarations []GeminiToolDeclaration + if len(funcDecls) > 0 { + declarations = append(declarations, GeminiToolDeclaration{ + FunctionDeclarations: funcDecls, + }) + } + if hasWebSearch { + declarations = append(declarations, GeminiToolDeclaration{ GoogleSearch: &GeminiGoogleSearch{ EnhancedContent: &GeminiEnhancedContent{ ImageSearch: &GeminiImageSearch{ @@ -744,10 +745,11 @@ func buildTools(tools []ClaudeTool) []GeminiToolDeclaration { }, }, }, - }} + }) + } + if len(declarations) == 0 { + return nil } - return []GeminiToolDeclaration{{ - FunctionDeclarations: funcDecls, - }} + return declarations } diff --git a/backend/internal/pkg/antigravity/request_transformer_test.go b/backend/internal/pkg/antigravity/request_transformer_test.go index 9e46295a8d..6fae5b7c56 100644 --- a/backend/internal/pkg/antigravity/request_transformer_test.go +++ b/backend/internal/pkg/antigravity/request_transformer_test.go @@ -263,6 +263,29 @@ func TestBuildTools_CustomTypeTools(t *testing.T) { } } +func TestBuildTools_PreservesWebSearchAlongsideFunctions(t *testing.T) { + tools := []ClaudeTool{ + { + Name: "get_weather", + Description: "Get weather information", + InputSchema: map[string]any{"type": "object"}, + }, + { + Type: "web_search_20250305", + Name: "web_search", + }, + } + + result := buildTools(tools) + require.Len(t, result, 2) + require.Len(t, result[0].FunctionDeclarations, 1) + require.Equal(t, "get_weather", result[0].FunctionDeclarations[0].Name) + require.NotNil(t, result[1].GoogleSearch) + require.NotNil(t, result[1].GoogleSearch.EnhancedContent) + require.NotNil(t, result[1].GoogleSearch.EnhancedContent.ImageSearch) + require.Equal(t, 5, result[1].GoogleSearch.EnhancedContent.ImageSearch.MaxResultCount) +} + func TestBuildGenerationConfig_ThinkingDynamicBudget(t *testing.T) { tests := []struct { name string @@ -400,3 +423,36 @@ func TestTransformClaudeToGeminiWithOptions_PreservesBillingHeaderSystemBlock(t }) } } + +func TestTransformClaudeToGeminiWithOptions_PreservesWebSearchAlongsideFunctions(t *testing.T) { + claudeReq := &ClaudeRequest{ + Model: "claude-3-5-sonnet-latest", + Messages: []ClaudeMessage{ + { + Role: "user", + Content: json.RawMessage(`[{"type":"text","text":"hello"}]`), + }, + }, + Tools: []ClaudeTool{ + { + Name: "get_weather", + Description: "Get weather information", + InputSchema: map[string]any{"type": "object"}, + }, + { + Type: "web_search_20250305", + Name: "web_search", + }, + }, + } + + body, err := TransformClaudeToGeminiWithOptions(claudeReq, "project-1", "gemini-2.5-flash", DefaultTransformOptions()) + require.NoError(t, err) + + var req V1InternalRequest + require.NoError(t, json.Unmarshal(body, &req)) + require.Len(t, req.Request.Tools, 2) + require.Len(t, req.Request.Tools[0].FunctionDeclarations, 1) + require.Equal(t, "get_weather", req.Request.Tools[0].FunctionDeclarations[0].Name) + require.NotNil(t, req.Request.Tools[1].GoogleSearch) +} diff --git a/backend/internal/pkg/apicompat/chatcompletions_responses_test.go b/backend/internal/pkg/apicompat/chatcompletions_responses_test.go index f54a4a027f..c140449a00 100644 --- a/backend/internal/pkg/apicompat/chatcompletions_responses_test.go +++ b/backend/internal/pkg/apicompat/chatcompletions_responses_test.go @@ -181,6 +181,50 @@ func TestChatCompletionsToResponses_ImageURL(t *testing.T) { assert.Equal(t, "data:image/png;base64,abc123", parts[1].ImageURL) } +func TestChatCompletionsToResponses_EmptyBase64ImageURLSkipped(t *testing.T) { + content := `[{"type":"text","text":"Describe this"},{"type":"image_url","image_url":{"url":"data:image/png;base64,"}}]` + req := &ChatCompletionsRequest{ + Model: "gpt-4o", + Messages: []ChatMessage{ + {Role: "user", Content: json.RawMessage(content)}, + }, + } + resp, err := ChatCompletionsToResponses(req) + require.NoError(t, err) + + var items []ResponsesInputItem + require.NoError(t, json.Unmarshal(resp.Input, &items)) + require.Len(t, items, 1) + + var parts []ResponsesContentPart + require.NoError(t, json.Unmarshal(items[0].Content, &parts)) + require.Len(t, parts, 1) + assert.Equal(t, "input_text", parts[0].Type) + assert.Equal(t, "Describe this", parts[0].Text) +} + +func TestChatCompletionsToResponses_WhitespaceOnlyBase64ImageURLSkipped(t *testing.T) { + content := `[{"type":"text","text":"Describe this"},{"type":"image_url","image_url":{"url":"data:image/png;base64, "}}]` + req := &ChatCompletionsRequest{ + Model: "gpt-4o", + Messages: []ChatMessage{ + {Role: "user", Content: json.RawMessage(content)}, + }, + } + resp, err := ChatCompletionsToResponses(req) + require.NoError(t, err) + + var items []ResponsesInputItem + require.NoError(t, json.Unmarshal(resp.Input, &items)) + require.Len(t, items, 1) + + var parts []ResponsesContentPart + require.NoError(t, json.Unmarshal(items[0].Content, &parts)) + require.Len(t, parts, 1) + assert.Equal(t, "input_text", parts[0].Type) + assert.Equal(t, "Describe this", parts[0].Text) +} + func TestChatCompletionsToResponses_SystemArrayContent(t *testing.T) { req := &ChatCompletionsRequest{ Model: "gpt-4o", @@ -876,3 +920,182 @@ func TestChatCompletionsStreamRoundTrip(t *testing.T) { assert.Equal(t, "resp_rt", c.ID) } } + +// --------------------------------------------------------------------------- +// BufferedResponseAccumulator tests +// --------------------------------------------------------------------------- + +func TestBufferedResponseAccumulator_TextOnly(t *testing.T) { + acc := NewBufferedResponseAccumulator() + + acc.ProcessEvent(&ResponsesStreamEvent{Type: "response.output_text.delta", Delta: "Hello"}) + acc.ProcessEvent(&ResponsesStreamEvent{Type: "response.output_text.delta", Delta: ", world!"}) + + assert.True(t, acc.HasContent()) + + output := acc.BuildOutput() + require.Len(t, output, 1) + assert.Equal(t, "message", output[0].Type) + assert.Equal(t, "assistant", output[0].Role) + require.Len(t, output[0].Content, 1) + assert.Equal(t, "output_text", output[0].Content[0].Type) + assert.Equal(t, "Hello, world!", output[0].Content[0].Text) +} + +func TestBufferedResponseAccumulator_ToolCalls(t *testing.T) { + acc := NewBufferedResponseAccumulator() + + // Add function call at output_index=1 + acc.ProcessEvent(&ResponsesStreamEvent{ + Type: "response.output_item.added", + OutputIndex: 1, + Item: &ResponsesOutput{ + Type: "function_call", + CallID: "call_abc", + Name: "get_weather", + }, + }) + acc.ProcessEvent(&ResponsesStreamEvent{ + Type: "response.function_call_arguments.delta", + OutputIndex: 1, + Delta: `{"city":`, + }) + acc.ProcessEvent(&ResponsesStreamEvent{ + Type: "response.function_call_arguments.delta", + OutputIndex: 1, + Delta: `"NYC"}`, + }) + + assert.True(t, acc.HasContent()) + + output := acc.BuildOutput() + require.Len(t, output, 1) + assert.Equal(t, "function_call", output[0].Type) + assert.Equal(t, "call_abc", output[0].CallID) + assert.Equal(t, "get_weather", output[0].Name) + assert.Equal(t, `{"city":"NYC"}`, output[0].Arguments) +} + +func TestBufferedResponseAccumulator_Reasoning(t *testing.T) { + acc := NewBufferedResponseAccumulator() + + acc.ProcessEvent(&ResponsesStreamEvent{Type: "response.reasoning_summary_text.delta", Delta: "Step 1: "}) + acc.ProcessEvent(&ResponsesStreamEvent{Type: "response.reasoning_summary_text.delta", Delta: "think about it"}) + + assert.True(t, acc.HasContent()) + + output := acc.BuildOutput() + require.Len(t, output, 1) + assert.Equal(t, "reasoning", output[0].Type) + require.Len(t, output[0].Summary, 1) + assert.Equal(t, "summary_text", output[0].Summary[0].Type) + assert.Equal(t, "Step 1: think about it", output[0].Summary[0].Text) +} + +func TestBufferedResponseAccumulator_Mixed(t *testing.T) { + acc := NewBufferedResponseAccumulator() + + // Reasoning first + acc.ProcessEvent(&ResponsesStreamEvent{Type: "response.reasoning_summary_text.delta", Delta: "I thought about it."}) + + // Then text + acc.ProcessEvent(&ResponsesStreamEvent{Type: "response.output_text.delta", Delta: "The answer is 42."}) + + // Then a tool call + acc.ProcessEvent(&ResponsesStreamEvent{ + Type: "response.output_item.added", + OutputIndex: 2, + Item: &ResponsesOutput{ + Type: "function_call", + CallID: "call_1", + Name: "verify", + }, + }) + acc.ProcessEvent(&ResponsesStreamEvent{ + Type: "response.function_call_arguments.delta", + OutputIndex: 2, + Delta: `{}`, + }) + + assert.True(t, acc.HasContent()) + + output := acc.BuildOutput() + // Order: reasoning → message → function_calls + require.Len(t, output, 3) + assert.Equal(t, "reasoning", output[0].Type) + assert.Equal(t, "message", output[1].Type) + assert.Equal(t, "function_call", output[2].Type) + assert.Equal(t, "The answer is 42.", output[1].Content[0].Text) + assert.Equal(t, "verify", output[2].Name) +} + +func TestBufferedResponseAccumulator_SupplementEmptyOutput(t *testing.T) { + acc := NewBufferedResponseAccumulator() + acc.ProcessEvent(&ResponsesStreamEvent{Type: "response.output_text.delta", Delta: "Hello"}) + + resp := &ResponsesResponse{ + ID: "resp_1", + Status: "completed", + Output: nil, // empty output + Usage: &ResponsesUsage{InputTokens: 10, OutputTokens: 5}, + } + + acc.SupplementResponseOutput(resp) + + require.Len(t, resp.Output, 1) + assert.Equal(t, "message", resp.Output[0].Type) + assert.Equal(t, "Hello", resp.Output[0].Content[0].Text) + // Usage should be untouched + assert.Equal(t, 10, resp.Usage.InputTokens) +} + +func TestBufferedResponseAccumulator_NoSupplementWhenOutputExists(t *testing.T) { + acc := NewBufferedResponseAccumulator() + acc.ProcessEvent(&ResponsesStreamEvent{Type: "response.output_text.delta", Delta: "from deltas"}) + + resp := &ResponsesResponse{ + ID: "resp_2", + Status: "completed", + Output: []ResponsesOutput{ + { + Type: "message", + Content: []ResponsesContentPart{ + {Type: "output_text", Text: "from terminal event"}, + }, + }, + }, + } + + acc.SupplementResponseOutput(resp) + + // Output should NOT be overwritten + require.Len(t, resp.Output, 1) + assert.Equal(t, "from terminal event", resp.Output[0].Content[0].Text) +} + +func TestBufferedResponseAccumulator_EmptyDeltas(t *testing.T) { + acc := NewBufferedResponseAccumulator() + + // Process events with empty delta — should not accumulate + acc.ProcessEvent(&ResponsesStreamEvent{Type: "response.output_text.delta", Delta: ""}) + acc.ProcessEvent(&ResponsesStreamEvent{Type: "response.created"}) + + assert.False(t, acc.HasContent()) + + resp := &ResponsesResponse{ID: "resp_3", Status: "completed"} + acc.SupplementResponseOutput(resp) + assert.Nil(t, resp.Output) +} + +func TestBufferedResponseAccumulator_IgnoresNonFunctionCallItems(t *testing.T) { + acc := NewBufferedResponseAccumulator() + + // output_item.added with type "message" should be ignored + acc.ProcessEvent(&ResponsesStreamEvent{ + Type: "response.output_item.added", + OutputIndex: 0, + Item: &ResponsesOutput{Type: "message"}, + }) + + assert.False(t, acc.HasContent()) +} diff --git a/backend/internal/pkg/apicompat/chatcompletions_to_responses.go b/backend/internal/pkg/apicompat/chatcompletions_to_responses.go index c9a61ecc51..c272540624 100644 --- a/backend/internal/pkg/apicompat/chatcompletions_to_responses.go +++ b/backend/internal/pkg/apicompat/chatcompletions_to_responses.go @@ -340,7 +340,7 @@ func convertChatContentPartsToResponses(parts []ChatContentPart) []ResponsesCont }) } case "image_url": - if p.ImageURL != nil && p.ImageURL.URL != "" { + if p.ImageURL != nil && p.ImageURL.URL != "" && !isEmptyBase64DataURI(p.ImageURL.URL) { responseParts = append(responseParts, ResponsesContentPart{ Type: "input_image", ImageURL: p.ImageURL.URL, @@ -351,6 +351,22 @@ func convertChatContentPartsToResponses(parts []ChatContentPart) []ResponsesCont return responseParts } +func isEmptyBase64DataURI(raw string) bool { + if !strings.HasPrefix(raw, "data:") { + return false + } + rest := strings.TrimPrefix(raw, "data:") + semicolonIdx := strings.Index(rest, ";") + if semicolonIdx < 0 { + return false + } + rest = rest[semicolonIdx+1:] + if !strings.HasPrefix(rest, "base64,") { + return false + } + return strings.TrimSpace(strings.TrimPrefix(rest, "base64,")) == "" +} + func flattenChatContentParts(parts []ChatContentPart) string { var textParts []string for _, p := range parts { diff --git a/backend/internal/pkg/apicompat/responses_to_chatcompletions.go b/backend/internal/pkg/apicompat/responses_to_chatcompletions.go index 688a68ebf8..61b3bf9cde 100644 --- a/backend/internal/pkg/apicompat/responses_to_chatcompletions.go +++ b/backend/internal/pkg/apicompat/responses_to_chatcompletions.go @@ -5,6 +5,7 @@ import ( "encoding/hex" "encoding/json" "fmt" + "strings" "time" ) @@ -372,3 +373,119 @@ func generateChatCmplID() string { _, _ = rand.Read(b) return "chatcmpl-" + hex.EncodeToString(b) } + +// --------------------------------------------------------------------------- +// BufferedResponseAccumulator: accumulates SSE delta events for non-streaming +// paths where the terminal event may have empty output. +// --------------------------------------------------------------------------- + +type bufferedFuncCall struct { + CallID string + Name string + Args strings.Builder +} + +// BufferedResponseAccumulator collects content from Responses SSE delta events +// so that non-streaming handlers can reconstruct output when the terminal event +// (response.completed / response.done) carries an empty output array. +type BufferedResponseAccumulator struct { + text strings.Builder + reasoning strings.Builder + funcCalls []bufferedFuncCall + outputIndexToFuncIdx map[int]int +} + +// NewBufferedResponseAccumulator returns an initialised accumulator. +func NewBufferedResponseAccumulator() *BufferedResponseAccumulator { + return &BufferedResponseAccumulator{ + outputIndexToFuncIdx: make(map[int]int), + } +} + +// ProcessEvent inspects a single Responses SSE event and accumulates any +// content it carries. Only delta events that contribute to the final output +// are handled; all other event types are silently ignored. +func (a *BufferedResponseAccumulator) ProcessEvent(event *ResponsesStreamEvent) { + switch event.Type { + case "response.output_text.delta": + if event.Delta != "" { + _, _ = a.text.WriteString(event.Delta) + } + case "response.output_item.added": + if event.Item != nil && event.Item.Type == "function_call" { + idx := len(a.funcCalls) + a.outputIndexToFuncIdx[event.OutputIndex] = idx + a.funcCalls = append(a.funcCalls, bufferedFuncCall{ + CallID: event.Item.CallID, + Name: event.Item.Name, + }) + } + case "response.function_call_arguments.delta": + if event.Delta != "" { + if idx, ok := a.outputIndexToFuncIdx[event.OutputIndex]; ok { + _, _ = a.funcCalls[idx].Args.WriteString(event.Delta) + } + } + case "response.reasoning_summary_text.delta": + if event.Delta != "" { + _, _ = a.reasoning.WriteString(event.Delta) + } + } +} + +// HasContent reports whether any content has been accumulated. +func (a *BufferedResponseAccumulator) HasContent() bool { + return a.text.Len() > 0 || len(a.funcCalls) > 0 || a.reasoning.Len() > 0 +} + +// BuildOutput constructs a []ResponsesOutput from the accumulated delta +// content. The order matches what ResponsesToChatCompletions expects: +// reasoning → message → function_calls. +func (a *BufferedResponseAccumulator) BuildOutput() []ResponsesOutput { + var out []ResponsesOutput + + if a.reasoning.Len() > 0 { + out = append(out, ResponsesOutput{ + Type: "reasoning", + Summary: []ResponsesSummary{{ + Type: "summary_text", + Text: a.reasoning.String(), + }}, + }) + } + + if a.text.Len() > 0 { + out = append(out, ResponsesOutput{ + Type: "message", + Role: "assistant", + Content: []ResponsesContentPart{{ + Type: "output_text", + Text: a.text.String(), + }}, + }) + } + + for i := range a.funcCalls { + out = append(out, ResponsesOutput{ + Type: "function_call", + CallID: a.funcCalls[i].CallID, + Name: a.funcCalls[i].Name, + Arguments: a.funcCalls[i].Args.String(), + }) + } + + return out +} + +// SupplementResponseOutput fills resp.Output from accumulated delta content +// when the terminal event delivered an empty output array. If resp.Output is +// already populated, this is a no-op (preserves backward compatibility). +func (a *BufferedResponseAccumulator) SupplementResponseOutput(resp *ResponsesResponse) { + if resp == nil || len(resp.Output) > 0 { + return + } + if !a.HasContent() { + return + } + resp.Output = a.BuildOutput() +} diff --git a/backend/internal/pkg/apicompat/types.go b/backend/internal/pkg/apicompat/types.go index d9546cb008..e0d1a53e8a 100644 --- a/backend/internal/pkg/apicompat/types.go +++ b/backend/internal/pkg/apicompat/types.go @@ -28,7 +28,7 @@ type AnthropicRequest struct { // AnthropicOutputConfig controls output generation parameters. type AnthropicOutputConfig struct { - Effort string `json:"effort,omitempty"` // "low" | "medium" | "high" + Effort string `json:"effort,omitempty"` // "low" | "medium" | "high" | "max" } // AnthropicThinking configures extended thinking in the Anthropic API. @@ -168,7 +168,7 @@ type ResponsesRequest struct { // ResponsesReasoning configures reasoning effort in the Responses API. type ResponsesReasoning struct { - Effort string `json:"effort"` // "low" | "medium" | "high" + Effort string `json:"effort"` // "low" | "medium" | "high" | "xhigh" Summary string `json:"summary,omitempty"` // "auto" | "concise" | "detailed" } @@ -347,7 +347,7 @@ type ChatCompletionsRequest struct { StreamOptions *ChatStreamOptions `json:"stream_options,omitempty"` Tools []ChatTool `json:"tools,omitempty"` ToolChoice json.RawMessage `json:"tool_choice,omitempty"` - ReasoningEffort string `json:"reasoning_effort,omitempty"` // "low" | "medium" | "high" + ReasoningEffort string `json:"reasoning_effort,omitempty"` // "low" | "medium" | "high" | "xhigh" ServiceTier string `json:"service_tier,omitempty"` Stop json.RawMessage `json:"stop,omitempty"` // string or []string diff --git a/backend/internal/pkg/pagination/pagination.go b/backend/internal/pkg/pagination/pagination.go index c162588ae6..ce8e74b8ce 100644 --- a/backend/internal/pkg/pagination/pagination.go +++ b/backend/internal/pkg/pagination/pagination.go @@ -1,10 +1,19 @@ // Package pagination provides types and helpers for paginated responses. package pagination +import "strings" + +const ( + SortOrderAsc = "asc" + SortOrderDesc = "desc" +) + // PaginationParams 分页参数 type PaginationParams struct { - Page int - PageSize int + Page int + PageSize int + SortBy string + SortOrder string } // PaginationResult 分页结果 @@ -18,8 +27,9 @@ type PaginationResult struct { // DefaultPagination 默认分页参数 func DefaultPagination() PaginationParams { return PaginationParams{ - Page: 1, - PageSize: 20, + Page: 1, + PageSize: 20, + SortOrder: SortOrderDesc, } } @@ -36,8 +46,32 @@ func (p PaginationParams) Limit() int { if p.PageSize < 1 { return 20 } - if p.PageSize > 100 { - return 100 + if p.PageSize > 1000 { + return 1000 } return p.PageSize } + +// NormalizeSortOrder normalizes sort order to asc/desc and falls back to defaultOrder. +func NormalizeSortOrder(order string, defaultOrder string) string { + switch strings.ToLower(strings.TrimSpace(defaultOrder)) { + case SortOrderAsc: + defaultOrder = SortOrderAsc + default: + defaultOrder = SortOrderDesc + } + + switch strings.ToLower(strings.TrimSpace(order)) { + case SortOrderAsc: + return SortOrderAsc + case SortOrderDesc: + return SortOrderDesc + default: + return defaultOrder + } +} + +// NormalizedSortOrder returns the normalized sort order using defaultOrder as fallback. +func (p PaginationParams) NormalizedSortOrder(defaultOrder string) string { + return NormalizeSortOrder(p.SortOrder, defaultOrder) +} diff --git a/backend/internal/pkg/pagination/pagination_test.go b/backend/internal/pkg/pagination/pagination_test.go new file mode 100644 index 0000000000..9a3b069d90 --- /dev/null +++ b/backend/internal/pkg/pagination/pagination_test.go @@ -0,0 +1,71 @@ +package pagination + +import "testing" + +func TestNormalizeSortOrder(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + input string + defaultOrder string + want string + }{ + {name: "asc", input: "asc", defaultOrder: "desc", want: "asc"}, + {name: "uppercase asc", input: "ASC", defaultOrder: "desc", want: "asc"}, + {name: "desc", input: "desc", defaultOrder: "asc", want: "desc"}, + {name: "trim spaces", input: " desc ", defaultOrder: "asc", want: "desc"}, + {name: "invalid falls back", input: "sideways", defaultOrder: "asc", want: "asc"}, + {name: "empty falls back", input: "", defaultOrder: "desc", want: "desc"}, + {name: "invalid default falls back to desc", input: "", defaultOrder: "wat", want: "desc"}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + if got := NormalizeSortOrder(tt.input, tt.defaultOrder); got != tt.want { + t.Fatalf("NormalizeSortOrder(%q, %q) = %q, want %q", tt.input, tt.defaultOrder, got, tt.want) + } + }) + } +} + +func TestPaginationParamsNormalizedSortOrder(t *testing.T) { + t.Parallel() + + params := PaginationParams{SortOrder: "ASC"} + if got := params.NormalizedSortOrder("desc"); got != "asc" { + t.Fatalf("NormalizedSortOrder = %q, want asc", got) + } + + params = PaginationParams{SortOrder: "bad"} + if got := params.NormalizedSortOrder("asc"); got != "asc" { + t.Fatalf("NormalizedSortOrder invalid fallback = %q, want asc", got) + } +} + +func TestPaginationParamsLimit(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + pageSize int + want int + }{ + {name: "non-positive falls back to default", pageSize: 0, want: 20}, + {name: "negative falls back to default", pageSize: -1, want: 20}, + {name: "normal value keeps", pageSize: 50, want: 50}, + {name: "max value keeps", pageSize: 1000, want: 1000}, + {name: "beyond max clamps to 1000", pageSize: 1500, want: 1000}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + p := PaginationParams{PageSize: tt.pageSize} + if got := p.Limit(); got != tt.want { + t.Fatalf("Limit() for PageSize=%d = %d, want %d", tt.pageSize, got, tt.want) + } + }) + } +} diff --git a/backend/internal/repository/account_repo.go b/backend/internal/repository/account_repo.go index 94bfb09d58..24115c33d7 100644 --- a/backend/internal/repository/account_repo.go +++ b/backend/internal/repository/account_repo.go @@ -468,16 +468,61 @@ func (r *accountRepository) ListWithFilters(ctx context.Context, params paginati } if status != "" { switch status { + case service.StatusActive: + q = q.Where( + dbaccount.StatusEQ(status), + dbaccount.SchedulableEQ(true), + dbaccount.Or( + dbaccount.RateLimitResetAtIsNil(), + dbaccount.RateLimitResetAtLTE(time.Now()), + ), + dbpredicate.Account(func(s *entsql.Selector) { + col := s.C("temp_unschedulable_until") + s.Where(entsql.Or( + entsql.IsNull(col), + entsql.LTE(col, entsql.Expr("NOW()")), + )) + }), + ) case "rate_limited": - q = q.Where(dbaccount.RateLimitResetAtGT(time.Now())) + q = q.Where( + dbaccount.StatusEQ(service.StatusActive), + dbaccount.RateLimitResetAtGT(time.Now()), + dbpredicate.Account(func(s *entsql.Selector) { + col := s.C("temp_unschedulable_until") + s.Where(entsql.Or( + entsql.IsNull(col), + entsql.LTE(col, entsql.Expr("NOW()")), + )) + }), + ) case "temp_unschedulable": - q = q.Where(dbpredicate.Account(func(s *entsql.Selector) { - col := s.C("temp_unschedulable_until") - s.Where(entsql.And( - entsql.Not(entsql.IsNull(col)), - entsql.GT(col, entsql.Expr("NOW()")), - )) - })) + q = q.Where( + dbaccount.StatusEQ(service.StatusActive), + dbpredicate.Account(func(s *entsql.Selector) { + col := s.C("temp_unschedulable_until") + s.Where(entsql.And( + entsql.Not(entsql.IsNull(col)), + entsql.GT(col, entsql.Expr("NOW()")), + )) + }), + ) + case "unschedulable": + q = q.Where( + dbaccount.StatusEQ(service.StatusActive), + dbaccount.SchedulableEQ(false), + dbaccount.Or( + dbaccount.RateLimitResetAtIsNil(), + dbaccount.RateLimitResetAtLTE(time.Now()), + ), + dbpredicate.Account(func(s *entsql.Selector) { + col := s.C("temp_unschedulable_until") + s.Where(entsql.Or( + entsql.IsNull(col), + entsql.LTE(col, entsql.Expr("NOW()")), + )) + }), + ) default: q = q.Where(dbaccount.StatusEQ(status)) } @@ -510,11 +555,14 @@ func (r *accountRepository) ListWithFilters(ctx context.Context, params paginati return nil, nil, err } - accounts, err := q. + accountsQuery := q. Offset(params.Offset()). - Limit(params.Limit()). - Order(dbent.Desc(dbaccount.FieldID)). - All(ctx) + Limit(params.Limit()) + for _, order := range accountListOrder(params) { + accountsQuery = accountsQuery.Order(order) + } + + accounts, err := accountsQuery.All(ctx) if err != nil { return nil, nil, err } @@ -526,6 +574,50 @@ func (r *accountRepository) ListWithFilters(ctx context.Context, params paginati return outAccounts, paginationResultFromTotal(int64(total), params), nil } +func accountListOrder(params pagination.PaginationParams) []func(*entsql.Selector) { + sortBy := strings.ToLower(strings.TrimSpace(params.SortBy)) + sortOrder := params.NormalizedSortOrder(pagination.SortOrderAsc) + + field := dbaccount.FieldName + defaultOrder := true + switch sortBy { + case "", "name": + field = dbaccount.FieldName + case "id": + field = dbaccount.FieldID + defaultOrder = false + case "status": + field = dbaccount.FieldStatus + defaultOrder = false + case "schedulable": + field = dbaccount.FieldSchedulable + defaultOrder = false + case "priority": + field = dbaccount.FieldPriority + defaultOrder = false + case "rate_multiplier": + field = dbaccount.FieldRateMultiplier + defaultOrder = false + case "last_used_at": + field = dbaccount.FieldLastUsedAt + defaultOrder = false + case "expires_at": + field = dbaccount.FieldExpiresAt + defaultOrder = false + case "created_at": + field = dbaccount.FieldCreatedAt + defaultOrder = false + } + + if sortOrder == pagination.SortOrderDesc { + return []func(*entsql.Selector){dbent.Desc(field), dbent.Desc(dbaccount.FieldID)} + } + if defaultOrder { + return []func(*entsql.Selector){dbent.Asc(dbaccount.FieldName), dbent.Asc(dbaccount.FieldID)} + } + return []func(*entsql.Selector){dbent.Asc(field), dbent.Asc(dbaccount.FieldID)} +} + func (r *accountRepository) ListByGroup(ctx context.Context, groupID int64) ([]service.Account, error) { accounts, err := r.queryAccountsByGroup(ctx, groupID, accountGroupQueryOptions{ status: service.StatusActive, diff --git a/backend/internal/repository/account_repo_integration_test.go b/backend/internal/repository/account_repo_integration_test.go index 8da30c92a8..b249bb61b7 100644 --- a/backend/internal/repository/account_repo_integration_test.go +++ b/backend/internal/repository/account_repo_integration_test.go @@ -255,6 +255,101 @@ func (s *AccountRepoSuite) TestListWithFilters() { s.Require().Equal(service.StatusDisabled, accounts[0].Status) }, }, + { + name: "filter_by_status_active_excludes_runtime_blocked_accounts", + setup: func(client *dbent.Client) { + mustCreateAccount(s.T(), client, &service.Account{Name: "active-normal", Status: service.StatusActive}) + rateLimited := mustCreateAccount(s.T(), client, &service.Account{Name: "active-rate-limited", Status: service.StatusActive}) + err := client.Account.UpdateOneID(rateLimited.ID). + SetRateLimitResetAt(time.Now().Add(10 * time.Minute)). + Exec(context.Background()) + s.Require().NoError(err) + tempUnsched := mustCreateAccount(s.T(), client, &service.Account{Name: "active-temp-unsched", Status: service.StatusActive}) + err = client.Account.UpdateOneID(tempUnsched.ID). + SetTempUnschedulableUntil(time.Now().Add(15 * time.Minute)). + Exec(context.Background()) + s.Require().NoError(err) + unsched := mustCreateAccount(s.T(), client, &service.Account{Name: "active-unsched", Status: service.StatusActive}) + err = client.Account.UpdateOneID(unsched.ID). + SetSchedulable(false). + Exec(context.Background()) + s.Require().NoError(err) + }, + status: service.StatusActive, + wantCount: 1, + validate: func(accounts []service.Account) { + s.Require().Equal("active-normal", accounts[0].Name) + }, + }, + { + name: "filter_by_status_unschedulable_excludes_rate_limited_and_temp_unschedulable", + setup: func(client *dbent.Client) { + mustCreateAccount(s.T(), client, &service.Account{Name: "active-normal", Status: service.StatusActive, Schedulable: true}) + unsched := mustCreateAccount(s.T(), client, &service.Account{Name: "active-unsched", Status: service.StatusActive}) + err := client.Account.UpdateOneID(unsched.ID). + SetSchedulable(false). + Exec(context.Background()) + s.Require().NoError(err) + rateLimited := mustCreateAccount(s.T(), client, &service.Account{Name: "active-rate-limited", Status: service.StatusActive}) + err = client.Account.UpdateOneID(rateLimited.ID). + SetSchedulable(false). + SetRateLimitResetAt(time.Now().Add(10 * time.Minute)). + Exec(context.Background()) + s.Require().NoError(err) + tempUnsched := mustCreateAccount(s.T(), client, &service.Account{Name: "active-temp-unsched", Status: service.StatusActive}) + err = client.Account.UpdateOneID(tempUnsched.ID). + SetSchedulable(false). + SetTempUnschedulableUntil(time.Now().Add(15 * time.Minute)). + Exec(context.Background()) + s.Require().NoError(err) + }, + status: "unschedulable", + wantCount: 1, + validate: func(accounts []service.Account) { + s.Require().Equal("active-unsched", accounts[0].Name) + }, + }, + { + name: "filter_by_status_rate_limited_excludes_temp_unschedulable", + setup: func(client *dbent.Client) { + rateLimited := mustCreateAccount(s.T(), client, &service.Account{Name: "active-rate-limited", Status: service.StatusActive}) + err := client.Account.UpdateOneID(rateLimited.ID). + SetRateLimitResetAt(time.Now().Add(10 * time.Minute)). + Exec(context.Background()) + s.Require().NoError(err) + tempUnsched := mustCreateAccount(s.T(), client, &service.Account{Name: "active-temp-unsched", Status: service.StatusActive}) + err = client.Account.UpdateOneID(tempUnsched.ID). + SetRateLimitResetAt(time.Now().Add(20 * time.Minute)). + SetTempUnschedulableUntil(time.Now().Add(15 * time.Minute)). + Exec(context.Background()) + s.Require().NoError(err) + }, + status: "rate_limited", + wantCount: 1, + validate: func(accounts []service.Account) { + s.Require().Equal("active-rate-limited", accounts[0].Name) + }, + }, + { + name: "filter_by_status_temp_unschedulable_excludes_manually_unschedulable", + setup: func(client *dbent.Client) { + tempUnsched := mustCreateAccount(s.T(), client, &service.Account{Name: "active-temp-unsched", Status: service.StatusActive, Schedulable: true}) + err := client.Account.UpdateOneID(tempUnsched.ID). + SetTempUnschedulableUntil(time.Now().Add(15 * time.Minute)). + Exec(context.Background()) + s.Require().NoError(err) + unsched := mustCreateAccount(s.T(), client, &service.Account{Name: "active-unsched", Status: service.StatusActive}) + err = client.Account.UpdateOneID(unsched.ID). + SetSchedulable(false). + Exec(context.Background()) + s.Require().NoError(err) + }, + status: "temp_unschedulable", + wantCount: 1, + validate: func(accounts []service.Account) { + s.Require().Equal("active-temp-unsched", accounts[0].Name) + }, + }, { name: "filter_by_search", setup: func(client *dbent.Client) { diff --git a/backend/internal/repository/account_repo_sort_integration_test.go b/backend/internal/repository/account_repo_sort_integration_test.go new file mode 100644 index 0000000000..098dde7bbd --- /dev/null +++ b/backend/internal/repository/account_repo_sort_integration_test.go @@ -0,0 +1,35 @@ +//go:build integration + +package repository + +import ( + "github.com/Wei-Shaw/sub2api/internal/pkg/pagination" + "github.com/Wei-Shaw/sub2api/internal/service" +) + +func (s *AccountRepoSuite) TestList_DefaultSortByNameAsc() { + mustCreateAccount(s.T(), s.client, &service.Account{Name: "z-account"}) + mustCreateAccount(s.T(), s.client, &service.Account{Name: "a-account"}) + + accounts, _, err := s.repo.List(s.ctx, pagination.PaginationParams{Page: 1, PageSize: 10}) + s.Require().NoError(err) + s.Require().Len(accounts, 2) + s.Require().Equal("a-account", accounts[0].Name) + s.Require().Equal("z-account", accounts[1].Name) +} + +func (s *AccountRepoSuite) TestListWithFilters_SortByPriorityDesc() { + mustCreateAccount(s.T(), s.client, &service.Account{Name: "low-priority", Priority: 10}) + mustCreateAccount(s.T(), s.client, &service.Account{Name: "high-priority", Priority: 90}) + + accounts, _, err := s.repo.ListWithFilters(s.ctx, pagination.PaginationParams{ + Page: 1, + PageSize: 10, + SortBy: "priority", + SortOrder: "desc", + }, "", "", "", "", 0, "") + s.Require().NoError(err) + s.Require().Len(accounts, 2) + s.Require().Equal("high-priority", accounts[0].Name) + s.Require().Equal("low-priority", accounts[1].Name) +} diff --git a/backend/internal/repository/announcement_repo.go b/backend/internal/repository/announcement_repo.go index 53dc335f8c..afe1fb25c9 100644 --- a/backend/internal/repository/announcement_repo.go +++ b/backend/internal/repository/announcement_repo.go @@ -2,12 +2,15 @@ package repository import ( "context" + "strings" "time" dbent "github.com/Wei-Shaw/sub2api/ent" "github.com/Wei-Shaw/sub2api/ent/announcement" "github.com/Wei-Shaw/sub2api/internal/pkg/pagination" "github.com/Wei-Shaw/sub2api/internal/service" + + entsql "entgo.io/ent/dialect/sql" ) type announcementRepository struct { @@ -128,11 +131,14 @@ func (r *announcementRepository) List( return nil, nil, err } - items, err := q. + itemsQuery := q. Offset(params.Offset()). - Limit(params.Limit()). - Order(dbent.Desc(announcement.FieldID)). - All(ctx) + Limit(params.Limit()) + for _, order := range announcementListOrders(params) { + itemsQuery = itemsQuery.Order(order) + } + + items, err := itemsQuery.All(ctx) if err != nil { return nil, nil, err } @@ -141,6 +147,56 @@ func (r *announcementRepository) List( return out, paginationResultFromTotal(int64(total), params), nil } +func announcementListOrder(params pagination.PaginationParams) (string, string) { + sortBy := strings.ToLower(strings.TrimSpace(params.SortBy)) + sortOrder := params.NormalizedSortOrder(pagination.SortOrderDesc) + + switch sortBy { + case "title": + return announcement.FieldTitle, sortOrder + case "status": + return announcement.FieldStatus, sortOrder + case "notify_mode": + return announcement.FieldNotifyMode, sortOrder + case "starts_at": + return announcement.FieldStartsAt, sortOrder + case "ends_at": + return announcement.FieldEndsAt, sortOrder + case "id": + return announcement.FieldID, sortOrder + case "", "created_at": + return announcement.FieldCreatedAt, sortOrder + default: + return announcement.FieldCreatedAt, pagination.SortOrderDesc + } +} + +func announcementListOrders(params pagination.PaginationParams) []func(*entsql.Selector) { + field, sortOrder := announcementListOrder(params) + + if sortOrder == pagination.SortOrderAsc { + if field == announcement.FieldID { + return []func(*entsql.Selector){ + dbent.Asc(field), + } + } + return []func(*entsql.Selector){ + dbent.Asc(field), + dbent.Asc(announcement.FieldID), + } + } + + if field == announcement.FieldID { + return []func(*entsql.Selector){ + dbent.Desc(field), + } + } + return []func(*entsql.Selector){ + dbent.Desc(field), + dbent.Desc(announcement.FieldID), + } +} + func (r *announcementRepository) ListActive(ctx context.Context, now time.Time) ([]service.Announcement, error) { q := r.client.Announcement.Query(). Where( diff --git a/backend/internal/repository/announcement_repo_sort_test.go b/backend/internal/repository/announcement_repo_sort_test.go new file mode 100644 index 0000000000..e47f98dcfc --- /dev/null +++ b/backend/internal/repository/announcement_repo_sort_test.go @@ -0,0 +1,63 @@ +package repository + +import ( + "testing" + + "github.com/Wei-Shaw/sub2api/internal/pkg/pagination" +) + +func TestAnnouncementListOrder(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + params pagination.PaginationParams + wantBy string + want string + }{ + { + name: "default created_at desc", + params: pagination.PaginationParams{}, + wantBy: "created_at", + want: "desc", + }, + { + name: "title asc", + params: pagination.PaginationParams{ + SortBy: "title", + SortOrder: "ASC", + }, + wantBy: "title", + want: "asc", + }, + { + name: "status desc", + params: pagination.PaginationParams{ + SortBy: "status", + SortOrder: "desc", + }, + wantBy: "status", + want: "desc", + }, + { + name: "invalid falls back", + params: pagination.PaginationParams{ + SortBy: "sideways", + SortOrder: "wat", + }, + wantBy: "created_at", + want: "desc", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + gotBy, gotOrder := announcementListOrder(tt.params) + if gotBy != tt.wantBy || gotOrder != tt.want { + t.Fatalf("announcementListOrder(%+v) = (%q, %q), want (%q, %q)", tt.params, gotBy, gotOrder, tt.wantBy, tt.want) + } + }) + } +} diff --git a/backend/internal/repository/api_key_repo.go b/backend/internal/repository/api_key_repo.go index b3b12e8113..7fd988550f 100644 --- a/backend/internal/repository/api_key_repo.go +++ b/backend/internal/repository/api_key_repo.go @@ -4,6 +4,7 @@ import ( "context" "database/sql" "fmt" + "strings" "time" dbent "github.com/Wei-Shaw/sub2api/ent" @@ -14,6 +15,8 @@ import ( "github.com/Wei-Shaw/sub2api/internal/service" "github.com/Wei-Shaw/sub2api/internal/pkg/pagination" + + entsql "entgo.io/ent/dialect/sql" ) type apiKeyRepository struct { @@ -164,6 +167,7 @@ func (r *apiKeyRepository) GetByKeyForAuth(ctx context.Context, key string) (*se group.FieldSupportedModelScopes, group.FieldAllowMessagesDispatch, group.FieldDefaultMappedModel, + group.FieldMessagesDispatchModelConfig, ) }). Only(ctx) @@ -309,12 +313,15 @@ func (r *apiKeyRepository) ListByUserID(ctx context.Context, userID int64, param return nil, nil, err } - keys, err := q. + keysQuery := q. WithGroup(). Offset(params.Offset()). - Limit(params.Limit()). - Order(dbent.Desc(apikey.FieldID)). - All(ctx) + Limit(params.Limit()) + for _, order := range apiKeyListOrder(params) { + keysQuery = keysQuery.Order(order) + } + + keys, err := keysQuery.All(ctx) if err != nil { return nil, nil, err } @@ -359,12 +366,15 @@ func (r *apiKeyRepository) ListByGroupID(ctx context.Context, groupID int64, par return nil, nil, err } - keys, err := q. + keysQuery := q. WithUser(). Offset(params.Offset()). - Limit(params.Limit()). - Order(dbent.Desc(apikey.FieldID)). - All(ctx) + Limit(params.Limit()) + for _, order := range apiKeyListOrder(params) { + keysQuery = keysQuery.Order(order) + } + + keys, err := keysQuery.All(ctx) if err != nil { return nil, nil, err } @@ -377,6 +387,32 @@ func (r *apiKeyRepository) ListByGroupID(ctx context.Context, groupID int64, par return outKeys, paginationResultFromTotal(int64(total), params), nil } +func apiKeyListOrder(params pagination.PaginationParams) []func(*entsql.Selector) { + sortBy := strings.ToLower(strings.TrimSpace(params.SortBy)) + sortOrder := params.NormalizedSortOrder(pagination.SortOrderDesc) + + var field string + switch sortBy { + case "name": + field = apikey.FieldName + case "status": + field = apikey.FieldStatus + case "expires_at": + field = apikey.FieldExpiresAt + case "last_used_at": + field = apikey.FieldLastUsedAt + case "created_at": + field = apikey.FieldCreatedAt + default: + field = apikey.FieldID + } + + if sortOrder == pagination.SortOrderAsc { + return []func(*entsql.Selector){dbent.Asc(field), dbent.Asc(apikey.FieldID)} + } + return []func(*entsql.Selector){dbent.Desc(field), dbent.Desc(apikey.FieldID)} +} + // SearchAPIKeys searches API keys by user ID and/or keyword (name) func (r *apiKeyRepository) SearchAPIKeys(ctx context.Context, userID int64, keyword string, limit int) ([]service.APIKey, error) { q := r.activeQuery() @@ -654,6 +690,7 @@ func groupEntityToService(g *dbent.Group) *service.Group { RequireOAuthOnly: g.RequireOauthOnly, RequirePrivacySet: g.RequirePrivacySet, DefaultMappedModel: g.DefaultMappedModel, + MessagesDispatchModelConfig: g.MessagesDispatchModelConfig, CreatedAt: g.CreatedAt, UpdatedAt: g.UpdatedAt, } diff --git a/backend/internal/repository/api_key_repo_integration_test.go b/backend/internal/repository/api_key_repo_integration_test.go index 7d5c18260b..e926ed8602 100644 --- a/backend/internal/repository/api_key_repo_integration_test.go +++ b/backend/internal/repository/api_key_repo_integration_test.go @@ -86,6 +86,45 @@ func (s *APIKeyRepoSuite) TestGetByKey_NotFound() { s.Require().Error(err, "expected error for non-existent key") } +func (s *APIKeyRepoSuite) TestGetByKeyForAuth_PreservesMessagesDispatchModelConfig() { + user := s.mustCreateUser("getbykey-auth-dispatch@test.com") + group, err := s.client.Group.Create(). + SetName("g-auth-dispatch"). + SetPlatform(service.PlatformOpenAI). + SetStatus(service.StatusActive). + SetSubscriptionType(service.SubscriptionTypeStandard). + SetRateMultiplier(1). + SetAllowMessagesDispatch(true). + SetDefaultMappedModel("gpt-5.4"). + SetMessagesDispatchModelConfig(service.OpenAIMessagesDispatchModelConfig{ + OpusMappedModel: "gpt-5.4-nano", + SonnetMappedModel: "gpt-5.3-codex", + HaikuMappedModel: "gpt-5.4-mini", + ExactModelMappings: map[string]string{ + "claude-sonnet-4.5": "gpt-5.4-nano", + }, + }). + Save(s.ctx) + s.Require().NoError(err) + + key := &service.APIKey{ + UserID: user.ID, + Key: "sk-getbykey-auth-dispatch", + Name: "Dispatch Key", + GroupID: &group.ID, + Status: service.StatusActive, + } + s.Require().NoError(s.repo.Create(s.ctx, key)) + + got, err := s.repo.GetByKeyForAuth(s.ctx, key.Key) + s.Require().NoError(err) + s.Require().NotNil(got.Group) + s.Require().True(got.Group.AllowMessagesDispatch) + s.Require().Equal("gpt-5.4", got.Group.DefaultMappedModel) + s.Require().Equal("gpt-5.4-nano", got.Group.MessagesDispatchModelConfig.OpusMappedModel) + s.Require().Equal("gpt-5.4-nano", got.Group.MessagesDispatchModelConfig.ExactModelMappings["claude-sonnet-4.5"]) +} + // --- Update --- func (s *APIKeyRepoSuite) TestUpdate() { diff --git a/backend/internal/repository/api_key_repo_messages_dispatch_unit_test.go b/backend/internal/repository/api_key_repo_messages_dispatch_unit_test.go new file mode 100644 index 0000000000..aba62ead2e --- /dev/null +++ b/backend/internal/repository/api_key_repo_messages_dispatch_unit_test.go @@ -0,0 +1,74 @@ +package repository + +import ( + "context" + "testing" + + dbent "github.com/Wei-Shaw/sub2api/ent" + "github.com/Wei-Shaw/sub2api/internal/service" + "github.com/stretchr/testify/require" +) + +func TestGroupEntityToService_PreservesMessagesDispatchModelConfig(t *testing.T) { + group := &dbent.Group{ + ID: 1, + Name: "openai-dispatch", + Platform: service.PlatformOpenAI, + Status: service.StatusActive, + SubscriptionType: service.SubscriptionTypeStandard, + RateMultiplier: 1, + AllowMessagesDispatch: true, + DefaultMappedModel: "gpt-5.4", + MessagesDispatchModelConfig: service.OpenAIMessagesDispatchModelConfig{ + OpusMappedModel: "gpt-5.4-nano", + SonnetMappedModel: "gpt-5.3-codex", + HaikuMappedModel: "gpt-5.4-mini", + ExactModelMappings: map[string]string{ + "claude-sonnet-4.5": "gpt-5.4-nano", + }, + }, + } + + got := groupEntityToService(group) + require.NotNil(t, got) + require.Equal(t, group.MessagesDispatchModelConfig, got.MessagesDispatchModelConfig) +} + +func TestAPIKeyRepository_GetByKeyForAuth_PreservesMessagesDispatchModelConfig_SQLite(t *testing.T) { + repo, client := newAPIKeyRepoSQLite(t) + ctx := context.Background() + user := mustCreateAPIKeyRepoUser(t, ctx, client, "getbykey-auth-dispatch-unit@test.com") + + group, err := client.Group.Create(). + SetName("g-auth-dispatch-unit"). + SetPlatform(service.PlatformOpenAI). + SetStatus(service.StatusActive). + SetSubscriptionType(service.SubscriptionTypeStandard). + SetRateMultiplier(1). + SetAllowMessagesDispatch(true). + SetDefaultMappedModel("gpt-5.4"). + SetMessagesDispatchModelConfig(service.OpenAIMessagesDispatchModelConfig{ + OpusMappedModel: "gpt-5.4-nano", + SonnetMappedModel: "gpt-5.3-codex", + HaikuMappedModel: "gpt-5.4-mini", + ExactModelMappings: map[string]string{ + "claude-sonnet-4.5": "gpt-5.4-nano", + }, + }). + Save(ctx) + require.NoError(t, err) + + key := &service.APIKey{ + UserID: user.ID, + Key: "sk-getbykey-auth-dispatch-unit", + Name: "Dispatch Key Unit", + GroupID: &group.ID, + Status: service.StatusActive, + } + require.NoError(t, repo.Create(ctx, key)) + + got, err := repo.GetByKeyForAuth(ctx, key.Key) + require.NoError(t, err) + require.NotNil(t, got.Group) + require.Equal(t, group.MessagesDispatchModelConfig, got.Group.MessagesDispatchModelConfig) +} diff --git a/backend/internal/repository/api_key_repo_sort_integration_test.go b/backend/internal/repository/api_key_repo_sort_integration_test.go new file mode 100644 index 0000000000..69812882fe --- /dev/null +++ b/backend/internal/repository/api_key_repo_sort_integration_test.go @@ -0,0 +1,25 @@ +//go:build integration + +package repository + +import ( + "github.com/Wei-Shaw/sub2api/internal/pkg/pagination" + "github.com/Wei-Shaw/sub2api/internal/service" +) + +func (s *APIKeyRepoSuite) TestListByUserID_SortByNameAsc() { + user := s.mustCreateUser("sort-name@example.com") + s.mustCreateApiKey(user.ID, "sk-z", "z-key", nil) + s.mustCreateApiKey(user.ID, "sk-a", "a-key", nil) + + keys, _, err := s.repo.ListByUserID(s.ctx, user.ID, pagination.PaginationParams{ + Page: 1, + PageSize: 10, + SortBy: "name", + SortOrder: "asc", + }, service.APIKeyListFilters{}) + s.Require().NoError(err) + s.Require().Len(keys, 2) + s.Require().Equal("a-key", keys[0].Name) + s.Require().Equal("z-key", keys[1].Name) +} diff --git a/backend/internal/repository/channel_repo.go b/backend/internal/repository/channel_repo.go index baad31f7b3..5aad7718cb 100644 --- a/backend/internal/repository/channel_repo.go +++ b/backend/internal/repository/channel_repo.go @@ -188,8 +188,8 @@ func (r *channelRepository) List(ctx context.Context, params pagination.Paginati // 查询 channel 列表 dataQuery := fmt.Sprintf( `SELECT c.id, c.name, c.description, c.status, c.model_mapping, c.billing_model_source, c.restrict_models, c.features, c.created_at, c.updated_at - FROM channels c WHERE %s ORDER BY c.id ASC LIMIT $%d OFFSET $%d`, - whereClause, argIdx, argIdx+1, + FROM channels c WHERE %s ORDER BY %s LIMIT $%d OFFSET $%d`, + whereClause, channelListOrderBy(params), argIdx, argIdx+1, ) args = append(args, pageSize, offset) @@ -246,6 +246,31 @@ func (r *channelRepository) List(ctx context.Context, params pagination.Paginati return channels, paginationResult, nil } +func channelListOrderBy(params pagination.PaginationParams) string { + sortBy := strings.ToLower(strings.TrimSpace(params.SortBy)) + sortOrder := strings.ToUpper(params.NormalizedSortOrder(pagination.SortOrderAsc)) + + var column string + switch sortBy { + case "": + column = "c.id" + sortOrder = "ASC" + case "id": + column = "c.id" + case "name": + column = "c.name" + case "status": + column = "c.status" + case "created_at": + column = "c.created_at" + default: + column = "c.id" + sortOrder = "ASC" + } + + return fmt.Sprintf("%s %s, c.id %s", column, sortOrder, sortOrder) +} + func (r *channelRepository) ListAll(ctx context.Context) ([]service.Channel, error) { rows, err := r.db.QueryContext(ctx, `SELECT id, name, description, status, model_mapping, billing_model_source, restrict_models, features, created_at, updated_at FROM channels ORDER BY id`, diff --git a/backend/internal/repository/channel_repo_test.go b/backend/internal/repository/channel_repo_test.go index 5a59948d25..e761866d70 100644 --- a/backend/internal/repository/channel_repo_test.go +++ b/backend/internal/repository/channel_repo_test.go @@ -8,6 +8,7 @@ import ( "fmt" "testing" + "github.com/Wei-Shaw/sub2api/internal/pkg/pagination" "github.com/lib/pq" "github.com/stretchr/testify/require" ) @@ -225,3 +226,12 @@ func TestIsUniqueViolation(t *testing.T) { }) } } + +func TestChannelListOrderBy_AllowsDescendingIDSort(t *testing.T) { + params := pagination.PaginationParams{ + SortBy: "id", + SortOrder: "desc", + } + + require.Equal(t, "c.id DESC, c.id DESC", channelListOrderBy(params)) +} diff --git a/backend/internal/repository/group_repo.go b/backend/internal/repository/group_repo.go index a075b586c4..c17e3365d3 100644 --- a/backend/internal/repository/group_repo.go +++ b/backend/internal/repository/group_repo.go @@ -5,6 +5,7 @@ import ( "database/sql" "errors" "fmt" + "sort" "strings" dbent "github.com/Wei-Shaw/sub2api/ent" @@ -14,6 +15,8 @@ import ( "github.com/Wei-Shaw/sub2api/internal/pkg/pagination" "github.com/Wei-Shaw/sub2api/internal/service" "github.com/lib/pq" + + entsql "entgo.io/ent/dialect/sql" ) type sqlExecutor interface { @@ -40,6 +43,7 @@ func (r *groupRepository) Create(ctx context.Context, groupIn *service.Group) er SetDescription(groupIn.Description). SetPlatform(groupIn.Platform). SetRateMultiplier(groupIn.RateMultiplier). + SetSortOrder(groupIn.SortOrder). SetIsExclusive(groupIn.IsExclusive). SetStatus(groupIn.Status). SetSubscriptionType(groupIn.SubscriptionType). @@ -58,7 +62,8 @@ func (r *groupRepository) Create(ctx context.Context, groupIn *service.Group) er SetAllowMessagesDispatch(groupIn.AllowMessagesDispatch). SetRequireOauthOnly(groupIn.RequireOAuthOnly). SetRequirePrivacySet(groupIn.RequirePrivacySet). - SetDefaultMappedModel(groupIn.DefaultMappedModel) + SetDefaultMappedModel(groupIn.DefaultMappedModel). + SetMessagesDispatchModelConfig(groupIn.MessagesDispatchModelConfig) // 设置模型路由配置 if groupIn.ModelRouting != nil { @@ -124,7 +129,8 @@ func (r *groupRepository) Update(ctx context.Context, groupIn *service.Group) er SetAllowMessagesDispatch(groupIn.AllowMessagesDispatch). SetRequireOauthOnly(groupIn.RequireOAuthOnly). SetRequirePrivacySet(groupIn.RequirePrivacySet). - SetDefaultMappedModel(groupIn.DefaultMappedModel) + SetDefaultMappedModel(groupIn.DefaultMappedModel). + SetMessagesDispatchModelConfig(groupIn.MessagesDispatchModelConfig) // 显式处理可空字段:nil 需要 clear,非 nil 需要 set。 if groupIn.DailyLimitUSD != nil { @@ -231,11 +237,18 @@ func (r *groupRepository) ListWithFilters(ctx context.Context, params pagination return nil, nil, err } - groups, err := q. + if strings.EqualFold(strings.TrimSpace(params.SortBy), "account_count") { + return r.listWithAccountCountSort(ctx, q, params, total) + } + + groupsQuery := q. Offset(params.Offset()). - Limit(params.Limit()). - Order(dbent.Asc(group.FieldSortOrder), dbent.Asc(group.FieldID)). - All(ctx) + Limit(params.Limit()) + for _, order := range groupListOrder(params) { + groupsQuery = groupsQuery.Order(order) + } + + groups, err := groupsQuery.All(ctx) if err != nil { return nil, nil, err } @@ -261,6 +274,104 @@ func (r *groupRepository) ListWithFilters(ctx context.Context, params pagination return outGroups, paginationResultFromTotal(int64(total), params), nil } +func (r *groupRepository) listWithAccountCountSort(ctx context.Context, q *dbent.GroupQuery, params pagination.PaginationParams, total int) ([]service.Group, *pagination.PaginationResult, error) { + groups, err := q. + Order(dbent.Asc(group.FieldSortOrder), dbent.Asc(group.FieldID)). + All(ctx) + if err != nil { + return nil, nil, err + } + + groupIDs := make([]int64, 0, len(groups)) + outGroups := make([]service.Group, 0, len(groups)) + for i := range groups { + g := groupEntityToService(groups[i]) + outGroups = append(outGroups, *g) + groupIDs = append(groupIDs, g.ID) + } + + counts, err := r.loadAccountCounts(ctx, groupIDs) + if err != nil { + return nil, nil, err + } + for i := range outGroups { + c := counts[outGroups[i].ID] + outGroups[i].AccountCount = c.Total + outGroups[i].ActiveAccountCount = c.Active + outGroups[i].RateLimitedAccountCount = c.RateLimited + } + + sortOrder := params.NormalizedSortOrder(pagination.SortOrderDesc) + sort.SliceStable(outGroups, func(i, j int) bool { + if outGroups[i].AccountCount == outGroups[j].AccountCount { + if outGroups[i].SortOrder == outGroups[j].SortOrder { + return outGroups[i].ID < outGroups[j].ID + } + return outGroups[i].SortOrder < outGroups[j].SortOrder + } + if sortOrder == pagination.SortOrderAsc { + return outGroups[i].AccountCount < outGroups[j].AccountCount + } + return outGroups[i].AccountCount > outGroups[j].AccountCount + }) + + return paginateSlice(outGroups, params), paginationResultFromTotal(int64(total), params), nil +} + +func groupListOrder(params pagination.PaginationParams) []func(*entsql.Selector) { + sortBy := strings.ToLower(strings.TrimSpace(params.SortBy)) + sortOrder := params.NormalizedSortOrder(pagination.SortOrderAsc) + + var field string + tieField := group.FieldID + defaultOrder := true + switch sortBy { + case "", "sort_order": + field = group.FieldSortOrder + case "name": + field = group.FieldName + defaultOrder = false + case "platform": + field = group.FieldPlatform + defaultOrder = false + case "billing_type", "subscription_type": + field = group.FieldSubscriptionType + defaultOrder = false + case "rate_multiplier": + field = group.FieldRateMultiplier + defaultOrder = false + case "is_exclusive": + field = group.FieldIsExclusive + defaultOrder = false + case "status": + field = group.FieldStatus + defaultOrder = false + case "created_at": + field = group.FieldCreatedAt + defaultOrder = false + case "id": + field = group.FieldID + defaultOrder = false + tieField = "" + default: + field = group.FieldSortOrder + } + + if sortOrder == pagination.SortOrderDesc && sortBy != "" { + if tieField == "" { + return []func(*entsql.Selector){dbent.Desc(field)} + } + return []func(*entsql.Selector){dbent.Desc(field), dbent.Desc(tieField)} + } + if defaultOrder { + return []func(*entsql.Selector){dbent.Asc(group.FieldSortOrder), dbent.Asc(group.FieldID)} + } + if tieField == "" { + return []func(*entsql.Selector){dbent.Asc(field)} + } + return []func(*entsql.Selector){dbent.Asc(field), dbent.Asc(tieField)} +} + func (r *groupRepository) ListActive(ctx context.Context) ([]service.Group, error) { groups, err := r.client.Group.Query(). Where(group.StatusEQ(service.StatusActive)). diff --git a/backend/internal/repository/group_repo_integration_test.go b/backend/internal/repository/group_repo_integration_test.go index eccf5ceacb..f91dae43eb 100644 --- a/backend/internal/repository/group_repo_integration_test.go +++ b/backend/internal/repository/group_repo_integration_test.go @@ -113,6 +113,33 @@ func (s *GroupRepoSuite) TestUpdate() { s.Require().Equal("updated", got.Name) } +func (s *GroupRepoSuite) TestGetByID_PreservesMessagesDispatchModelConfig() { + group := &service.Group{ + Name: "openai-dispatch", + Platform: service.PlatformOpenAI, + RateMultiplier: 1.0, + IsExclusive: false, + Status: service.StatusActive, + SubscriptionType: service.SubscriptionTypeStandard, + AllowMessagesDispatch: true, + DefaultMappedModel: "gpt-5.4", + MessagesDispatchModelConfig: service.OpenAIMessagesDispatchModelConfig{ + OpusMappedModel: "gpt-5.4", + SonnetMappedModel: "gpt-5.3-codex", + HaikuMappedModel: "gpt-5.4-mini", + ExactModelMappings: map[string]string{ + "claude-sonnet-4.5": "gpt-5.4-nano", + }, + }, + } + + s.Require().NoError(s.repo.Create(s.ctx, group)) + + got, err := s.repo.GetByID(s.ctx, group.ID) + s.Require().NoError(err) + s.Require().Equal(group.MessagesDispatchModelConfig, got.MessagesDispatchModelConfig) +} + func (s *GroupRepoSuite) TestDelete() { group := &service.Group{ Name: "to-delete", diff --git a/backend/internal/repository/group_repo_sort_integration_test.go b/backend/internal/repository/group_repo_sort_integration_test.go new file mode 100644 index 0000000000..85b2efcc61 --- /dev/null +++ b/backend/internal/repository/group_repo_sort_integration_test.go @@ -0,0 +1,50 @@ +//go:build integration + +package repository + +import ( + "github.com/Wei-Shaw/sub2api/internal/pkg/pagination" + "github.com/Wei-Shaw/sub2api/internal/service" +) + +func (s *GroupRepoSuite) TestList_DefaultSortBySortOrderAsc() { + g1 := &service.Group{Name: "g1", Platform: service.PlatformAnthropic, RateMultiplier: 1, Status: service.StatusActive, SubscriptionType: service.SubscriptionTypeStandard, SortOrder: 20} + g2 := &service.Group{Name: "g2", Platform: service.PlatformAnthropic, RateMultiplier: 1, Status: service.StatusActive, SubscriptionType: service.SubscriptionTypeStandard, SortOrder: 10} + s.Require().NoError(s.repo.Create(s.ctx, g1)) + s.Require().NoError(s.repo.Create(s.ctx, g2)) + + groups, _, err := s.repo.List(s.ctx, pagination.PaginationParams{Page: 1, PageSize: 100}) + s.Require().NoError(err) + s.Require().GreaterOrEqual(len(groups), 2) + indexByID := make(map[int64]int, len(groups)) + for i, g := range groups { + indexByID[g.ID] = i + } + s.Require().Contains(indexByID, g1.ID) + s.Require().Contains(indexByID, g2.ID) + // g2 has SortOrder=10, g1 has SortOrder=20; ascending means g2 comes first + s.Require().Less(indexByID[g2.ID], indexByID[g1.ID]) +} + +func (s *GroupRepoSuite) TestList_SortBySortOrderDesc() { + g1 := &service.Group{Name: "g1", Platform: service.PlatformAnthropic, RateMultiplier: 1, Status: service.StatusActive, SubscriptionType: service.SubscriptionTypeStandard, SortOrder: 40} + g2 := &service.Group{Name: "g2", Platform: service.PlatformAnthropic, RateMultiplier: 1, Status: service.StatusActive, SubscriptionType: service.SubscriptionTypeStandard, SortOrder: 50} + s.Require().NoError(s.repo.Create(s.ctx, g1)) + s.Require().NoError(s.repo.Create(s.ctx, g2)) + + groups, _, err := s.repo.List(s.ctx, pagination.PaginationParams{ + Page: 1, + PageSize: 10, + SortBy: "sort_order", + SortOrder: "desc", + }) + s.Require().NoError(err) + s.Require().GreaterOrEqual(len(groups), 2) + indexByID := make(map[int64]int, len(groups)) + for i, group := range groups { + indexByID[group.ID] = i + } + s.Require().Contains(indexByID, g1.ID) + s.Require().Contains(indexByID, g2.ID) + s.Require().Less(indexByID[g2.ID], indexByID[g1.ID]) +} diff --git a/backend/internal/repository/integration_harness_test.go b/backend/internal/repository/integration_harness_test.go index fb9c26c4aa..5857fbcb2a 100644 --- a/backend/internal/repository/integration_harness_test.go +++ b/backend/internal/repository/integration_harness_test.go @@ -332,6 +332,10 @@ func (h prefixHook) prefixCmd(cmd redisclient.Cmder) { "hgetall", "hget", "hset", "hdel", "hincrbyfloat", "exists", "zadd", "zcard", "zrange", "zrangebyscore", "zrem", "zremrangebyscore", "zrevrange", "zrevrangebyscore", "zscore": prefixOne(1) + case "mget": + for i := 1; i < len(args); i++ { + prefixOne(i) + } case "del", "unlink": for i := 1; i < len(args); i++ { prefixOne(i) diff --git a/backend/internal/repository/pagination.go b/backend/internal/repository/pagination.go index ff08c34be1..87c42a5981 100644 --- a/backend/internal/repository/pagination.go +++ b/backend/internal/repository/pagination.go @@ -14,3 +14,22 @@ func paginationResultFromTotal(total int64, params pagination.PaginationParams) Pages: pages, } } + +func paginateSlice[T any](items []T, params pagination.PaginationParams) []T { + if len(items) == 0 { + return []T{} + } + + offset := params.Offset() + if offset >= len(items) { + return []T{} + } + + limit := params.Limit() + end := offset + limit + if end > len(items) { + end = len(items) + } + + return items[offset:end] +} diff --git a/backend/internal/repository/promo_code_repo.go b/backend/internal/repository/promo_code_repo.go index 95ce687a2a..d9c76bb3d4 100644 --- a/backend/internal/repository/promo_code_repo.go +++ b/backend/internal/repository/promo_code_repo.go @@ -2,12 +2,15 @@ package repository import ( "context" + "strings" dbent "github.com/Wei-Shaw/sub2api/ent" "github.com/Wei-Shaw/sub2api/ent/promocode" "github.com/Wei-Shaw/sub2api/ent/promocodeusage" "github.com/Wei-Shaw/sub2api/internal/pkg/pagination" "github.com/Wei-Shaw/sub2api/internal/service" + + entsql "entgo.io/ent/dialect/sql" ) type promoCodeRepository struct { @@ -137,11 +140,14 @@ func (r *promoCodeRepository) ListWithFilters(ctx context.Context, params pagina return nil, nil, err } - codes, err := q. + codesQuery := q. Offset(params.Offset()). - Limit(params.Limit()). - Order(dbent.Desc(promocode.FieldID)). - All(ctx) + Limit(params.Limit()) + for _, order := range promoCodeListOrder(params) { + codesQuery = codesQuery.Order(order) + } + + codes, err := codesQuery.All(ctx) if err != nil { return nil, nil, err } @@ -151,6 +157,32 @@ func (r *promoCodeRepository) ListWithFilters(ctx context.Context, params pagina return outCodes, paginationResultFromTotal(int64(total), params), nil } +func promoCodeListOrder(params pagination.PaginationParams) []func(*entsql.Selector) { + sortBy := strings.ToLower(strings.TrimSpace(params.SortBy)) + sortOrder := params.NormalizedSortOrder(pagination.SortOrderDesc) + + var field string + switch sortBy { + case "bonus_amount": + field = promocode.FieldBonusAmount + case "status": + field = promocode.FieldStatus + case "expires_at": + field = promocode.FieldExpiresAt + case "created_at": + field = promocode.FieldCreatedAt + case "code": + field = promocode.FieldCode + default: + field = promocode.FieldID + } + + if sortOrder == pagination.SortOrderAsc { + return []func(*entsql.Selector){dbent.Asc(field), dbent.Asc(promocode.FieldID)} + } + return []func(*entsql.Selector){dbent.Desc(field), dbent.Desc(promocode.FieldID)} +} + func (r *promoCodeRepository) CreateUsage(ctx context.Context, usage *service.PromoCodeUsage) error { client := clientFromContext(ctx, r.client) created, err := client.PromoCodeUsage.Create(). diff --git a/backend/internal/repository/proxy_repo.go b/backend/internal/repository/proxy_repo.go index 07c2a20498..60b2f069ec 100644 --- a/backend/internal/repository/proxy_repo.go +++ b/backend/internal/repository/proxy_repo.go @@ -3,12 +3,16 @@ package repository import ( "context" "database/sql" + "sort" + "strings" dbent "github.com/Wei-Shaw/sub2api/ent" "github.com/Wei-Shaw/sub2api/ent/proxy" "github.com/Wei-Shaw/sub2api/internal/service" "github.com/Wei-Shaw/sub2api/internal/pkg/pagination" + + entsql "entgo.io/ent/dialect/sql" ) type sqlQuerier interface { @@ -135,11 +139,14 @@ func (r *proxyRepository) ListWithFilters(ctx context.Context, params pagination return nil, nil, err } - proxies, err := q. + proxiesQuery := q. Offset(params.Offset()). - Limit(params.Limit()). - Order(dbent.Desc(proxy.FieldID)). - All(ctx) + Limit(params.Limit()) + for _, order := range proxyListOrder(params) { + proxiesQuery = proxiesQuery.Order(order) + } + + proxies, err := proxiesQuery.All(ctx) if err != nil { return nil, nil, err } @@ -170,22 +177,58 @@ func (r *proxyRepository) ListWithFiltersAndAccountCount(ctx context.Context, pa return nil, nil, err } - proxies, err := q. + if strings.EqualFold(strings.TrimSpace(params.SortBy), "account_count") { + return r.listWithAccountCountSort(ctx, q, params, total) + } + + proxiesQuery := q. Offset(params.Offset()). - Limit(params.Limit()). + Limit(params.Limit()) + for _, order := range proxyListOrder(params) { + proxiesQuery = proxiesQuery.Order(order) + } + + proxies, err := proxiesQuery.All(ctx) + if err != nil { + return nil, nil, err + } + + return r.buildProxyWithAccountCountResult(ctx, proxies, params, int64(total)) +} + +func (r *proxyRepository) listWithAccountCountSort(ctx context.Context, q *dbent.ProxyQuery, params pagination.PaginationParams, total int) ([]service.ProxyWithAccountCount, *pagination.PaginationResult, error) { + proxies, err := q. Order(dbent.Desc(proxy.FieldID)). All(ctx) if err != nil { return nil, nil, err } - // Get account counts + result, _, err := r.buildProxyWithAccountCountResult(ctx, proxies, params, int64(total)) + if err != nil { + return nil, nil, err + } + + sortOrder := params.NormalizedSortOrder(pagination.SortOrderDesc) + sort.SliceStable(result, func(i, j int) bool { + if result[i].AccountCount == result[j].AccountCount { + return result[i].ID > result[j].ID + } + if sortOrder == pagination.SortOrderAsc { + return result[i].AccountCount < result[j].AccountCount + } + return result[i].AccountCount > result[j].AccountCount + }) + + return paginateSlice(result, params), paginationResultFromTotal(int64(total), params), nil +} + +func (r *proxyRepository) buildProxyWithAccountCountResult(ctx context.Context, proxies []*dbent.Proxy, params pagination.PaginationParams, total int64) ([]service.ProxyWithAccountCount, *pagination.PaginationResult, error) { counts, err := r.GetAccountCountsForProxies(ctx) if err != nil { return nil, nil, err } - // Build result with account counts result := make([]service.ProxyWithAccountCount, 0, len(proxies)) for i := range proxies { proxyOut := proxyEntityToService(proxies[i]) @@ -198,7 +241,31 @@ func (r *proxyRepository) ListWithFiltersAndAccountCount(ctx context.Context, pa }) } - return result, paginationResultFromTotal(int64(total), params), nil + return result, paginationResultFromTotal(total, params), nil +} + +func proxyListOrder(params pagination.PaginationParams) []func(*entsql.Selector) { + sortBy := strings.ToLower(strings.TrimSpace(params.SortBy)) + sortOrder := params.NormalizedSortOrder(pagination.SortOrderDesc) + + var field string + switch sortBy { + case "name": + field = proxy.FieldName + case "protocol": + field = proxy.FieldProtocol + case "status": + field = proxy.FieldStatus + case "created_at": + field = proxy.FieldCreatedAt + default: + field = proxy.FieldID + } + + if sortOrder == pagination.SortOrderAsc { + return []func(*entsql.Selector){dbent.Asc(field), dbent.Asc(proxy.FieldID)} + } + return []func(*entsql.Selector){dbent.Desc(field), dbent.Desc(proxy.FieldID)} } func (r *proxyRepository) ListActive(ctx context.Context) ([]service.Proxy, error) { diff --git a/backend/internal/repository/proxy_repo_sort_integration_test.go b/backend/internal/repository/proxy_repo_sort_integration_test.go new file mode 100644 index 0000000000..fe1c2873ba --- /dev/null +++ b/backend/internal/repository/proxy_repo_sort_integration_test.go @@ -0,0 +1,28 @@ +//go:build integration + +package repository + +import ( + "github.com/Wei-Shaw/sub2api/internal/pkg/pagination" + "github.com/Wei-Shaw/sub2api/internal/service" +) + +func (s *ProxyRepoSuite) TestListWithFiltersAndAccountCount_SortByAccountCountDesc() { + p1 := s.mustCreateProxy(&service.Proxy{Name: "p1", Protocol: "http", Host: "127.0.0.1", Port: 8080, Status: service.StatusActive}) + p2 := s.mustCreateProxy(&service.Proxy{Name: "p2", Protocol: "http", Host: "127.0.0.1", Port: 8081, Status: service.StatusActive}) + s.mustInsertAccount("a1", &p1.ID) + s.mustInsertAccount("a2", &p1.ID) + s.mustInsertAccount("a3", &p2.ID) + + proxies, _, err := s.repo.ListWithFiltersAndAccountCount(s.ctx, pagination.PaginationParams{ + Page: 1, + PageSize: 10, + SortBy: "account_count", + SortOrder: "desc", + }, "", "", "") + s.Require().NoError(err) + s.Require().Len(proxies, 2) + s.Require().Equal(p1.ID, proxies[0].ID) + s.Require().Equal(int64(2), proxies[0].AccountCount) + s.Require().Equal(p2.ID, proxies[1].ID) +} diff --git a/backend/internal/repository/redeem_code_repo.go b/backend/internal/repository/redeem_code_repo.go index 934a309568..07975970ef 100644 --- a/backend/internal/repository/redeem_code_repo.go +++ b/backend/internal/repository/redeem_code_repo.go @@ -2,6 +2,7 @@ package repository import ( "context" + "strings" "time" dbent "github.com/Wei-Shaw/sub2api/ent" @@ -9,6 +10,8 @@ import ( "github.com/Wei-Shaw/sub2api/ent/user" "github.com/Wei-Shaw/sub2api/internal/pkg/pagination" "github.com/Wei-Shaw/sub2api/internal/service" + + entsql "entgo.io/ent/dialect/sql" ) type redeemCodeRepository struct { @@ -120,13 +123,16 @@ func (r *redeemCodeRepository) ListWithFilters(ctx context.Context, params pagin return nil, nil, err } - codes, err := q. + codesQuery := q. WithUser(). WithGroup(). Offset(params.Offset()). - Limit(params.Limit()). - Order(dbent.Desc(redeemcode.FieldID)). - All(ctx) + Limit(params.Limit()) + for _, order := range redeemCodeListOrder(params) { + codesQuery = codesQuery.Order(order) + } + + codes, err := codesQuery.All(ctx) if err != nil { return nil, nil, err } @@ -136,6 +142,34 @@ func (r *redeemCodeRepository) ListWithFilters(ctx context.Context, params pagin return outCodes, paginationResultFromTotal(int64(total), params), nil } +func redeemCodeListOrder(params pagination.PaginationParams) []func(*entsql.Selector) { + sortBy := strings.ToLower(strings.TrimSpace(params.SortBy)) + sortOrder := params.NormalizedSortOrder(pagination.SortOrderDesc) + + var field string + switch sortBy { + case "type": + field = redeemcode.FieldType + case "value": + field = redeemcode.FieldValue + case "status": + field = redeemcode.FieldStatus + case "used_at": + field = redeemcode.FieldUsedAt + case "created_at": + field = redeemcode.FieldCreatedAt + case "code": + field = redeemcode.FieldCode + default: + field = redeemcode.FieldID + } + + if sortOrder == pagination.SortOrderAsc { + return []func(*entsql.Selector){dbent.Asc(field), dbent.Asc(redeemcode.FieldID)} + } + return []func(*entsql.Selector){dbent.Desc(field), dbent.Desc(redeemcode.FieldID)} +} + func (r *redeemCodeRepository) Update(ctx context.Context, code *service.RedeemCode) error { up := r.client.RedeemCode.UpdateOneID(code.ID). SetCode(code.Code). diff --git a/backend/internal/repository/redeem_code_repo_sort_integration_test.go b/backend/internal/repository/redeem_code_repo_sort_integration_test.go new file mode 100644 index 0000000000..30d32f4cf9 --- /dev/null +++ b/backend/internal/repository/redeem_code_repo_sort_integration_test.go @@ -0,0 +1,24 @@ +//go:build integration + +package repository + +import ( + "github.com/Wei-Shaw/sub2api/internal/pkg/pagination" + "github.com/Wei-Shaw/sub2api/internal/service" +) + +func (s *RedeemCodeRepoSuite) TestListWithFilters_SortByValueAsc() { + s.Require().NoError(s.repo.Create(s.ctx, &service.RedeemCode{Code: "VALUE-20", Type: service.RedeemTypeBalance, Value: 20, Status: service.StatusUnused})) + s.Require().NoError(s.repo.Create(s.ctx, &service.RedeemCode{Code: "VALUE-10", Type: service.RedeemTypeBalance, Value: 10, Status: service.StatusUnused})) + + codes, _, err := s.repo.ListWithFilters(s.ctx, pagination.PaginationParams{ + Page: 1, + PageSize: 10, + SortBy: "value", + SortOrder: "asc", + }, "", "", "") + s.Require().NoError(err) + s.Require().Len(codes, 2) + s.Require().Equal("VALUE-10", codes[0].Code) + s.Require().Equal("VALUE-20", codes[1].Code) +} diff --git a/backend/internal/repository/scheduler_cache.go b/backend/internal/repository/scheduler_cache.go index 4f447e4fea..e9be8c7a6e 100644 --- a/backend/internal/repository/scheduler_cache.go +++ b/backend/internal/repository/scheduler_cache.go @@ -15,19 +15,39 @@ const ( schedulerBucketSetKey = "sched:buckets" schedulerOutboxWatermarkKey = "sched:outbox:watermark" schedulerAccountPrefix = "sched:acc:" + schedulerAccountMetaPrefix = "sched:meta:" schedulerActivePrefix = "sched:active:" schedulerReadyPrefix = "sched:ready:" schedulerVersionPrefix = "sched:ver:" schedulerSnapshotPrefix = "sched:" schedulerLockPrefix = "sched:lock:" + + defaultSchedulerSnapshotMGetChunkSize = 128 + defaultSchedulerSnapshotWriteChunkSize = 256 ) type schedulerCache struct { - rdb *redis.Client + rdb *redis.Client + mgetChunkSize int + writeChunkSize int } func NewSchedulerCache(rdb *redis.Client) service.SchedulerCache { - return &schedulerCache{rdb: rdb} + return newSchedulerCacheWithChunkSizes(rdb, defaultSchedulerSnapshotMGetChunkSize, defaultSchedulerSnapshotWriteChunkSize) +} + +func newSchedulerCacheWithChunkSizes(rdb *redis.Client, mgetChunkSize, writeChunkSize int) service.SchedulerCache { + if mgetChunkSize <= 0 { + mgetChunkSize = defaultSchedulerSnapshotMGetChunkSize + } + if writeChunkSize <= 0 { + writeChunkSize = defaultSchedulerSnapshotWriteChunkSize + } + return &schedulerCache{ + rdb: rdb, + mgetChunkSize: mgetChunkSize, + writeChunkSize: writeChunkSize, + } } func (c *schedulerCache) GetSnapshot(ctx context.Context, bucket service.SchedulerBucket) ([]*service.Account, bool, error) { @@ -65,9 +85,9 @@ func (c *schedulerCache) GetSnapshot(ctx context.Context, bucket service.Schedul keys := make([]string, 0, len(ids)) for _, id := range ids { - keys = append(keys, schedulerAccountKey(id)) + keys = append(keys, schedulerAccountMetaKey(id)) } - values, err := c.rdb.MGet(ctx, keys...).Result() + values, err := c.mgetChunked(ctx, keys) if err != nil { return nil, false, err } @@ -100,14 +120,11 @@ func (c *schedulerCache) SetSnapshot(ctx context.Context, bucket service.Schedul versionStr := strconv.FormatInt(version, 10) snapshotKey := schedulerSnapshotKey(bucket, versionStr) - pipe := c.rdb.Pipeline() - for _, account := range accounts { - payload, err := json.Marshal(account) - if err != nil { - return err - } - pipe.Set(ctx, schedulerAccountKey(strconv.FormatInt(account.ID, 10)), payload, 0) + if err := c.writeAccounts(ctx, accounts); err != nil { + return err } + + pipe := c.rdb.Pipeline() if len(accounts) > 0 { // 使用序号作为 score,保持数据库返回的排序语义。 members := make([]redis.Z, 0, len(accounts)) @@ -117,7 +134,13 @@ func (c *schedulerCache) SetSnapshot(ctx context.Context, bucket service.Schedul Member: strconv.FormatInt(account.ID, 10), }) } - pipe.ZAdd(ctx, snapshotKey, members...) + for start := 0; start < len(members); start += c.writeChunkSize { + end := start + c.writeChunkSize + if end > len(members) { + end = len(members) + } + pipe.ZAdd(ctx, snapshotKey, members[start:end]...) + } } else { pipe.Del(ctx, snapshotKey) } @@ -151,20 +174,15 @@ func (c *schedulerCache) SetAccount(ctx context.Context, account *service.Accoun if account == nil || account.ID <= 0 { return nil } - payload, err := json.Marshal(account) - if err != nil { - return err - } - key := schedulerAccountKey(strconv.FormatInt(account.ID, 10)) - return c.rdb.Set(ctx, key, payload, 0).Err() + return c.writeAccounts(ctx, []service.Account{*account}) } func (c *schedulerCache) DeleteAccount(ctx context.Context, accountID int64) error { if accountID <= 0 { return nil } - key := schedulerAccountKey(strconv.FormatInt(accountID, 10)) - return c.rdb.Del(ctx, key).Err() + id := strconv.FormatInt(accountID, 10) + return c.rdb.Del(ctx, schedulerAccountKey(id), schedulerAccountMetaKey(id)).Err() } func (c *schedulerCache) UpdateLastUsed(ctx context.Context, updates map[int64]time.Time) error { @@ -179,7 +197,7 @@ func (c *schedulerCache) UpdateLastUsed(ctx context.Context, updates map[int64]t ids = append(ids, id) } - values, err := c.rdb.MGet(ctx, keys...).Result() + values, err := c.mgetChunked(ctx, keys) if err != nil { return err } @@ -198,7 +216,12 @@ func (c *schedulerCache) UpdateLastUsed(ctx context.Context, updates map[int64]t if err != nil { return err } + metaPayload, err := json.Marshal(buildSchedulerMetadataAccount(*account)) + if err != nil { + return err + } pipe.Set(ctx, keys[i], updated, 0) + pipe.Set(ctx, schedulerAccountMetaKey(strconv.FormatInt(ids[i], 10)), metaPayload, 0) } _, err = pipe.Exec(ctx) return err @@ -256,6 +279,10 @@ func schedulerAccountKey(id string) string { return schedulerAccountPrefix + id } +func schedulerAccountMetaKey(id string) string { + return schedulerAccountMetaPrefix + id +} + func ptrTime(t time.Time) *time.Time { return &t } @@ -276,3 +303,138 @@ func decodeCachedAccount(val any) (*service.Account, error) { } return &account, nil } + +func (c *schedulerCache) writeAccounts(ctx context.Context, accounts []service.Account) error { + if len(accounts) == 0 { + return nil + } + + pipe := c.rdb.Pipeline() + pending := 0 + flush := func() error { + if pending == 0 { + return nil + } + if _, err := pipe.Exec(ctx); err != nil { + return err + } + pipe = c.rdb.Pipeline() + pending = 0 + return nil + } + + for _, account := range accounts { + fullPayload, err := json.Marshal(account) + if err != nil { + return err + } + metaPayload, err := json.Marshal(buildSchedulerMetadataAccount(account)) + if err != nil { + return err + } + + id := strconv.FormatInt(account.ID, 10) + pipe.Set(ctx, schedulerAccountKey(id), fullPayload, 0) + pipe.Set(ctx, schedulerAccountMetaKey(id), metaPayload, 0) + pending++ + if pending >= c.writeChunkSize { + if err := flush(); err != nil { + return err + } + } + } + + return flush() +} + +func (c *schedulerCache) mgetChunked(ctx context.Context, keys []string) ([]any, error) { + if len(keys) == 0 { + return []any{}, nil + } + + out := make([]any, 0, len(keys)) + chunkSize := c.mgetChunkSize + if chunkSize <= 0 { + chunkSize = defaultSchedulerSnapshotMGetChunkSize + } + for start := 0; start < len(keys); start += chunkSize { + end := start + chunkSize + if end > len(keys) { + end = len(keys) + } + part, err := c.rdb.MGet(ctx, keys[start:end]...).Result() + if err != nil { + return nil, err + } + out = append(out, part...) + } + return out, nil +} + +func buildSchedulerMetadataAccount(account service.Account) service.Account { + return service.Account{ + ID: account.ID, + Name: account.Name, + Platform: account.Platform, + Type: account.Type, + Concurrency: account.Concurrency, + LoadFactor: account.LoadFactor, + Priority: account.Priority, + RateMultiplier: account.RateMultiplier, + Status: account.Status, + LastUsedAt: account.LastUsedAt, + ExpiresAt: account.ExpiresAt, + AutoPauseOnExpired: account.AutoPauseOnExpired, + Schedulable: account.Schedulable, + RateLimitedAt: account.RateLimitedAt, + RateLimitResetAt: account.RateLimitResetAt, + OverloadUntil: account.OverloadUntil, + TempUnschedulableUntil: account.TempUnschedulableUntil, + TempUnschedulableReason: account.TempUnschedulableReason, + SessionWindowStart: account.SessionWindowStart, + SessionWindowEnd: account.SessionWindowEnd, + SessionWindowStatus: account.SessionWindowStatus, + Credentials: filterSchedulerCredentials(account.Credentials), + Extra: filterSchedulerExtra(account.Extra), + } +} + +func filterSchedulerCredentials(credentials map[string]any) map[string]any { + if len(credentials) == 0 { + return nil + } + keys := []string{"model_mapping", "api_key", "project_id", "oauth_type"} + filtered := make(map[string]any) + for _, key := range keys { + if value, ok := credentials[key]; ok && value != nil { + filtered[key] = value + } + } + if len(filtered) == 0 { + return nil + } + return filtered +} + +func filterSchedulerExtra(extra map[string]any) map[string]any { + if len(extra) == 0 { + return nil + } + keys := []string{ + "mixed_scheduling", + "window_cost_limit", + "window_cost_sticky_reserve", + "max_sessions", + "session_idle_timeout_minutes", + } + filtered := make(map[string]any) + for _, key := range keys { + if value, ok := extra[key]; ok && value != nil { + filtered[key] = value + } + } + if len(filtered) == 0 { + return nil + } + return filtered +} diff --git a/backend/internal/repository/scheduler_cache_integration_test.go b/backend/internal/repository/scheduler_cache_integration_test.go new file mode 100644 index 0000000000..134a6a0753 --- /dev/null +++ b/backend/internal/repository/scheduler_cache_integration_test.go @@ -0,0 +1,88 @@ +//go:build integration + +package repository + +import ( + "context" + "strings" + "testing" + "time" + + "github.com/Wei-Shaw/sub2api/internal/service" + "github.com/stretchr/testify/require" +) + +func TestSchedulerCacheSnapshotUsesSlimMetadataButKeepsFullAccount(t *testing.T) { + ctx := context.Background() + rdb := testRedis(t) + cache := NewSchedulerCache(rdb) + + bucket := service.SchedulerBucket{GroupID: 2, Platform: service.PlatformGemini, Mode: service.SchedulerModeSingle} + now := time.Now().UTC().Truncate(time.Second) + limitReset := now.Add(10 * time.Minute) + overloadUntil := now.Add(2 * time.Minute) + tempUnschedUntil := now.Add(3 * time.Minute) + windowEnd := now.Add(5 * time.Hour) + + account := service.Account{ + ID: 101, + Name: "gemini-heavy", + Platform: service.PlatformGemini, + Type: service.AccountTypeOAuth, + Status: service.StatusActive, + Schedulable: true, + Concurrency: 3, + Priority: 7, + LastUsedAt: &now, + Credentials: map[string]any{ + "api_key": "gemini-api-key", + "access_token": "secret-access-token", + "project_id": "proj-1", + "oauth_type": "ai_studio", + "model_mapping": map[string]any{"gemini-2.5-pro": "gemini-2.5-pro"}, + "huge_blob": strings.Repeat("x", 4096), + }, + Extra: map[string]any{ + "mixed_scheduling": true, + "window_cost_limit": 12.5, + "window_cost_sticky_reserve": 8.0, + "max_sessions": 4, + "session_idle_timeout_minutes": 11, + "unused_large_field": strings.Repeat("y", 4096), + }, + RateLimitResetAt: &limitReset, + OverloadUntil: &overloadUntil, + TempUnschedulableUntil: &tempUnschedUntil, + SessionWindowStart: &now, + SessionWindowEnd: &windowEnd, + SessionWindowStatus: "active", + } + + require.NoError(t, cache.SetSnapshot(ctx, bucket, []service.Account{account})) + + snapshot, hit, err := cache.GetSnapshot(ctx, bucket) + require.NoError(t, err) + require.True(t, hit) + require.Len(t, snapshot, 1) + + got := snapshot[0] + require.NotNil(t, got) + require.Equal(t, "gemini-api-key", got.GetCredential("api_key")) + require.Equal(t, "proj-1", got.GetCredential("project_id")) + require.Equal(t, "ai_studio", got.GetCredential("oauth_type")) + require.NotEmpty(t, got.GetModelMapping()) + require.Empty(t, got.GetCredential("access_token")) + require.Empty(t, got.GetCredential("huge_blob")) + require.Equal(t, true, got.Extra["mixed_scheduling"]) + require.Equal(t, 12.5, got.GetWindowCostLimit()) + require.Equal(t, 8.0, got.GetWindowCostStickyReserve()) + require.Equal(t, 4, got.GetMaxSessions()) + require.Equal(t, 11, got.GetSessionIdleTimeoutMinutes()) + require.Nil(t, got.Extra["unused_large_field"]) + + full, err := cache.GetAccount(ctx, account.ID) + require.NoError(t, err) + require.NotNil(t, full) + require.Equal(t, "secret-access-token", full.GetCredential("access_token")) + require.Equal(t, strings.Repeat("x", 4096), full.GetCredential("huge_blob")) +} diff --git a/backend/internal/repository/usage_log_repo.go b/backend/internal/repository/usage_log_repo.go index e6c77d6a92..7e671a784b 100644 --- a/backend/internal/repository/usage_log_repo.go +++ b/backend/internal/repository/usage_log_repo.go @@ -3771,7 +3771,7 @@ func (r *usageLogRepository) listUsageLogsWithPagination(ctx context.Context, wh limitPos := len(args) + 1 offsetPos := len(args) + 2 listArgs := append(append([]any{}, args...), params.Limit(), params.Offset()) - query := fmt.Sprintf("SELECT %s FROM usage_logs %s ORDER BY id DESC LIMIT $%d OFFSET $%d", usageLogSelectColumns, whereClause, limitPos, offsetPos) + query := fmt.Sprintf("SELECT %s FROM usage_logs %s ORDER BY %s LIMIT $%d OFFSET $%d", usageLogSelectColumns, whereClause, usageLogOrderBy(params), limitPos, offsetPos) logs, err := r.queryUsageLogs(ctx, query, listArgs...) if err != nil { return nil, nil, err @@ -3786,7 +3786,7 @@ func (r *usageLogRepository) listUsageLogsWithFastPagination(ctx context.Context limitPos := len(args) + 1 offsetPos := len(args) + 2 listArgs := append(append([]any{}, args...), limit+1, offset) - query := fmt.Sprintf("SELECT %s FROM usage_logs %s ORDER BY id DESC LIMIT $%d OFFSET $%d", usageLogSelectColumns, whereClause, limitPos, offsetPos) + query := fmt.Sprintf("SELECT %s FROM usage_logs %s ORDER BY %s LIMIT $%d OFFSET $%d", usageLogSelectColumns, whereClause, usageLogOrderBy(params), limitPos, offsetPos) logs, err := r.queryUsageLogs(ctx, query, listArgs...) if err != nil { @@ -3808,6 +3808,26 @@ func (r *usageLogRepository) listUsageLogsWithFastPagination(ctx context.Context return logs, paginationResultFromTotal(total, params), nil } +func usageLogOrderBy(params pagination.PaginationParams) string { + sortBy := strings.ToLower(strings.TrimSpace(params.SortBy)) + sortOrder := strings.ToUpper(params.NormalizedSortOrder(pagination.SortOrderDesc)) + + var column string + switch sortBy { + case "model": + column = "COALESCE(NULLIF(TRIM(requested_model), ''), model)" + case "created_at": + column = "created_at" + default: + column = "id" + } + + if column == "id" { + return fmt.Sprintf("id %s", sortOrder) + } + return fmt.Sprintf("%s %s, id %s", column, sortOrder, sortOrder) +} + func (r *usageLogRepository) queryUsageLogs(ctx context.Context, query string, args ...any) (logs []service.UsageLog, err error) { rows, err := r.sql.QueryContext(ctx, query, args...) if err != nil { diff --git a/backend/internal/repository/usage_log_repo_request_type_test.go b/backend/internal/repository/usage_log_repo_request_type_test.go index ce0c5f0040..b9cb6a1330 100644 --- a/backend/internal/repository/usage_log_repo_request_type_test.go +++ b/backend/internal/repository/usage_log_repo_request_type_test.go @@ -330,6 +330,15 @@ func TestUsageLogRepositoryGetStatsWithFiltersRequestTypePriority(t *testing.T) "total_account_cost", "avg_duration_ms", }).AddRow(int64(1), int64(2), int64(3), int64(4), 1.2, 1.0, 1.2, 20.0)) + mock.ExpectQuery("SELECT COALESCE\\(NULLIF\\(TRIM\\(inbound_endpoint\\), ''\\), 'unknown'\\) AS endpoint"). + WithArgs(sqlmock.AnyArg(), sqlmock.AnyArg(), requestType). + WillReturnRows(sqlmock.NewRows([]string{"endpoint", "requests", "total_tokens", "cost", "actual_cost"})) + mock.ExpectQuery("SELECT COALESCE\\(NULLIF\\(TRIM\\(upstream_endpoint\\), ''\\), 'unknown'\\) AS endpoint"). + WithArgs(sqlmock.AnyArg(), sqlmock.AnyArg(), requestType). + WillReturnRows(sqlmock.NewRows([]string{"endpoint", "requests", "total_tokens", "cost", "actual_cost"})) + mock.ExpectQuery("SELECT CONCAT\\("). + WithArgs(sqlmock.AnyArg(), sqlmock.AnyArg(), requestType). + WillReturnRows(sqlmock.NewRows([]string{"endpoint", "requests", "total_tokens", "cost", "actual_cost"})) stats, err := repo.GetStatsWithFilters(context.Background(), filters) require.NoError(t, err) diff --git a/backend/internal/repository/usage_log_repo_sort_integration_test.go b/backend/internal/repository/usage_log_repo_sort_integration_test.go new file mode 100644 index 0000000000..4c69f97538 --- /dev/null +++ b/backend/internal/repository/usage_log_repo_sort_integration_test.go @@ -0,0 +1,61 @@ +//go:build integration + +package repository + +import ( + "time" + + "github.com/Wei-Shaw/sub2api/internal/pkg/pagination" + "github.com/Wei-Shaw/sub2api/internal/pkg/usagestats" + "github.com/Wei-Shaw/sub2api/internal/service" + "github.com/google/uuid" +) + +func (s *UsageLogRepoSuite) TestListWithFilters_SortByModelAsc() { + user := mustCreateUser(s.T(), s.client, &service.User{Email: "usage-sort@example.com"}) + apiKey := mustCreateApiKey(s.T(), s.client, &service.APIKey{UserID: user.ID, Key: "sk-usage-sort", Name: "k"}) + account := mustCreateAccount(s.T(), s.client, &service.Account{Name: "usage-sort-account"}) + + first := &service.UsageLog{ + UserID: user.ID, + APIKeyID: apiKey.ID, + AccountID: account.ID, + RequestID: uuid.New().String(), + Model: "z-model", + RequestedModel: "z-model", + InputTokens: 10, + OutputTokens: 20, + TotalCost: 0.5, + ActualCost: 0.5, + CreatedAt: time.Now(), + } + _, err := s.repo.Create(s.ctx, first) + s.Require().NoError(err) + + second := &service.UsageLog{ + UserID: user.ID, + APIKeyID: apiKey.ID, + AccountID: account.ID, + RequestID: uuid.New().String(), + Model: "a-model", + RequestedModel: "a-model", + InputTokens: 10, + OutputTokens: 20, + TotalCost: 0.5, + ActualCost: 0.5, + CreatedAt: time.Now().Add(time.Second), + } + _, err = s.repo.Create(s.ctx, second) + s.Require().NoError(err) + + logs, _, err := s.repo.ListWithFilters(s.ctx, pagination.PaginationParams{ + Page: 1, + PageSize: 10, + SortBy: "model", + SortOrder: "asc", + }, usagestats.UsageLogFilters{UserID: user.ID}) + s.Require().NoError(err) + s.Require().Len(logs, 2) + s.Require().Equal("a-model", logs[0].RequestedModel) + s.Require().Equal("z-model", logs[1].RequestedModel) +} diff --git a/backend/internal/repository/user_repo.go b/backend/internal/repository/user_repo.go index 06c79113e2..d5a13607ae 100644 --- a/backend/internal/repository/user_repo.go +++ b/backend/internal/repository/user_repo.go @@ -17,6 +17,8 @@ import ( "github.com/Wei-Shaw/sub2api/ent/usersubscription" "github.com/Wei-Shaw/sub2api/internal/pkg/pagination" "github.com/Wei-Shaw/sub2api/internal/service" + + entsql "entgo.io/ent/dialect/sql" ) type userRepository struct { @@ -224,11 +226,14 @@ func (r *userRepository) ListWithFilters(ctx context.Context, params pagination. return nil, nil, err } - users, err := q. + usersQuery := q. Offset(params.Offset()). - Limit(params.Limit()). - Order(dbent.Desc(dbuser.FieldID)). - All(ctx) + Limit(params.Limit()) + for _, order := range userListOrder(params) { + usersQuery = usersQuery.Order(order) + } + + users, err := usersQuery.All(ctx) if err != nil { return nil, nil, err } @@ -281,6 +286,50 @@ func (r *userRepository) ListWithFilters(ctx context.Context, params pagination. return outUsers, paginationResultFromTotal(int64(total), params), nil } +func userListOrder(params pagination.PaginationParams) []func(*entsql.Selector) { + sortBy := strings.ToLower(strings.TrimSpace(params.SortBy)) + sortOrder := params.NormalizedSortOrder(pagination.SortOrderDesc) + + var field string + defaultField := true + switch sortBy { + case "email": + field = dbuser.FieldEmail + defaultField = false + case "username": + field = dbuser.FieldUsername + defaultField = false + case "role": + field = dbuser.FieldRole + defaultField = false + case "balance": + field = dbuser.FieldBalance + defaultField = false + case "concurrency": + field = dbuser.FieldConcurrency + defaultField = false + case "status": + field = dbuser.FieldStatus + defaultField = false + case "created_at": + field = dbuser.FieldCreatedAt + defaultField = false + default: + field = dbuser.FieldID + } + + if sortOrder == pagination.SortOrderAsc { + if defaultField && field == dbuser.FieldID { + return []func(*entsql.Selector){dbent.Asc(dbuser.FieldID)} + } + return []func(*entsql.Selector){dbent.Asc(field), dbent.Asc(dbuser.FieldID)} + } + if defaultField && field == dbuser.FieldID { + return []func(*entsql.Selector){dbent.Desc(dbuser.FieldID)} + } + return []func(*entsql.Selector){dbent.Desc(field), dbent.Desc(dbuser.FieldID)} +} + // filterUsersByAttributes returns user IDs that match ALL the given attribute filters func (r *userRepository) filterUsersByAttributes(ctx context.Context, attrs map[int64]string) ([]int64, error) { if len(attrs) == 0 { diff --git a/backend/internal/repository/user_repo_sort_integration_test.go b/backend/internal/repository/user_repo_sort_integration_test.go new file mode 100644 index 0000000000..ab84b0e93b --- /dev/null +++ b/backend/internal/repository/user_repo_sort_integration_test.go @@ -0,0 +1,39 @@ +//go:build integration + +package repository + +import ( + "testing" + + "github.com/Wei-Shaw/sub2api/internal/pkg/pagination" + "github.com/Wei-Shaw/sub2api/internal/service" +) + +func (s *UserRepoSuite) TestListWithFilters_SortByEmailAsc() { + s.mustCreateUser(&service.User{Email: "z-last@example.com", Username: "z-user"}) + s.mustCreateUser(&service.User{Email: "a-first@example.com", Username: "a-user"}) + + users, _, err := s.repo.ListWithFilters(s.ctx, pagination.PaginationParams{ + Page: 1, + PageSize: 10, + SortBy: "email", + SortOrder: "asc", + }, service.UserListFilters{}) + s.Require().NoError(err) + s.Require().Len(users, 2) + s.Require().Equal("a-first@example.com", users[0].Email) + s.Require().Equal("z-last@example.com", users[1].Email) +} + +func (s *UserRepoSuite) TestList_DefaultSortByNewestFirst() { + first := s.mustCreateUser(&service.User{Email: "first@example.com"}) + second := s.mustCreateUser(&service.User{Email: "second@example.com"}) + + users, _, err := s.repo.List(s.ctx, pagination.PaginationParams{Page: 1, PageSize: 10}) + s.Require().NoError(err) + s.Require().Len(users, 2) + s.Require().Equal(second.ID, users[0].ID) + s.Require().Equal(first.ID, users[1].ID) +} + +func TestUserRepoSortSuiteSmoke(_ *testing.T) {} diff --git a/backend/internal/repository/wire.go b/backend/internal/repository/wire.go index 657e3ed66c..d3adb4a0aa 100644 --- a/backend/internal/repository/wire.go +++ b/backend/internal/repository/wire.go @@ -47,6 +47,21 @@ func ProvideSessionLimitCache(rdb *redis.Client, cfg *config.Config) service.Ses return NewSessionLimitCache(rdb, defaultIdleTimeoutMinutes) } +// ProvideSchedulerCache 创建调度快照缓存,并注入快照分块参数。 +func ProvideSchedulerCache(rdb *redis.Client, cfg *config.Config) service.SchedulerCache { + mgetChunkSize := defaultSchedulerSnapshotMGetChunkSize + writeChunkSize := defaultSchedulerSnapshotWriteChunkSize + if cfg != nil { + if cfg.Gateway.Scheduling.SnapshotMGetChunkSize > 0 { + mgetChunkSize = cfg.Gateway.Scheduling.SnapshotMGetChunkSize + } + if cfg.Gateway.Scheduling.SnapshotWriteChunkSize > 0 { + writeChunkSize = cfg.Gateway.Scheduling.SnapshotWriteChunkSize + } + } + return newSchedulerCacheWithChunkSizes(rdb, mgetChunkSize, writeChunkSize) +} + // ProviderSet is the Wire provider set for all repositories var ProviderSet = wire.NewSet( NewUserRepository, @@ -92,7 +107,7 @@ var ProviderSet = wire.NewSet( NewRedeemCache, NewUpdateCache, NewGeminiTokenCache, - NewSchedulerCache, + ProvideSchedulerCache, NewSchedulerOutboxRepository, NewProxyLatencyCache, NewTotpCache, diff --git a/backend/internal/server/api_contract_test.go b/backend/internal/server/api_contract_test.go index 1e355ca540..af4a140992 100644 --- a/backend/internal/server/api_contract_test.go +++ b/backend/internal/server/api_contract_test.go @@ -461,6 +461,28 @@ func TestAPIContracts(t *testing.T) { service.SettingKeyTurnstileSiteKey: "site-key", service.SettingKeyTurnstileSecretKey: "secret-key", + service.SettingKeyOIDCConnectEnabled: "false", + service.SettingKeyOIDCConnectProviderName: "OIDC", + service.SettingKeyOIDCConnectClientID: "", + service.SettingKeyOIDCConnectIssuerURL: "", + service.SettingKeyOIDCConnectDiscoveryURL: "", + service.SettingKeyOIDCConnectAuthorizeURL: "", + service.SettingKeyOIDCConnectTokenURL: "", + service.SettingKeyOIDCConnectUserInfoURL: "", + service.SettingKeyOIDCConnectJWKSURL: "", + service.SettingKeyOIDCConnectScopes: "openid email profile", + service.SettingKeyOIDCConnectRedirectURL: "", + service.SettingKeyOIDCConnectFrontendRedirectURL: "/auth/oidc/callback", + service.SettingKeyOIDCConnectTokenAuthMethod: "client_secret_post", + service.SettingKeyOIDCConnectUsePKCE: "false", + service.SettingKeyOIDCConnectValidateIDToken: "true", + service.SettingKeyOIDCConnectAllowedSigningAlgs: "RS256,ES256,PS256", + service.SettingKeyOIDCConnectClockSkewSeconds: "120", + service.SettingKeyOIDCConnectRequireEmailVerified: "false", + service.SettingKeyOIDCConnectUserInfoEmailPath: "", + service.SettingKeyOIDCConnectUserInfoIDPath: "", + service.SettingKeyOIDCConnectUserInfoUsernamePath: "", + service.SettingKeySiteName: "Sub2API", service.SettingKeySiteLogo: "", service.SettingKeySiteSubtitle: "Subtitle", @@ -468,8 +490,10 @@ func TestAPIContracts(t *testing.T) { service.SettingKeyContactInfo: "support", service.SettingKeyDocURL: "https://docs.example.com", - service.SettingKeyDefaultConcurrency: "5", - service.SettingKeyDefaultBalance: "1.25", + service.SettingKeyDefaultConcurrency: "5", + service.SettingKeyDefaultBalance: "1.25", + service.SettingKeyTableDefaultPageSize: "20", + service.SettingKeyTablePageSizeOptions: "[10,20,50,100]", service.SettingKeyOpsMonitoringEnabled: "false", service.SettingKeyOpsRealtimeMonitoringEnabled: "true", @@ -502,10 +526,32 @@ func TestAPIContracts(t *testing.T) { "turnstile_enabled": true, "turnstile_site_key": "site-key", "turnstile_secret_key_configured": true, - "linuxdo_connect_enabled": false, + "linuxdo_connect_enabled": false, "linuxdo_connect_client_id": "", "linuxdo_connect_client_secret_configured": false, "linuxdo_connect_redirect_url": "", + "oidc_connect_enabled": false, + "oidc_connect_provider_name": "OIDC", + "oidc_connect_client_id": "", + "oidc_connect_client_secret_configured": false, + "oidc_connect_issuer_url": "", + "oidc_connect_discovery_url": "", + "oidc_connect_authorize_url": "", + "oidc_connect_token_url": "", + "oidc_connect_userinfo_url": "", + "oidc_connect_jwks_url": "", + "oidc_connect_scopes": "openid email profile", + "oidc_connect_redirect_url": "", + "oidc_connect_frontend_redirect_url": "/auth/oidc/callback", + "oidc_connect_token_auth_method": "client_secret_post", + "oidc_connect_use_pkce": false, + "oidc_connect_validate_id_token": true, + "oidc_connect_allowed_signing_algs": "RS256,ES256,PS256", + "oidc_connect_clock_skew_seconds": 120, + "oidc_connect_require_email_verified": false, + "oidc_connect_userinfo_email_path": "", + "oidc_connect_userinfo_id_path": "", + "oidc_connect_userinfo_username_path": "", "ops_monitoring_enabled": false, "ops_realtime_monitoring_enabled": true, "ops_query_mode_default": "auto", @@ -531,10 +577,13 @@ func TestAPIContracts(t *testing.T) { "hide_ccs_import_button": false, "purchase_subscription_enabled": false, "purchase_subscription_url": "", + "table_default_page_size": 20, + "table_page_size_options": [10, 20, 50, 100], "min_claude_code_version": "", "max_claude_code_version": "", "allow_ungrouped_key_scheduling": false, "backend_mode_enabled": false, + "enable_cch_signing": false, "enable_fingerprint_unification": true, "enable_metadata_passthrough": false, "custom_menu_items": [], diff --git a/backend/internal/server/routes/auth.go b/backend/internal/server/routes/auth.go index a6c0ecf568..c143b030fc 100644 --- a/backend/internal/server/routes/auth.go +++ b/backend/internal/server/routes/auth.go @@ -70,6 +70,14 @@ func RegisterAuthRoutes( }), h.Auth.CompleteLinuxDoOAuthRegistration, ) + auth.GET("/oauth/oidc/start", h.Auth.OIDCOAuthStart) + auth.GET("/oauth/oidc/callback", h.Auth.OIDCOAuthCallback) + auth.POST("/oauth/oidc/complete-registration", + rateLimiter.LimitWithOptions("oauth-oidc-complete", 10, time.Minute, middleware.RateLimitOptions{ + FailureMode: middleware.RateLimitFailClose, + }), + h.Auth.CompleteOIDCOAuthRegistration, + ) } // 公开设置(无需认证) diff --git a/backend/internal/service/admin_service.go b/backend/internal/service/admin_service.go index 8032f8717a..97b42c2458 100644 --- a/backend/internal/service/admin_service.go +++ b/backend/internal/service/admin_service.go @@ -21,13 +21,13 @@ import ( // AdminService interface defines admin management operations type AdminService interface { // User management - ListUsers(ctx context.Context, page, pageSize int, filters UserListFilters) ([]User, int64, error) + ListUsers(ctx context.Context, page, pageSize int, filters UserListFilters, sortBy, sortOrder string) ([]User, int64, error) GetUser(ctx context.Context, id int64) (*User, error) CreateUser(ctx context.Context, input *CreateUserInput) (*User, error) UpdateUser(ctx context.Context, id int64, input *UpdateUserInput) (*User, error) DeleteUser(ctx context.Context, id int64) error UpdateUserBalance(ctx context.Context, userID int64, balance float64, operation string, notes string) (*User, error) - GetUserAPIKeys(ctx context.Context, userID int64, page, pageSize int) ([]APIKey, int64, error) + GetUserAPIKeys(ctx context.Context, userID int64, page, pageSize int, sortBy, sortOrder string) ([]APIKey, int64, error) GetUserUsageStats(ctx context.Context, userID int64, period string) (any, error) // GetUserBalanceHistory returns paginated balance/concurrency change records for a user. // codeType is optional - pass empty string to return all types. @@ -35,7 +35,7 @@ type AdminService interface { GetUserBalanceHistory(ctx context.Context, userID int64, page, pageSize int, codeType string) ([]RedeemCode, int64, float64, error) // Group management - ListGroups(ctx context.Context, page, pageSize int, platform, status, search string, isExclusive *bool) ([]Group, int64, error) + ListGroups(ctx context.Context, page, pageSize int, platform, status, search string, isExclusive *bool, sortBy, sortOrder string) ([]Group, int64, error) GetAllGroups(ctx context.Context) ([]Group, error) GetAllGroupsByPlatform(ctx context.Context, platform string) ([]Group, error) GetGroup(ctx context.Context, id int64) (*Group, error) @@ -55,7 +55,7 @@ type AdminService interface { ReplaceUserGroup(ctx context.Context, userID, oldGroupID, newGroupID int64) (*ReplaceUserGroupResult, error) // Account management - ListAccounts(ctx context.Context, page, pageSize int, platform, accountType, status, search string, groupID int64, privacyMode string) ([]Account, int64, error) + ListAccounts(ctx context.Context, page, pageSize int, platform, accountType, status, search string, groupID int64, privacyMode string, sortBy, sortOrder string) ([]Account, int64, error) GetAccount(ctx context.Context, id int64) (*Account, error) GetAccountsByIDs(ctx context.Context, ids []int64) ([]*Account, error) CreateAccount(ctx context.Context, input *CreateAccountInput) (*Account, error) @@ -77,8 +77,8 @@ type AdminService interface { CheckMixedChannelRisk(ctx context.Context, currentAccountID int64, currentAccountPlatform string, groupIDs []int64) error // Proxy management - ListProxies(ctx context.Context, page, pageSize int, protocol, status, search string) ([]Proxy, int64, error) - ListProxiesWithAccountCount(ctx context.Context, page, pageSize int, protocol, status, search string) ([]ProxyWithAccountCount, int64, error) + ListProxies(ctx context.Context, page, pageSize int, protocol, status, search string, sortBy, sortOrder string) ([]Proxy, int64, error) + ListProxiesWithAccountCount(ctx context.Context, page, pageSize int, protocol, status, search string, sortBy, sortOrder string) ([]ProxyWithAccountCount, int64, error) GetAllProxies(ctx context.Context) ([]Proxy, error) GetAllProxiesWithAccountCount(ctx context.Context) ([]ProxyWithAccountCount, error) GetProxy(ctx context.Context, id int64) (*Proxy, error) @@ -93,7 +93,7 @@ type AdminService interface { CheckProxyQuality(ctx context.Context, id int64) (*ProxyQualityCheckResult, error) // Redeem code management - ListRedeemCodes(ctx context.Context, page, pageSize int, codeType, status, search string) ([]RedeemCode, int64, error) + ListRedeemCodes(ctx context.Context, page, pageSize int, codeType, status, search string, sortBy, sortOrder string) ([]RedeemCode, int64, error) GetRedeemCode(ctx context.Context, id int64) (*RedeemCode, error) GenerateRedeemCodes(ctx context.Context, input *GenerateRedeemCodesInput) ([]RedeemCode, error) DeleteRedeemCode(ctx context.Context, id int64) error @@ -152,10 +152,11 @@ type CreateGroupInput struct { // 支持的模型系列(仅 antigravity 平台使用) SupportedModelScopes []string // OpenAI Messages 调度配置(仅 openai 平台使用) - AllowMessagesDispatch bool - DefaultMappedModel string - RequireOAuthOnly bool - RequirePrivacySet bool + AllowMessagesDispatch bool + DefaultMappedModel string + RequireOAuthOnly bool + RequirePrivacySet bool + MessagesDispatchModelConfig OpenAIMessagesDispatchModelConfig // 从指定分组复制账号(创建分组后在同一事务内绑定) CopyAccountsFromGroupIDs []int64 } @@ -186,10 +187,11 @@ type UpdateGroupInput struct { // 支持的模型系列(仅 antigravity 平台使用) SupportedModelScopes *[]string // OpenAI Messages 调度配置(仅 openai 平台使用) - AllowMessagesDispatch *bool - DefaultMappedModel *string - RequireOAuthOnly *bool - RequirePrivacySet *bool + AllowMessagesDispatch *bool + DefaultMappedModel *string + RequireOAuthOnly *bool + RequirePrivacySet *bool + MessagesDispatchModelConfig *OpenAIMessagesDispatchModelConfig // 从指定分组复制账号(同步操作:先清空当前分组的账号绑定,再绑定源分组的账号) CopyAccountsFromGroupIDs []int64 } @@ -483,8 +485,8 @@ func NewAdminService( } // User management implementations -func (s *adminServiceImpl) ListUsers(ctx context.Context, page, pageSize int, filters UserListFilters) ([]User, int64, error) { - params := pagination.PaginationParams{Page: page, PageSize: pageSize} +func (s *adminServiceImpl) ListUsers(ctx context.Context, page, pageSize int, filters UserListFilters, sortBy, sortOrder string) ([]User, int64, error) { + params := pagination.PaginationParams{Page: page, PageSize: pageSize, SortBy: sortBy, SortOrder: sortOrder} users, result, err := s.userRepo.ListWithFilters(ctx, params, filters) if err != nil { return nil, 0, err @@ -751,8 +753,8 @@ func (s *adminServiceImpl) UpdateUserBalance(ctx context.Context, userID int64, return user, nil } -func (s *adminServiceImpl) GetUserAPIKeys(ctx context.Context, userID int64, page, pageSize int) ([]APIKey, int64, error) { - params := pagination.PaginationParams{Page: page, PageSize: pageSize} +func (s *adminServiceImpl) GetUserAPIKeys(ctx context.Context, userID int64, page, pageSize int, sortBy, sortOrder string) ([]APIKey, int64, error) { + params := pagination.PaginationParams{Page: page, PageSize: pageSize, SortBy: sortBy, SortOrder: sortOrder} keys, result, err := s.apiKeyRepo.ListByUserID(ctx, userID, params, APIKeyListFilters{}) if err != nil { return nil, 0, err @@ -787,8 +789,8 @@ func (s *adminServiceImpl) GetUserBalanceHistory(ctx context.Context, userID int } // Group management implementations -func (s *adminServiceImpl) ListGroups(ctx context.Context, page, pageSize int, platform, status, search string, isExclusive *bool) ([]Group, int64, error) { - params := pagination.PaginationParams{Page: page, PageSize: pageSize} +func (s *adminServiceImpl) ListGroups(ctx context.Context, page, pageSize int, platform, status, search string, isExclusive *bool, sortBy, sortOrder string) ([]Group, int64, error) { + params := pagination.PaginationParams{Page: page, PageSize: pageSize, SortBy: sortBy, SortOrder: sortOrder} groups, result, err := s.groupRepo.ListWithFilters(ctx, params, platform, status, search, isExclusive) if err != nil { return nil, 0, err @@ -908,7 +910,9 @@ func (s *adminServiceImpl) CreateGroup(ctx context.Context, input *CreateGroupIn RequireOAuthOnly: input.RequireOAuthOnly, RequirePrivacySet: input.RequirePrivacySet, DefaultMappedModel: input.DefaultMappedModel, + MessagesDispatchModelConfig: normalizeOpenAIMessagesDispatchModelConfig(input.MessagesDispatchModelConfig), } + sanitizeGroupMessagesDispatchFields(group) if err := s.groupRepo.Create(ctx, group); err != nil { return nil, err } @@ -1135,6 +1139,10 @@ func (s *adminServiceImpl) UpdateGroup(ctx context.Context, id int64, input *Upd if input.DefaultMappedModel != nil { group.DefaultMappedModel = *input.DefaultMappedModel } + if input.MessagesDispatchModelConfig != nil { + group.MessagesDispatchModelConfig = normalizeOpenAIMessagesDispatchModelConfig(*input.MessagesDispatchModelConfig) + } + sanitizeGroupMessagesDispatchFields(group) if err := s.groupRepo.Update(ctx, group); err != nil { return nil, err @@ -1456,8 +1464,8 @@ func (s *adminServiceImpl) ReplaceUserGroup(ctx context.Context, userID, oldGrou } // Account management implementations -func (s *adminServiceImpl) ListAccounts(ctx context.Context, page, pageSize int, platform, accountType, status, search string, groupID int64, privacyMode string) ([]Account, int64, error) { - params := pagination.PaginationParams{Page: page, PageSize: pageSize} +func (s *adminServiceImpl) ListAccounts(ctx context.Context, page, pageSize int, platform, accountType, status, search string, groupID int64, privacyMode string, sortBy, sortOrder string) ([]Account, int64, error) { + params := pagination.PaginationParams{Page: page, PageSize: pageSize, SortBy: sortBy, SortOrder: sortOrder} accounts, result, err := s.accountRepo.ListWithFilters(ctx, params, platform, accountType, status, search, groupID, privacyMode) if err != nil { return nil, 0, err @@ -1885,8 +1893,8 @@ func (s *adminServiceImpl) SetAccountSchedulable(ctx context.Context, id int64, } // Proxy management implementations -func (s *adminServiceImpl) ListProxies(ctx context.Context, page, pageSize int, protocol, status, search string) ([]Proxy, int64, error) { - params := pagination.PaginationParams{Page: page, PageSize: pageSize} +func (s *adminServiceImpl) ListProxies(ctx context.Context, page, pageSize int, protocol, status, search string, sortBy, sortOrder string) ([]Proxy, int64, error) { + params := pagination.PaginationParams{Page: page, PageSize: pageSize, SortBy: sortBy, SortOrder: sortOrder} proxies, result, err := s.proxyRepo.ListWithFilters(ctx, params, protocol, status, search) if err != nil { return nil, 0, err @@ -1894,8 +1902,8 @@ func (s *adminServiceImpl) ListProxies(ctx context.Context, page, pageSize int, return proxies, result.Total, nil } -func (s *adminServiceImpl) ListProxiesWithAccountCount(ctx context.Context, page, pageSize int, protocol, status, search string) ([]ProxyWithAccountCount, int64, error) { - params := pagination.PaginationParams{Page: page, PageSize: pageSize} +func (s *adminServiceImpl) ListProxiesWithAccountCount(ctx context.Context, page, pageSize int, protocol, status, search string, sortBy, sortOrder string) ([]ProxyWithAccountCount, int64, error) { + params := pagination.PaginationParams{Page: page, PageSize: pageSize, SortBy: sortBy, SortOrder: sortOrder} proxies, result, err := s.proxyRepo.ListWithFiltersAndAccountCount(ctx, params, protocol, status, search) if err != nil { return nil, 0, err @@ -2032,8 +2040,8 @@ func (s *adminServiceImpl) CheckProxyExists(ctx context.Context, host string, po } // Redeem code management implementations -func (s *adminServiceImpl) ListRedeemCodes(ctx context.Context, page, pageSize int, codeType, status, search string) ([]RedeemCode, int64, error) { - params := pagination.PaginationParams{Page: page, PageSize: pageSize} +func (s *adminServiceImpl) ListRedeemCodes(ctx context.Context, page, pageSize int, codeType, status, search string, sortBy, sortOrder string) ([]RedeemCode, int64, error) { + params := pagination.PaginationParams{Page: page, PageSize: pageSize, SortBy: sortBy, SortOrder: sortOrder} codes, result, err := s.redeemCodeRepo.ListWithFilters(ctx, params, codeType, status, search) if err != nil { return nil, 0, err diff --git a/backend/internal/service/admin_service_group_test.go b/backend/internal/service/admin_service_group_test.go index 536be0b583..a4c6d0caba 100644 --- a/backend/internal/service/admin_service_group_test.go +++ b/backend/internal/service/admin_service_group_test.go @@ -10,6 +10,11 @@ import ( "github.com/stretchr/testify/require" ) +func ptrString[T ~string](v T) *string { + s := string(v) + return &s +} + // groupRepoStubForAdmin 用于测试 AdminService 的 GroupRepository Stub type groupRepoStubForAdmin struct { created *Group // 记录 Create 调用的参数 @@ -120,6 +125,22 @@ func (s *groupRepoStubForAdmin) UpdateSortOrders(_ context.Context, _ []GroupSor return nil } +func TestAdminService_ListGroups_PassesSortParams(t *testing.T) { + repo := &groupRepoStubForAdmin{ + listWithFiltersGroups: []Group{{ID: 1, Name: "g1"}}, + } + svc := &adminServiceImpl{groupRepo: repo} + + _, _, err := svc.ListGroups(context.Background(), 3, 25, PlatformOpenAI, StatusActive, "needle", nil, "account_count", "ASC") + require.NoError(t, err) + require.Equal(t, pagination.PaginationParams{ + Page: 3, + PageSize: 25, + SortBy: "account_count", + SortOrder: "ASC", + }, repo.listWithFiltersParams) +} + // TestAdminService_CreateGroup_WithImagePricing 测试创建分组时 ImagePrice 字段正确传递 func TestAdminService_CreateGroup_WithImagePricing(t *testing.T) { repo := &groupRepoStubForAdmin{} @@ -245,6 +266,116 @@ func TestAdminService_UpdateGroup_PartialImagePricing(t *testing.T) { require.Nil(t, repo.updated.ImagePrice4K) } +func TestAdminService_CreateGroup_NormalizesMessagesDispatchModelConfig(t *testing.T) { + repo := &groupRepoStubForAdmin{} + svc := &adminServiceImpl{groupRepo: repo} + + group, err := svc.CreateGroup(context.Background(), &CreateGroupInput{ + Name: "dispatch-group", + Description: "dispatch config", + Platform: PlatformOpenAI, + RateMultiplier: 1.0, + MessagesDispatchModelConfig: OpenAIMessagesDispatchModelConfig{ + OpusMappedModel: " gpt-5.4-high ", + SonnetMappedModel: " gpt-5.3-codex ", + HaikuMappedModel: " gpt-5.4-mini-medium ", + ExactModelMappings: map[string]string{ + " claude-sonnet-4-5-20250929 ": " gpt-5.2-high ", + }, + }, + }) + require.NoError(t, err) + require.NotNil(t, group) + require.NotNil(t, repo.created) + require.Equal(t, OpenAIMessagesDispatchModelConfig{ + OpusMappedModel: "gpt-5.4", + SonnetMappedModel: "gpt-5.3-codex", + HaikuMappedModel: "gpt-5.4-mini", + ExactModelMappings: map[string]string{ + "claude-sonnet-4-5-20250929": "gpt-5.2", + }, + }, repo.created.MessagesDispatchModelConfig) +} + +func TestAdminService_UpdateGroup_NormalizesMessagesDispatchModelConfig(t *testing.T) { + existingGroup := &Group{ + ID: 1, + Name: "existing-group", + Platform: PlatformOpenAI, + Status: StatusActive, + } + repo := &groupRepoStubForAdmin{getByID: existingGroup} + svc := &adminServiceImpl{groupRepo: repo} + + group, err := svc.UpdateGroup(context.Background(), 1, &UpdateGroupInput{ + MessagesDispatchModelConfig: &OpenAIMessagesDispatchModelConfig{ + SonnetMappedModel: " gpt-5.4-medium ", + ExactModelMappings: map[string]string{ + " claude-haiku-4-5-20251001 ": " gpt-5.4-mini-high ", + }, + }, + }) + require.NoError(t, err) + require.NotNil(t, group) + require.NotNil(t, repo.updated) + require.Equal(t, OpenAIMessagesDispatchModelConfig{ + SonnetMappedModel: "gpt-5.4", + ExactModelMappings: map[string]string{ + "claude-haiku-4-5-20251001": "gpt-5.4-mini", + }, + }, repo.updated.MessagesDispatchModelConfig) +} + +func TestAdminService_CreateGroup_ClearsMessagesDispatchFieldsForNonOpenAIPlatform(t *testing.T) { + repo := &groupRepoStubForAdmin{} + svc := &adminServiceImpl{groupRepo: repo} + + group, err := svc.CreateGroup(context.Background(), &CreateGroupInput{ + Name: "anthropic-group", + Description: "non-openai", + Platform: PlatformAnthropic, + RateMultiplier: 1.0, + AllowMessagesDispatch: true, + DefaultMappedModel: "gpt-5.4", + MessagesDispatchModelConfig: OpenAIMessagesDispatchModelConfig{ + OpusMappedModel: "gpt-5.4", + }, + }) + require.NoError(t, err) + require.NotNil(t, group) + require.NotNil(t, repo.created) + require.False(t, repo.created.AllowMessagesDispatch) + require.Empty(t, repo.created.DefaultMappedModel) + require.Equal(t, OpenAIMessagesDispatchModelConfig{}, repo.created.MessagesDispatchModelConfig) +} + +func TestAdminService_UpdateGroup_ClearsMessagesDispatchFieldsWhenPlatformChangesAwayFromOpenAI(t *testing.T) { + existingGroup := &Group{ + ID: 1, + Name: "existing-openai-group", + Platform: PlatformOpenAI, + Status: StatusActive, + AllowMessagesDispatch: true, + DefaultMappedModel: "gpt-5.4", + MessagesDispatchModelConfig: OpenAIMessagesDispatchModelConfig{ + SonnetMappedModel: "gpt-5.3-codex", + }, + } + repo := &groupRepoStubForAdmin{getByID: existingGroup} + svc := &adminServiceImpl{groupRepo: repo} + + group, err := svc.UpdateGroup(context.Background(), 1, &UpdateGroupInput{ + Platform: PlatformAnthropic, + }) + require.NoError(t, err) + require.NotNil(t, group) + require.NotNil(t, repo.updated) + require.Equal(t, PlatformAnthropic, repo.updated.Platform) + require.False(t, repo.updated.AllowMessagesDispatch) + require.Empty(t, repo.updated.DefaultMappedModel) + require.Equal(t, OpenAIMessagesDispatchModelConfig{}, repo.updated.MessagesDispatchModelConfig) +} + func TestAdminService_ListGroups_WithSearch(t *testing.T) { // 测试: // 1. search 参数正常传递到 repository 层 @@ -258,7 +389,7 @@ func TestAdminService_ListGroups_WithSearch(t *testing.T) { } svc := &adminServiceImpl{groupRepo: repo} - groups, total, err := svc.ListGroups(context.Background(), 1, 20, "", "", "alpha", nil) + groups, total, err := svc.ListGroups(context.Background(), 1, 20, "", "", "alpha", nil, "", "") require.NoError(t, err) require.Equal(t, int64(1), total) require.Equal(t, []Group{{ID: 1, Name: "alpha"}}, groups) @@ -276,7 +407,7 @@ func TestAdminService_ListGroups_WithSearch(t *testing.T) { } svc := &adminServiceImpl{groupRepo: repo} - groups, total, err := svc.ListGroups(context.Background(), 2, 10, "", "", "", nil) + groups, total, err := svc.ListGroups(context.Background(), 2, 10, "", "", "", nil, "", "") require.NoError(t, err) require.Empty(t, groups) require.Equal(t, int64(0), total) @@ -295,7 +426,7 @@ func TestAdminService_ListGroups_WithSearch(t *testing.T) { } svc := &adminServiceImpl{groupRepo: repo} - groups, total, err := svc.ListGroups(context.Background(), 3, 50, PlatformAntigravity, StatusActive, "beta", &isExclusive) + groups, total, err := svc.ListGroups(context.Background(), 3, 50, PlatformAntigravity, StatusActive, "beta", &isExclusive, "", "") require.NoError(t, err) require.Equal(t, int64(42), total) require.Equal(t, []Group{{ID: 2, Name: "beta"}}, groups) diff --git a/backend/internal/service/admin_service_list_users_test.go b/backend/internal/service/admin_service_list_users_test.go index 37f348dfbd..ceeb52c294 100644 --- a/backend/internal/service/admin_service_list_users_test.go +++ b/backend/internal/service/admin_service_list_users_test.go @@ -13,11 +13,13 @@ import ( type userRepoStubForListUsers struct { userRepoStub - users []User - err error + users []User + err error + listWithFiltersParams pagination.PaginationParams } func (s *userRepoStubForListUsers) ListWithFilters(_ context.Context, params pagination.PaginationParams, _ UserListFilters) ([]User, *pagination.PaginationResult, error) { + s.listWithFiltersParams = params if s.err != nil { return nil, nil, s.err } @@ -103,7 +105,7 @@ func TestAdminService_ListUsers_BatchRateFallbackToSingle(t *testing.T) { userGroupRateRepo: rateRepo, } - users, total, err := svc.ListUsers(context.Background(), 1, 20, UserListFilters{}) + users, total, err := svc.ListUsers(context.Background(), 1, 20, UserListFilters{}, "", "") require.NoError(t, err) require.Equal(t, int64(2), total) require.Len(t, users, 2) @@ -112,3 +114,19 @@ func TestAdminService_ListUsers_BatchRateFallbackToSingle(t *testing.T) { require.Equal(t, 1.1, users[0].GroupRates[11]) require.Equal(t, 2.2, users[1].GroupRates[22]) } + +func TestAdminService_ListUsers_PassesSortParams(t *testing.T) { + userRepo := &userRepoStubForListUsers{ + users: []User{{ID: 1, Email: "a@example.com"}}, + } + svc := &adminServiceImpl{userRepo: userRepo} + + _, _, err := svc.ListUsers(context.Background(), 2, 50, UserListFilters{}, "email", "ASC") + require.NoError(t, err) + require.Equal(t, pagination.PaginationParams{ + Page: 2, + PageSize: 50, + SortBy: "email", + SortOrder: "ASC", + }, userRepo.listWithFiltersParams) +} diff --git a/backend/internal/service/admin_service_search_test.go b/backend/internal/service/admin_service_search_test.go index eb213e6af6..595e99e344 100644 --- a/backend/internal/service/admin_service_search_test.go +++ b/backend/internal/service/admin_service_search_test.go @@ -170,13 +170,13 @@ func TestAdminService_ListAccounts_WithSearch(t *testing.T) { } svc := &adminServiceImpl{accountRepo: repo} - accounts, total, err := svc.ListAccounts(context.Background(), 1, 20, PlatformGemini, AccountTypeOAuth, StatusActive, "acc", 0, "") + accounts, total, err := svc.ListAccounts(context.Background(), 1, 20, PlatformGemini, AccountTypeOAuth, StatusActive, "acc", 0, "", "name", "ASC") require.NoError(t, err) require.Equal(t, int64(10), total) require.Equal(t, []Account{{ID: 1, Name: "acc"}}, accounts) require.Equal(t, 1, repo.listWithFiltersCalls) - require.Equal(t, pagination.PaginationParams{Page: 1, PageSize: 20}, repo.listWithFiltersParams) + require.Equal(t, pagination.PaginationParams{Page: 1, PageSize: 20, SortBy: "name", SortOrder: "ASC"}, repo.listWithFiltersParams) require.Equal(t, PlatformGemini, repo.listWithFiltersPlatform) require.Equal(t, AccountTypeOAuth, repo.listWithFiltersType) require.Equal(t, StatusActive, repo.listWithFiltersStatus) @@ -192,7 +192,7 @@ func TestAdminService_ListAccounts_WithPrivacyMode(t *testing.T) { } svc := &adminServiceImpl{accountRepo: repo} - accounts, total, err := svc.ListAccounts(context.Background(), 1, 20, PlatformOpenAI, AccountTypeOAuth, StatusActive, "acc2", 0, PrivacyModeCFBlocked) + accounts, total, err := svc.ListAccounts(context.Background(), 1, 20, PlatformOpenAI, AccountTypeOAuth, StatusActive, "acc2", 0, PrivacyModeCFBlocked, "", "") require.NoError(t, err) require.Equal(t, int64(1), total) require.Equal(t, []Account{{ID: 2, Name: "acc2"}}, accounts) @@ -208,13 +208,13 @@ func TestAdminService_ListProxies_WithSearch(t *testing.T) { } svc := &adminServiceImpl{proxyRepo: repo} - proxies, total, err := svc.ListProxies(context.Background(), 3, 50, "http", StatusActive, "p1") + proxies, total, err := svc.ListProxies(context.Background(), 3, 50, "http", StatusActive, "p1", "name", "ASC") require.NoError(t, err) require.Equal(t, int64(7), total) require.Equal(t, []Proxy{{ID: 2, Name: "p1"}}, proxies) require.Equal(t, 1, repo.listWithFiltersCalls) - require.Equal(t, pagination.PaginationParams{Page: 3, PageSize: 50}, repo.listWithFiltersParams) + require.Equal(t, pagination.PaginationParams{Page: 3, PageSize: 50, SortBy: "name", SortOrder: "ASC"}, repo.listWithFiltersParams) require.Equal(t, "http", repo.listWithFiltersProtocol) require.Equal(t, StatusActive, repo.listWithFiltersStatus) require.Equal(t, "p1", repo.listWithFiltersSearch) @@ -229,13 +229,13 @@ func TestAdminService_ListProxiesWithAccountCount_WithSearch(t *testing.T) { } svc := &adminServiceImpl{proxyRepo: repo} - proxies, total, err := svc.ListProxiesWithAccountCount(context.Background(), 2, 10, "socks5", StatusDisabled, "p2") + proxies, total, err := svc.ListProxiesWithAccountCount(context.Background(), 2, 10, "socks5", StatusDisabled, "p2", "account_count", "DESC") require.NoError(t, err) require.Equal(t, int64(9), total) require.Equal(t, []ProxyWithAccountCount{{Proxy: Proxy{ID: 3, Name: "p2"}, AccountCount: 5}}, proxies) require.Equal(t, 1, repo.listWithFiltersAndAccountCountCalls) - require.Equal(t, pagination.PaginationParams{Page: 2, PageSize: 10}, repo.listWithFiltersAndAccountCountParams) + require.Equal(t, pagination.PaginationParams{Page: 2, PageSize: 10, SortBy: "account_count", SortOrder: "DESC"}, repo.listWithFiltersAndAccountCountParams) require.Equal(t, "socks5", repo.listWithFiltersAndAccountCountProtocol) require.Equal(t, StatusDisabled, repo.listWithFiltersAndAccountCountStatus) require.Equal(t, "p2", repo.listWithFiltersAndAccountCountSearch) @@ -250,13 +250,13 @@ func TestAdminService_ListRedeemCodes_WithSearch(t *testing.T) { } svc := &adminServiceImpl{redeemCodeRepo: repo} - codes, total, err := svc.ListRedeemCodes(context.Background(), 1, 20, RedeemTypeBalance, StatusUnused, "ABC") + codes, total, err := svc.ListRedeemCodes(context.Background(), 1, 20, RedeemTypeBalance, StatusUnused, "ABC", "value", "ASC") require.NoError(t, err) require.Equal(t, int64(3), total) require.Equal(t, []RedeemCode{{ID: 4, Code: "ABC"}}, codes) require.Equal(t, 1, repo.listWithFiltersCalls) - require.Equal(t, pagination.PaginationParams{Page: 1, PageSize: 20}, repo.listWithFiltersParams) + require.Equal(t, pagination.PaginationParams{Page: 1, PageSize: 20, SortBy: "value", SortOrder: "ASC"}, repo.listWithFiltersParams) require.Equal(t, RedeemTypeBalance, repo.listWithFiltersType) require.Equal(t, StatusUnused, repo.listWithFiltersStatus) require.Equal(t, "ABC", repo.listWithFiltersSearch) diff --git a/backend/internal/service/api_key_auth_cache.go b/backend/internal/service/api_key_auth_cache.go index ad6ba0e930..c2e96df13a 100644 --- a/backend/internal/service/api_key_auth_cache.go +++ b/backend/internal/service/api_key_auth_cache.go @@ -4,6 +4,7 @@ import "time" // APIKeyAuthSnapshot API Key 认证缓存快照(仅包含认证所需字段) type APIKeyAuthSnapshot struct { + Version int `json:"version"` APIKeyID int64 `json:"api_key_id"` UserID int64 `json:"user_id"` GroupID *int64 `json:"group_id,omitempty"` @@ -63,8 +64,9 @@ type APIKeyAuthGroupSnapshot struct { SupportedModelScopes []string `json:"supported_model_scopes,omitempty"` // OpenAI Messages 调度配置(仅 openai 平台使用) - AllowMessagesDispatch bool `json:"allow_messages_dispatch"` - DefaultMappedModel string `json:"default_mapped_model,omitempty"` + AllowMessagesDispatch bool `json:"allow_messages_dispatch"` + DefaultMappedModel string `json:"default_mapped_model,omitempty"` + MessagesDispatchModelConfig OpenAIMessagesDispatchModelConfig `json:"messages_dispatch_model_config,omitempty"` } // APIKeyAuthCacheEntry 缓存条目,支持负缓存 diff --git a/backend/internal/service/api_key_auth_cache_impl.go b/backend/internal/service/api_key_auth_cache_impl.go index 64a70e8cf4..8069ed4fe2 100644 --- a/backend/internal/service/api_key_auth_cache_impl.go +++ b/backend/internal/service/api_key_auth_cache_impl.go @@ -13,6 +13,8 @@ import ( "github.com/dgraph-io/ristretto" ) +const apiKeyAuthSnapshotVersion = 3 + type apiKeyAuthCacheConfig struct { l1Size int l1TTL time.Duration @@ -192,6 +194,9 @@ func (s *APIKeyService) applyAuthCacheEntry(key string, entry *APIKeyAuthCacheEn if entry.Snapshot == nil { return nil, false, nil } + if entry.Snapshot.Version != apiKeyAuthSnapshotVersion { + return nil, false, nil + } return s.snapshotToAPIKey(key, entry.Snapshot), true, nil } @@ -200,6 +205,7 @@ func (s *APIKeyService) snapshotFromAPIKey(apiKey *APIKey) *APIKeyAuthSnapshot { return nil } snapshot := &APIKeyAuthSnapshot{ + Version: apiKeyAuthSnapshotVersion, APIKeyID: apiKey.ID, UserID: apiKey.UserID, GroupID: apiKey.GroupID, @@ -243,6 +249,7 @@ func (s *APIKeyService) snapshotFromAPIKey(apiKey *APIKey) *APIKeyAuthSnapshot { SupportedModelScopes: apiKey.Group.SupportedModelScopes, AllowMessagesDispatch: apiKey.Group.AllowMessagesDispatch, DefaultMappedModel: apiKey.Group.DefaultMappedModel, + MessagesDispatchModelConfig: apiKey.Group.MessagesDispatchModelConfig, } } return snapshot @@ -298,6 +305,7 @@ func (s *APIKeyService) snapshotToAPIKey(key string, snapshot *APIKeyAuthSnapsho SupportedModelScopes: snapshot.Group.SupportedModelScopes, AllowMessagesDispatch: snapshot.Group.AllowMessagesDispatch, DefaultMappedModel: snapshot.Group.DefaultMappedModel, + MessagesDispatchModelConfig: snapshot.Group.MessagesDispatchModelConfig, } } s.compileAPIKeyIPRules(apiKey) diff --git a/backend/internal/service/api_key_service_cache_test.go b/backend/internal/service/api_key_service_cache_test.go index 357f8deff7..3c2f7dbb5c 100644 --- a/backend/internal/service/api_key_service_cache_test.go +++ b/backend/internal/service/api_key_service_cache_test.go @@ -188,6 +188,7 @@ func TestAPIKeyService_GetByKey_UsesL2Cache(t *testing.T) { groupID := int64(9) cacheEntry := &APIKeyAuthCacheEntry{ Snapshot: &APIKeyAuthSnapshot{ + Version: apiKeyAuthSnapshotVersion, APIKeyID: 1, UserID: 2, GroupID: &groupID, @@ -226,6 +227,129 @@ func TestAPIKeyService_GetByKey_UsesL2Cache(t *testing.T) { require.Equal(t, map[string][]int64{"claude-opus-*": {1, 2}}, apiKey.Group.ModelRouting) } +func TestAPIKeyService_SnapshotRoundTrip_PreservesMessagesDispatchModelConfig(t *testing.T) { + svc := NewAPIKeyService(nil, nil, nil, nil, nil, nil, &config.Config{}) + groupID := int64(9) + apiKey := &APIKey{ + ID: 1, + UserID: 2, + GroupID: &groupID, + Key: "k-roundtrip", + Status: StatusActive, + User: &User{ + ID: 2, + Status: StatusActive, + Role: RoleUser, + Balance: 10, + Concurrency: 3, + }, + Group: &Group{ + ID: groupID, + Name: "openai", + Platform: PlatformOpenAI, + Status: StatusActive, + SubscriptionType: SubscriptionTypeStandard, + RateMultiplier: 1, + AllowMessagesDispatch: true, + DefaultMappedModel: "gpt-5.4", + MessagesDispatchModelConfig: OpenAIMessagesDispatchModelConfig{ + OpusMappedModel: "gpt-5.4-nano", + SonnetMappedModel: "gpt-5.3-codex", + HaikuMappedModel: "gpt-5.4-mini", + ExactModelMappings: map[string]string{ + "claude-sonnet-4.5": "gpt-5.4-nano", + }, + }, + }, + } + + snapshot := svc.snapshotFromAPIKey(apiKey) + roundTrip := svc.snapshotToAPIKey(apiKey.Key, snapshot) + + require.NotNil(t, roundTrip) + require.NotNil(t, roundTrip.Group) + require.Equal(t, apiKey.Group.MessagesDispatchModelConfig, roundTrip.Group.MessagesDispatchModelConfig) +} + +func TestAPIKeyService_GetByKey_IgnoresLegacyAuthCacheSnapshotWithoutMessagesDispatchConfig(t *testing.T) { + cache := &authCacheStub{} + var repoCalls int32 + repo := &authRepoStub{ + getByKeyForAuth: func(ctx context.Context, key string) (*APIKey, error) { + atomic.AddInt32(&repoCalls, 1) + groupID := int64(9) + return &APIKey{ + ID: 1, + UserID: 2, + GroupID: &groupID, + Status: StatusActive, + User: &User{ + ID: 2, + Status: StatusActive, + Role: RoleUser, + Balance: 10, + Concurrency: 3, + }, + Group: &Group{ + ID: groupID, + Name: "openai", + Platform: PlatformOpenAI, + Status: StatusActive, + Hydrated: true, + SubscriptionType: SubscriptionTypeStandard, + RateMultiplier: 1, + AllowMessagesDispatch: true, + DefaultMappedModel: "gpt-5.4", + MessagesDispatchModelConfig: OpenAIMessagesDispatchModelConfig{ + OpusMappedModel: "gpt-5.4-nano", + }, + }, + }, nil + }, + } + cfg := &config.Config{ + APIKeyAuth: config.APIKeyAuthCacheConfig{ + L2TTLSeconds: 60, + }, + } + svc := NewAPIKeyService(repo, nil, nil, nil, nil, cache, cfg) + + groupID := int64(9) + cache.getAuthCache = func(ctx context.Context, key string) (*APIKeyAuthCacheEntry, error) { + return &APIKeyAuthCacheEntry{ + Snapshot: &APIKeyAuthSnapshot{ + APIKeyID: 1, + UserID: 2, + GroupID: &groupID, + Status: StatusActive, + User: APIKeyAuthUserSnapshot{ + ID: 2, + Status: StatusActive, + Role: RoleUser, + Balance: 10, + Concurrency: 3, + }, + Group: &APIKeyAuthGroupSnapshot{ + ID: groupID, + Name: "openai", + Platform: PlatformOpenAI, + Status: StatusActive, + SubscriptionType: SubscriptionTypeStandard, + RateMultiplier: 1, + AllowMessagesDispatch: true, + DefaultMappedModel: "gpt-5.4", + }, + }, + }, nil + } + + apiKey, err := svc.GetByKey(context.Background(), "k-legacy") + require.NoError(t, err) + require.Equal(t, int32(1), atomic.LoadInt32(&repoCalls)) + require.NotNil(t, apiKey.Group) + require.Equal(t, "gpt-5.4-nano", apiKey.Group.MessagesDispatchModelConfig.OpusMappedModel) +} + func TestAPIKeyService_GetByKey_NegativeCache(t *testing.T) { cache := &authCacheStub{} repo := &authRepoStub{ diff --git a/backend/internal/service/auth_service.go b/backend/internal/service/auth_service.go index 6e524fb91c..fd28cd4235 100644 --- a/backend/internal/service/auth_service.go +++ b/backend/internal/service/auth_service.go @@ -833,7 +833,8 @@ func randomHexString(byteLength int) (string, error) { func isReservedEmail(email string) bool { normalized := strings.ToLower(strings.TrimSpace(email)) - return strings.HasSuffix(normalized, LinuxDoConnectSyntheticEmailDomain) + return strings.HasSuffix(normalized, LinuxDoConnectSyntheticEmailDomain) || + strings.HasSuffix(normalized, OIDCConnectSyntheticEmailDomain) } // GenerateToken 生成JWT access token diff --git a/backend/internal/service/channel_service.go b/backend/internal/service/channel_service.go index ab22ed1337..ce684e5ea7 100644 --- a/backend/internal/service/channel_service.go +++ b/backend/internal/service/channel_service.go @@ -281,7 +281,6 @@ 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...) @@ -309,7 +308,6 @@ 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] @@ -318,8 +316,6 @@ func populateChannelCache(channels []Channel, groupPlatforms map[int64]string) * } } - // 通配符条目保持配置顺序(最先匹配到优先) - return cache } @@ -484,7 +480,10 @@ func (s *ChannelService) ResolveChannelMapping(ctx context.Context, groupID int6 // 返回 true 表示模型被限制(不在允许列表中)。 // 如果渠道未启用模型限制或分组无渠道关联,返回 false。 func (s *ChannelService) IsModelRestricted(ctx context.Context, groupID int64, model string) bool { - lk, _ := s.lookupGroupChannel(ctx, groupID) + 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) + } if lk == nil { return false } @@ -804,7 +803,6 @@ 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/domain_constants.go b/backend/internal/service/domain_constants.go index 52df52d656..68d7da3b45 100644 --- a/backend/internal/service/domain_constants.go +++ b/backend/internal/service/domain_constants.go @@ -71,6 +71,9 @@ const ( // LinuxDoConnectSyntheticEmailDomain 是 LinuxDo Connect 用户的合成邮箱后缀(RFC 保留域名)。 const LinuxDoConnectSyntheticEmailDomain = "@linuxdo-connect.invalid" +// OIDCConnectSyntheticEmailDomain 是 OIDC 用户的合成邮箱后缀(RFC 保留域名)。 +const OIDCConnectSyntheticEmailDomain = "@oidc-connect.invalid" + // Setting keys const ( // 注册设置 @@ -105,6 +108,30 @@ const ( SettingKeyLinuxDoConnectClientSecret = "linuxdo_connect_client_secret" SettingKeyLinuxDoConnectRedirectURL = "linuxdo_connect_redirect_url" + // Generic OIDC OAuth 登录设置 + SettingKeyOIDCConnectEnabled = "oidc_connect_enabled" + SettingKeyOIDCConnectProviderName = "oidc_connect_provider_name" + SettingKeyOIDCConnectClientID = "oidc_connect_client_id" + SettingKeyOIDCConnectClientSecret = "oidc_connect_client_secret" + SettingKeyOIDCConnectIssuerURL = "oidc_connect_issuer_url" + SettingKeyOIDCConnectDiscoveryURL = "oidc_connect_discovery_url" + SettingKeyOIDCConnectAuthorizeURL = "oidc_connect_authorize_url" + SettingKeyOIDCConnectTokenURL = "oidc_connect_token_url" + SettingKeyOIDCConnectUserInfoURL = "oidc_connect_userinfo_url" + SettingKeyOIDCConnectJWKSURL = "oidc_connect_jwks_url" + SettingKeyOIDCConnectScopes = "oidc_connect_scopes" + SettingKeyOIDCConnectRedirectURL = "oidc_connect_redirect_url" + SettingKeyOIDCConnectFrontendRedirectURL = "oidc_connect_frontend_redirect_url" + SettingKeyOIDCConnectTokenAuthMethod = "oidc_connect_token_auth_method" + SettingKeyOIDCConnectUsePKCE = "oidc_connect_use_pkce" + SettingKeyOIDCConnectValidateIDToken = "oidc_connect_validate_id_token" + SettingKeyOIDCConnectAllowedSigningAlgs = "oidc_connect_allowed_signing_algs" + SettingKeyOIDCConnectClockSkewSeconds = "oidc_connect_clock_skew_seconds" + SettingKeyOIDCConnectRequireEmailVerified = "oidc_connect_require_email_verified" + SettingKeyOIDCConnectUserInfoEmailPath = "oidc_connect_userinfo_email_path" + SettingKeyOIDCConnectUserInfoIDPath = "oidc_connect_userinfo_id_path" + SettingKeyOIDCConnectUserInfoUsernamePath = "oidc_connect_userinfo_username_path" + // OEM设置 SettingKeySiteName = "site_name" // 网站名称 SettingKeySiteLogo = "site_logo" // 网站Logo (base64) @@ -116,6 +143,8 @@ const ( SettingKeyHideCcsImportButton = "hide_ccs_import_button" // 是否隐藏 API Keys 页面的导入 CCS 按钮 SettingKeyPurchaseSubscriptionEnabled = "purchase_subscription_enabled" // 是否展示"购买订阅"页面入口 SettingKeyPurchaseSubscriptionURL = "purchase_subscription_url" // "购买订阅"页面 URL(作为 iframe src) + SettingKeyTableDefaultPageSize = "table_default_page_size" // 表格默认每页条数 + SettingKeyTablePageSizeOptions = "table_page_size_options" // 表格可选每页条数(JSON 数组) SettingKeyCustomMenuItems = "custom_menu_items" // 自定义菜单项(JSON 数组) SettingKeyCustomEndpoints = "custom_endpoints" // 自定义端点列表(JSON 数组) @@ -218,6 +247,8 @@ const ( SettingKeyEnableFingerprintUnification = "enable_fingerprint_unification" // SettingKeyEnableMetadataPassthrough 是否透传客户端原始 metadata.user_id(默认 false) SettingKeyEnableMetadataPassthrough = "enable_metadata_passthrough" + // SettingKeyEnableCCHSigning 是否对 billing header 中的 cch 进行 xxHash64 签名(默认 false) + SettingKeyEnableCCHSigning = "enable_cch_signing" ) // AdminAPIKeyPrefix is the prefix for admin API keys (distinct from user "sk-" keys). diff --git a/backend/internal/service/gateway_anthropic_apikey_passthrough_test.go b/backend/internal/service/gateway_anthropic_apikey_passthrough_test.go index 6e19db322d..5be1f73328 100644 --- a/backend/internal/service/gateway_anthropic_apikey_passthrough_test.go +++ b/backend/internal/service/gateway_anthropic_apikey_passthrough_test.go @@ -761,7 +761,16 @@ func TestGatewayService_AnthropicOAuth_ForwardPreservesBillingHeaderSystemBlock( system := gjson.GetBytes(upstream.lastBody, "system") require.True(t, system.Exists()) - require.Contains(t, system.Raw, "x-anthropic-billing-header keep") + require.True(t, system.IsArray(), "system should be an array") + require.Equal(t, claudeCodeSystemPrompt, system.Array()[0].Get("text").String()) + require.Equal(t, "ephemeral", system.Array()[0].Get("cache_control.type").String()) + + // 原始 system prompt 应迁移至 messages 中 + messages := gjson.GetBytes(upstream.lastBody, "messages") + require.True(t, messages.IsArray()) + firstMsg := messages.Array()[0] + require.Equal(t, "user", firstMsg.Get("role").String()) + require.Contains(t, firstMsg.Get("content.0.text").String(), "x-anthropic-billing-header keep") }) } } diff --git a/backend/internal/service/gateway_billing_header.go b/backend/internal/service/gateway_billing_header.go new file mode 100644 index 0000000000..91fbfd8fdd --- /dev/null +++ b/backend/internal/service/gateway_billing_header.go @@ -0,0 +1,73 @@ +package service + +import ( + "fmt" + "regexp" + "strings" + + "github.com/cespare/xxhash/v2" + "github.com/tidwall/gjson" + "github.com/tidwall/sjson" +) + +// ccVersionInBillingRe matches the semver part of cc_version (X.Y.Z), preserving +// the trailing message-derived suffix (e.g. ".c02") if present. +var ccVersionInBillingRe = regexp.MustCompile(`cc_version=\d+\.\d+\.\d+`) + +// cchPlaceholderRe matches the cch=00000 placeholder in billing header text, +// scoped to x-anthropic-billing-header to avoid touching user content. +var cchPlaceholderRe = regexp.MustCompile(`(x-anthropic-billing-header:[^"]*?\bcch=)(00000)(;)`) + +const cchSeed uint64 = 0x6E52736AC806831E + +// syncBillingHeaderVersion rewrites cc_version in x-anthropic-billing-header +// system text blocks to match the version extracted from userAgent. +// Only touches system array blocks whose text starts with "x-anthropic-billing-header". +func syncBillingHeaderVersion(body []byte, userAgent string) []byte { + version := ExtractCLIVersion(userAgent) + if version == "" { + return body + } + + systemResult := gjson.GetBytes(body, "system") + if !systemResult.Exists() || !systemResult.IsArray() { + return body + } + + replacement := "cc_version=" + version + idx := 0 + systemResult.ForEach(func(_, item gjson.Result) bool { + text := item.Get("text") + if text.Exists() && text.Type == gjson.String && + strings.HasPrefix(text.String(), "x-anthropic-billing-header") { + newText := ccVersionInBillingRe.ReplaceAllString(text.String(), replacement) + if newText != text.String() { + if updated, err := sjson.SetBytes(body, fmt.Sprintf("system.%d.text", idx), newText); err == nil { + body = updated + } + } + } + idx++ + return true + }) + + return body +} + +// signBillingHeaderCCH computes the xxHash64-based CCH signature for the request +// body and replaces the cch=00000 placeholder with the computed 5-hex-char hash. +// The body must contain the placeholder when this function is called. +func signBillingHeaderCCH(body []byte) []byte { + if !cchPlaceholderRe.Match(body) { + return body + } + cch := fmt.Sprintf("%05x", xxHash64Seeded(body, cchSeed)&0xFFFFF) + return cchPlaceholderRe.ReplaceAll(body, []byte("${1}"+cch+"${3}")) +} + +// xxHash64Seeded computes xxHash64 of data with a custom seed. +func xxHash64Seeded(data []byte, seed uint64) uint64 { + d := xxhash.NewWithSeed(seed) + _, _ = d.Write(data) + return d.Sum64() +} diff --git a/backend/internal/service/gateway_billing_header_test.go b/backend/internal/service/gateway_billing_header_test.go new file mode 100644 index 0000000000..ffc4091c63 --- /dev/null +++ b/backend/internal/service/gateway_billing_header_test.go @@ -0,0 +1,165 @@ +package service + +import ( + "fmt" + "testing" + + "github.com/cespare/xxhash/v2" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "github.com/tidwall/gjson" +) + +func TestSyncBillingHeaderVersion(t *testing.T) { + tests := []struct { + name string + body string + userAgent string + wantSub string // substring expected in result + unchanged bool // expect body to remain the same + }{ + { + name: "replaces cc_version preserving message-derived suffix", + body: `{"system":[{"type":"text","text":"x-anthropic-billing-header: cc_version=2.1.81.df2; cc_entrypoint=cli; cch=00000;"},{"type":"text","text":"You are Claude Code.","cache_control":{"type":"ephemeral"}}],"messages":[]}`, + userAgent: "claude-cli/2.1.22 (external, cli)", + wantSub: "cc_version=2.1.22.df2", + }, + { + name: "no billing header in system", + body: `{"system":[{"type":"text","text":"You are Claude Code."}],"messages":[]}`, + userAgent: "claude-cli/2.1.22", + unchanged: true, + }, + { + name: "no system field", + body: `{"messages":[]}`, + userAgent: "claude-cli/2.1.22", + unchanged: true, + }, + { + name: "user-agent without version", + body: `{"system":[{"type":"text","text":"x-anthropic-billing-header: cc_version=2.1.81; cc_entrypoint=cli; cch=00000;"}],"messages":[]}`, + userAgent: "Mozilla/5.0", + unchanged: true, + }, + { + name: "empty user-agent", + body: `{"system":[{"type":"text","text":"x-anthropic-billing-header: cc_version=2.1.81; cc_entrypoint=cli; cch=00000;"}],"messages":[]}`, + userAgent: "", + unchanged: true, + }, + { + name: "version already matches", + body: `{"system":[{"type":"text","text":"x-anthropic-billing-header: cc_version=2.1.22; cc_entrypoint=cli; cch=00000;"}],"messages":[]}`, + userAgent: "claude-cli/2.1.22", + unchanged: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + result := syncBillingHeaderVersion([]byte(tt.body), tt.userAgent) + if tt.unchanged { + assert.Equal(t, tt.body, string(result), "body should remain unchanged") + } else { + assert.Contains(t, string(result), tt.wantSub) + // Ensure old semver is gone + assert.NotContains(t, string(result), "cc_version=2.1.81") + } + }) + } +} + +func TestSignBillingHeaderCCH(t *testing.T) { + t.Run("replaces placeholder with hash", func(t *testing.T) { + body := []byte(`{"system":[{"type":"text","text":"x-anthropic-billing-header: cc_version=2.1.63.a43; cc_entrypoint=cli; cch=00000;"}],"messages":[{"role":"user","content":[{"type":"text","text":"hello"}]}]}`) + result := signBillingHeaderCCH(body) + + // Should not have the placeholder anymore + assert.NotContains(t, string(result), "cch=00000") + + // Should have a 5 hex-char cch value + billingText := gjson.GetBytes(result, "system.0.text").String() + require.Contains(t, billingText, "cch=") + assert.Regexp(t, `cch=[0-9a-f]{5};`, billingText) + }) + + t.Run("no placeholder - body unchanged", func(t *testing.T) { + body := []byte(`{"system":[{"type":"text","text":"x-anthropic-billing-header: cc_version=2.1.63; cc_entrypoint=cli; cch=abcde;"}],"messages":[]}`) + result := signBillingHeaderCCH(body) + assert.Equal(t, string(body), string(result)) + }) + + t.Run("no billing header - body unchanged", func(t *testing.T) { + body := []byte(`{"system":[{"type":"text","text":"You are Claude Code."}],"messages":[]}`) + result := signBillingHeaderCCH(body) + assert.Equal(t, string(body), string(result)) + }) + + t.Run("cch=00000 in user content is not touched", func(t *testing.T) { + body := []byte(`{"system":[{"type":"text","text":"x-anthropic-billing-header: cc_version=2.1.63; cc_entrypoint=cli; cch=00000;"}],"messages":[{"role":"user","content":[{"type":"text","text":"keep literal cch=00000 in this message"}]}]}`) + result := signBillingHeaderCCH(body) + + // Billing header should be signed + billingText := gjson.GetBytes(result, "system.0.text").String() + assert.NotContains(t, billingText, "cch=00000") + + // User message should keep its literal cch=00000 + userText := gjson.GetBytes(result, "messages.0.content.0.text").String() + assert.Contains(t, userText, "cch=00000") + }) + + t.Run("signing is deterministic", func(t *testing.T) { + body := []byte(`{"system":[{"type":"text","text":"x-anthropic-billing-header: cc_version=2.1.63; cc_entrypoint=cli; cch=00000;"}],"messages":[{"role":"user","content":"hi"}]}`) + r1 := signBillingHeaderCCH(body) + body2 := []byte(`{"system":[{"type":"text","text":"x-anthropic-billing-header: cc_version=2.1.63; cc_entrypoint=cli; cch=00000;"}],"messages":[{"role":"user","content":"hi"}]}`) + r2 := signBillingHeaderCCH(body2) + assert.Equal(t, string(r1), string(r2)) + }) + + t.Run("matches reference algorithm", func(t *testing.T) { + // Verify: signBillingHeaderCCH(body) produces cch = xxHash64(body_with_placeholder, seed) & 0xFFFFF + body := []byte(`{"system":[{"type":"text","text":"x-anthropic-billing-header: cc_version=2.1.63.a43; cc_entrypoint=cli; cch=00000;"}],"messages":[{"role":"user","content":[{"type":"text","text":"hello"}]}]}`) + expectedCCH := fmt.Sprintf("%05x", xxHash64Seeded(body, cchSeed)&0xFFFFF) + + result := signBillingHeaderCCH(body) + billingText := gjson.GetBytes(result, "system.0.text").String() + assert.Contains(t, billingText, "cch="+expectedCCH+";") + }) +} + +func TestXXHash64Seeded(t *testing.T) { + t.Run("matches cespare/xxhash for seed 0", func(t *testing.T) { + inputs := []string{"", "a", "hello world", "The quick brown fox jumps over the lazy dog"} + for _, s := range inputs { + data := []byte(s) + expected := xxhash.Sum64(data) + got := xxHash64Seeded(data, 0) + assert.Equal(t, expected, got, "mismatch for input %q", s) + } + }) + + t.Run("large input matches cespare", func(t *testing.T) { + data := make([]byte, 256) + for i := range data { + data[i] = byte(i) + } + expected := xxhash.Sum64(data) + got := xxHash64Seeded(data, 0) + assert.Equal(t, expected, got) + }) + + t.Run("deterministic with custom seed", func(t *testing.T) { + data := []byte("hello world") + h1 := xxHash64Seeded(data, cchSeed) + h2 := xxHash64Seeded(data, cchSeed) + assert.Equal(t, h1, h2) + }) + + t.Run("different seeds produce different results", func(t *testing.T) { + data := []byte("test data for hashing") + h1 := xxHash64Seeded(data, 0) + h2 := xxHash64Seeded(data, cchSeed) + assert.NotEqual(t, h1, h2) + }) +} diff --git a/backend/internal/service/gateway_prompt_test.go b/backend/internal/service/gateway_prompt_test.go index 356536b063..e27e18aaa7 100644 --- a/backend/internal/service/gateway_prompt_test.go +++ b/backend/internal/service/gateway_prompt_test.go @@ -278,3 +278,148 @@ func TestInjectClaudeCodePrompt(t *testing.T) { }) } } + +func TestRewriteSystemForNonClaudeCode(t *testing.T) { + tests := []struct { + name string + body string + system any + wantSystemText string // system array 第一个 block 的 text + wantMessagesLen int // messages 数组长度 + wantFirstMsgRole string // 第一条消息的 role + wantFirstMsgText string // 第一条消息的 content[0].text + wantAckMsgText string // 第二条消息的 content[0].text + }{ + { + name: "nil system - no messages injected", + body: `{"model":"claude-3","messages":[{"role":"user","content":"hello"}]}`, + system: nil, + wantSystemText: claudeCodeSystemPrompt, + wantMessagesLen: 1, // 原始 1 条消息,不注入 + }, + { + name: "empty string system - no messages injected", + body: `{"model":"claude-3","messages":[{"role":"user","content":"hello"}]}`, + system: "", + wantSystemText: claudeCodeSystemPrompt, + wantMessagesLen: 1, + }, + { + name: "custom string system - migrated to messages", + body: `{"model":"claude-3","messages":[{"role":"user","content":"hello"}]}`, + system: "You are a personal assistant running inside OpenClaw.", + wantSystemText: claudeCodeSystemPrompt, + wantMessagesLen: 3, // instruction + ack + original + wantFirstMsgRole: "user", + wantFirstMsgText: "[System Instructions]\nYou are a personal assistant running inside OpenClaw.", + wantAckMsgText: "Understood. I will follow these instructions.", + }, + { + name: "system equals Claude Code prompt - no messages injected", + body: `{"model":"claude-3","messages":[{"role":"user","content":"hello"}]}`, + system: claudeCodeSystemPrompt, + wantSystemText: claudeCodeSystemPrompt, + wantMessagesLen: 1, + }, + { + name: "array system with custom blocks - text joined and migrated", + body: `{"model":"claude-3","messages":[{"role":"user","content":"hello"}]}`, + system: []any{ + map[string]any{"type": "text", "text": "First instruction"}, + map[string]any{"type": "text", "text": "Second instruction"}, + }, + wantSystemText: claudeCodeSystemPrompt, + wantMessagesLen: 3, + wantFirstMsgRole: "user", + wantFirstMsgText: "[System Instructions]\nFirst instruction\n\nSecond instruction", + wantAckMsgText: "Understood. I will follow these instructions.", + }, + { + name: "empty array system - no messages injected", + body: `{"model":"claude-3","messages":[{"role":"user","content":"hello"}]}`, + system: []any{}, + wantSystemText: claudeCodeSystemPrompt, + wantMessagesLen: 1, + }, + { + name: "json.RawMessage string system", + body: `{"model":"claude-3","system":"Custom prompt","messages":[{"role":"user","content":"hello"}]}`, + system: json.RawMessage(`"Custom prompt"`), + wantSystemText: claudeCodeSystemPrompt, + wantMessagesLen: 3, + wantFirstMsgRole: "user", + wantFirstMsgText: "[System Instructions]\nCustom prompt", + wantAckMsgText: "Understood. I will follow these instructions.", + }, + { + name: "json.RawMessage nil system", + body: `{"model":"claude-3","messages":[{"role":"user","content":"hello"}]}`, + system: json.RawMessage(nil), + wantSystemText: claudeCodeSystemPrompt, + wantMessagesLen: 1, + }, + { + name: "multiple original messages preserved", + body: `{"model":"claude-3","messages":[{"role":"user","content":"msg1"},{"role":"assistant","content":"resp1"},{"role":"user","content":"msg2"}]}`, + system: "Be helpful", + wantSystemText: claudeCodeSystemPrompt, + wantMessagesLen: 5, // 2 injected + 3 original + wantFirstMsgRole: "user", + wantFirstMsgText: "[System Instructions]\nBe helpful", + wantAckMsgText: "Understood. I will follow these instructions.", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + result := rewriteSystemForNonClaudeCode([]byte(tt.body), tt.system) + + var parsed map[string]any + err := json.Unmarshal(result, &parsed) + require.NoError(t, err) + + // system 应为 array 格式: [{type: "text", text: "...", cache_control: {type: "ephemeral"}}] + systemArr, ok := parsed["system"].([]any) + require.True(t, ok, "system should be an array, got %T", parsed["system"]) + require.Len(t, systemArr, 1, "system array should have exactly 1 block") + systemBlock, ok := systemArr[0].(map[string]any) + require.True(t, ok) + require.Equal(t, "text", systemBlock["type"]) + require.Equal(t, tt.wantSystemText, systemBlock["text"]) + cc, ok := systemBlock["cache_control"].(map[string]any) + require.True(t, ok, "system block should have cache_control") + require.Equal(t, "ephemeral", cc["type"]) + + // 检查 messages + messages, ok := parsed["messages"].([]any) + require.True(t, ok, "messages should be an array") + require.Len(t, messages, tt.wantMessagesLen) + + if tt.wantFirstMsgRole != "" && len(messages) >= 2 { + // 检查注入的 instruction 消息 + firstMsg, ok := messages[0].(map[string]any) + require.True(t, ok) + require.Equal(t, tt.wantFirstMsgRole, firstMsg["role"]) + + firstContent, ok := firstMsg["content"].([]any) + require.True(t, ok) + require.Len(t, firstContent, 1) + firstBlock, ok := firstContent[0].(map[string]any) + require.True(t, ok) + require.Equal(t, tt.wantFirstMsgText, firstBlock["text"]) + + // 检查注入的 ack 消息 + ackMsg, ok := messages[1].(map[string]any) + require.True(t, ok) + require.Equal(t, "assistant", ackMsg["role"]) + + ackContent, ok := ackMsg["content"].([]any) + require.True(t, ok) + require.Len(t, ackContent, 1) + ackBlock, ok := ackContent[0].(map[string]any) + require.True(t, ok) + require.Equal(t, tt.wantAckMsgText, ackBlock["text"]) + } + }) + } +} diff --git a/backend/internal/service/gateway_service.go b/backend/internal/service/gateway_service.go index e8e4343f8a..629bdbba16 100644 --- a/backend/internal/service/gateway_service.go +++ b/backend/internal/service/gateway_service.go @@ -1191,12 +1191,20 @@ func (s *GatewayService) SelectAccountForModelWithExclusions(ctx context.Context // anthropic/gemini 分组支持混合调度(包含启用了 mixed_scheduling 的 antigravity 账户) // 注意:强制平台模式不走混合调度 if (platform == PlatformAnthropic || platform == PlatformGemini) && !hasForcePlatform { - return s.selectAccountWithMixedScheduling(ctx, groupID, sessionHash, requestedModel, excludedIDs, platform) + account, err := s.selectAccountWithMixedScheduling(ctx, groupID, sessionHash, requestedModel, excludedIDs, platform) + if err != nil { + return nil, err + } + return s.hydrateSelectedAccount(ctx, account) } // antigravity 分组、强制平台模式或无分组使用单平台选择 // 注意:强制平台模式也必须遵守分组限制,不再回退到全平台查询 - return s.selectAccountForModelWithPlatform(ctx, groupID, sessionHash, requestedModel, excludedIDs, platform) + account, err := s.selectAccountForModelWithPlatform(ctx, groupID, sessionHash, requestedModel, excludedIDs, platform) + if err != nil { + return nil, err + } + return s.hydrateSelectedAccount(ctx, account) } // SelectAccountWithLoadAwareness selects account with load-awareness and wait plan. @@ -1272,11 +1280,7 @@ func (s *GatewayService) SelectAccountWithLoadAwareness(ctx context.Context, gro localExcluded[account.ID] = struct{}{} // 排除此账号 continue // 重新选择 } - return &AccountSelectionResult{ - Account: account, - Acquired: true, - ReleaseFunc: result.ReleaseFunc, - }, nil + return s.newSelectionResult(ctx, account, true, result.ReleaseFunc, nil) } // 对于等待计划的情况,也需要先检查会话限制 @@ -1288,26 +1292,20 @@ func (s *GatewayService) SelectAccountWithLoadAwareness(ctx context.Context, gro if stickyAccountID > 0 && stickyAccountID == account.ID && s.concurrencyService != nil { waitingCount, _ := s.concurrencyService.GetAccountWaitingCount(ctx, account.ID) if waitingCount < cfg.StickySessionMaxWaiting { - return &AccountSelectionResult{ - Account: account, - WaitPlan: &AccountWaitPlan{ - AccountID: account.ID, - MaxConcurrency: account.Concurrency, - Timeout: cfg.StickySessionWaitTimeout, - MaxWaiting: cfg.StickySessionMaxWaiting, - }, - }, nil + return s.newSelectionResult(ctx, account, false, nil, &AccountWaitPlan{ + AccountID: account.ID, + MaxConcurrency: account.Concurrency, + Timeout: cfg.StickySessionWaitTimeout, + MaxWaiting: cfg.StickySessionMaxWaiting, + }) } } - return &AccountSelectionResult{ - Account: account, - WaitPlan: &AccountWaitPlan{ - AccountID: account.ID, - MaxConcurrency: account.Concurrency, - Timeout: cfg.FallbackWaitTimeout, - MaxWaiting: cfg.FallbackMaxWaiting, - }, - }, nil + return s.newSelectionResult(ctx, account, false, nil, &AccountWaitPlan{ + AccountID: account.ID, + MaxConcurrency: account.Concurrency, + Timeout: cfg.FallbackWaitTimeout, + MaxWaiting: cfg.FallbackMaxWaiting, + }) } } @@ -1453,11 +1451,7 @@ func (s *GatewayService) SelectAccountWithLoadAwareness(ctx context.Context, gro if s.debugModelRoutingEnabled() { logger.LegacyPrintf("service.gateway", "[ModelRoutingDebug] routed sticky hit: group_id=%v model=%s session=%s account=%d", derefGroupID(groupID), requestedModel, shortSessionHash(sessionHash), stickyAccountID) } - return &AccountSelectionResult{ - Account: stickyAccount, - Acquired: true, - ReleaseFunc: result.ReleaseFunc, - }, nil + return s.newSelectionResult(ctx, stickyAccount, true, result.ReleaseFunc, nil) } } @@ -1568,11 +1562,7 @@ func (s *GatewayService) SelectAccountWithLoadAwareness(ctx context.Context, gro if s.debugModelRoutingEnabled() { logger.LegacyPrintf("service.gateway", "[ModelRoutingDebug] routed select: group_id=%v model=%s session=%s account=%d", derefGroupID(groupID), requestedModel, shortSessionHash(sessionHash), item.account.ID) } - return &AccountSelectionResult{ - Account: item.account, - Acquired: true, - ReleaseFunc: result.ReleaseFunc, - }, nil + return s.newSelectionResult(ctx, item.account, true, result.ReleaseFunc, nil) } } @@ -1585,15 +1575,12 @@ func (s *GatewayService) SelectAccountWithLoadAwareness(ctx context.Context, gro if s.debugModelRoutingEnabled() { logger.LegacyPrintf("service.gateway", "[ModelRoutingDebug] routed wait: group_id=%v model=%s session=%s account=%d", derefGroupID(groupID), requestedModel, shortSessionHash(sessionHash), item.account.ID) } - return &AccountSelectionResult{ - Account: item.account, - WaitPlan: &AccountWaitPlan{ - AccountID: item.account.ID, - MaxConcurrency: item.account.Concurrency, - Timeout: cfg.StickySessionWaitTimeout, - MaxWaiting: cfg.StickySessionMaxWaiting, - }, - }, nil + return s.newSelectionResult(ctx, item.account, false, nil, &AccountWaitPlan{ + AccountID: item.account.ID, + MaxConcurrency: item.account.Concurrency, + Timeout: cfg.StickySessionWaitTimeout, + MaxWaiting: cfg.StickySessionMaxWaiting, + }) } // 所有路由账号会话限制都已满,继续到 Layer 2 回退 } @@ -1627,11 +1614,10 @@ func (s *GatewayService) SelectAccountWithLoadAwareness(ctx context.Context, gro if !s.checkAndRegisterSession(ctx, account, sessionHash) { result.ReleaseFunc() // 释放槽位,继续到 Layer 2 } else { - return &AccountSelectionResult{ - Account: account, - Acquired: true, - ReleaseFunc: result.ReleaseFunc, - }, nil + if s.cache != nil { + _ = s.cache.RefreshSessionTTL(ctx, derefGroupID(groupID), sessionHash, stickySessionTTL) + } + return s.newSelectionResult(ctx, account, true, result.ReleaseFunc, nil) } } @@ -1641,15 +1627,12 @@ func (s *GatewayService) SelectAccountWithLoadAwareness(ctx context.Context, gro if !s.checkAndRegisterSession(ctx, account, sessionHash) { // 会话限制已满,继续到 Layer 2 } else { - return &AccountSelectionResult{ - Account: account, - WaitPlan: &AccountWaitPlan{ - AccountID: accountID, - MaxConcurrency: account.Concurrency, - Timeout: cfg.StickySessionWaitTimeout, - MaxWaiting: cfg.StickySessionMaxWaiting, - }, - }, nil + return s.newSelectionResult(ctx, account, false, nil, &AccountWaitPlan{ + AccountID: accountID, + MaxConcurrency: account.Concurrency, + Timeout: cfg.StickySessionWaitTimeout, + MaxWaiting: cfg.StickySessionMaxWaiting, + }) } } } @@ -1708,7 +1691,9 @@ func (s *GatewayService) SelectAccountWithLoadAwareness(ctx context.Context, gro loadMap, err := s.concurrencyService.GetAccountsLoadBatch(ctx, accountLoads) if err != nil { - if result, ok := s.tryAcquireByLegacyOrder(ctx, candidates, groupID, sessionHash, preferOAuth); ok { + if result, ok, legacyErr := s.tryAcquireByLegacyOrder(ctx, candidates, groupID, sessionHash, preferOAuth); legacyErr != nil { + return nil, legacyErr + } else if ok { return result, nil } } else { @@ -1747,11 +1732,7 @@ func (s *GatewayService) SelectAccountWithLoadAwareness(ctx context.Context, gro if sessionHash != "" && s.cache != nil { _ = s.cache.SetSessionAccountID(ctx, derefGroupID(groupID), sessionHash, selected.account.ID, stickySessionTTL) } - return &AccountSelectionResult{ - Account: selected.account, - Acquired: true, - ReleaseFunc: result.ReleaseFunc, - }, nil + return s.newSelectionResult(ctx, selected.account, true, result.ReleaseFunc, nil) } } @@ -1774,20 +1755,17 @@ func (s *GatewayService) SelectAccountWithLoadAwareness(ctx context.Context, gro if !s.checkAndRegisterSession(ctx, acc, sessionHash) { continue // 会话限制已满,尝试下一个账号 } - return &AccountSelectionResult{ - Account: acc, - WaitPlan: &AccountWaitPlan{ - AccountID: acc.ID, - MaxConcurrency: acc.Concurrency, - Timeout: cfg.FallbackWaitTimeout, - MaxWaiting: cfg.FallbackMaxWaiting, - }, - }, nil + return s.newSelectionResult(ctx, acc, false, nil, &AccountWaitPlan{ + AccountID: acc.ID, + MaxConcurrency: acc.Concurrency, + Timeout: cfg.FallbackWaitTimeout, + MaxWaiting: cfg.FallbackMaxWaiting, + }) } return nil, ErrNoAvailableAccounts } -func (s *GatewayService) tryAcquireByLegacyOrder(ctx context.Context, candidates []*Account, groupID *int64, sessionHash string, preferOAuth bool) (*AccountSelectionResult, bool) { +func (s *GatewayService) tryAcquireByLegacyOrder(ctx context.Context, candidates []*Account, groupID *int64, sessionHash string, preferOAuth bool) (*AccountSelectionResult, bool, error) { ordered := append([]*Account(nil), candidates...) sortAccountsByPriorityAndLastUsed(ordered, preferOAuth) @@ -1802,15 +1780,15 @@ func (s *GatewayService) tryAcquireByLegacyOrder(ctx context.Context, candidates if sessionHash != "" && s.cache != nil { _ = s.cache.SetSessionAccountID(ctx, derefGroupID(groupID), sessionHash, acc.ID, stickySessionTTL) } - return &AccountSelectionResult{ - Account: acc, - Acquired: true, - ReleaseFunc: result.ReleaseFunc, - }, true + selection, err := s.newSelectionResult(ctx, acc, true, result.ReleaseFunc, nil) + if err != nil { + return nil, false, err + } + return selection, true, nil } } - return nil, false + return nil, false, nil } func (s *GatewayService) schedulingConfig() config.GatewaySchedulingConfig { @@ -2425,6 +2403,33 @@ func (s *GatewayService) getSchedulableAccount(ctx context.Context, accountID in return s.accountRepo.GetByID(ctx, accountID) } +func (s *GatewayService) hydrateSelectedAccount(ctx context.Context, account *Account) (*Account, error) { + if account == nil || s.schedulerSnapshot == nil { + return account, nil + } + hydrated, err := s.schedulerSnapshot.GetAccount(ctx, account.ID) + if err != nil { + return nil, err + } + if hydrated == nil { + return nil, fmt.Errorf("selected gateway account %d not found during hydration", account.ID) + } + return hydrated, nil +} + +func (s *GatewayService) newSelectionResult(ctx context.Context, account *Account, acquired bool, release func(), waitPlan *AccountWaitPlan) (*AccountSelectionResult, error) { + hydrated, err := s.hydrateSelectedAccount(ctx, account) + if err != nil { + return nil, err + } + return &AccountSelectionResult{ + Account: hydrated, + Acquired: acquired, + ReleaseFunc: release, + WaitPlan: waitPlan, + }, nil +} + // filterByMinPriority 过滤出优先级最小的账号集合 func filterByMinPriority(accounts []accountWithLoad) []accountWithLoad { if len(accounts) == 0 { @@ -3712,6 +3717,86 @@ func injectClaudeCodePrompt(body []byte, system any) []byte { return result } +// rewriteSystemForNonClaudeCode 将非 Claude Code 客户端的 system prompt 迁移至 messages, +// system 字段仅保留 Claude Code 标识提示词。 +// Anthropic 基于 system 参数内容检测第三方应用,仅前置追加 Claude Code 提示词 +// 无法通过检测,因为后续内容仍为非 Claude Code 格式。 +// 策略:将原始 system prompt 提取并注入为 user/assistant 消息对,system 仅保留 Claude Code 标识。 +func rewriteSystemForNonClaudeCode(body []byte, system any) []byte { + system = normalizeSystemParam(system) + + // 1. 提取原始 system prompt 文本 + var originalSystemText string + switch v := system.(type) { + case string: + originalSystemText = strings.TrimSpace(v) + case []any: + var parts []string + for _, item := range v { + if m, ok := item.(map[string]any); ok { + if text, ok := m["text"].(string); ok && strings.TrimSpace(text) != "" { + parts = append(parts, text) + } + } + } + originalSystemText = strings.Join(parts, "\n\n") + } + + // 2. 将 system 替换为 Claude Code 标准提示词(array 格式,与真实 Claude Code 一致) + // 真实 Claude Code 始终以 [{type: "text", text: "...", cache_control: {type: "ephemeral"}}] 发送 system。 + // 使用 string 格式会被 Anthropic 检测为第三方应用。 + claudeCodeSystemBlock := []map[string]any{ + { + "type": "text", + "text": claudeCodeSystemPrompt, + "cache_control": map[string]string{"type": "ephemeral"}, + }, + } + out, ok := setJSONValueBytes(body, "system", claudeCodeSystemBlock) + if !ok { + logger.LegacyPrintf("service.gateway", "Warning: failed to set Claude Code system prompt") + return body + } + + // 3. 将原始 system prompt 作为 user/assistant 消息对注入到 messages 开头 + // 模型仍通过 messages 接收完整指令,保留客户端功能 + ccPromptTrimmed := strings.TrimSpace(claudeCodeSystemPrompt) + if originalSystemText != "" && originalSystemText != ccPromptTrimmed && !hasClaudeCodePrefix(originalSystemText) { + instrMsg, err1 := json.Marshal(map[string]any{ + "role": "user", + "content": []map[string]any{ + {"type": "text", "text": "[System Instructions]\n" + originalSystemText}, + }, + }) + ackMsg, err2 := json.Marshal(map[string]any{ + "role": "assistant", + "content": []map[string]any{ + {"type": "text", "text": "Understood. I will follow these instructions."}, + }, + }) + if err1 != nil || err2 != nil { + logger.LegacyPrintf("service.gateway", "Warning: failed to marshal system-to-messages injection") + return out + } + + // 重建 messages 数组:[instruction, ack, ...originalMessages] + items := [][]byte{instrMsg, ackMsg} + messagesResult := gjson.GetBytes(out, "messages") + if messagesResult.IsArray() { + messagesResult.ForEach(func(_, msg gjson.Result) bool { + items = append(items, []byte(msg.Raw)) + return true + }) + } + + if next, setOk := setJSONRawBytes(out, "messages", buildJSONArrayRaw(items)); setOk { + out = next + } + } + + return out +} + type cacheControlPath struct { path string log string @@ -3873,7 +3958,7 @@ func (s *GatewayService) Forward(ctx context.Context, c *gin.Context, account *A // Beta policy: evaluate once; block check + cache filter set for buildUpstreamRequest. // Always overwrite the cache to prevent stale values from a previous retry with a different account. if account.Platform == PlatformAnthropic && c != nil { - policy := s.evaluateBetaPolicy(ctx, c.GetHeader("anthropic-beta"), account) + policy := s.evaluateBetaPolicy(ctx, c.GetHeader("anthropic-beta"), account, parsed.Model) if policy.blockErr != nil { return nil, policy.blockErr } @@ -3903,19 +3988,24 @@ func (s *GatewayService) Forward(ctx context.Context, c *gin.Context, account *A shouldMimicClaudeCode := account.IsOAuth() && !isClaudeCode if shouldMimicClaudeCode { - // 智能注入 Claude Code 系统提示词(仅 OAuth/SetupToken 账号需要) + // 非 Claude Code 客户端:将 system 替换为 Claude Code 标识,原始 system 迁移至 messages // 条件:1) OAuth/SetupToken 账号 2) 不是 Claude Code 客户端 3) 不是 Haiku 模型 4) system 中还没有 Claude Code 提示词 + systemRewritten := false if !strings.Contains(strings.ToLower(reqModel), "haiku") && !systemIncludesClaudeCodePrompt(parsed.System) { - body = injectClaudeCodePrompt(body, parsed.System) + body = rewriteSystemForNonClaudeCode(body, parsed.System) + systemRewritten = true } - normalizeOpts := claudeOAuthNormalizeOptions{stripSystemCacheControl: true} + // system 被重写时保留 CC prompt 的 cache_control: ephemeral(匹配真实 Claude Code 行为); + // 未重写时(haiku / 已含 CC 前缀)剥离客户端 cache_control,与原有行为一致。 + // 两种情况下 enforceCacheControlLimit 都会兜底处理上限。 + normalizeOpts := claudeOAuthNormalizeOptions{stripSystemCacheControl: !systemRewritten} if s.identityService != nil { fp, err := s.identityService.GetOrCreateFingerprint(ctx, account.ID, c.Request.Header) if err == nil && fp != nil { // metadata 透传开启时跳过 metadata 注入 - _, mimicMPT := s.settingService.GetGatewayForwardingSettings(ctx) + _, mimicMPT, _ := s.settingService.GetGatewayForwardingSettings(ctx) if !mimicMPT { if metadataUserID := s.buildOAuthMetadataUserID(parsed, account, fp); metadataUserID != "" { normalizeOpts.injectMetadata = true @@ -5461,9 +5551,9 @@ func (s *GatewayService) buildUpstreamRequest(ctx context.Context, c *gin.Contex // OAuth账号:应用统一指纹和metadata重写(受设置开关控制) var fingerprint *Fingerprint - enableFP, enableMPT := true, false + enableFP, enableMPT, enableCCH := true, false, false if s.settingService != nil { - enableFP, enableMPT = s.settingService.GetGatewayForwardingSettings(ctx) + enableFP, enableMPT, enableCCH = s.settingService.GetGatewayForwardingSettings(ctx) } if account.IsOAuth() && s.identityService != nil { // 1. 获取或创建指纹(包含随机生成的ClientID) @@ -5490,6 +5580,15 @@ func (s *GatewayService) buildUpstreamRequest(ctx context.Context, c *gin.Contex } } + // 同步 billing header cc_version 与实际发送的 User-Agent 版本 + if fingerprint != nil { + body = syncBillingHeaderVersion(body, fingerprint.UserAgent) + } + // CCH 签名:将 cch=00000 占位符替换为 xxHash64 签名(需在所有 body 修改之后) + if enableCCH { + body = signBillingHeaderCCH(body) + } + req, err := http.NewRequestWithContext(ctx, "POST", targetURL, bytes.NewReader(body)) if err != nil { return nil, err @@ -5530,9 +5629,8 @@ func (s *GatewayService) buildUpstreamRequest(ctx context.Context, c *gin.Contex } // Build effective drop set: merge static defaults with dynamic beta policy filter rules - policyFilterSet := s.getBetaPolicyFilterSet(ctx, c, account) + policyFilterSet := s.getBetaPolicyFilterSet(ctx, c, account, modelID) effectiveDropSet := mergeDropSets(policyFilterSet) - effectiveDropWithClaudeCodeSet := mergeDropSets(policyFilterSet, claude.BetaClaudeCode) // 处理 anthropic-beta header(OAuth 账号需要包含 oauth beta) if tokenType == "oauth" { @@ -5543,11 +5641,16 @@ func (s *GatewayService) buildUpstreamRequest(ctx context.Context, c *gin.Contex applyClaudeCodeMimicHeaders(req, reqStream) incomingBeta := getHeaderRaw(req.Header, "anthropic-beta") - // Match real Claude CLI traffic (per mitmproxy reports): - // messages requests typically use only oauth + interleaved-thinking. - // Also drop claude-code beta if a downstream client added it. + // Claude Code OAuth credentials are scoped to Claude Code. + // Non-haiku models MUST include claude-code beta for Anthropic to recognize + // this as a legitimate Claude Code request; without it, the request is + // rejected as third-party ("out of extra usage"). + // Haiku models are exempt from third-party detection and don't need it. requiredBetas := []string{claude.BetaOAuth, claude.BetaInterleavedThinking} - setHeaderRaw(req.Header, "anthropic-beta", mergeAnthropicBetaDropping(requiredBetas, incomingBeta, effectiveDropWithClaudeCodeSet)) + if !strings.Contains(strings.ToLower(modelID), "haiku") { + requiredBetas = []string{claude.BetaClaudeCode, claude.BetaOAuth, claude.BetaInterleavedThinking} + } + setHeaderRaw(req.Header, "anthropic-beta", mergeAnthropicBetaDropping(requiredBetas, incomingBeta, effectiveDropSet)) } else { // Claude Code 客户端:尽量透传原始 header,仅补齐 oauth beta clientBetaHeader := getHeaderRaw(req.Header, "anthropic-beta") @@ -5770,7 +5873,7 @@ type betaPolicyResult struct { } // evaluateBetaPolicy loads settings once and evaluates all rules against the given request. -func (s *GatewayService) evaluateBetaPolicy(ctx context.Context, betaHeader string, account *Account) betaPolicyResult { +func (s *GatewayService) evaluateBetaPolicy(ctx context.Context, betaHeader string, account *Account, model string) betaPolicyResult { if s.settingService == nil { return betaPolicyResult{} } @@ -5785,10 +5888,11 @@ func (s *GatewayService) evaluateBetaPolicy(ctx context.Context, betaHeader stri if !betaPolicyScopeMatches(rule.Scope, isOAuth, isBedrock) { continue } - switch rule.Action { + effectiveAction, effectiveErrMsg := resolveRuleAction(rule, model) + switch effectiveAction { case BetaPolicyActionBlock: if result.blockErr == nil && betaHeader != "" && containsBetaToken(betaHeader, rule.BetaToken) { - msg := rule.ErrorMessage + msg := effectiveErrMsg if msg == "" { msg = "beta feature " + rule.BetaToken + " is not allowed" } @@ -5830,7 +5934,7 @@ const betaPolicyFilterSetKey = "betaPolicyFilterSet" // In the /v1/messages path, Forward() evaluates the policy first and caches the result; // buildUpstreamRequest reuses it (zero extra DB calls). In the count_tokens path, this // evaluates on demand (one DB call). -func (s *GatewayService) getBetaPolicyFilterSet(ctx context.Context, c *gin.Context, account *Account) map[string]struct{} { +func (s *GatewayService) getBetaPolicyFilterSet(ctx context.Context, c *gin.Context, account *Account, model string) map[string]struct{} { if c != nil { if v, ok := c.Get(betaPolicyFilterSetKey); ok { if fs, ok := v.(map[string]struct{}); ok { @@ -5838,7 +5942,7 @@ func (s *GatewayService) getBetaPolicyFilterSet(ctx context.Context, c *gin.Cont } } } - return s.evaluateBetaPolicy(ctx, "", account).filterSet + return s.evaluateBetaPolicy(ctx, "", account, model).filterSet } // betaPolicyScopeMatches checks whether a rule's scope matches the current account type. @@ -5857,6 +5961,33 @@ func betaPolicyScopeMatches(scope string, isOAuth bool, isBedrock bool) bool { } } +// matchModelWhitelist checks if a model matches any pattern in the whitelist. +// Reuses matchModelPattern from group.go which supports exact and wildcard prefix matching. +func matchModelWhitelist(model string, whitelist []string) bool { + for _, pattern := range whitelist { + if matchModelPattern(pattern, model) { + return true + } + } + return false +} + +// resolveRuleAction determines the effective action and error message for a rule given the request model. +// When ModelWhitelist is empty, the rule's primary Action/ErrorMessage applies unconditionally. +// When non-empty, Action applies to matching models; FallbackAction/FallbackErrorMessage applies to others. +func resolveRuleAction(rule BetaPolicyRule, model string) (action, errorMessage string) { + if len(rule.ModelWhitelist) == 0 { + return rule.Action, rule.ErrorMessage + } + if matchModelWhitelist(model, rule.ModelWhitelist) { + return rule.Action, rule.ErrorMessage + } + if rule.FallbackAction != "" { + return rule.FallbackAction, rule.FallbackErrorMessage + } + return BetaPolicyActionPass, "" // default fallback: pass (fail-open) +} + // droppedBetaSet returns claude.DroppedBetas as a set, with optional extra tokens. func droppedBetaSet(extra ...string) map[string]struct{} { m := make(map[string]struct{}, len(defaultDroppedBetasSet)+len(extra)) @@ -5903,7 +6034,7 @@ func (s *GatewayService) resolveBedrockBetaTokensForRequest( modelID string, ) ([]string, error) { // 1. 对原始 header 中的 beta token 做 block 检查(快速失败) - policy := s.evaluateBetaPolicy(ctx, betaHeader, account) + policy := s.evaluateBetaPolicy(ctx, betaHeader, account, modelID) if policy.blockErr != nil { return nil, policy.blockErr } @@ -5915,7 +6046,7 @@ func (s *GatewayService) resolveBedrockBetaTokensForRequest( // 例如:管理员 block 了 interleaved-thinking,客户端不在 header 中带该 token, // 但请求体中包含 thinking 字段 → autoInjectBedrockBetaTokens 会自动补齐 → // 如果不做此检查,block 规则会被绕过。 - if blockErr := s.checkBetaPolicyBlockForTokens(ctx, betaTokens, account); blockErr != nil { + if blockErr := s.checkBetaPolicyBlockForTokens(ctx, betaTokens, account, modelID); blockErr != nil { return nil, blockErr } @@ -5924,7 +6055,7 @@ func (s *GatewayService) resolveBedrockBetaTokensForRequest( // checkBetaPolicyBlockForTokens 检查 token 列表中是否有被管理员 block 规则命中的 token。 // 用于补充 evaluateBetaPolicy 对 header 的检查,覆盖 body 自动注入的 token。 -func (s *GatewayService) checkBetaPolicyBlockForTokens(ctx context.Context, tokens []string, account *Account) *BetaBlockedError { +func (s *GatewayService) checkBetaPolicyBlockForTokens(ctx context.Context, tokens []string, account *Account, model string) *BetaBlockedError { if s.settingService == nil || len(tokens) == 0 { return nil } @@ -5936,14 +6067,15 @@ func (s *GatewayService) checkBetaPolicyBlockForTokens(ctx context.Context, toke isBedrock := account.IsBedrock() tokenSet := buildBetaTokenSet(tokens) for _, rule := range settings.Rules { - if rule.Action != BetaPolicyActionBlock { + effectiveAction, effectiveErrMsg := resolveRuleAction(rule, model) + if effectiveAction != BetaPolicyActionBlock { continue } if !betaPolicyScopeMatches(rule.Scope, isOAuth, isBedrock) { continue } if _, present := tokenSet[rule.BetaToken]; present { - msg := rule.ErrorMessage + msg := effectiveErrMsg if msg == "" { msg = "beta feature " + rule.BetaToken + " is not allowed" } @@ -8355,9 +8487,9 @@ func (s *GatewayService) buildCountTokensRequest(ctx context.Context, c *gin.Con // OAuth 账号:应用统一指纹和重写 userID(受设置开关控制) // 如果启用了会话ID伪装,会在重写后替换 session 部分为固定值 - ctEnableFP, ctEnableMPT := true, false + ctEnableFP, ctEnableMPT, ctEnableCCH := true, false, false if s.settingService != nil { - ctEnableFP, ctEnableMPT = s.settingService.GetGatewayForwardingSettings(ctx) + ctEnableFP, ctEnableMPT, ctEnableCCH = s.settingService.GetGatewayForwardingSettings(ctx) } var ctFingerprint *Fingerprint if account.IsOAuth() && s.identityService != nil { @@ -8375,6 +8507,14 @@ func (s *GatewayService) buildCountTokensRequest(ctx context.Context, c *gin.Con } } + // 同步 billing header cc_version 与实际发送的 User-Agent 版本 + if ctFingerprint != nil && ctEnableFP { + body = syncBillingHeaderVersion(body, ctFingerprint.UserAgent) + } + if ctEnableCCH { + body = signBillingHeaderCCH(body) + } + req, err := http.NewRequestWithContext(ctx, "POST", targetURL, bytes.NewReader(body)) if err != nil { return nil, err @@ -8415,7 +8555,7 @@ func (s *GatewayService) buildCountTokensRequest(ctx context.Context, c *gin.Con } // Build effective drop set for count_tokens: merge static defaults with dynamic beta policy filter rules - ctEffectiveDropSet := mergeDropSets(s.getBetaPolicyFilterSet(ctx, c, account)) + ctEffectiveDropSet := mergeDropSets(s.getBetaPolicyFilterSet(ctx, c, account, modelID)) // OAuth 账号:处理 anthropic-beta header if tokenType == "oauth" { diff --git a/backend/internal/service/gemini_messages_compat_service.go b/backend/internal/service/gemini_messages_compat_service.go index b35ebce5c9..5a9490f36d 100644 --- a/backend/internal/service/gemini_messages_compat_service.go +++ b/backend/internal/service/gemini_messages_compat_service.go @@ -137,7 +137,7 @@ func (s *GeminiMessagesCompatService) SelectAccountForModelWithExclusions(ctx co _ = s.cache.SetSessionAccountID(ctx, derefGroupID(groupID), cacheKey, selected.ID, geminiStickySessionTTL) } - return selected, nil + return s.hydrateSelectedAccount(ctx, selected) } // resolvePlatformAndSchedulingMode 解析目标平台和调度模式。 @@ -416,6 +416,20 @@ func (s *GeminiMessagesCompatService) getSchedulableAccount(ctx context.Context, return s.accountRepo.GetByID(ctx, accountID) } +func (s *GeminiMessagesCompatService) hydrateSelectedAccount(ctx context.Context, account *Account) (*Account, error) { + if account == nil || s.schedulerSnapshot == nil { + return account, nil + } + hydrated, err := s.schedulerSnapshot.GetAccount(ctx, account.ID) + if err != nil { + return nil, err + } + if hydrated == nil { + return nil, fmt.Errorf("selected gemini account %d not found during hydration", account.ID) + } + return hydrated, nil +} + func (s *GeminiMessagesCompatService) listSchedulableAccountsOnce(ctx context.Context, groupID *int64, platform string, hasForcePlatform bool) ([]Account, error) { if s.schedulerSnapshot != nil { accounts, _, err := s.schedulerSnapshot.ListSchedulableAccounts(ctx, groupID, platform, hasForcePlatform) @@ -546,7 +560,7 @@ func (s *GeminiMessagesCompatService) SelectAccountForAIStudioEndpoints(ctx cont if selected == nil { return nil, errors.New("no available Gemini accounts") } - return selected, nil + return s.hydrateSelectedAccount(ctx, selected) } func (s *GeminiMessagesCompatService) Forward(ctx context.Context, c *gin.Context, account *Account, body []byte) (*ForwardResult, error) { @@ -612,7 +626,8 @@ func (s *GeminiMessagesCompatService) Forward(ctx context.Context, c *gin.Contex fullURL += "?alt=sse" } - upstreamReq, err := http.NewRequestWithContext(ctx, http.MethodPost, fullURL, bytes.NewReader(geminiReq)) + restGeminiReq := normalizeGeminiRequestForAIStudio(geminiReq) + upstreamReq, err := http.NewRequestWithContext(ctx, http.MethodPost, fullURL, bytes.NewReader(restGeminiReq)) if err != nil { return nil, "", err } @@ -685,7 +700,8 @@ func (s *GeminiMessagesCompatService) Forward(ctx context.Context, c *gin.Contex fullURL += "?alt=sse" } - upstreamReq, err := http.NewRequestWithContext(ctx, http.MethodPost, fullURL, bytes.NewReader(geminiReq)) + restGeminiReq := normalizeGeminiRequestForAIStudio(geminiReq) + upstreamReq, err := http.NewRequestWithContext(ctx, http.MethodPost, fullURL, bytes.NewReader(restGeminiReq)) if err != nil { return nil, "", err } @@ -3184,12 +3200,17 @@ func convertClaudeToolsToGeminiTools(tools any) []any { return nil } + hasWebSearch := false funcDecls := make([]any, 0, len(arr)) for _, t := range arr { tm, ok := t.(map[string]any) if !ok { continue } + if isClaudeWebSearchToolMap(tm) { + hasWebSearch = true + continue + } var name, desc string var params any @@ -3233,13 +3254,75 @@ func convertClaudeToolsToGeminiTools(tools any) []any { }) } - if len(funcDecls) == 0 { + out := make([]any, 0, 2) + if len(funcDecls) > 0 { + out = append(out, map[string]any{ + "functionDeclarations": funcDecls, + }) + } + if hasWebSearch { + out = append(out, map[string]any{ + "googleSearch": map[string]any{}, + }) + } + if len(out) == 0 { return nil } - return []any{ - map[string]any{ - "functionDeclarations": funcDecls, - }, + return out +} + +func normalizeGeminiRequestForAIStudio(body []byte) []byte { + var payload map[string]any + if err := json.Unmarshal(body, &payload); err != nil { + return body + } + + tools, ok := payload["tools"].([]any) + if !ok || len(tools) == 0 { + return body + } + + modified := false + for _, rawTool := range tools { + tool, ok := rawTool.(map[string]any) + if !ok { + continue + } + googleSearch, ok := tool["googleSearch"] + if !ok { + continue + } + if _, exists := tool["google_search"]; exists { + continue + } + tool["google_search"] = googleSearch + delete(tool, "googleSearch") + modified = true + } + + if !modified { + return body + } + + normalized, err := json.Marshal(payload) + if err != nil { + return body + } + return normalized +} + +func isClaudeWebSearchToolMap(tool map[string]any) bool { + toolType, _ := tool["type"].(string) + if strings.HasPrefix(toolType, "web_search") || toolType == "google_search" { + return true + } + + name, _ := tool["name"].(string) + switch strings.TrimSpace(name) { + case "web_search", "google_search", "web_search_20250305": + return true + default: + return false } } diff --git a/backend/internal/service/gemini_messages_compat_service_test.go b/backend/internal/service/gemini_messages_compat_service_test.go index f659f0e60a..c2adf45ded 100644 --- a/backend/internal/service/gemini_messages_compat_service_test.go +++ b/backend/internal/service/gemini_messages_compat_service_test.go @@ -164,6 +164,35 @@ func TestConvertClaudeToolsToGeminiTools_CustomType(t *testing.T) { } } +func TestConvertClaudeToolsToGeminiTools_PreservesWebSearchAlongsideFunctions(t *testing.T) { + tools := []any{ + map[string]any{ + "name": "get_weather", + "description": "Get weather info", + "input_schema": map[string]any{"type": "object"}, + }, + map[string]any{ + "type": "web_search_20250305", + "name": "web_search", + }, + } + + result := convertClaudeToolsToGeminiTools(tools) + require.Len(t, result, 2) + + functionDecl, ok := result[0].(map[string]any) + require.True(t, ok) + funcDecls, ok := functionDecl["functionDeclarations"].([]any) + require.True(t, ok) + require.Len(t, funcDecls, 1) + + searchDecl, ok := result[1].(map[string]any) + require.True(t, ok) + googleSearch, ok := searchDecl["googleSearch"].(map[string]any) + require.True(t, ok) + require.Empty(t, googleSearch) +} + func TestGeminiHandleNativeNonStreamingResponse_DebugDisabledDoesNotEmitHeaderLogs(t *testing.T) { gin.SetMode(gin.TestMode) logSink, restore := captureStructuredLog(t) @@ -232,6 +261,53 @@ func TestGeminiMessagesCompatServiceForward_PreservesRequestedModelAndMappedUpst require.Contains(t, httpStub.lastReq.URL.String(), "/models/claude-sonnet-4-20250514:") } +func TestGeminiMessagesCompatServiceForward_NormalizesWebSearchToolForAIStudio(t *testing.T) { + gin.SetMode(gin.TestMode) + w := httptest.NewRecorder() + c, _ := gin.CreateTestContext(w) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages", nil) + + httpStub := &geminiCompatHTTPUpstreamStub{ + response: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"x-request-id": []string{"gemini-req-2"}}, + Body: io.NopCloser(strings.NewReader(`{"candidates":[{"content":{"parts":[{"text":"hello"}]}}],"usageMetadata":{"promptTokenCount":10,"candidatesTokenCount":5}}`)), + }, + } + svc := &GeminiMessagesCompatService{httpUpstream: httpStub, cfg: &config.Config{}} + account := &Account{ + ID: 1, + Type: AccountTypeAPIKey, + Credentials: map[string]any{ + "api_key": "test-key", + }, + } + body := []byte(`{"model":"claude-sonnet-4","max_tokens":16,"messages":[{"role":"user","content":"hello"}],"tools":[{"name":"get_weather","description":"Get weather info","input_schema":{"type":"object"}},{"type":"web_search_20250305","name":"web_search"}]}`) + + result, err := svc.Forward(context.Background(), c, account, body) + require.NoError(t, err) + require.NotNil(t, result) + require.NotNil(t, httpStub.lastReq) + + postedBody, err := io.ReadAll(httpStub.lastReq.Body) + require.NoError(t, err) + + var posted map[string]any + require.NoError(t, json.Unmarshal(postedBody, &posted)) + tools, ok := posted["tools"].([]any) + require.True(t, ok) + require.Len(t, tools, 2) + + searchTool, ok := tools[1].(map[string]any) + require.True(t, ok) + _, hasSnake := searchTool["google_search"] + _, hasCamel := searchTool["googleSearch"] + require.True(t, hasSnake) + require.False(t, hasCamel) + _, hasFuncDecl := searchTool["functionDeclarations"] + require.False(t, hasFuncDecl) +} + func TestConvertClaudeMessagesToGeminiGenerateContent_AddsThoughtSignatureForToolUse(t *testing.T) { claudeReq := map[string]any{ "model": "claude-haiku-4-5-20251001", diff --git a/backend/internal/service/group.go b/backend/internal/service/group.go index d59af9e1c0..1226261357 100644 --- a/backend/internal/service/group.go +++ b/backend/internal/service/group.go @@ -3,8 +3,12 @@ package service import ( "strings" "time" + + "github.com/Wei-Shaw/sub2api/internal/domain" ) +type OpenAIMessagesDispatchModelConfig = domain.OpenAIMessagesDispatchModelConfig + type Group struct { ID int64 Name string @@ -49,10 +53,11 @@ type Group struct { SortOrder int // OpenAI Messages 调度配置(仅 openai 平台使用) - AllowMessagesDispatch bool - RequireOAuthOnly bool // 仅允许非 apikey 类型账号关联(OpenAI/Antigravity/Anthropic/Gemini) - RequirePrivacySet bool // 调度时仅允许 privacy 已成功设置的账号(OpenAI/Antigravity/Anthropic/Gemini) - DefaultMappedModel string + AllowMessagesDispatch bool + RequireOAuthOnly bool // 仅允许非 apikey 类型账号关联(OpenAI/Antigravity/Anthropic/Gemini) + RequirePrivacySet bool // 调度时仅允许 privacy 已成功设置的账号(OpenAI/Antigravity/Anthropic/Gemini) + DefaultMappedModel string + MessagesDispatchModelConfig OpenAIMessagesDispatchModelConfig CreatedAt time.Time UpdatedAt time.Time diff --git a/backend/internal/service/oauth_refresh_api.go b/backend/internal/service/oauth_refresh_api.go index 5dbba63858..571e9ecdb7 100644 --- a/backend/internal/service/oauth_refresh_api.go +++ b/backend/internal/service/oauth_refresh_api.go @@ -5,6 +5,8 @@ import ( "fmt" "log/slog" "strconv" + "strings" + "sync" "time" ) @@ -17,7 +19,7 @@ type OAuthRefreshExecutor interface { CacheKey(account *Account) string } -const refreshLockTTL = 30 * time.Second +const defaultRefreshLockTTL = 60 * time.Second // OAuthRefreshResult 统一刷新结果 type OAuthRefreshResult struct { @@ -28,20 +30,39 @@ type OAuthRefreshResult struct { } // OAuthRefreshAPI 统一的 OAuth Token 刷新入口 -// 封装分布式锁、DB 重读、已刷新检查等通用逻辑 +// 封装分布式锁、进程内互斥锁、DB 重读、已刷新检查、竞争恢复等通用逻辑 type OAuthRefreshAPI struct { accountRepo AccountRepository - tokenCache GeminiTokenCache // 可选,nil = 无锁 + tokenCache GeminiTokenCache // 可选,nil = 无分布式锁 + lockTTL time.Duration + localLocks sync.Map // key: cacheKey string -> value: *sync.Mutex } // NewOAuthRefreshAPI 创建统一刷新 API -func NewOAuthRefreshAPI(accountRepo AccountRepository, tokenCache GeminiTokenCache) *OAuthRefreshAPI { +// 可选传入 lockTTL 覆盖默认的 60s 分布式锁 TTL +func NewOAuthRefreshAPI(accountRepo AccountRepository, tokenCache GeminiTokenCache, lockTTL ...time.Duration) *OAuthRefreshAPI { + ttl := defaultRefreshLockTTL + if len(lockTTL) > 0 && lockTTL[0] > 0 { + ttl = lockTTL[0] + } return &OAuthRefreshAPI{ accountRepo: accountRepo, tokenCache: tokenCache, + lockTTL: ttl, } } +// getLocalLock 返回指定 cacheKey 的进程内互斥锁 +func (api *OAuthRefreshAPI) getLocalLock(cacheKey string) *sync.Mutex { + actual, _ := api.localLocks.LoadOrStore(cacheKey, &sync.Mutex{}) + mu, ok := actual.(*sync.Mutex) + if !ok { + mu = &sync.Mutex{} + api.localLocks.Store(cacheKey, mu) + } + return mu +} + // RefreshIfNeeded 在分布式锁保护下按需刷新 OAuth token // // 流程: @@ -59,12 +80,17 @@ func (api *OAuthRefreshAPI) RefreshIfNeeded( ) (*OAuthRefreshResult, error) { cacheKey := executor.CacheKey(account) + // 0. 获取进程内互斥锁(防止同一进程内的并发刷新竞争) + localMu := api.getLocalLock(cacheKey) + localMu.Lock() + defer localMu.Unlock() + // 1. 获取分布式锁 lockAcquired := false if api.tokenCache != nil { - acquired, lockErr := api.tokenCache.AcquireRefreshLock(ctx, cacheKey, refreshLockTTL) + acquired, lockErr := api.tokenCache.AcquireRefreshLock(ctx, cacheKey, api.lockTTL) if lockErr != nil { - // Redis 错误,降级为无锁刷新 + // Redis 错误,降级为无锁刷新(进程内互斥锁仍生效) slog.Warn("oauth_refresh_lock_failed_degraded", "account_id", account.ID, "cache_key", cacheKey, @@ -102,6 +128,19 @@ func (api *OAuthRefreshAPI) RefreshIfNeeded( // 4. 执行平台特定刷新逻辑 newCredentials, refreshErr := executor.Refresh(ctx, freshAccount) if refreshErr != nil { + // 竞争恢复:invalid_grant 可能是另一个 worker 已消费了旧 refresh_token + // 重新读取 DB,如果 refresh_token 已更新则说明是竞争,返回成功 + if isInvalidGrantError(refreshErr) { + if recoveredAccount, recovered := api.tryRecoverFromRefreshRace(ctx, freshAccount); recovered { + slog.Info("oauth_refresh_race_recovered", + "account_id", freshAccount.ID, + "platform", freshAccount.Platform, + ) + return &OAuthRefreshResult{ + Account: recoveredAccount, + }, nil + } + } return nil, refreshErr } @@ -126,6 +165,33 @@ func (api *OAuthRefreshAPI) RefreshIfNeeded( }, nil } +// isInvalidGrantError 检查错误是否为 invalid_grant +func isInvalidGrantError(err error) bool { + return err != nil && strings.Contains(strings.ToLower(err.Error()), "invalid_grant") +} + +// tryRecoverFromRefreshRace 在 invalid_grant 错误后尝试竞争恢复 +// 重新读取 DB,如果 refresh_token 已改变(说明另一个 worker 成功刷新),则返回更新后的 account +func (api *OAuthRefreshAPI) tryRecoverFromRefreshRace(ctx context.Context, usedAccount *Account) (*Account, bool) { + if api.accountRepo == nil { + return nil, false + } + reReadAccount, err := api.accountRepo.GetByID(ctx, usedAccount.ID) + if err != nil || reReadAccount == nil { + return nil, false + } + usedRT := usedAccount.GetCredential("refresh_token") + currentRT := reReadAccount.GetCredential("refresh_token") + if usedRT == "" || currentRT == "" { + return nil, false + } + // refresh_token 不同 → 另一个 worker 已成功刷新 + if usedRT != currentRT { + return reReadAccount, true + } + return nil, false +} + // MergeCredentials 将旧 credentials 中不存在于新 map 的字段保留到新 map 中 func MergeCredentials(oldCreds, newCreds map[string]any) map[string]any { if newCreds == nil { diff --git a/backend/internal/service/oauth_refresh_api_test.go b/backend/internal/service/oauth_refresh_api_test.go index c3b38ddf6e..4a60723b8b 100644 --- a/backend/internal/service/oauth_refresh_api_test.go +++ b/backend/internal/service/oauth_refresh_api_test.go @@ -5,6 +5,7 @@ package service import ( "context" "errors" + "sync" "testing" "time" @@ -385,6 +386,224 @@ func TestBuildClaudeAccountCredentials_Minimal(t *testing.T) { require.False(t, hasScope, "scope should not be set when empty") } +// refreshAPIAccountRepoWithRace supports returning a different account on subsequent GetByID calls +// to simulate race conditions where another worker has refreshed the token. +type refreshAPIAccountRepoWithRace struct { + refreshAPIAccountRepo + raceAccount *Account // returned on 2nd+ GetByID call + getByIDCalls int +} + +func (r *refreshAPIAccountRepoWithRace) GetByID(_ context.Context, _ int64) (*Account, error) { + r.getByIDCalls++ + if r.getByIDCalls > 1 && r.raceAccount != nil { + return r.raceAccount, nil + } + if r.getByIDErr != nil { + return nil, r.getByIDErr + } + return r.account, nil +} + +// ========== Race recovery tests ========== + +func TestRefreshIfNeeded_InvalidGrantRaceRecovered(t *testing.T) { + // Account with old refresh token + account := &Account{ + ID: 10, + Platform: PlatformAnthropic, + Type: AccountTypeOAuth, + Credentials: map[string]any{"refresh_token": "old-rt", "access_token": "old-at"}, + } + // After race, DB has new refresh token from another worker + racedAccount := &Account{ + ID: 10, + Platform: PlatformAnthropic, + Type: AccountTypeOAuth, + Credentials: map[string]any{"refresh_token": "new-rt", "access_token": "new-at"}, + } + repo := &refreshAPIAccountRepoWithRace{ + refreshAPIAccountRepo: refreshAPIAccountRepo{account: account}, + raceAccount: racedAccount, + } + cache := &refreshAPICacheStub{lockResult: true} + executor := &refreshAPIExecutorStub{ + needsRefresh: true, + err: errors.New("invalid_grant: refresh token not found or invalid"), + } + + api := NewOAuthRefreshAPI(repo, cache) + result, err := api.RefreshIfNeeded(context.Background(), account, executor, 3*time.Minute) + + require.NoError(t, err, "race-recovered invalid_grant should not return error") + require.False(t, result.Refreshed) + require.False(t, result.LockHeld) + require.NotNil(t, result.Account) + require.Equal(t, "new-rt", result.Account.GetCredential("refresh_token")) + require.Equal(t, 0, repo.updateCalls) // no DB update needed, another worker did it +} + +func TestRefreshIfNeeded_InvalidGrantGenuine(t *testing.T) { + // Account with revoked refresh token - DB still has the same token + account := &Account{ + ID: 11, + Platform: PlatformAnthropic, + Type: AccountTypeOAuth, + Credentials: map[string]any{"refresh_token": "revoked-rt", "access_token": "old-at"}, + } + repo := &refreshAPIAccountRepoWithRace{ + refreshAPIAccountRepo: refreshAPIAccountRepo{account: account}, + raceAccount: account, // same refresh_token on re-read + } + cache := &refreshAPICacheStub{lockResult: true} + executor := &refreshAPIExecutorStub{ + needsRefresh: true, + err: errors.New("invalid_grant: refresh token revoked"), + } + + api := NewOAuthRefreshAPI(repo, cache) + result, err := api.RefreshIfNeeded(context.Background(), account, executor, 3*time.Minute) + + require.Error(t, err, "genuine invalid_grant should propagate error") + require.Nil(t, result) + require.Contains(t, err.Error(), "invalid_grant") +} + +func TestRefreshIfNeeded_InvalidGrantDBRereadFailsOnRecovery(t *testing.T) { + account := &Account{ + ID: 12, + Platform: PlatformAnthropic, + Type: AccountTypeOAuth, + Credentials: map[string]any{"refresh_token": "old-rt"}, + } + repo := &refreshAPIAccountRepoWithRace{ + refreshAPIAccountRepo: refreshAPIAccountRepo{account: account}, + raceAccount: nil, // GetByID returns nil on recovery attempt + } + cache := &refreshAPICacheStub{lockResult: true} + executor := &refreshAPIExecutorStub{ + needsRefresh: true, + err: errors.New("invalid_grant"), + } + + api := NewOAuthRefreshAPI(repo, cache) + result, err := api.RefreshIfNeeded(context.Background(), account, executor, 3*time.Minute) + + require.Error(t, err, "should propagate error when recovery DB re-read fails") + require.Nil(t, result) +} + +func TestRefreshIfNeeded_LocalMutexSerializesConcurrent(t *testing.T) { + // Test that two goroutines for the same account are serialized by the local mutex. + // The first goroutine refreshes successfully; the second sees NeedsRefresh=false. + refreshed := &Account{ + ID: 20, + Platform: PlatformAnthropic, + Type: AccountTypeOAuth, + Credentials: map[string]any{"refresh_token": "new-rt", "access_token": "new-at"}, + } + callCount := 0 + repo := &refreshAPIAccountRepo{account: &Account{ + ID: 20, + Platform: PlatformAnthropic, + Type: AccountTypeOAuth, + Credentials: map[string]any{"refresh_token": "old-rt"}, + }} + + // After first refresh, NeedsRefresh should return false + // We simulate this by using an executor that decrements needsRefresh after first call + var mu sync.Mutex + dynamicExecutor := &dynamicRefreshExecutor{ + canRefresh: true, + cacheKey: "test:mutex:anthropic", + refreshFunc: func(_ context.Context, _ *Account) (map[string]any, error) { + mu.Lock() + callCount++ + mu.Unlock() + time.Sleep(50 * time.Millisecond) // slow refresh + return map[string]any{"access_token": "new-at"}, nil + }, + needsRefreshFunc: func() bool { + mu.Lock() + defer mu.Unlock() + return callCount == 0 // only first call needs refresh + }, + } + + _ = refreshed + + api := NewOAuthRefreshAPI(repo, nil) // no distributed lock, only local mutex + + var wg sync.WaitGroup + results := make([]*OAuthRefreshResult, 2) + errs := make([]error, 2) + + for i := 0; i < 2; i++ { + wg.Add(1) + go func(idx int) { + defer wg.Done() + results[idx], errs[idx] = api.RefreshIfNeeded(context.Background(), repo.account, dynamicExecutor, 3*time.Minute) + }(i) + } + wg.Wait() + + require.NoError(t, errs[0]) + require.NoError(t, errs[1]) + + // Only one goroutine should have actually called Refresh + mu.Lock() + require.Equal(t, 1, callCount, "only one refresh call should have been made") + mu.Unlock() +} + +// dynamicRefreshExecutor is a test helper with function-based NeedsRefresh and Refresh. +type dynamicRefreshExecutor struct { + canRefresh bool + cacheKey string + needsRefreshFunc func() bool + refreshFunc func(context.Context, *Account) (map[string]any, error) +} + +func (e *dynamicRefreshExecutor) CanRefresh(_ *Account) bool { return e.canRefresh } + +func (e *dynamicRefreshExecutor) NeedsRefresh(_ *Account, _ time.Duration) bool { + return e.needsRefreshFunc() +} + +func (e *dynamicRefreshExecutor) Refresh(ctx context.Context, account *Account) (map[string]any, error) { + return e.refreshFunc(ctx, account) +} + +func (e *dynamicRefreshExecutor) CacheKey(_ *Account) string { + return e.cacheKey +} + +// ========== NewOAuthRefreshAPI TTL tests ========== + +func TestNewOAuthRefreshAPI_DefaultTTL(t *testing.T) { + api := NewOAuthRefreshAPI(nil, nil) + require.Equal(t, defaultRefreshLockTTL, api.lockTTL) +} + +func TestNewOAuthRefreshAPI_CustomTTL(t *testing.T) { + api := NewOAuthRefreshAPI(nil, nil, 90*time.Second) + require.Equal(t, 90*time.Second, api.lockTTL) +} + +func TestNewOAuthRefreshAPI_ZeroTTLUsesDefault(t *testing.T) { + api := NewOAuthRefreshAPI(nil, nil, 0) + require.Equal(t, defaultRefreshLockTTL, api.lockTTL) +} + +// ========== isInvalidGrantError tests ========== + +func TestIsInvalidGrantError(t *testing.T) { + require.True(t, isInvalidGrantError(errors.New("invalid_grant: token revoked"))) + require.True(t, isInvalidGrantError(errors.New("INVALID_GRANT"))) + require.False(t, isInvalidGrantError(errors.New("invalid_client"))) + require.False(t, isInvalidGrantError(nil)) +} + // ========== BackgroundRefreshPolicy tests ========== func TestBackgroundRefreshPolicy_DefaultSkips(t *testing.T) { diff --git a/backend/internal/service/openai_codex_instructions_template.go b/backend/internal/service/openai_codex_instructions_template.go new file mode 100644 index 0000000000..5588c73cbb --- /dev/null +++ b/backend/internal/service/openai_codex_instructions_template.go @@ -0,0 +1,55 @@ +package service + +import ( + "bytes" + "fmt" + "strings" + "text/template" +) + +type forcedCodexInstructionsTemplateData struct { + ExistingInstructions string + OriginalModel string + NormalizedModel string + BillingModel string + UpstreamModel string +} + +func applyForcedCodexInstructionsTemplate( + reqBody map[string]any, + templateText string, + data forcedCodexInstructionsTemplateData, +) (bool, error) { + rendered, err := renderForcedCodexInstructionsTemplate(templateText, data) + if err != nil { + return false, err + } + if rendered == "" { + return false, nil + } + + existing, _ := reqBody["instructions"].(string) + if strings.TrimSpace(existing) == rendered { + return false, nil + } + + reqBody["instructions"] = rendered + return true, nil +} + +func renderForcedCodexInstructionsTemplate( + templateText string, + data forcedCodexInstructionsTemplateData, +) (string, error) { + tmpl, err := template.New("forced_codex_instructions").Option("missingkey=zero").Parse(templateText) + if err != nil { + return "", fmt.Errorf("parse forced codex instructions template: %w", err) + } + + var buf bytes.Buffer + if err := tmpl.Execute(&buf, data); err != nil { + return "", fmt.Errorf("render forced codex instructions template: %w", err) + } + + return strings.TrimSpace(buf.String()), nil +} diff --git a/backend/internal/service/openai_codex_transform.go b/backend/internal/service/openai_codex_transform.go index 21b4874eb3..4ec038e068 100644 --- a/backend/internal/service/openai_codex_transform.go +++ b/backend/internal/service/openai_codex_transform.go @@ -275,6 +275,13 @@ func normalizeCodexModel(model string) string { return "gpt-5.1" } +func normalizeOpenAIModelForUpstream(account *Account, model string) string { + if account == nil || account.Type == AccountTypeOAuth { + return normalizeCodexModel(model) + } + return strings.TrimSpace(model) +} + func SupportsVerbosity(model string) bool { if !strings.HasPrefix(model, "gpt-") { return true diff --git a/backend/internal/service/openai_compat_model_test.go b/backend/internal/service/openai_compat_model_test.go index 32c646d490..4396c15fd0 100644 --- a/backend/internal/service/openai_compat_model_test.go +++ b/backend/internal/service/openai_compat_model_test.go @@ -6,9 +6,12 @@ import ( "io" "net/http" "net/http/httptest" + "os" + "path/filepath" "strings" "testing" + "github.com/Wei-Shaw/sub2api/internal/config" "github.com/Wei-Shaw/sub2api/internal/pkg/apicompat" "github.com/gin-gonic/gin" "github.com/stretchr/testify/require" @@ -127,3 +130,101 @@ func TestForwardAsAnthropic_NormalizesRoutingAndEffortForGpt54XHigh(t *testing.T t.Logf("upstream body: %s", string(upstream.lastBody)) t.Logf("response body: %s", rec.Body.String()) } + +func TestForwardAsAnthropic_ForcedCodexInstructionsTemplatePrependsRenderedInstructions(t *testing.T) { + t.Parallel() + gin.SetMode(gin.TestMode) + + templateDir := t.TempDir() + templatePath := filepath.Join(templateDir, "codex-instructions.md.tmpl") + require.NoError(t, os.WriteFile(templatePath, []byte("server-prefix\n\n{{ .ExistingInstructions }}"), 0o644)) + + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + body := []byte(`{"model":"gpt-5.4","max_tokens":16,"system":"client-system","messages":[{"role":"user","content":"hello"}],"stream":false}`) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages", bytes.NewReader(body)) + c.Request.Header.Set("Content-Type", "application/json") + + upstreamBody := strings.Join([]string{ + `data: {"type":"response.completed","response":{"id":"resp_1","object":"response","model":"gpt-5.4","status":"completed","output":[{"type":"message","id":"msg_1","role":"assistant","status":"completed","content":[{"type":"output_text","text":"ok"}]}],"usage":{"input_tokens":5,"output_tokens":2,"total_tokens":7}}}`, + "", + "data: [DONE]", + "", + }, "\n") + upstream := &httpUpstreamRecorder{resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"text/event-stream"}, "x-request-id": []string{"rid_forced"}}, + Body: io.NopCloser(strings.NewReader(upstreamBody)), + }} + + svc := &OpenAIGatewayService{ + cfg: &config.Config{Gateway: config.GatewayConfig{ + ForcedCodexInstructionsTemplateFile: templatePath, + ForcedCodexInstructionsTemplate: "server-prefix\n\n{{ .ExistingInstructions }}", + }}, + httpUpstream: upstream, + } + account := &Account{ + ID: 1, + Name: "openai-oauth", + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Concurrency: 1, + Credentials: map[string]any{ + "access_token": "oauth-token", + "chatgpt_account_id": "chatgpt-acc", + }, + } + + result, err := svc.ForwardAsAnthropic(context.Background(), c, account, body, "", "gpt-5.1") + require.NoError(t, err) + require.NotNil(t, result) + require.Equal(t, "server-prefix\n\nclient-system", gjson.GetBytes(upstream.lastBody, "instructions").String()) +} + +func TestForwardAsAnthropic_ForcedCodexInstructionsTemplateUsesCachedTemplateContent(t *testing.T) { + t.Parallel() + gin.SetMode(gin.TestMode) + + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + body := []byte(`{"model":"gpt-5.4","max_tokens":16,"system":"client-system","messages":[{"role":"user","content":"hello"}],"stream":false}`) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages", bytes.NewReader(body)) + c.Request.Header.Set("Content-Type", "application/json") + + upstreamBody := strings.Join([]string{ + `data: {"type":"response.completed","response":{"id":"resp_1","object":"response","model":"gpt-5.4","status":"completed","output":[{"type":"message","id":"msg_1","role":"assistant","status":"completed","content":[{"type":"output_text","text":"ok"}]}],"usage":{"input_tokens":5,"output_tokens":2,"total_tokens":7}}}`, + "", + "data: [DONE]", + "", + }, "\n") + upstream := &httpUpstreamRecorder{resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"text/event-stream"}, "x-request-id": []string{"rid_forced_cached"}}, + Body: io.NopCloser(strings.NewReader(upstreamBody)), + }} + + svc := &OpenAIGatewayService{ + cfg: &config.Config{Gateway: config.GatewayConfig{ + ForcedCodexInstructionsTemplateFile: "/path/that/should/not/be/read.tmpl", + ForcedCodexInstructionsTemplate: "cached-prefix\n\n{{ .ExistingInstructions }}", + }}, + httpUpstream: upstream, + } + account := &Account{ + ID: 1, + Name: "openai-oauth", + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Concurrency: 1, + Credentials: map[string]any{ + "access_token": "oauth-token", + "chatgpt_account_id": "chatgpt-acc", + }, + } + + result, err := svc.ForwardAsAnthropic(context.Background(), c, account, body, "", "gpt-5.1") + require.NoError(t, err) + require.NotNil(t, result) + require.Equal(t, "cached-prefix\n\nclient-system", gjson.GetBytes(upstream.lastBody, "instructions").String()) +} diff --git a/backend/internal/service/openai_content_session_seed.go b/backend/internal/service/openai_content_session_seed.go new file mode 100644 index 0000000000..7c2ba25140 --- /dev/null +++ b/backend/internal/service/openai_content_session_seed.go @@ -0,0 +1,107 @@ +package service + +import ( + "encoding/json" + "strings" + + "github.com/tidwall/gjson" +) + +// contentSessionSeedPrefix prevents collisions between content-derived seeds +// and explicit session IDs (e.g. "sess-xxx" or "compat_cc_xxx"). +const contentSessionSeedPrefix = "compat_cs_" + +// deriveOpenAIContentSessionSeed builds a stable session seed from an +// OpenAI-format request body. Only fields constant across conversation turns +// are included: model, tools/functions definitions, system/developer prompts, +// instructions (Responses API), and the first user message. +// Supports both Chat Completions (messages) and Responses API (input). +func deriveOpenAIContentSessionSeed(body []byte) string { + if len(body) == 0 { + return "" + } + + var b strings.Builder + + if model := gjson.GetBytes(body, "model").String(); model != "" { + _, _ = b.WriteString("model=") + _, _ = b.WriteString(model) + } + + if tools := gjson.GetBytes(body, "tools"); tools.Exists() && tools.IsArray() && tools.Raw != "[]" { + _, _ = b.WriteString("|tools=") + _, _ = b.WriteString(normalizeCompatSeedJSON(json.RawMessage(tools.Raw))) + } + + if funcs := gjson.GetBytes(body, "functions"); funcs.Exists() && funcs.IsArray() && funcs.Raw != "[]" { + _, _ = b.WriteString("|functions=") + _, _ = b.WriteString(normalizeCompatSeedJSON(json.RawMessage(funcs.Raw))) + } + + if instr := gjson.GetBytes(body, "instructions").String(); instr != "" { + _, _ = b.WriteString("|instructions=") + _, _ = b.WriteString(instr) + } + + firstUserCaptured := false + + msgs := gjson.GetBytes(body, "messages") + if msgs.Exists() && msgs.IsArray() { + msgs.ForEach(func(_, msg gjson.Result) bool { + role := msg.Get("role").String() + switch role { + case "system", "developer": + _, _ = b.WriteString("|system=") + if c := msg.Get("content"); c.Exists() { + _, _ = b.WriteString(normalizeCompatSeedJSON(json.RawMessage(c.Raw))) + } + case "user": + if !firstUserCaptured { + _, _ = b.WriteString("|first_user=") + if c := msg.Get("content"); c.Exists() { + _, _ = b.WriteString(normalizeCompatSeedJSON(json.RawMessage(c.Raw))) + } + firstUserCaptured = true + } + } + return true + }) + } else if inp := gjson.GetBytes(body, "input"); inp.Exists() { + if inp.Type == gjson.String { + _, _ = b.WriteString("|input=") + _, _ = b.WriteString(inp.String()) + } else if inp.IsArray() { + inp.ForEach(func(_, item gjson.Result) bool { + role := item.Get("role").String() + switch role { + case "system", "developer": + _, _ = b.WriteString("|system=") + if c := item.Get("content"); c.Exists() { + _, _ = b.WriteString(normalizeCompatSeedJSON(json.RawMessage(c.Raw))) + } + case "user": + if !firstUserCaptured { + _, _ = b.WriteString("|first_user=") + if c := item.Get("content"); c.Exists() { + _, _ = b.WriteString(normalizeCompatSeedJSON(json.RawMessage(c.Raw))) + } + firstUserCaptured = true + } + } + if !firstUserCaptured && item.Get("type").String() == "input_text" { + _, _ = b.WriteString("|first_user=") + if text := item.Get("text").String(); text != "" { + _, _ = b.WriteString(text) + } + firstUserCaptured = true + } + return true + }) + } + } + + if b.Len() == 0 { + return "" + } + return contentSessionSeedPrefix + b.String() +} diff --git a/backend/internal/service/openai_content_session_seed_test.go b/backend/internal/service/openai_content_session_seed_test.go new file mode 100644 index 0000000000..65a0bf1808 --- /dev/null +++ b/backend/internal/service/openai_content_session_seed_test.go @@ -0,0 +1,218 @@ +package service + +import ( + "testing" + + "github.com/stretchr/testify/require" +) + +func TestDeriveOpenAIContentSessionSeed_EmptyInputs(t *testing.T) { + require.Empty(t, deriveOpenAIContentSessionSeed(nil)) + require.Empty(t, deriveOpenAIContentSessionSeed([]byte{})) + require.Empty(t, deriveOpenAIContentSessionSeed([]byte(`{}`))) +} + +func TestDeriveOpenAIContentSessionSeed_ModelOnly(t *testing.T) { + seed := deriveOpenAIContentSessionSeed([]byte(`{"model":"gpt-5.4"}`)) + require.Contains(t, seed, contentSessionSeedPrefix) + require.Contains(t, seed, "model=gpt-5.4") +} + +func TestDeriveOpenAIContentSessionSeed_ChatCompletions_StableAcrossTurns(t *testing.T) { + turn1 := []byte(`{ + "model": "gpt-5.4", + "messages": [ + {"role": "system", "content": "You are helpful."}, + {"role": "user", "content": "Hello"} + ] + }`) + turn2 := []byte(`{ + "model": "gpt-5.4", + "messages": [ + {"role": "system", "content": "You are helpful."}, + {"role": "user", "content": "Hello"}, + {"role": "assistant", "content": "Hi there!"}, + {"role": "user", "content": "How are you?"} + ] + }`) + s1 := deriveOpenAIContentSessionSeed(turn1) + s2 := deriveOpenAIContentSessionSeed(turn2) + require.Equal(t, s1, s2, "seed should be stable across later turns") + require.NotEmpty(t, s1) +} + +func TestDeriveOpenAIContentSessionSeed_ChatCompletions_DifferentFirstUserDiffers(t *testing.T) { + req1 := []byte(`{"model":"gpt-5.4","messages":[{"role":"user","content":"Question A"}]}`) + req2 := []byte(`{"model":"gpt-5.4","messages":[{"role":"user","content":"Question B"}]}`) + s1 := deriveOpenAIContentSessionSeed(req1) + s2 := deriveOpenAIContentSessionSeed(req2) + require.NotEqual(t, s1, s2) +} + +func TestDeriveOpenAIContentSessionSeed_ChatCompletions_DifferentSystemDiffers(t *testing.T) { + req1 := []byte(`{"model":"gpt-5.4","messages":[{"role":"system","content":"A"},{"role":"user","content":"Hi"}]}`) + req2 := []byte(`{"model":"gpt-5.4","messages":[{"role":"system","content":"B"},{"role":"user","content":"Hi"}]}`) + s1 := deriveOpenAIContentSessionSeed(req1) + s2 := deriveOpenAIContentSessionSeed(req2) + require.NotEqual(t, s1, s2) +} + +func TestDeriveOpenAIContentSessionSeed_ChatCompletions_DifferentModelDiffers(t *testing.T) { + req1 := []byte(`{"model":"gpt-5.4","messages":[{"role":"user","content":"Hi"}]}`) + req2 := []byte(`{"model":"gpt-4o","messages":[{"role":"user","content":"Hi"}]}`) + s1 := deriveOpenAIContentSessionSeed(req1) + s2 := deriveOpenAIContentSessionSeed(req2) + require.NotEqual(t, s1, s2) +} + +func TestDeriveOpenAIContentSessionSeed_ChatCompletions_WithTools(t *testing.T) { + withTools := []byte(`{ + "model": "gpt-5.4", + "tools": [{"type":"function","function":{"name":"get_weather"}}], + "messages": [{"role": "user", "content": "Hello"}] + }`) + withoutTools := []byte(`{ + "model": "gpt-5.4", + "messages": [{"role": "user", "content": "Hello"}] + }`) + s1 := deriveOpenAIContentSessionSeed(withTools) + s2 := deriveOpenAIContentSessionSeed(withoutTools) + require.NotEqual(t, s1, s2, "tools should affect the seed") + require.Contains(t, s1, "|tools=") +} + +func TestDeriveOpenAIContentSessionSeed_ChatCompletions_WithFunctions(t *testing.T) { + body := []byte(`{ + "model": "gpt-5.4", + "functions": [{"name":"get_weather","parameters":{}}], + "messages": [{"role": "user", "content": "Hello"}] + }`) + seed := deriveOpenAIContentSessionSeed(body) + require.Contains(t, seed, "|functions=") +} + +func TestDeriveOpenAIContentSessionSeed_ChatCompletions_DeveloperRole(t *testing.T) { + body := []byte(`{ + "model": "gpt-5.4", + "messages": [ + {"role": "developer", "content": "You are helpful."}, + {"role": "user", "content": "Hello"} + ] + }`) + seed := deriveOpenAIContentSessionSeed(body) + require.Contains(t, seed, "|system=") + require.Contains(t, seed, "|first_user=") +} + +func TestDeriveOpenAIContentSessionSeed_ChatCompletions_StructuredContent(t *testing.T) { + body := []byte(`{ + "model": "gpt-5.4", + "messages": [ + {"role": "user", "content": [{"type":"text","text":"Hello"}]} + ] + }`) + seed := deriveOpenAIContentSessionSeed(body) + require.NotEmpty(t, seed) + require.Contains(t, seed, "|first_user=") +} + +func TestDeriveOpenAIContentSessionSeed_ResponsesAPI_InputString(t *testing.T) { + body := []byte(`{"model":"gpt-5.4","input":"Hello, how are you?"}`) + seed := deriveOpenAIContentSessionSeed(body) + require.Contains(t, seed, "|input=Hello, how are you?") +} + +func TestDeriveOpenAIContentSessionSeed_ResponsesAPI_InputArray(t *testing.T) { + body := []byte(`{ + "model": "gpt-5.4", + "input": [ + {"role": "system", "content": "You are helpful."}, + {"role": "user", "content": "Hello"} + ] + }`) + seed := deriveOpenAIContentSessionSeed(body) + require.Contains(t, seed, "|system=") + require.Contains(t, seed, "|first_user=") +} + +func TestDeriveOpenAIContentSessionSeed_ResponsesAPI_WithInstructions(t *testing.T) { + body := []byte(`{ + "model": "gpt-5.4", + "instructions": "You are a coding assistant.", + "input": "Write a hello world" + }`) + seed := deriveOpenAIContentSessionSeed(body) + require.Contains(t, seed, "|instructions=You are a coding assistant.") + require.Contains(t, seed, "|input=Write a hello world") +} + +func TestDeriveOpenAIContentSessionSeed_Deterministic(t *testing.T) { + body := []byte(`{ + "model": "gpt-5.4", + "messages": [ + {"role": "system", "content": "You are helpful."}, + {"role": "user", "content": "Hello"} + ] + }`) + s1 := deriveOpenAIContentSessionSeed(body) + s2 := deriveOpenAIContentSessionSeed(body) + require.Equal(t, s1, s2, "seed must be deterministic") +} + +func TestDeriveOpenAIContentSessionSeed_PrefixPresent(t *testing.T) { + body := []byte(`{"model":"gpt-5.4","messages":[{"role":"user","content":"Hi"}]}`) + seed := deriveOpenAIContentSessionSeed(body) + require.True(t, len(seed) > len(contentSessionSeedPrefix)) + require.Equal(t, contentSessionSeedPrefix, seed[:len(contentSessionSeedPrefix)]) +} + +func TestDeriveOpenAIContentSessionSeed_EmptyToolsIgnored(t *testing.T) { + body := []byte(`{"model":"gpt-5.4","tools":[],"messages":[{"role":"user","content":"Hi"}]}`) + seed := deriveOpenAIContentSessionSeed(body) + require.NotContains(t, seed, "|tools=") +} + +func TestDeriveOpenAIContentSessionSeed_MessagesPreferredOverInput(t *testing.T) { + body := []byte(`{ + "model": "gpt-5.4", + "messages": [{"role": "user", "content": "from messages"}], + "input": "from input" + }`) + seed := deriveOpenAIContentSessionSeed(body) + require.Contains(t, seed, "|first_user=") + require.NotContains(t, seed, "|input=") +} + +func TestDeriveOpenAIContentSessionSeed_JSONCanonicalisation(t *testing.T) { + compact := []byte(`{"model":"gpt-5.4","tools":[{"type":"function","function":{"name":"get_weather","description":"Get weather"}}],"messages":[{"role":"user","content":"Hi"}]}`) + spaced := []byte(`{ + "model": "gpt-5.4", + "tools": [ + { "type" : "function", "function": { "description": "Get weather", "name": "get_weather" } } + ], + "messages": [ { "role": "user", "content": "Hi" } ] + }`) + s1 := deriveOpenAIContentSessionSeed(compact) + s2 := deriveOpenAIContentSessionSeed(spaced) + require.Equal(t, s1, s2, "different formatting of identical JSON should produce the same seed") +} + +func TestDeriveOpenAIContentSessionSeed_ResponsesAPI_InputTextTypedItem(t *testing.T) { + body := []byte(`{ + "model": "gpt-5.4", + "input": [{"type": "input_text", "text": "Hello world"}] + }`) + seed := deriveOpenAIContentSessionSeed(body) + require.Contains(t, seed, "|first_user=") + require.Contains(t, seed, "Hello world") +} + +func TestDeriveOpenAIContentSessionSeed_ResponsesAPI_TypedMessageItem(t *testing.T) { + body := []byte(`{ + "model": "gpt-5.4", + "input": [{"type": "message", "role": "user", "content": "Hello from typed message"}] + }`) + seed := deriveOpenAIContentSessionSeed(body) + require.Contains(t, seed, "|first_user=") + require.Contains(t, seed, "Hello from typed message") +} diff --git a/backend/internal/service/openai_gateway_chat_completions.go b/backend/internal/service/openai_gateway_chat_completions.go index 3cada2ebdd..25451b2b6b 100644 --- a/backend/internal/service/openai_gateway_chat_completions.go +++ b/backend/internal/service/openai_gateway_chat_completions.go @@ -46,7 +46,7 @@ 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. billingModel := resolveOpenAIForwardModel(account, originalModel, defaultMappedModel) - upstreamModel := normalizeCodexModel(billingModel) + upstreamModel := normalizeOpenAIModelForUpstream(account, billingModel) promptCacheKey = strings.TrimSpace(promptCacheKey) compatPromptCacheInjected := false @@ -259,6 +259,7 @@ func (s *OpenAIGatewayService) handleChatBufferedStreamingResponse( var finalResponse *apicompat.ResponsesResponse var usage OpenAIUsage + acc := apicompat.NewBufferedResponseAccumulator() for scanner.Scan() { line := scanner.Text() @@ -276,7 +277,11 @@ func (s *OpenAIGatewayService) handleChatBufferedStreamingResponse( continue } - if (event.Type == "response.completed" || event.Type == "response.incomplete" || event.Type == "response.failed") && + // Accumulate delta content for fallback when terminal output is empty. + acc.ProcessEvent(&event) + + if (event.Type == "response.completed" || event.Type == "response.done" || + event.Type == "response.incomplete" || event.Type == "response.failed") && event.Response != nil { finalResponse = event.Response if event.Response.Usage != nil { @@ -305,6 +310,10 @@ func (s *OpenAIGatewayService) handleChatBufferedStreamingResponse( return nil, fmt.Errorf("upstream stream ended without terminal event") } + // When the terminal event has an empty output array, reconstruct from + // accumulated delta events so the client receives the full content. + acc.SupplementResponseOutput(finalResponse) + chatResp := apicompat.ResponsesToChatCompletions(finalResponse, originalModel) if s.responseHeaderFilter != nil { diff --git a/backend/internal/service/openai_gateway_messages.go b/backend/internal/service/openai_gateway_messages.go index dd416269f4..7a4862d335 100644 --- a/backend/internal/service/openai_gateway_messages.go +++ b/backend/internal/service/openai_gateway_messages.go @@ -62,7 +62,7 @@ func (s *OpenAIGatewayService) ForwardAsAnthropic( // 3. Model mapping billingModel := resolveOpenAIForwardModel(account, normalizedModel, defaultMappedModel) - upstreamModel := normalizeCodexModel(billingModel) + upstreamModel := normalizeOpenAIModelForUpstream(account, billingModel) responsesReq.Model = upstreamModel logger.L().Debug("openai messages: model mapping applied", @@ -86,6 +86,24 @@ func (s *OpenAIGatewayService) ForwardAsAnthropic( return nil, fmt.Errorf("unmarshal for codex transform: %w", err) } codexResult := applyCodexOAuthTransform(reqBody, false, false) + forcedTemplateText := "" + if s.cfg != nil { + forcedTemplateText = s.cfg.Gateway.ForcedCodexInstructionsTemplate + } + templateUpstreamModel := upstreamModel + if codexResult.NormalizedModel != "" { + templateUpstreamModel = codexResult.NormalizedModel + } + existingInstructions, _ := reqBody["instructions"].(string) + if _, err := applyForcedCodexInstructionsTemplate(reqBody, forcedTemplateText, forcedCodexInstructionsTemplateData{ + ExistingInstructions: strings.TrimSpace(existingInstructions), + OriginalModel: originalModel, + NormalizedModel: normalizedModel, + BillingModel: billingModel, + UpstreamModel: templateUpstreamModel, + }); err != nil { + return nil, err + } if codexResult.NormalizedModel != "" { upstreamModel = codexResult.NormalizedModel } diff --git a/backend/internal/service/openai_gateway_service.go b/backend/internal/service/openai_gateway_service.go index 831d940f53..32596f0fe0 100644 --- a/backend/internal/service/openai_gateway_service.go +++ b/backend/internal/service/openai_gateway_service.go @@ -21,6 +21,7 @@ import ( "time" "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/Wei-Shaw/sub2api/internal/pkg/apicompat" "github.com/Wei-Shaw/sub2api/internal/pkg/logger" "github.com/Wei-Shaw/sub2api/internal/pkg/openai" "github.com/Wei-Shaw/sub2api/internal/util/responseheaders" @@ -1120,6 +1121,7 @@ func (s *OpenAIGatewayService) ExtractSessionID(c *gin.Context, body []byte) str // 1. Header: session_id // 2. Header: conversation_id // 3. Body: prompt_cache_key (opencode) +// 4. Body: content-based fallback (model + system + tools + first user message) func (s *OpenAIGatewayService) GenerateSessionHash(c *gin.Context, body []byte) string { if c == nil { return "" @@ -1132,6 +1134,9 @@ func (s *OpenAIGatewayService) GenerateSessionHash(c *gin.Context, body []byte) if sessionID == "" && len(body) > 0 { sessionID = strings.TrimSpace(gjson.GetBytes(body, "prompt_cache_key").String()) } + if sessionID == "" && len(body) > 0 { + sessionID = deriveOpenAIContentSessionSeed(body) + } if sessionID == "" { return "" } @@ -1238,7 +1243,7 @@ func (s *OpenAIGatewayService) selectAccountForModelWithExclusions(ctx context.C _ = s.setStickySessionAccountID(ctx, groupID, sessionHash, selected.ID, openaiStickySessionTTL) } - return selected, nil + return s.hydrateSelectedAccount(ctx, selected) } // tryStickySessionHit 尝试从粘性会话获取账号。 @@ -1403,35 +1408,25 @@ func (s *OpenAIGatewayService) SelectAccountWithLoadAwareness(ctx context.Contex } result, err := s.tryAcquireAccountSlot(ctx, account.ID, account.Concurrency) if err == nil && result.Acquired { - return &AccountSelectionResult{ - Account: account, - Acquired: true, - ReleaseFunc: result.ReleaseFunc, - }, nil + return s.newSelectionResult(ctx, account, true, result.ReleaseFunc, nil) } if stickyAccountID > 0 && stickyAccountID == account.ID && s.concurrencyService != nil { waitingCount, _ := s.concurrencyService.GetAccountWaitingCount(ctx, account.ID) if waitingCount < cfg.StickySessionMaxWaiting { - return &AccountSelectionResult{ - Account: account, - WaitPlan: &AccountWaitPlan{ - AccountID: account.ID, - MaxConcurrency: account.Concurrency, - Timeout: cfg.StickySessionWaitTimeout, - MaxWaiting: cfg.StickySessionMaxWaiting, - }, - }, nil + return s.newSelectionResult(ctx, account, false, nil, &AccountWaitPlan{ + AccountID: account.ID, + MaxConcurrency: account.Concurrency, + Timeout: cfg.StickySessionWaitTimeout, + MaxWaiting: cfg.StickySessionMaxWaiting, + }) } } - return &AccountSelectionResult{ - Account: account, - WaitPlan: &AccountWaitPlan{ - AccountID: account.ID, - MaxConcurrency: account.Concurrency, - Timeout: cfg.FallbackWaitTimeout, - MaxWaiting: cfg.FallbackMaxWaiting, - }, - }, nil + return s.newSelectionResult(ctx, account, false, nil, &AccountWaitPlan{ + AccountID: account.ID, + MaxConcurrency: account.Concurrency, + Timeout: cfg.FallbackWaitTimeout, + MaxWaiting: cfg.FallbackMaxWaiting, + }) } accounts, err := s.listSchedulableAccounts(ctx, groupID) @@ -1471,24 +1466,17 @@ func (s *OpenAIGatewayService) SelectAccountWithLoadAwareness(ctx context.Contex result, err := s.tryAcquireAccountSlot(ctx, accountID, account.Concurrency) if err == nil && result.Acquired { _ = s.refreshStickySessionTTL(ctx, groupID, sessionHash, openaiStickySessionTTL) - return &AccountSelectionResult{ - Account: account, - Acquired: true, - ReleaseFunc: result.ReleaseFunc, - }, nil + return s.newSelectionResult(ctx, account, true, result.ReleaseFunc, nil) } waitingCount, _ := s.concurrencyService.GetAccountWaitingCount(ctx, accountID) if waitingCount < cfg.StickySessionMaxWaiting { - return &AccountSelectionResult{ - Account: account, - WaitPlan: &AccountWaitPlan{ - AccountID: accountID, - MaxConcurrency: account.Concurrency, - Timeout: cfg.StickySessionWaitTimeout, - MaxWaiting: cfg.StickySessionMaxWaiting, - }, - }, nil + return s.newSelectionResult(ctx, account, false, nil, &AccountWaitPlan{ + AccountID: accountID, + MaxConcurrency: account.Concurrency, + Timeout: cfg.StickySessionWaitTimeout, + MaxWaiting: cfg.StickySessionMaxWaiting, + }) } } } @@ -1547,11 +1535,7 @@ func (s *OpenAIGatewayService) SelectAccountWithLoadAwareness(ctx context.Contex if sessionHash != "" { _ = s.setStickySessionAccountID(ctx, groupID, sessionHash, fresh.ID, openaiStickySessionTTL) } - return &AccountSelectionResult{ - Account: fresh, - Acquired: true, - ReleaseFunc: result.ReleaseFunc, - }, nil + return s.newSelectionResult(ctx, fresh, true, result.ReleaseFunc, nil) } } } else { @@ -1604,11 +1588,7 @@ func (s *OpenAIGatewayService) SelectAccountWithLoadAwareness(ctx context.Contex if sessionHash != "" { _ = s.setStickySessionAccountID(ctx, groupID, sessionHash, fresh.ID, openaiStickySessionTTL) } - return &AccountSelectionResult{ - Account: fresh, - Acquired: true, - ReleaseFunc: result.ReleaseFunc, - }, nil + return s.newSelectionResult(ctx, fresh, true, result.ReleaseFunc, nil) } } } @@ -1624,15 +1604,12 @@ func (s *OpenAIGatewayService) SelectAccountWithLoadAwareness(ctx context.Contex if needsUpstreamCheck && s.isUpstreamModelRestrictedByChannel(ctx, *groupID, fresh, requestedModel) { continue } - return &AccountSelectionResult{ - Account: fresh, - WaitPlan: &AccountWaitPlan{ - AccountID: fresh.ID, - MaxConcurrency: fresh.Concurrency, - Timeout: cfg.FallbackWaitTimeout, - MaxWaiting: cfg.FallbackMaxWaiting, - }, - }, nil + return s.newSelectionResult(ctx, fresh, false, nil, &AccountWaitPlan{ + AccountID: fresh.ID, + MaxConcurrency: fresh.Concurrency, + Timeout: cfg.FallbackWaitTimeout, + MaxWaiting: cfg.FallbackMaxWaiting, + }) } return nil, ErrNoAvailableAccounts @@ -1727,6 +1704,33 @@ func (s *OpenAIGatewayService) getSchedulableAccount(ctx context.Context, accoun return account, nil } +func (s *OpenAIGatewayService) hydrateSelectedAccount(ctx context.Context, account *Account) (*Account, error) { + if account == nil || s.schedulerSnapshot == nil { + return account, nil + } + hydrated, err := s.schedulerSnapshot.GetAccount(ctx, account.ID) + if err != nil { + return nil, err + } + if hydrated == nil { + return nil, fmt.Errorf("selected openai account %d not found during hydration", account.ID) + } + return hydrated, nil +} + +func (s *OpenAIGatewayService) newSelectionResult(ctx context.Context, account *Account, acquired bool, release func(), waitPlan *AccountWaitPlan) (*AccountSelectionResult, error) { + hydrated, err := s.hydrateSelectedAccount(ctx, account) + if err != nil { + return nil, err + } + return &AccountSelectionResult{ + Account: hydrated, + Acquired: acquired, + ReleaseFunc: release, + WaitPlan: waitPlan, + }, nil +} + func (s *OpenAIGatewayService) schedulingConfig() config.GatewaySchedulingConfig { if s.cfg != nil { return s.cfg.Gateway.Scheduling @@ -1937,9 +1941,11 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco } upstreamModel := billingModel - // 针对所有 OpenAI 账号执行 Codex 模型名规范化,确保上游识别一致。 + // OpenAI OAuth 账号走 ChatGPT internal Codex endpoint,需要将模型名规范化为 + // 上游可识别的 Codex/GPT 系列。API Key 账号则应保留原始/映射后的模型名, + // 以兼容自定义 base_url 的 OpenAI-compatible 上游。 if model, ok := reqBody["model"].(string); ok { - upstreamModel = normalizeCodexModel(model) + upstreamModel = normalizeOpenAIModelForUpstream(account, 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) @@ -2045,6 +2051,11 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco } } + if sanitizeEmptyBase64InputImagesInOpenAIRequestBodyMap(reqBody) { + bodyModified = true + disablePatch() + } + // Re-serialize body only if modified if bodyModified { serializedByPatch := false @@ -2472,6 +2483,14 @@ func (s *OpenAIGatewayService) forwardOpenAIPassthrough( reqStream = gjson.GetBytes(body, "stream").Bool() } + sanitizedBody, sanitized, err := sanitizeEmptyBase64InputImagesInOpenAIBody(body) + if err != nil { + return nil, err + } + if sanitized { + body = sanitizedBody + } + logger.LegacyPrintf("service.openai_gateway", "[OpenAI 自动透传] 命中自动透传分支: account=%d name=%s type=%s model=%s stream=%v", account.ID, @@ -2544,7 +2563,11 @@ func (s *OpenAIGatewayService) forwardOpenAIPassthrough( defer func() { _ = resp.Body.Close() }() if resp.StatusCode >= 400 { - // 透传模式不做 failover(避免改变原始上游语义),按上游原样返回错误响应。 + // 透传模式默认保持原样代理;但 429/529 属于网关必须兜底的 + // 上游容量类错误,应先触发多账号 failover 以维持基础 SLA。 + if shouldFailoverOpenAIPassthroughResponse(resp.StatusCode) { + return nil, s.handleFailoverErrorResponsePassthrough(ctx, resp, c, account, body) + } return nil, s.handleErrorResponsePassthrough(ctx, resp, c, account, body) } @@ -2727,6 +2750,58 @@ func (s *OpenAIGatewayService) buildUpstreamRequestOpenAIPassthrough( return req, nil } +func shouldFailoverOpenAIPassthroughResponse(statusCode int) bool { + switch statusCode { + case http.StatusTooManyRequests, 529: + return true + default: + return false + } +} + +func (s *OpenAIGatewayService) handleFailoverErrorResponsePassthrough( + ctx context.Context, + resp *http.Response, + c *gin.Context, + account *Account, + requestBody []byte, +) error { + body, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20)) + + upstreamMsg := strings.TrimSpace(extractUpstreamErrorMessage(body)) + upstreamMsg = sanitizeUpstreamErrorMessage(upstreamMsg) + upstreamDetail := "" + if s.cfg != nil && s.cfg.Gateway.LogUpstreamErrorBody { + maxBytes := s.cfg.Gateway.LogUpstreamErrorBodyMaxBytes + if maxBytes <= 0 { + maxBytes = 2048 + } + upstreamDetail = truncateString(string(body), maxBytes) + } + setOpsUpstreamError(c, resp.StatusCode, upstreamMsg, upstreamDetail) + logOpenAIInstructionsRequiredDebug(ctx, c, account, resp.StatusCode, upstreamMsg, requestBody, body) + if s.rateLimitService != nil { + _ = s.rateLimitService.HandleUpstreamError(ctx, account, resp.StatusCode, resp.Header, body) + } + appendOpsUpstreamError(c, OpsUpstreamErrorEvent{ + Platform: account.Platform, + AccountID: account.ID, + AccountName: account.Name, + UpstreamStatusCode: resp.StatusCode, + UpstreamRequestID: resp.Header.Get("x-request-id"), + Passthrough: true, + Kind: "failover", + Message: upstreamMsg, + Detail: upstreamDetail, + UpstreamResponseBody: upstreamDetail, + }) + return &UpstreamFailoverError{ + StatusCode: resp.StatusCode, + ResponseBody: body, + ResponseHeaders: resp.Header.Clone(), + } +} + func (s *OpenAIGatewayService) handleErrorResponsePassthrough( ctx context.Context, resp *http.Response, @@ -2948,6 +3023,14 @@ func (s *OpenAIGatewayService) handleNonStreamingResponsePassthrough( return nil, err } + // Detect SSE responses from upstream and convert to JSON. + // Some upstreams (e.g. other sub2api instances) may return SSE even when + // stream=false was requested. Without this conversion the client would + // receive raw SSE text or a terminal event with empty output. + if isEventStreamResponse(resp.Header) { + return s.handlePassthroughSSEToJSON(resp, c, body) + } + usage := &OpenAIUsage{} usageParsed := false if len(body) > 0 { @@ -2971,6 +3054,56 @@ func (s *OpenAIGatewayService) handleNonStreamingResponsePassthrough( return usage, nil } +// handlePassthroughSSEToJSON converts an SSE response body into a JSON +// response for the passthrough path. It mirrors handleSSEToJSON but skips +// model replacement (passthrough does not remap models). +func (s *OpenAIGatewayService) handlePassthroughSSEToJSON(resp *http.Response, c *gin.Context, body []byte) (*OpenAIUsage, error) { + bodyText := string(body) + finalResponse, ok := extractCodexFinalResponse(bodyText) + + usage := &OpenAIUsage{} + if ok { + if parsedUsage, parsed := extractOpenAIUsageFromJSONBytes(finalResponse); parsed { + *usage = parsedUsage + } + // When the terminal event has an empty output array, reconstruct + // output from accumulated delta events so the client gets full content. + if len(gjson.GetBytes(finalResponse, "output").Array()) == 0 { + if outputJSON, reconstructed := reconstructResponseOutputFromSSE(bodyText); reconstructed { + if patched, err := sjson.SetRawBytes(finalResponse, "output", outputJSON); err == nil { + finalResponse = patched + } + } + } + body = finalResponse + // Correct tool calls in final response + body = s.correctToolCallsInResponseBody(body) + } else { + terminalType, terminalPayload, terminalOK := extractOpenAISSETerminalEvent(bodyText) + if terminalOK && terminalType == "response.failed" { + msg := extractOpenAISSEErrorMessage(terminalPayload) + if msg == "" { + msg = "Upstream compact response failed" + } + return nil, s.writeOpenAINonStreamingProtocolError(resp, c, msg) + } + usage = s.parseSSEUsageFromBody(bodyText) + } + + writeOpenAIPassthroughResponseHeaders(c.Writer.Header(), resp.Header, s.responseHeaderFilter) + + contentType := "application/json; charset=utf-8" + if !ok { + contentType = resp.Header.Get("Content-Type") + if contentType == "" { + contentType = "text/event-stream" + } + } + c.Data(resp.StatusCode, contentType, body) + + return usage, nil +} + func writeOpenAIPassthroughResponseHeaders(dst http.Header, src http.Header, filter *responseheaders.CompiledHeaderFilter) { if dst == nil || src == nil { return @@ -3799,10 +3932,21 @@ func (s *OpenAIGatewayService) handleNonStreamingResponse(ctx context.Context, r return nil, err } + // Detect SSE responses for ALL account types via Content-Type header. + // Some OpenAI-compatible upstreams (including other sub2api instances) + // may return SSE even when stream=false was requested. + if isEventStreamResponse(resp.Header) { + return s.handleSSEToJSON(resp, c, body, originalModel, mappedModel) + } + // For OAuth accounts, also fall back to a body-content heuristic because + // the upstream may omit the Content-Type header while still sending SSE. + // This heuristic is NOT applied to API-key accounts to avoid false + // positives on JSON responses that coincidentally contain "data:" or + // "event:" in their text content. if account.Type == AccountTypeOAuth { bodyLooksLikeSSE := bytes.Contains(body, []byte("data:")) || bytes.Contains(body, []byte("event:")) - if isEventStreamResponse(resp.Header) || bodyLooksLikeSSE { - return s.handleOAuthSSEToJSON(resp, c, body, originalModel, mappedModel) + if bodyLooksLikeSSE { + return s.handleSSEToJSON(resp, c, body, originalModel, mappedModel) } } @@ -3836,7 +3980,7 @@ func isEventStreamResponse(header http.Header) bool { return strings.Contains(contentType, "text/event-stream") } -func (s *OpenAIGatewayService) handleOAuthSSEToJSON(resp *http.Response, c *gin.Context, body []byte, originalModel, mappedModel string) (*OpenAIUsage, error) { +func (s *OpenAIGatewayService) handleSSEToJSON(resp *http.Response, c *gin.Context, body []byte, originalModel, mappedModel string) (*OpenAIUsage, error) { bodyText := string(body) finalResponse, ok := extractCodexFinalResponse(bodyText) @@ -3845,6 +3989,16 @@ func (s *OpenAIGatewayService) handleOAuthSSEToJSON(resp *http.Response, c *gin. if parsedUsage, parsed := extractOpenAIUsageFromJSONBytes(finalResponse); parsed { *usage = parsedUsage } + // When the terminal event has an empty output array, reconstruct + // output from accumulated delta events so the client gets full content. + // gjson Array() returns empty slice for null, missing, or empty arrays. + if len(gjson.GetBytes(finalResponse, "output").Array()) == 0 { + if outputJSON, reconstructed := reconstructResponseOutputFromSSE(bodyText); reconstructed { + if patched, err := sjson.SetRawBytes(finalResponse, "output", outputJSON); err == nil { + finalResponse = patched + } + } + } body = finalResponse if originalModel != mappedModel { body = s.replaceModelInResponseBody(body, mappedModel, originalModel) @@ -3946,6 +4100,34 @@ func extractCodexFinalResponse(body string) ([]byte, bool) { return nil, false } +// reconstructResponseOutputFromSSE scans raw SSE body text for delta events and +// returns a JSON-encoded output array reconstructed from accumulated deltas. +// Returns (nil, false) if no content was found in deltas. +func reconstructResponseOutputFromSSE(bodyText string) ([]byte, bool) { + acc := apicompat.NewBufferedResponseAccumulator() + lines := strings.Split(bodyText, "\n") + for _, line := range lines { + data, ok := extractOpenAISSEDataLine(line) + if !ok || data == "" || data == "[DONE]" { + continue + } + var event apicompat.ResponsesStreamEvent + if err := json.Unmarshal([]byte(data), &event); err != nil { + continue + } + acc.ProcessEvent(&event) + } + if !acc.HasContent() { + return nil, false + } + output := acc.BuildOutput() + outputJSON, err := json.Marshal(output) + if err != nil { + return nil, false + } + return outputJSON, true +} + func (s *OpenAIGatewayService) parseSSEUsageFromBody(body string) *OpenAIUsage { usage := &OpenAIUsage{} lines := strings.Split(body, "\n") @@ -4857,6 +5039,123 @@ func normalizeOpenAIServiceTier(raw string) *string { } } +func sanitizeEmptyBase64InputImagesInOpenAIBody(body []byte) ([]byte, bool, error) { + if len(body) == 0 || !bytes.Contains(body, []byte(`"image_url"`)) || !bytes.Contains(body, []byte(`base64,`)) { + return body, false, nil + } + + var reqBody map[string]any + if err := json.Unmarshal(body, &reqBody); err != nil { + return body, false, fmt.Errorf("sanitize request body: %w", err) + } + if !sanitizeEmptyBase64InputImagesInOpenAIRequestBodyMap(reqBody) { + return body, false, nil + } + normalized, err := json.Marshal(reqBody) + if err != nil { + return body, false, fmt.Errorf("serialize sanitized request body: %w", err) + } + return normalized, true, nil +} + +func sanitizeEmptyBase64InputImagesInOpenAIRequestBodyMap(reqBody map[string]any) bool { + if reqBody == nil { + return false + } + input, ok := reqBody["input"] + if !ok { + return false + } + normalizedInput, changed := sanitizeEmptyBase64InputImagesInOpenAIInput(input) + if !changed { + return false + } + reqBody["input"] = normalizedInput + return true +} + +func sanitizeEmptyBase64InputImagesInOpenAIInput(input any) (any, bool) { + items, ok := input.([]any) + if !ok { + return input, false + } + + normalizedItems := make([]any, 0, len(items)) + changed := false + for _, item := range items { + itemMap, ok := item.(map[string]any) + if !ok { + normalizedItems = append(normalizedItems, item) + continue + } + if shouldDropEmptyBase64InputImagePart(itemMap) { + changed = true + continue + } + content, ok := itemMap["content"] + if !ok { + normalizedItems = append(normalizedItems, itemMap) + continue + } + parts, ok := content.([]any) + if !ok { + normalizedItems = append(normalizedItems, itemMap) + continue + } + + normalizedParts := make([]any, 0, len(parts)) + itemChanged := false + for _, part := range parts { + if shouldDropEmptyBase64InputImagePart(part) { + changed = true + itemChanged = true + continue + } + normalizedParts = append(normalizedParts, part) + } + if itemChanged { + if len(normalizedParts) == 0 { + continue + } + itemMap["content"] = normalizedParts + } + normalizedItems = append(normalizedItems, itemMap) + } + if !changed { + return input, false + } + return normalizedItems, true +} + +func shouldDropEmptyBase64InputImagePart(part any) bool { + partMap, ok := part.(map[string]any) + if !ok { + return false + } + typeValue, _ := partMap["type"].(string) + if strings.TrimSpace(typeValue) != "input_image" { + return false + } + imageURL, _ := partMap["image_url"].(string) + return isEmptyBase64DataURI(imageURL) +} + +func isEmptyBase64DataURI(raw string) bool { + if !strings.HasPrefix(raw, "data:") { + return false + } + rest := strings.TrimPrefix(raw, "data:") + semicolonIdx := strings.Index(rest, ";") + if semicolonIdx < 0 { + return false + } + rest = rest[semicolonIdx+1:] + if !strings.HasPrefix(rest, "base64,") { + return false + } + return strings.TrimSpace(strings.TrimPrefix(rest, "base64,")) == "" +} + func getOpenAIRequestBodyMap(c *gin.Context, body []byte) (map[string]any, error) { if c != nil { if cached, ok := c.Get(OpenAIParsedRequestBodyKey); ok { diff --git a/backend/internal/service/openai_gateway_service_hotpath_test.go b/backend/internal/service/openai_gateway_service_hotpath_test.go index f73c06c5e1..234dee00cf 100644 --- a/backend/internal/service/openai_gateway_service_hotpath_test.go +++ b/backend/internal/service/openai_gateway_service_hotpath_test.go @@ -1,6 +1,7 @@ package service import ( + "encoding/json" "net/http/httptest" "testing" @@ -139,3 +140,61 @@ func TestGetOpenAIRequestBodyMap_WriteBackContextCache(t *testing.T) { require.True(t, ok) require.Equal(t, got, cachedMap) } + +func TestSanitizeEmptyBase64InputImagesInOpenAIRequestBodyMap(t *testing.T) { + var reqBody map[string]any + require.NoError(t, json.Unmarshal([]byte(`{ + "model":"gpt-5.4", + "input":[ + {"role":"user","content":[ + {"type":"input_text","text":"Describe this"}, + {"type":"input_image","image_url":"data:image/png;base64, "}, + {"type":"input_image","image_url":"data:image/png;base64,abc123"} + ]}, + {"role":"user","content":[ + {"type":"input_image","image_url":"data:image/png;base64,"} + ]}, + {"type":"input_image","image_url":"data:image/png;base64,"}, + {"type":"input_image","image_url":"data:image/png;base64,top-level-valid"} + ] + }`), &reqBody)) + + require.True(t, sanitizeEmptyBase64InputImagesInOpenAIRequestBodyMap(reqBody)) + + normalized, err := json.Marshal(reqBody) + require.NoError(t, err) + require.JSONEq(t, `{ + "model":"gpt-5.4", + "input":[ + {"role":"user","content":[ + {"type":"input_text","text":"Describe this"}, + {"type":"input_image","image_url":"data:image/png;base64,abc123"} + ]}, + {"type":"input_image","image_url":"data:image/png;base64,top-level-valid"} + ] + }`, string(normalized)) +} + +func TestSanitizeEmptyBase64InputImagesInOpenAIBody(t *testing.T) { + body, changed, err := sanitizeEmptyBase64InputImagesInOpenAIBody([]byte(`{ + "model":"gpt-5.4", + "stream":true, + "input":[ + {"role":"user","content":[ + {"type":"input_text","text":"Describe this"}, + {"type":"input_image","image_url":"data:image/png;base64,"} + ]} + ] + }`)) + require.NoError(t, err) + require.True(t, changed) + require.JSONEq(t, `{ + "model":"gpt-5.4", + "stream":true, + "input":[ + {"role":"user","content":[ + {"type":"input_text","text":"Describe this"} + ]} + ] + }`, string(body)) +} diff --git a/backend/internal/service/openai_gateway_service_test.go b/backend/internal/service/openai_gateway_service_test.go index 9e2f33f22a..cf2d875fcd 100644 --- a/backend/internal/service/openai_gateway_service_test.go +++ b/backend/internal/service/openai_gateway_service_test.go @@ -237,6 +237,60 @@ func TestOpenAIGatewayService_GenerateSessionHashWithFallback(t *testing.T) { require.Equal(t, "", empty) } +func TestOpenAIGatewayService_GenerateSessionHash_ContentFallback(t *testing.T) { + gin.SetMode(gin.TestMode) + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/openai/v1/chat/completions", nil) + + svc := &OpenAIGatewayService{} + + body := []byte(`{"model":"gpt-5.4","messages":[{"role":"system","content":"You are helpful."},{"role":"user","content":"Hello"}]}`) + + hash := svc.GenerateSessionHash(c, body) + require.NotEmpty(t, hash, "content-based fallback should produce a hash") + + hash2 := svc.GenerateSessionHash(c, body) + require.Equal(t, hash, hash2, "same content should produce same hash") + + bodyExtended := []byte(`{"model":"gpt-5.4","messages":[{"role":"system","content":"You are helpful."},{"role":"user","content":"Hello"},{"role":"assistant","content":"Hi!"},{"role":"user","content":"How are you?"}]}`) + hashExtended := svc.GenerateSessionHash(c, bodyExtended) + require.Equal(t, hash, hashExtended, "hash should be stable across later turns") + + bodyDifferent := []byte(`{"model":"gpt-5.4","messages":[{"role":"user","content":"Different question"}]}`) + hashDifferent := svc.GenerateSessionHash(c, bodyDifferent) + require.NotEqual(t, hash, hashDifferent, "different content should produce different hash") +} + +func TestOpenAIGatewayService_GenerateSessionHash_ExplicitSignalWinsOverContent(t *testing.T) { + gin.SetMode(gin.TestMode) + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/openai/v1/chat/completions", nil) + + svc := &OpenAIGatewayService{} + body := []byte(`{"model":"gpt-5.4","messages":[{"role":"user","content":"Hello"}]}`) + + contentHash := svc.GenerateSessionHash(c, body) + require.NotEmpty(t, contentHash) + + c.Request.Header.Set("session_id", "explicit-session") + explicitHash := svc.GenerateSessionHash(c, body) + require.NotEmpty(t, explicitHash) + require.NotEqual(t, contentHash, explicitHash, "explicit session_id should override content fallback") +} + +func TestOpenAIGatewayService_GenerateSessionHash_EmptyBodyStillEmpty(t *testing.T) { + gin.SetMode(gin.TestMode) + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/openai/v1/chat/completions", nil) + + svc := &OpenAIGatewayService{} + require.Empty(t, svc.GenerateSessionHash(c, []byte(`{}`))) + require.Empty(t, svc.GenerateSessionHash(c, nil)) +} + func (c stubConcurrencyCache) GetAccountWaitingCount(ctx context.Context, accountID int64) (int, error) { if c.waitCounts != nil { if count, ok := c.waitCounts[accountID]; ok { @@ -1797,7 +1851,7 @@ func TestExtractCodexFinalResponse_SampleReplay(t *testing.T) { require.Contains(t, string(finalResp), `"input_tokens":11`) } -func TestHandleOAuthSSEToJSON_CompletedEventReturnsJSON(t *testing.T) { +func TestHandleSSEToJSON_CompletedEventReturnsJSON(t *testing.T) { gin.SetMode(gin.TestMode) rec := httptest.NewRecorder() c, _ := gin.CreateTestContext(rec) @@ -1814,7 +1868,7 @@ func TestHandleOAuthSSEToJSON_CompletedEventReturnsJSON(t *testing.T) { `data: [DONE]`, }, "\n")) - usage, err := svc.handleOAuthSSEToJSON(resp, c, body, "gpt-4o", "gpt-4o") + usage, err := svc.handleSSEToJSON(resp, c, body, "gpt-4o", "gpt-4o") require.NoError(t, err) require.NotNil(t, usage) require.Equal(t, 7, usage.InputTokens) @@ -1826,7 +1880,7 @@ func TestHandleOAuthSSEToJSON_CompletedEventReturnsJSON(t *testing.T) { require.NotContains(t, rec.Body.String(), "data:") } -func TestHandleOAuthSSEToJSON_NoFinalResponseKeepsSSEBody(t *testing.T) { +func TestHandleSSEToJSON_NoFinalResponseKeepsSSEBody(t *testing.T) { gin.SetMode(gin.TestMode) rec := httptest.NewRecorder() c, _ := gin.CreateTestContext(rec) @@ -1842,7 +1896,7 @@ func TestHandleOAuthSSEToJSON_NoFinalResponseKeepsSSEBody(t *testing.T) { `data: [DONE]`, }, "\n")) - usage, err := svc.handleOAuthSSEToJSON(resp, c, body, "gpt-4o", "gpt-4o") + usage, err := svc.handleSSEToJSON(resp, c, body, "gpt-4o", "gpt-4o") require.NoError(t, err) require.NotNil(t, usage) require.Equal(t, 0, usage.InputTokens) @@ -1850,7 +1904,7 @@ func TestHandleOAuthSSEToJSON_NoFinalResponseKeepsSSEBody(t *testing.T) { require.Contains(t, rec.Body.String(), `data: {"type":"response.in_progress"`) } -func TestHandleOAuthSSEToJSON_ResponseFailedReturnsProtocolError(t *testing.T) { +func TestHandleSSEToJSON_ResponseFailedReturnsProtocolError(t *testing.T) { gin.SetMode(gin.TestMode) rec := httptest.NewRecorder() c, _ := gin.CreateTestContext(rec) @@ -1866,7 +1920,7 @@ func TestHandleOAuthSSEToJSON_ResponseFailedReturnsProtocolError(t *testing.T) { `data: [DONE]`, }, "\n")) - usage, err := svc.handleOAuthSSEToJSON(resp, c, body, "gpt-4o", "gpt-4o") + usage, err := svc.handleSSEToJSON(resp, c, body, "gpt-4o", "gpt-4o") require.Nil(t, usage) require.Error(t, err) require.Equal(t, http.StatusBadGateway, rec.Code) diff --git a/backend/internal/service/openai_messages_dispatch.go b/backend/internal/service/openai_messages_dispatch.go new file mode 100644 index 0000000000..f2c1ad3c1b --- /dev/null +++ b/backend/internal/service/openai_messages_dispatch.go @@ -0,0 +1,100 @@ +package service + +import "strings" + +const ( + defaultOpenAIMessagesDispatchOpusMappedModel = "gpt-5.4" + defaultOpenAIMessagesDispatchSonnetMappedModel = "gpt-5.3-codex" + defaultOpenAIMessagesDispatchHaikuMappedModel = "gpt-5.4-mini" +) + +func normalizeOpenAIMessagesDispatchMappedModel(model string) string { + model = NormalizeOpenAICompatRequestedModel(strings.TrimSpace(model)) + return strings.TrimSpace(model) +} + +func normalizeOpenAIMessagesDispatchModelConfig(cfg OpenAIMessagesDispatchModelConfig) OpenAIMessagesDispatchModelConfig { + out := OpenAIMessagesDispatchModelConfig{ + OpusMappedModel: normalizeOpenAIMessagesDispatchMappedModel(cfg.OpusMappedModel), + SonnetMappedModel: normalizeOpenAIMessagesDispatchMappedModel(cfg.SonnetMappedModel), + HaikuMappedModel: normalizeOpenAIMessagesDispatchMappedModel(cfg.HaikuMappedModel), + } + + if len(cfg.ExactModelMappings) > 0 { + out.ExactModelMappings = make(map[string]string, len(cfg.ExactModelMappings)) + for requestedModel, mappedModel := range cfg.ExactModelMappings { + requestedModel = strings.TrimSpace(requestedModel) + mappedModel = normalizeOpenAIMessagesDispatchMappedModel(mappedModel) + if requestedModel == "" || mappedModel == "" { + continue + } + out.ExactModelMappings[requestedModel] = mappedModel + } + if len(out.ExactModelMappings) == 0 { + out.ExactModelMappings = nil + } + } + + return out +} + +func claudeMessagesDispatchFamily(model string) string { + normalized := strings.ToLower(strings.TrimSpace(model)) + if !strings.HasPrefix(normalized, "claude") { + return "" + } + switch { + case strings.Contains(normalized, "opus"): + return "opus" + case strings.Contains(normalized, "sonnet"): + return "sonnet" + case strings.Contains(normalized, "haiku"): + return "haiku" + default: + return "" + } +} + +func (g *Group) ResolveMessagesDispatchModel(requestedModel string) string { + if g == nil { + return "" + } + requestedModel = strings.TrimSpace(requestedModel) + if requestedModel == "" { + return "" + } + + cfg := normalizeOpenAIMessagesDispatchModelConfig(g.MessagesDispatchModelConfig) + if mappedModel := strings.TrimSpace(cfg.ExactModelMappings[requestedModel]); mappedModel != "" { + return mappedModel + } + + switch claudeMessagesDispatchFamily(requestedModel) { + case "opus": + if mappedModel := strings.TrimSpace(cfg.OpusMappedModel); mappedModel != "" { + return mappedModel + } + return defaultOpenAIMessagesDispatchOpusMappedModel + case "sonnet": + if mappedModel := strings.TrimSpace(cfg.SonnetMappedModel); mappedModel != "" { + return mappedModel + } + return defaultOpenAIMessagesDispatchSonnetMappedModel + case "haiku": + if mappedModel := strings.TrimSpace(cfg.HaikuMappedModel); mappedModel != "" { + return mappedModel + } + return defaultOpenAIMessagesDispatchHaikuMappedModel + default: + return "" + } +} + +func sanitizeGroupMessagesDispatchFields(g *Group) { + if g == nil || g.Platform == PlatformOpenAI { + return + } + g.AllowMessagesDispatch = false + g.DefaultMappedModel = "" + g.MessagesDispatchModelConfig = OpenAIMessagesDispatchModelConfig{} +} diff --git a/backend/internal/service/openai_messages_dispatch_test.go b/backend/internal/service/openai_messages_dispatch_test.go new file mode 100644 index 0000000000..a625aaddd4 --- /dev/null +++ b/backend/internal/service/openai_messages_dispatch_test.go @@ -0,0 +1,27 @@ +package service + +import "testing" + +import "github.com/stretchr/testify/require" + +func TestNormalizeOpenAIMessagesDispatchModelConfig(t *testing.T) { + t.Parallel() + + cfg := normalizeOpenAIMessagesDispatchModelConfig(OpenAIMessagesDispatchModelConfig{ + OpusMappedModel: " gpt-5.4-high ", + SonnetMappedModel: "gpt-5.3-codex", + HaikuMappedModel: " gpt-5.4-mini-medium ", + ExactModelMappings: map[string]string{ + " claude-sonnet-4-5-20250929 ": " gpt-5.2-high ", + "": "gpt-5.4", + "claude-opus-4-6": " ", + }, + }) + + require.Equal(t, "gpt-5.4", cfg.OpusMappedModel) + require.Equal(t, "gpt-5.3-codex", cfg.SonnetMappedModel) + require.Equal(t, "gpt-5.4-mini", cfg.HaikuMappedModel) + require.Equal(t, map[string]string{ + "claude-sonnet-4-5-20250929": "gpt-5.2", + }, cfg.ExactModelMappings) +} diff --git a/backend/internal/service/openai_model_mapping_test.go b/backend/internal/service/openai_model_mapping_test.go index 5ce2602c1c..cda7e3698a 100644 --- a/backend/internal/service/openai_model_mapping_test.go +++ b/backend/internal/service/openai_model_mapping_test.go @@ -99,3 +99,39 @@ func TestNormalizeCodexModel(t *testing.T) { } } } + +func TestNormalizeOpenAIModelForUpstream(t *testing.T) { + tests := []struct { + name string + account *Account + model string + want string + }{ + { + name: "oauth keeps codex normalization behavior", + account: &Account{Type: AccountTypeOAuth}, + model: "gemini-3-flash-preview", + want: "gpt-5.1", + }, + { + name: "apikey preserves custom compatible model", + account: &Account{Type: AccountTypeAPIKey}, + model: "gemini-3-flash-preview", + want: "gemini-3-flash-preview", + }, + { + name: "apikey preserves official non codex model", + account: &Account{Type: AccountTypeAPIKey}, + model: "gpt-4.1", + want: "gpt-4.1", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if got := normalizeOpenAIModelForUpstream(tt.account, tt.model); got != tt.want { + t.Fatalf("normalizeOpenAIModelForUpstream(...) = %q, want %q", got, tt.want) + } + }) + } +} diff --git a/backend/internal/service/openai_oauth_passthrough_test.go b/backend/internal/service/openai_oauth_passthrough_test.go index 97fa218d92..69c9de42e3 100644 --- a/backend/internal/service/openai_oauth_passthrough_test.go +++ b/backend/internal/service/openai_oauth_passthrough_test.go @@ -48,6 +48,22 @@ func (u *httpUpstreamRecorder) DoWithTLS(req *http.Request, proxyURL string, acc return u.Do(req, proxyURL, accountID, accountConcurrency) } +type openAIPassthroughFailoverRepo struct { + stubOpenAIAccountRepo + rateLimitCalls []time.Time + overloadCalls []time.Time +} + +func (r *openAIPassthroughFailoverRepo) SetRateLimited(_ context.Context, _ int64, resetAt time.Time) error { + r.rateLimitCalls = append(r.rateLimitCalls, resetAt) + return nil +} + +func (r *openAIPassthroughFailoverRepo) SetOverloaded(_ context.Context, _ int64, until time.Time) error { + r.overloadCalls = append(r.overloadCalls, until) + return nil +} + var structuredLogCaptureMu sync.Mutex type inMemoryLogSink struct { @@ -527,6 +543,8 @@ func TestOpenAIGatewayService_OAuthPassthrough_UpstreamErrorIncludesPassthroughF _, err := svc.Forward(context.Background(), c, account, originalBody) require.Error(t, err) + require.True(t, c.Writer.Written(), "非 429/529 的 passthrough 错误应继续原样写回客户端") + require.Equal(t, http.StatusBadRequest, rec.Code) // should append an upstream error event with passthrough=true v, ok := c.Get(OpsUpstreamErrorsKey) @@ -535,55 +553,145 @@ func TestOpenAIGatewayService_OAuthPassthrough_UpstreamErrorIncludesPassthroughF require.True(t, ok) require.NotEmpty(t, arr) require.True(t, arr[len(arr)-1].Passthrough) + require.Equal(t, "http_error", arr[len(arr)-1].Kind) } -func TestOpenAIGatewayService_OAuthPassthrough_429PersistsRateLimit(t *testing.T) { +func TestOpenAIGatewayService_OpenAIPassthrough_429And529TriggerFailover(t *testing.T) { gin.SetMode(gin.TestMode) - - rec := httptest.NewRecorder() - c, _ := gin.CreateTestContext(rec) - c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", bytes.NewReader(nil)) - c.Request.Header.Set("User-Agent", "codex_cli_rs/0.1.0") - originalBody := []byte(`{"model":"gpt-5.2","stream":false,"instructions":"local-test-instructions","input":[{"type":"text","text":"hi"}]}`) - resetAt := time.Now().Add(7 * 24 * time.Hour).Unix() - resp := &http.Response{ - StatusCode: http.StatusTooManyRequests, - Header: http.Header{ - "Content-Type": []string{"application/json"}, - "x-request-id": []string{"rid-rate-limit"}, + + newAccount := func(accountType string) *Account { + account := &Account{ + ID: 123, + Name: "acc", + Platform: PlatformOpenAI, + Type: accountType, + Concurrency: 1, + Extra: map[string]any{"openai_passthrough": true}, + Status: StatusActive, + Schedulable: true, + RateMultiplier: f64p(1), + } + switch accountType { + case AccountTypeOAuth: + account.Credentials = map[string]any{"access_token": "oauth-token", "chatgpt_account_id": "chatgpt-acc"} + case AccountTypeAPIKey: + account.Credentials = map[string]any{"api_key": "sk-test"} + } + return account + } + + testCases := []struct { + name string + accountType string + statusCode int + body string + assertRepo func(t *testing.T, repo *openAIPassthroughFailoverRepo, start time.Time) + }{ + { + name: "oauth_429_rate_limit", + accountType: AccountTypeOAuth, + statusCode: http.StatusTooManyRequests, + body: func() string { + resetAt := time.Now().Add(7 * 24 * time.Hour).Unix() + return fmt.Sprintf(`{"error":{"message":"The usage limit has been reached","type":"usage_limit_reached","resets_at":%d}}`, resetAt) + }(), + assertRepo: func(t *testing.T, repo *openAIPassthroughFailoverRepo, _ time.Time) { + require.Len(t, repo.rateLimitCalls, 1) + require.Empty(t, repo.overloadCalls) + require.True(t, time.Until(repo.rateLimitCalls[0]) > 24*time.Hour) + }, + }, + { + name: "oauth_529_overload", + accountType: AccountTypeOAuth, + statusCode: 529, + body: `{"error":{"message":"server overloaded","type":"server_error"}}`, + assertRepo: func(t *testing.T, repo *openAIPassthroughFailoverRepo, start time.Time) { + require.Empty(t, repo.rateLimitCalls) + require.Len(t, repo.overloadCalls, 1) + require.WithinDuration(t, start.Add(10*time.Minute), repo.overloadCalls[0], 5*time.Second) + }, + }, + { + name: "apikey_429_rate_limit", + accountType: AccountTypeAPIKey, + statusCode: http.StatusTooManyRequests, + body: func() string { + resetAt := time.Now().Add(7 * 24 * time.Hour).Unix() + return fmt.Sprintf(`{"error":{"message":"The usage limit has been reached","type":"usage_limit_reached","resets_at":%d}}`, resetAt) + }(), + assertRepo: func(t *testing.T, repo *openAIPassthroughFailoverRepo, _ time.Time) { + require.Len(t, repo.rateLimitCalls, 1) + require.Empty(t, repo.overloadCalls) + require.True(t, time.Until(repo.rateLimitCalls[0]) > 24*time.Hour) + }, + }, + { + name: "apikey_529_overload", + accountType: AccountTypeAPIKey, + statusCode: 529, + body: `{"error":{"message":"server overloaded","type":"server_error"}}`, + assertRepo: func(t *testing.T, repo *openAIPassthroughFailoverRepo, start time.Time) { + require.Empty(t, repo.rateLimitCalls) + require.Len(t, repo.overloadCalls, 1) + require.WithinDuration(t, start.Add(10*time.Minute), repo.overloadCalls[0], 5*time.Second) + }, }, - Body: io.NopCloser(strings.NewReader(fmt.Sprintf(`{"error":{"message":"The usage limit has been reached","type":"usage_limit_reached","resets_at":%d}}`, resetAt))), - } - upstream := &httpUpstreamRecorder{resp: resp} - repo := &openAIWSRateLimitSignalRepo{} - rateSvc := &RateLimitService{accountRepo: repo} - - svc := &OpenAIGatewayService{ - cfg: &config.Config{Gateway: config.GatewayConfig{ForceCodexCLI: false}}, - httpUpstream: upstream, - rateLimitService: rateSvc, } - account := &Account{ - ID: 123, - Name: "acc", - Platform: PlatformOpenAI, - Type: AccountTypeOAuth, - Concurrency: 1, - Credentials: map[string]any{"access_token": "oauth-token", "chatgpt_account_id": "chatgpt-acc"}, - Extra: map[string]any{"openai_passthrough": true}, - Status: StatusActive, - Schedulable: true, - RateMultiplier: f64p(1), - } + for _, tc := range testCases { + t.Run(tc.name, func(t *testing.T) { + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", bytes.NewReader(nil)) + c.Request.Header.Set("User-Agent", "codex_cli_rs/0.1.0") - _, err := svc.Forward(context.Background(), c, account, originalBody) - require.Error(t, err) - require.Equal(t, http.StatusTooManyRequests, rec.Code) - require.Contains(t, rec.Body.String(), "usage_limit_reached") - require.Len(t, repo.rateLimitCalls, 1) - require.WithinDuration(t, time.Unix(resetAt, 0), repo.rateLimitCalls[0], 2*time.Second) + resp := &http.Response{ + StatusCode: tc.statusCode, + Header: http.Header{ + "Content-Type": []string{"application/json"}, + "x-request-id": []string{"rid-failover"}, + }, + Body: io.NopCloser(strings.NewReader(tc.body)), + } + upstream := &httpUpstreamRecorder{resp: resp} + repo := &openAIPassthroughFailoverRepo{} + rateSvc := &RateLimitService{ + accountRepo: repo, + cfg: &config.Config{ + RateLimit: config.RateLimitConfig{OverloadCooldownMinutes: 10}, + }, + } + + svc := &OpenAIGatewayService{ + cfg: &config.Config{Gateway: config.GatewayConfig{ForceCodexCLI: false}}, + httpUpstream: upstream, + rateLimitService: rateSvc, + } + + account := newAccount(tc.accountType) + start := time.Now() + _, err := svc.Forward(context.Background(), c, account, originalBody) + require.Error(t, err) + + var failoverErr *UpstreamFailoverError + require.ErrorAs(t, err, &failoverErr) + require.Equal(t, tc.statusCode, failoverErr.StatusCode) + require.False(t, c.Writer.Written(), "429/529 passthrough 应返回 failover 错误给上层换号,而不是直接向客户端写响应") + + v, ok := c.Get(OpsUpstreamErrorsKey) + require.True(t, ok) + arr, ok := v.([]*OpsUpstreamErrorEvent) + require.True(t, ok) + require.NotEmpty(t, arr) + require.True(t, arr[len(arr)-1].Passthrough) + require.Equal(t, "failover", arr[len(arr)-1].Kind) + require.Equal(t, tc.statusCode, arr[len(arr)-1].UpstreamStatusCode) + + tc.assertRepo(t, repo, start) + }) + } } func TestOpenAIGatewayService_OAuthPassthrough_NonCodexUAFallbackToCodexUA(t *testing.T) { diff --git a/backend/internal/service/openai_ws_forwarder.go b/backend/internal/service/openai_ws_forwarder.go index 6d45baab36..83849bf35c 100644 --- a/backend/internal/service/openai_ws_forwarder.go +++ b/backend/internal/service/openai_ws_forwarder.go @@ -2515,7 +2515,7 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient( } normalized = next } - upstreamModel := normalizeCodexModel(account.GetMappedModel(originalModel)) + upstreamModel := normalizeOpenAIModelForUpstream(account, account.GetMappedModel(originalModel)) if upstreamModel != originalModel { next, setErr := applyPayloadMutation(normalized, "model", upstreamModel) if setErr != nil { @@ -2773,7 +2773,7 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient( mappedModel := "" var mappedModelBytes []byte if originalModel != "" { - mappedModel = normalizeCodexModel(account.GetMappedModel(originalModel)) + mappedModel = normalizeOpenAIModelForUpstream(account, account.GetMappedModel(originalModel)) needModelReplace = mappedModel != "" && mappedModel != originalModel if needModelReplace { mappedModelBytes = []byte(mappedModel) diff --git a/backend/internal/service/openai_ws_ratelimit_signal_test.go b/backend/internal/service/openai_ws_ratelimit_signal_test.go index ffe7915262..6313d0c08e 100644 --- a/backend/internal/service/openai_ws_ratelimit_signal_test.go +++ b/backend/internal/service/openai_ws_ratelimit_signal_test.go @@ -492,7 +492,7 @@ func TestAdminService_ListAccounts_ExhaustedCodexExtraReturnsRateLimitedAccount( } svc := &adminServiceImpl{accountRepo: repo} - accounts, total, err := svc.ListAccounts(context.Background(), 1, 20, PlatformOpenAI, AccountTypeOAuth, "", "", 0, "") + accounts, total, err := svc.ListAccounts(context.Background(), 1, 20, PlatformOpenAI, AccountTypeOAuth, "", "", 0, "", "", "") require.NoError(t, err) require.Equal(t, int64(1), total) require.Len(t, accounts, 1) diff --git a/backend/internal/service/ops_service.go b/backend/internal/service/ops_service.go index 29f0aa8b50..cd3974a00f 100644 --- a/backend/internal/service/ops_service.go +++ b/backend/internal/service/ops_service.go @@ -16,7 +16,7 @@ import ( var ErrOpsDisabled = infraerrors.NotFound("OPS_DISABLED", "Ops monitoring is disabled") const ( - opsMaxStoredRequestBodyBytes = 10 * 1024 + opsMaxStoredRequestBodyBytes = 256 * 1024 opsMaxStoredErrorBodyBytes = 20 * 1024 ) diff --git a/backend/internal/service/scheduler_snapshot_hydration_test.go b/backend/internal/service/scheduler_snapshot_hydration_test.go new file mode 100644 index 0000000000..5c0b289b7a --- /dev/null +++ b/backend/internal/service/scheduler_snapshot_hydration_test.go @@ -0,0 +1,159 @@ +//go:build unit + +package service + +import ( + "context" + "testing" + "time" +) + +type snapshotHydrationCache struct { + snapshot []*Account + accounts map[int64]*Account +} + +func (c *snapshotHydrationCache) GetSnapshot(ctx context.Context, bucket SchedulerBucket) ([]*Account, bool, error) { + return c.snapshot, true, nil +} + +func (c *snapshotHydrationCache) SetSnapshot(ctx context.Context, bucket SchedulerBucket, accounts []Account) error { + return nil +} + +func (c *snapshotHydrationCache) GetAccount(ctx context.Context, accountID int64) (*Account, error) { + if c.accounts == nil { + return nil, nil + } + return c.accounts[accountID], nil +} + +func (c *snapshotHydrationCache) SetAccount(ctx context.Context, account *Account) error { + return nil +} + +func (c *snapshotHydrationCache) DeleteAccount(ctx context.Context, accountID int64) error { + return nil +} + +func (c *snapshotHydrationCache) UpdateLastUsed(ctx context.Context, updates map[int64]time.Time) error { + return nil +} + +func (c *snapshotHydrationCache) TryLockBucket(ctx context.Context, bucket SchedulerBucket, ttl time.Duration) (bool, error) { + return true, nil +} + +func (c *snapshotHydrationCache) ListBuckets(ctx context.Context) ([]SchedulerBucket, error) { + return nil, nil +} + +func (c *snapshotHydrationCache) GetOutboxWatermark(ctx context.Context) (int64, error) { + return 0, nil +} + +func (c *snapshotHydrationCache) SetOutboxWatermark(ctx context.Context, id int64) error { + return nil +} + +func TestOpenAISelectAccountWithLoadAwareness_HydratesSelectedAccountFromSchedulerSnapshot(t *testing.T) { + cache := &snapshotHydrationCache{ + snapshot: []*Account{ + { + ID: 1, + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Status: StatusActive, + Schedulable: true, + Concurrency: 1, + Priority: 1, + Credentials: map[string]any{ + "model_mapping": map[string]any{ + "gpt-4": "gpt-4", + }, + }, + }, + }, + accounts: map[int64]*Account{ + 1: { + ID: 1, + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Status: StatusActive, + Schedulable: true, + Concurrency: 1, + Priority: 1, + Credentials: map[string]any{ + "api_key": "sk-live", + "model_mapping": map[string]any{"gpt-4": "gpt-4"}, + }, + }, + }, + } + + schedulerSnapshot := NewSchedulerSnapshotService(cache, nil, nil, nil, nil) + groupID := int64(2) + svc := &OpenAIGatewayService{ + schedulerSnapshot: schedulerSnapshot, + cache: &stubGatewayCache{}, + } + + selection, err := svc.SelectAccountWithLoadAwareness(context.Background(), &groupID, "", "gpt-4", nil) + if err != nil { + t.Fatalf("SelectAccountWithLoadAwareness error: %v", err) + } + if selection == nil || selection.Account == nil { + t.Fatalf("expected selected account") + } + if got := selection.Account.GetOpenAIApiKey(); got != "sk-live" { + t.Fatalf("expected hydrated api key, got %q", got) + } +} + +func TestGatewaySelectAccountWithLoadAwareness_HydratesSelectedAccountFromSchedulerSnapshot(t *testing.T) { + cache := &snapshotHydrationCache{ + snapshot: []*Account{ + { + ID: 9, + Platform: PlatformAnthropic, + Type: AccountTypeAPIKey, + Status: StatusActive, + Schedulable: true, + Concurrency: 1, + Priority: 1, + }, + }, + accounts: map[int64]*Account{ + 9: { + ID: 9, + Platform: PlatformAnthropic, + Type: AccountTypeAPIKey, + Status: StatusActive, + Schedulable: true, + Concurrency: 1, + Priority: 1, + Credentials: map[string]any{ + "api_key": "anthropic-live-key", + }, + }, + }, + } + + schedulerSnapshot := NewSchedulerSnapshotService(cache, nil, nil, nil, nil) + svc := &GatewayService{ + schedulerSnapshot: schedulerSnapshot, + cache: &mockGatewayCacheForPlatform{}, + cfg: testConfig(), + } + + result, err := svc.SelectAccountWithLoadAwareness(context.Background(), nil, "", "claude-3-5-sonnet-20241022", nil, "", 0) + if err != nil { + t.Fatalf("SelectAccountWithLoadAwareness error: %v", err) + } + if result == nil || result.Account == nil { + t.Fatalf("expected selected account") + } + if got := result.Account.GetCredential("api_key"); got != "anthropic-live-key" { + t.Fatalf("expected hydrated api key, got %q", got) + } +} diff --git a/backend/internal/service/setting_service.go b/backend/internal/service/setting_service.go index 080fa12244..e8ae795b0a 100644 --- a/backend/internal/service/setting_service.go +++ b/backend/internal/service/setting_service.go @@ -9,6 +9,7 @@ import ( "fmt" "log/slog" "net/url" + "sort" "strconv" "strings" "sync/atomic" @@ -16,6 +17,7 @@ import ( "github.com/Wei-Shaw/sub2api/internal/config" infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" + "github.com/imroc/req/v3" "golang.org/x/sync/singleflight" ) @@ -81,6 +83,7 @@ const backendModeDBTimeout = 5 * time.Second type cachedGatewayForwardingSettings struct { fingerprintUnification bool metadataPassthrough bool + cchSigning bool expiresAt int64 // unix nano } @@ -159,11 +162,15 @@ func (s *SettingService) GetPublicSettings(ctx context.Context) (*PublicSettings SettingKeyHideCcsImportButton, SettingKeyPurchaseSubscriptionEnabled, SettingKeyPurchaseSubscriptionURL, + SettingKeyTableDefaultPageSize, + SettingKeyTablePageSizeOptions, SettingKeyCustomMenuItems, SettingKeyCustomEndpoints, SettingKeyLinuxDoConnectEnabled, SettingKeyBackendModeEnabled, SettingPaymentEnabled, + SettingKeyOIDCConnectEnabled, + SettingKeyOIDCConnectProviderName, } settings, err := s.settingRepo.GetMultiple(ctx, keys) @@ -177,6 +184,19 @@ func (s *SettingService) GetPublicSettings(ctx context.Context) (*PublicSettings } else { linuxDoEnabled = s.cfg != nil && s.cfg.LinuxDo.Enabled } + oidcEnabled := false + if raw, ok := settings[SettingKeyOIDCConnectEnabled]; ok { + oidcEnabled = raw == "true" + } else { + oidcEnabled = s.cfg != nil && s.cfg.OIDC.Enabled + } + oidcProviderName := strings.TrimSpace(settings[SettingKeyOIDCConnectProviderName]) + if oidcProviderName == "" && s.cfg != nil { + oidcProviderName = strings.TrimSpace(s.cfg.OIDC.ProviderName) + } + if oidcProviderName == "" { + oidcProviderName = "OIDC" + } // Password reset requires email verification to be enabled emailVerifyEnabled := settings[SettingKeyEmailVerifyEnabled] == "true" @@ -184,6 +204,10 @@ func (s *SettingService) GetPublicSettings(ctx context.Context) (*PublicSettings registrationEmailSuffixWhitelist := ParseRegistrationEmailSuffixWhitelist( settings[SettingKeyRegistrationEmailSuffixWhitelist], ) + tableDefaultPageSize, tablePageSizeOptions := parseTablePreferences( + settings[SettingKeyTableDefaultPageSize], + settings[SettingKeyTablePageSizeOptions], + ) return &PublicSettings{ RegistrationEnabled: settings[SettingKeyRegistrationEnabled] == "true", @@ -205,11 +229,15 @@ func (s *SettingService) GetPublicSettings(ctx context.Context) (*PublicSettings HideCcsImportButton: settings[SettingKeyHideCcsImportButton] == "true", PurchaseSubscriptionEnabled: settings[SettingKeyPurchaseSubscriptionEnabled] == "true", PurchaseSubscriptionURL: strings.TrimSpace(settings[SettingKeyPurchaseSubscriptionURL]), + TableDefaultPageSize: tableDefaultPageSize, + TablePageSizeOptions: tablePageSizeOptions, CustomMenuItems: settings[SettingKeyCustomMenuItems], CustomEndpoints: settings[SettingKeyCustomEndpoints], LinuxDoOAuthEnabled: linuxDoEnabled, BackendModeEnabled: settings[SettingKeyBackendModeEnabled] == "true", PaymentEnabled: settings[SettingPaymentEnabled] == "true", + OIDCOAuthEnabled: oidcEnabled, + OIDCOAuthProviderName: oidcProviderName, }, nil } @@ -253,11 +281,15 @@ func (s *SettingService) GetPublicSettingsForInjection(ctx context.Context) (any HideCcsImportButton bool `json:"hide_ccs_import_button"` PurchaseSubscriptionEnabled bool `json:"purchase_subscription_enabled"` PurchaseSubscriptionURL string `json:"purchase_subscription_url,omitempty"` + TableDefaultPageSize int `json:"table_default_page_size"` + TablePageSizeOptions []int `json:"table_page_size_options"` CustomMenuItems json.RawMessage `json:"custom_menu_items"` CustomEndpoints json.RawMessage `json:"custom_endpoints"` LinuxDoOAuthEnabled bool `json:"linuxdo_oauth_enabled"` BackendModeEnabled bool `json:"backend_mode_enabled"` PaymentEnabled bool `json:"payment_enabled"` + OIDCOAuthEnabled bool `json:"oidc_oauth_enabled"` + OIDCOAuthProviderName string `json:"oidc_oauth_provider_name"` Version string `json:"version,omitempty"` }{ RegistrationEnabled: settings.RegistrationEnabled, @@ -279,11 +311,15 @@ func (s *SettingService) GetPublicSettingsForInjection(ctx context.Context) (any HideCcsImportButton: settings.HideCcsImportButton, PurchaseSubscriptionEnabled: settings.PurchaseSubscriptionEnabled, PurchaseSubscriptionURL: settings.PurchaseSubscriptionURL, + TableDefaultPageSize: settings.TableDefaultPageSize, + TablePageSizeOptions: settings.TablePageSizeOptions, CustomMenuItems: filterUserVisibleMenuItems(settings.CustomMenuItems), CustomEndpoints: safeRawJSONArray(settings.CustomEndpoints), LinuxDoOAuthEnabled: settings.LinuxDoOAuthEnabled, BackendModeEnabled: settings.BackendModeEnabled, PaymentEnabled: settings.PaymentEnabled, + OIDCOAuthEnabled: settings.OIDCOAuthEnabled, + OIDCOAuthProviderName: settings.OIDCOAuthProviderName, Version: s.version, }, nil } @@ -336,8 +372,8 @@ func safeRawJSONArray(raw string) json.RawMessage { return json.RawMessage("[]") } -// GetFrameSrcOrigins returns deduplicated http(s) origins from purchase_subscription_url -// and all custom_menu_items URLs. Used by the router layer for CSP frame-src injection. +// GetFrameSrcOrigins returns deduplicated http(s) origins from home_content URL, +// purchase_subscription_url, and all custom_menu_items URLs. Used by the router layer for CSP frame-src injection. func (s *SettingService) GetFrameSrcOrigins(ctx context.Context) ([]string, error) { settings, err := s.GetPublicSettings(ctx) if err != nil { @@ -356,6 +392,9 @@ func (s *SettingService) GetFrameSrcOrigins(ctx context.Context) ([]string, erro } } + // home content URL (when home_content is set to a URL for iframe embedding) + addOrigin(settings.HomeContent) + // purchase subscription URL if settings.PurchaseSubscriptionEnabled { addOrigin(settings.PurchaseSubscriptionURL) @@ -463,6 +502,32 @@ func (s *SettingService) UpdateSettings(ctx context.Context, settings *SystemSet updates[SettingKeyLinuxDoConnectClientSecret] = settings.LinuxDoConnectClientSecret } + // Generic OIDC OAuth 登录 + updates[SettingKeyOIDCConnectEnabled] = strconv.FormatBool(settings.OIDCConnectEnabled) + updates[SettingKeyOIDCConnectProviderName] = settings.OIDCConnectProviderName + updates[SettingKeyOIDCConnectClientID] = settings.OIDCConnectClientID + updates[SettingKeyOIDCConnectIssuerURL] = settings.OIDCConnectIssuerURL + updates[SettingKeyOIDCConnectDiscoveryURL] = settings.OIDCConnectDiscoveryURL + updates[SettingKeyOIDCConnectAuthorizeURL] = settings.OIDCConnectAuthorizeURL + updates[SettingKeyOIDCConnectTokenURL] = settings.OIDCConnectTokenURL + updates[SettingKeyOIDCConnectUserInfoURL] = settings.OIDCConnectUserInfoURL + updates[SettingKeyOIDCConnectJWKSURL] = settings.OIDCConnectJWKSURL + updates[SettingKeyOIDCConnectScopes] = settings.OIDCConnectScopes + updates[SettingKeyOIDCConnectRedirectURL] = settings.OIDCConnectRedirectURL + updates[SettingKeyOIDCConnectFrontendRedirectURL] = settings.OIDCConnectFrontendRedirectURL + updates[SettingKeyOIDCConnectTokenAuthMethod] = settings.OIDCConnectTokenAuthMethod + updates[SettingKeyOIDCConnectUsePKCE] = strconv.FormatBool(settings.OIDCConnectUsePKCE) + updates[SettingKeyOIDCConnectValidateIDToken] = strconv.FormatBool(settings.OIDCConnectValidateIDToken) + updates[SettingKeyOIDCConnectAllowedSigningAlgs] = settings.OIDCConnectAllowedSigningAlgs + updates[SettingKeyOIDCConnectClockSkewSeconds] = strconv.Itoa(settings.OIDCConnectClockSkewSeconds) + updates[SettingKeyOIDCConnectRequireEmailVerified] = strconv.FormatBool(settings.OIDCConnectRequireEmailVerified) + updates[SettingKeyOIDCConnectUserInfoEmailPath] = settings.OIDCConnectUserInfoEmailPath + updates[SettingKeyOIDCConnectUserInfoIDPath] = settings.OIDCConnectUserInfoIDPath + updates[SettingKeyOIDCConnectUserInfoUsernamePath] = settings.OIDCConnectUserInfoUsernamePath + if settings.OIDCConnectClientSecret != "" { + updates[SettingKeyOIDCConnectClientSecret] = settings.OIDCConnectClientSecret + } + // OEM设置 updates[SettingKeySiteName] = settings.SiteName updates[SettingKeySiteLogo] = settings.SiteLogo @@ -474,6 +539,16 @@ func (s *SettingService) UpdateSettings(ctx context.Context, settings *SystemSet updates[SettingKeyHideCcsImportButton] = strconv.FormatBool(settings.HideCcsImportButton) updates[SettingKeyPurchaseSubscriptionEnabled] = strconv.FormatBool(settings.PurchaseSubscriptionEnabled) updates[SettingKeyPurchaseSubscriptionURL] = strings.TrimSpace(settings.PurchaseSubscriptionURL) + tableDefaultPageSize, tablePageSizeOptions := normalizeTablePreferences( + settings.TableDefaultPageSize, + settings.TablePageSizeOptions, + ) + updates[SettingKeyTableDefaultPageSize] = strconv.Itoa(tableDefaultPageSize) + tablePageSizeOptionsJSON, err := json.Marshal(tablePageSizeOptions) + if err != nil { + return fmt.Errorf("marshal table page size options: %w", err) + } + updates[SettingKeyTablePageSizeOptions] = string(tablePageSizeOptionsJSON) updates[SettingKeyCustomMenuItems] = settings.CustomMenuItems updates[SettingKeyCustomEndpoints] = settings.CustomEndpoints @@ -518,6 +593,7 @@ func (s *SettingService) UpdateSettings(ctx context.Context, settings *SystemSet // Gateway forwarding behavior updates[SettingKeyEnableFingerprintUnification] = strconv.FormatBool(settings.EnableFingerprintUnification) updates[SettingKeyEnableMetadataPassthrough] = strconv.FormatBool(settings.EnableMetadataPassthrough) + updates[SettingKeyEnableCCHSigning] = strconv.FormatBool(settings.EnableCCHSigning) err = s.settingRepo.SetMultiple(ctx, updates) if err == nil { @@ -537,6 +613,7 @@ func (s *SettingService) UpdateSettings(ctx context.Context, settings *SystemSet gatewayForwardingCache.Store(&cachedGatewayForwardingSettings{ fingerprintUnification: settings.EnableFingerprintUnification, metadataPassthrough: settings.EnableMetadataPassthrough, + cchSigning: settings.EnableCCHSigning, expiresAt: time.Now().Add(gatewayForwardingCacheTTL).UnixNano(), }) if s.onUpdate != nil { @@ -643,20 +720,20 @@ func (s *SettingService) IsBackendModeEnabled(ctx context.Context) bool { // GetGatewayForwardingSettings returns cached gateway forwarding settings. // Uses in-process atomic.Value cache with 60s TTL, zero-lock hot path. -// Returns (fingerprintUnification, metadataPassthrough). -func (s *SettingService) GetGatewayForwardingSettings(ctx context.Context) (fingerprintUnification, metadataPassthrough bool) { +// Returns (fingerprintUnification, metadataPassthrough, cchSigning). +func (s *SettingService) GetGatewayForwardingSettings(ctx context.Context) (fingerprintUnification, metadataPassthrough, cchSigning bool) { if cached, ok := gatewayForwardingCache.Load().(*cachedGatewayForwardingSettings); ok && cached != nil { if time.Now().UnixNano() < cached.expiresAt { - return cached.fingerprintUnification, cached.metadataPassthrough + return cached.fingerprintUnification, cached.metadataPassthrough, cached.cchSigning } } type gwfResult struct { - fp, mp bool + fp, mp, cch bool } val, _, _ := gatewayForwardingSF.Do("gateway_forwarding", func() (any, error) { if cached, ok := gatewayForwardingCache.Load().(*cachedGatewayForwardingSettings); ok && cached != nil { if time.Now().UnixNano() < cached.expiresAt { - return gwfResult{cached.fingerprintUnification, cached.metadataPassthrough}, nil + return gwfResult{cached.fingerprintUnification, cached.metadataPassthrough, cached.cchSigning}, nil } } dbCtx, cancel := context.WithTimeout(context.WithoutCancel(ctx), gatewayForwardingDBTimeout) @@ -664,32 +741,36 @@ func (s *SettingService) GetGatewayForwardingSettings(ctx context.Context) (fing values, err := s.settingRepo.GetMultiple(dbCtx, []string{ SettingKeyEnableFingerprintUnification, SettingKeyEnableMetadataPassthrough, + SettingKeyEnableCCHSigning, }) if err != nil { slog.Warn("failed to get gateway forwarding settings", "error", err) gatewayForwardingCache.Store(&cachedGatewayForwardingSettings{ fingerprintUnification: true, metadataPassthrough: false, + cchSigning: false, expiresAt: time.Now().Add(gatewayForwardingErrorTTL).UnixNano(), }) - return gwfResult{true, false}, nil + return gwfResult{true, false, false}, nil } fp := true if v, ok := values[SettingKeyEnableFingerprintUnification]; ok && v != "" { fp = v == "true" } mp := values[SettingKeyEnableMetadataPassthrough] == "true" + cch := values[SettingKeyEnableCCHSigning] == "true" gatewayForwardingCache.Store(&cachedGatewayForwardingSettings{ fingerprintUnification: fp, metadataPassthrough: mp, + cchSigning: cch, expiresAt: time.Now().Add(gatewayForwardingCacheTTL).UnixNano(), }) - return gwfResult{fp, mp}, nil + return gwfResult{fp, mp, cch}, nil }) if r, ok := val.(gwfResult); ok { - return r.fp, r.mp + return r.fp, r.mp, r.cch } - return true, false // fail-open defaults + return true, false, false // fail-open defaults } // IsEmailVerifyEnabled 检查是否开启邮件验证 @@ -821,8 +902,12 @@ func (s *SettingService) InitializeDefaultSettings(ctx context.Context) error { SettingKeySiteLogo: "", SettingKeyPurchaseSubscriptionEnabled: "false", SettingKeyPurchaseSubscriptionURL: "", + SettingKeyTableDefaultPageSize: "20", + SettingKeyTablePageSizeOptions: "[10,20,50,100]", SettingKeyCustomMenuItems: "[]", SettingKeyCustomEndpoints: "[]", + SettingKeyOIDCConnectEnabled: "false", + SettingKeyOIDCConnectProviderName: "OIDC", SettingKeyDefaultConcurrency: strconv.Itoa(s.cfg.Default.UserConcurrency), SettingKeyDefaultBalance: strconv.FormatFloat(s.cfg.Default.UserBalance, 'f', 8, 64), SettingKeyDefaultSubscriptions: "[]", @@ -890,6 +975,10 @@ func (s *SettingService) parseSettings(settings map[string]string) *SystemSettin CustomEndpoints: settings[SettingKeyCustomEndpoints], BackendModeEnabled: settings[SettingKeyBackendModeEnabled] == "true", } + result.TableDefaultPageSize, result.TablePageSizeOptions = parseTablePreferences( + settings[SettingKeyTableDefaultPageSize], + settings[SettingKeyTablePageSizeOptions], + ) // 解析整数类型 if port, err := strconv.Atoi(settings[SettingKeySMTPPort]); err == nil { @@ -948,6 +1037,138 @@ func (s *SettingService) parseSettings(settings map[string]string) *SystemSettin } result.LinuxDoConnectClientSecretConfigured = result.LinuxDoConnectClientSecret != "" + // Generic OIDC 设置: + // - 兼容 config.yaml/env + // - 支持后台系统设置覆盖并持久化(存储于 DB) + oidcBase := config.OIDCConnectConfig{} + if s.cfg != nil { + oidcBase = s.cfg.OIDC + } + + if raw, ok := settings[SettingKeyOIDCConnectEnabled]; ok { + result.OIDCConnectEnabled = raw == "true" + } else { + result.OIDCConnectEnabled = oidcBase.Enabled + } + + if v, ok := settings[SettingKeyOIDCConnectProviderName]; ok && strings.TrimSpace(v) != "" { + result.OIDCConnectProviderName = strings.TrimSpace(v) + } else { + result.OIDCConnectProviderName = strings.TrimSpace(oidcBase.ProviderName) + } + if result.OIDCConnectProviderName == "" { + result.OIDCConnectProviderName = "OIDC" + } + + if v, ok := settings[SettingKeyOIDCConnectClientID]; ok && strings.TrimSpace(v) != "" { + result.OIDCConnectClientID = strings.TrimSpace(v) + } else { + result.OIDCConnectClientID = strings.TrimSpace(oidcBase.ClientID) + } + if v, ok := settings[SettingKeyOIDCConnectIssuerURL]; ok && strings.TrimSpace(v) != "" { + result.OIDCConnectIssuerURL = strings.TrimSpace(v) + } else { + result.OIDCConnectIssuerURL = strings.TrimSpace(oidcBase.IssuerURL) + } + if v, ok := settings[SettingKeyOIDCConnectDiscoveryURL]; ok && strings.TrimSpace(v) != "" { + result.OIDCConnectDiscoveryURL = strings.TrimSpace(v) + } else { + result.OIDCConnectDiscoveryURL = strings.TrimSpace(oidcBase.DiscoveryURL) + } + if v, ok := settings[SettingKeyOIDCConnectAuthorizeURL]; ok && strings.TrimSpace(v) != "" { + result.OIDCConnectAuthorizeURL = strings.TrimSpace(v) + } else { + result.OIDCConnectAuthorizeURL = strings.TrimSpace(oidcBase.AuthorizeURL) + } + if v, ok := settings[SettingKeyOIDCConnectTokenURL]; ok && strings.TrimSpace(v) != "" { + result.OIDCConnectTokenURL = strings.TrimSpace(v) + } else { + result.OIDCConnectTokenURL = strings.TrimSpace(oidcBase.TokenURL) + } + if v, ok := settings[SettingKeyOIDCConnectUserInfoURL]; ok && strings.TrimSpace(v) != "" { + result.OIDCConnectUserInfoURL = strings.TrimSpace(v) + } else { + result.OIDCConnectUserInfoURL = strings.TrimSpace(oidcBase.UserInfoURL) + } + if v, ok := settings[SettingKeyOIDCConnectJWKSURL]; ok && strings.TrimSpace(v) != "" { + result.OIDCConnectJWKSURL = strings.TrimSpace(v) + } else { + result.OIDCConnectJWKSURL = strings.TrimSpace(oidcBase.JWKSURL) + } + if v, ok := settings[SettingKeyOIDCConnectScopes]; ok && strings.TrimSpace(v) != "" { + result.OIDCConnectScopes = strings.TrimSpace(v) + } else { + result.OIDCConnectScopes = strings.TrimSpace(oidcBase.Scopes) + } + if v, ok := settings[SettingKeyOIDCConnectRedirectURL]; ok && strings.TrimSpace(v) != "" { + result.OIDCConnectRedirectURL = strings.TrimSpace(v) + } else { + result.OIDCConnectRedirectURL = strings.TrimSpace(oidcBase.RedirectURL) + } + if v, ok := settings[SettingKeyOIDCConnectFrontendRedirectURL]; ok && strings.TrimSpace(v) != "" { + result.OIDCConnectFrontendRedirectURL = strings.TrimSpace(v) + } else { + result.OIDCConnectFrontendRedirectURL = strings.TrimSpace(oidcBase.FrontendRedirectURL) + } + if v, ok := settings[SettingKeyOIDCConnectTokenAuthMethod]; ok && strings.TrimSpace(v) != "" { + result.OIDCConnectTokenAuthMethod = strings.ToLower(strings.TrimSpace(v)) + } else { + result.OIDCConnectTokenAuthMethod = strings.ToLower(strings.TrimSpace(oidcBase.TokenAuthMethod)) + } + if raw, ok := settings[SettingKeyOIDCConnectUsePKCE]; ok { + result.OIDCConnectUsePKCE = raw == "true" + } else { + result.OIDCConnectUsePKCE = oidcBase.UsePKCE + } + if raw, ok := settings[SettingKeyOIDCConnectValidateIDToken]; ok { + result.OIDCConnectValidateIDToken = raw == "true" + } else { + result.OIDCConnectValidateIDToken = oidcBase.ValidateIDToken + } + if v, ok := settings[SettingKeyOIDCConnectAllowedSigningAlgs]; ok && strings.TrimSpace(v) != "" { + result.OIDCConnectAllowedSigningAlgs = strings.TrimSpace(v) + } else { + result.OIDCConnectAllowedSigningAlgs = strings.TrimSpace(oidcBase.AllowedSigningAlgs) + } + clockSkewSet := false + if raw, ok := settings[SettingKeyOIDCConnectClockSkewSeconds]; ok && strings.TrimSpace(raw) != "" { + if parsed, err := strconv.Atoi(strings.TrimSpace(raw)); err == nil { + result.OIDCConnectClockSkewSeconds = parsed + clockSkewSet = true + } + } + if !clockSkewSet { + result.OIDCConnectClockSkewSeconds = oidcBase.ClockSkewSeconds + } + if !clockSkewSet && result.OIDCConnectClockSkewSeconds == 0 { + result.OIDCConnectClockSkewSeconds = 120 + } + if raw, ok := settings[SettingKeyOIDCConnectRequireEmailVerified]; ok { + result.OIDCConnectRequireEmailVerified = raw == "true" + } else { + result.OIDCConnectRequireEmailVerified = oidcBase.RequireEmailVerified + } + if v, ok := settings[SettingKeyOIDCConnectUserInfoEmailPath]; ok { + result.OIDCConnectUserInfoEmailPath = strings.TrimSpace(v) + } else { + result.OIDCConnectUserInfoEmailPath = strings.TrimSpace(oidcBase.UserInfoEmailPath) + } + if v, ok := settings[SettingKeyOIDCConnectUserInfoIDPath]; ok { + result.OIDCConnectUserInfoIDPath = strings.TrimSpace(v) + } else { + result.OIDCConnectUserInfoIDPath = strings.TrimSpace(oidcBase.UserInfoIDPath) + } + if v, ok := settings[SettingKeyOIDCConnectUserInfoUsernamePath]; ok { + result.OIDCConnectUserInfoUsernamePath = strings.TrimSpace(v) + } else { + result.OIDCConnectUserInfoUsernamePath = strings.TrimSpace(oidcBase.UserInfoUsernamePath) + } + result.OIDCConnectClientSecret = strings.TrimSpace(settings[SettingKeyOIDCConnectClientSecret]) + if result.OIDCConnectClientSecret == "" { + result.OIDCConnectClientSecret = strings.TrimSpace(oidcBase.ClientSecret) + } + result.OIDCConnectClientSecretConfigured = result.OIDCConnectClientSecret != "" + // Model fallback settings result.EnableModelFallback = settings[SettingKeyEnableModelFallback] == "true" result.FallbackModelAnthropic = s.getStringOrDefault(settings, SettingKeyFallbackModelAnthropic, "claude-3-5-sonnet-20241022") @@ -987,13 +1208,14 @@ func (s *SettingService) parseSettings(settings map[string]string) *SystemSettin // 分组隔离 result.AllowUngroupedKeyScheduling = settings[SettingKeyAllowUngroupedKeyScheduling] == "true" - // Gateway forwarding behavior (defaults: fingerprint=true, metadata_passthrough=false) + // Gateway forwarding behavior (defaults: fingerprint=true, metadata_passthrough=false, cch_signing=false) if v, ok := settings[SettingKeyEnableFingerprintUnification]; ok && v != "" { result.EnableFingerprintUnification = v == "true" } else { result.EnableFingerprintUnification = true // default: enabled (current behavior) } result.EnableMetadataPassthrough = settings[SettingKeyEnableMetadataPassthrough] == "true" + result.EnableCCHSigning = settings[SettingKeyEnableCCHSigning] == "true" return result } @@ -1032,6 +1254,50 @@ func parseDefaultSubscriptions(raw string) []DefaultSubscriptionSetting { return normalized } +func parseTablePreferences(defaultPageSizeRaw, optionsRaw string) (int, []int) { + defaultPageSize := 20 + if v, err := strconv.Atoi(strings.TrimSpace(defaultPageSizeRaw)); err == nil { + defaultPageSize = v + } + + var options []int + if strings.TrimSpace(optionsRaw) != "" { + _ = json.Unmarshal([]byte(optionsRaw), &options) + } + + return normalizeTablePreferences(defaultPageSize, options) +} + +func normalizeTablePreferences(defaultPageSize int, options []int) (int, []int) { + const minPageSize = 5 + const maxPageSize = 1000 + const fallbackPageSize = 20 + + seen := make(map[int]struct{}, len(options)) + normalizedOptions := make([]int, 0, len(options)) + for _, option := range options { + if option < minPageSize || option > maxPageSize { + continue + } + if _, ok := seen[option]; ok { + continue + } + seen[option] = struct{}{} + normalizedOptions = append(normalizedOptions, option) + } + sort.Ints(normalizedOptions) + + if defaultPageSize < minPageSize || defaultPageSize > maxPageSize { + defaultPageSize = fallbackPageSize + } + + if len(normalizedOptions) == 0 { + normalizedOptions = []int{10, 20, 50} + } + + return defaultPageSize, normalizedOptions +} + // getStringOrDefault 获取字符串值或默认值 func (s *SettingService) getStringOrDefault(settings map[string]string, key, defaultValue string) string { if value, ok := settings[key]; ok && value != "" { @@ -1319,6 +1585,282 @@ func (s *SettingService) SetOverloadCooldownSettings(ctx context.Context, settin return s.settingRepo.Set(ctx, SettingKeyOverloadCooldownSettings, string(data)) } +// GetOIDCConnectOAuthConfig 返回用于登录的“最终生效” OIDC 配置。 +// +// 优先级: +// - 若对应系统设置键存在,则覆盖 config.yaml/env 的值 +// - 否则回退到 config.yaml/env 的值 +func (s *SettingService) GetOIDCConnectOAuthConfig(ctx context.Context) (config.OIDCConnectConfig, error) { + if s == nil || s.cfg == nil { + return config.OIDCConnectConfig{}, infraerrors.ServiceUnavailable("CONFIG_NOT_READY", "config not loaded") + } + + effective := s.cfg.OIDC + + keys := []string{ + SettingKeyOIDCConnectEnabled, + SettingKeyOIDCConnectProviderName, + SettingKeyOIDCConnectClientID, + SettingKeyOIDCConnectClientSecret, + SettingKeyOIDCConnectIssuerURL, + SettingKeyOIDCConnectDiscoveryURL, + SettingKeyOIDCConnectAuthorizeURL, + SettingKeyOIDCConnectTokenURL, + SettingKeyOIDCConnectUserInfoURL, + SettingKeyOIDCConnectJWKSURL, + SettingKeyOIDCConnectScopes, + SettingKeyOIDCConnectRedirectURL, + SettingKeyOIDCConnectFrontendRedirectURL, + SettingKeyOIDCConnectTokenAuthMethod, + SettingKeyOIDCConnectUsePKCE, + SettingKeyOIDCConnectValidateIDToken, + SettingKeyOIDCConnectAllowedSigningAlgs, + SettingKeyOIDCConnectClockSkewSeconds, + SettingKeyOIDCConnectRequireEmailVerified, + SettingKeyOIDCConnectUserInfoEmailPath, + SettingKeyOIDCConnectUserInfoIDPath, + SettingKeyOIDCConnectUserInfoUsernamePath, + } + settings, err := s.settingRepo.GetMultiple(ctx, keys) + if err != nil { + return config.OIDCConnectConfig{}, fmt.Errorf("get oidc connect settings: %w", err) + } + + if raw, ok := settings[SettingKeyOIDCConnectEnabled]; ok { + effective.Enabled = raw == "true" + } + if v, ok := settings[SettingKeyOIDCConnectProviderName]; ok && strings.TrimSpace(v) != "" { + effective.ProviderName = strings.TrimSpace(v) + } + if v, ok := settings[SettingKeyOIDCConnectClientID]; ok && strings.TrimSpace(v) != "" { + effective.ClientID = strings.TrimSpace(v) + } + if v, ok := settings[SettingKeyOIDCConnectClientSecret]; ok && strings.TrimSpace(v) != "" { + effective.ClientSecret = strings.TrimSpace(v) + } + if v, ok := settings[SettingKeyOIDCConnectIssuerURL]; ok && strings.TrimSpace(v) != "" { + effective.IssuerURL = strings.TrimSpace(v) + } + if v, ok := settings[SettingKeyOIDCConnectDiscoveryURL]; ok && strings.TrimSpace(v) != "" { + effective.DiscoveryURL = strings.TrimSpace(v) + } + if v, ok := settings[SettingKeyOIDCConnectAuthorizeURL]; ok && strings.TrimSpace(v) != "" { + effective.AuthorizeURL = strings.TrimSpace(v) + } + if v, ok := settings[SettingKeyOIDCConnectTokenURL]; ok && strings.TrimSpace(v) != "" { + effective.TokenURL = strings.TrimSpace(v) + } + if v, ok := settings[SettingKeyOIDCConnectUserInfoURL]; ok && strings.TrimSpace(v) != "" { + effective.UserInfoURL = strings.TrimSpace(v) + } + if v, ok := settings[SettingKeyOIDCConnectJWKSURL]; ok && strings.TrimSpace(v) != "" { + effective.JWKSURL = strings.TrimSpace(v) + } + if v, ok := settings[SettingKeyOIDCConnectScopes]; ok && strings.TrimSpace(v) != "" { + effective.Scopes = strings.TrimSpace(v) + } + if v, ok := settings[SettingKeyOIDCConnectRedirectURL]; ok && strings.TrimSpace(v) != "" { + effective.RedirectURL = strings.TrimSpace(v) + } + if v, ok := settings[SettingKeyOIDCConnectFrontendRedirectURL]; ok && strings.TrimSpace(v) != "" { + effective.FrontendRedirectURL = strings.TrimSpace(v) + } + if v, ok := settings[SettingKeyOIDCConnectTokenAuthMethod]; ok && strings.TrimSpace(v) != "" { + effective.TokenAuthMethod = strings.ToLower(strings.TrimSpace(v)) + } + if raw, ok := settings[SettingKeyOIDCConnectUsePKCE]; ok { + effective.UsePKCE = raw == "true" + } + if raw, ok := settings[SettingKeyOIDCConnectValidateIDToken]; ok { + effective.ValidateIDToken = raw == "true" + } + if v, ok := settings[SettingKeyOIDCConnectAllowedSigningAlgs]; ok && strings.TrimSpace(v) != "" { + effective.AllowedSigningAlgs = strings.TrimSpace(v) + } + if raw, ok := settings[SettingKeyOIDCConnectClockSkewSeconds]; ok && strings.TrimSpace(raw) != "" { + if parsed, parseErr := strconv.Atoi(strings.TrimSpace(raw)); parseErr == nil { + effective.ClockSkewSeconds = parsed + } + } + if raw, ok := settings[SettingKeyOIDCConnectRequireEmailVerified]; ok { + effective.RequireEmailVerified = raw == "true" + } + if v, ok := settings[SettingKeyOIDCConnectUserInfoEmailPath]; ok { + effective.UserInfoEmailPath = strings.TrimSpace(v) + } + if v, ok := settings[SettingKeyOIDCConnectUserInfoIDPath]; ok { + effective.UserInfoIDPath = strings.TrimSpace(v) + } + if v, ok := settings[SettingKeyOIDCConnectUserInfoUsernamePath]; ok { + effective.UserInfoUsernamePath = strings.TrimSpace(v) + } + + if !effective.Enabled { + return config.OIDCConnectConfig{}, infraerrors.NotFound("OAUTH_DISABLED", "oauth login is disabled") + } + if strings.TrimSpace(effective.ProviderName) == "" { + effective.ProviderName = "OIDC" + } + if strings.TrimSpace(effective.ClientID) == "" { + return config.OIDCConnectConfig{}, infraerrors.InternalServer("OAUTH_CONFIG_INVALID", "oauth client id not configured") + } + if strings.TrimSpace(effective.IssuerURL) == "" { + return config.OIDCConnectConfig{}, infraerrors.InternalServer("OAUTH_CONFIG_INVALID", "oauth issuer url not configured") + } + if strings.TrimSpace(effective.RedirectURL) == "" { + return config.OIDCConnectConfig{}, infraerrors.InternalServer("OAUTH_CONFIG_INVALID", "oauth redirect url not configured") + } + if strings.TrimSpace(effective.FrontendRedirectURL) == "" { + return config.OIDCConnectConfig{}, infraerrors.InternalServer("OAUTH_CONFIG_INVALID", "oauth frontend redirect url not configured") + } + if !scopesContainOpenID(effective.Scopes) { + return config.OIDCConnectConfig{}, infraerrors.InternalServer("OAUTH_CONFIG_INVALID", "oauth scopes must contain openid") + } + if effective.ClockSkewSeconds < 0 || effective.ClockSkewSeconds > 600 { + return config.OIDCConnectConfig{}, infraerrors.InternalServer("OAUTH_CONFIG_INVALID", "oauth clock skew must be between 0 and 600") + } + + if err := config.ValidateAbsoluteHTTPURL(effective.IssuerURL); err != nil { + return config.OIDCConnectConfig{}, infraerrors.InternalServer("OAUTH_CONFIG_INVALID", "oauth issuer url invalid") + } + + discoveryURL := strings.TrimSpace(effective.DiscoveryURL) + if discoveryURL == "" { + discoveryURL = oidcDefaultDiscoveryURL(effective.IssuerURL) + effective.DiscoveryURL = discoveryURL + } + if discoveryURL != "" { + if err := config.ValidateAbsoluteHTTPURL(discoveryURL); err != nil { + return config.OIDCConnectConfig{}, infraerrors.InternalServer("OAUTH_CONFIG_INVALID", "oauth discovery url invalid") + } + } + + needsDiscovery := strings.TrimSpace(effective.AuthorizeURL) == "" || + strings.TrimSpace(effective.TokenURL) == "" || + (effective.ValidateIDToken && strings.TrimSpace(effective.JWKSURL) == "") + if needsDiscovery && discoveryURL != "" { + metadata, resolveErr := oidcResolveProviderMetadata(ctx, discoveryURL) + if resolveErr != nil { + return config.OIDCConnectConfig{}, infraerrors.InternalServer("OAUTH_CONFIG_INVALID", "oauth discovery resolve failed").WithCause(resolveErr) + } + if strings.TrimSpace(effective.AuthorizeURL) == "" { + effective.AuthorizeURL = strings.TrimSpace(metadata.AuthorizationEndpoint) + } + if strings.TrimSpace(effective.TokenURL) == "" { + effective.TokenURL = strings.TrimSpace(metadata.TokenEndpoint) + } + if strings.TrimSpace(effective.UserInfoURL) == "" { + effective.UserInfoURL = strings.TrimSpace(metadata.UserInfoEndpoint) + } + if strings.TrimSpace(effective.JWKSURL) == "" { + effective.JWKSURL = strings.TrimSpace(metadata.JWKSURI) + } + } + + if strings.TrimSpace(effective.AuthorizeURL) == "" { + return config.OIDCConnectConfig{}, infraerrors.InternalServer("OAUTH_CONFIG_INVALID", "oauth authorize url not configured") + } + if strings.TrimSpace(effective.TokenURL) == "" { + return config.OIDCConnectConfig{}, infraerrors.InternalServer("OAUTH_CONFIG_INVALID", "oauth token url not configured") + } + if err := config.ValidateAbsoluteHTTPURL(effective.AuthorizeURL); err != nil { + return config.OIDCConnectConfig{}, infraerrors.InternalServer("OAUTH_CONFIG_INVALID", "oauth authorize url invalid") + } + if err := config.ValidateAbsoluteHTTPURL(effective.TokenURL); err != nil { + return config.OIDCConnectConfig{}, infraerrors.InternalServer("OAUTH_CONFIG_INVALID", "oauth token url invalid") + } + if v := strings.TrimSpace(effective.UserInfoURL); v != "" { + if err := config.ValidateAbsoluteHTTPURL(v); err != nil { + return config.OIDCConnectConfig{}, infraerrors.InternalServer("OAUTH_CONFIG_INVALID", "oauth userinfo url invalid") + } + } + if effective.ValidateIDToken { + if strings.TrimSpace(effective.JWKSURL) == "" { + return config.OIDCConnectConfig{}, infraerrors.InternalServer("OAUTH_CONFIG_INVALID", "oauth jwks url not configured") + } + if strings.TrimSpace(effective.AllowedSigningAlgs) == "" { + return config.OIDCConnectConfig{}, infraerrors.InternalServer("OAUTH_CONFIG_INVALID", "oauth signing algs not configured") + } + } + if v := strings.TrimSpace(effective.JWKSURL); v != "" { + if err := config.ValidateAbsoluteHTTPURL(v); err != nil { + return config.OIDCConnectConfig{}, infraerrors.InternalServer("OAUTH_CONFIG_INVALID", "oauth jwks url invalid") + } + } + if err := config.ValidateAbsoluteHTTPURL(effective.RedirectURL); err != nil { + return config.OIDCConnectConfig{}, infraerrors.InternalServer("OAUTH_CONFIG_INVALID", "oauth redirect url invalid") + } + if err := config.ValidateFrontendRedirectURL(effective.FrontendRedirectURL); err != nil { + return config.OIDCConnectConfig{}, infraerrors.InternalServer("OAUTH_CONFIG_INVALID", "oauth frontend redirect url invalid") + } + + method := strings.ToLower(strings.TrimSpace(effective.TokenAuthMethod)) + switch method { + case "", "client_secret_post", "client_secret_basic": + if strings.TrimSpace(effective.ClientSecret) == "" { + return config.OIDCConnectConfig{}, infraerrors.InternalServer("OAUTH_CONFIG_INVALID", "oauth client secret not configured") + } + case "none": + if !effective.UsePKCE { + return config.OIDCConnectConfig{}, infraerrors.InternalServer("OAUTH_CONFIG_INVALID", "oauth pkce must be enabled when token_auth_method=none") + } + default: + return config.OIDCConnectConfig{}, infraerrors.InternalServer("OAUTH_CONFIG_INVALID", "oauth token_auth_method invalid") + } + + return effective, nil +} + +func scopesContainOpenID(scopes string) bool { + for _, scope := range strings.Fields(strings.ToLower(strings.TrimSpace(scopes))) { + if scope == "openid" { + return true + } + } + return false +} + +type oidcProviderMetadata struct { + AuthorizationEndpoint string `json:"authorization_endpoint"` + TokenEndpoint string `json:"token_endpoint"` + UserInfoEndpoint string `json:"userinfo_endpoint"` + JWKSURI string `json:"jwks_uri"` +} + +func oidcDefaultDiscoveryURL(issuerURL string) string { + issuerURL = strings.TrimSpace(issuerURL) + if issuerURL == "" { + return "" + } + return strings.TrimRight(issuerURL, "/") + "/.well-known/openid-configuration" +} + +func oidcResolveProviderMetadata(ctx context.Context, discoveryURL string) (*oidcProviderMetadata, error) { + discoveryURL = strings.TrimSpace(discoveryURL) + if discoveryURL == "" { + return nil, fmt.Errorf("discovery url is empty") + } + + resp, err := req.C(). + SetTimeout(15*time.Second). + R(). + SetContext(ctx). + SetHeader("Accept", "application/json"). + Get(discoveryURL) + if err != nil { + return nil, fmt.Errorf("request discovery document: %w", err) + } + if !resp.IsSuccessState() { + return nil, fmt.Errorf("discovery request failed: status=%d", resp.StatusCode) + } + + metadata := &oidcProviderMetadata{} + if err := json.Unmarshal(resp.Bytes(), metadata); err != nil { + return nil, fmt.Errorf("parse discovery document: %w", err) + } + return metadata, nil +} + // GetStreamTimeoutSettings 获取流超时处理配置 func (s *SettingService) GetStreamTimeoutSettings(ctx context.Context) (*StreamTimeoutSettings, error) { value, err := s.settingRepo.GetValue(ctx, SettingKeyStreamTimeoutSettings) @@ -1531,6 +2073,18 @@ func (s *SettingService) SetBetaPolicySettings(ctx context.Context, settings *Be if !validScopes[rule.Scope] { return fmt.Errorf("rule[%d]: invalid scope %q", i, rule.Scope) } + // Validate model_whitelist patterns + for j, pattern := range rule.ModelWhitelist { + trimmed := strings.TrimSpace(pattern) + if trimmed == "" { + return fmt.Errorf("rule[%d]: model_whitelist[%d] cannot be empty", i, j) + } + settings.Rules[i].ModelWhitelist[j] = trimmed + } + // Validate fallback_action + if rule.FallbackAction != "" && !validActions[rule.FallbackAction] { + return fmt.Errorf("rule[%d]: invalid fallback_action %q", i, rule.FallbackAction) + } } data, err := json.Marshal(settings) diff --git a/backend/internal/service/setting_service_oidc_config_test.go b/backend/internal/service/setting_service_oidc_config_test.go new file mode 100644 index 0000000000..3809b332bd --- /dev/null +++ b/backend/internal/service/setting_service_oidc_config_test.go @@ -0,0 +1,103 @@ +//go:build unit + +package service + +import ( + "context" + "fmt" + "net/http" + "net/http/httptest" + "testing" + + "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/stretchr/testify/require" +) + +type settingOIDCRepoStub struct { + values map[string]string +} + +func (s *settingOIDCRepoStub) Get(ctx context.Context, key string) (*Setting, error) { + panic("unexpected Get call") +} + +func (s *settingOIDCRepoStub) GetValue(ctx context.Context, key string) (string, error) { + panic("unexpected GetValue call") +} + +func (s *settingOIDCRepoStub) Set(ctx context.Context, key, value string) error { + panic("unexpected Set call") +} + +func (s *settingOIDCRepoStub) GetMultiple(ctx context.Context, keys []string) (map[string]string, error) { + out := make(map[string]string, len(keys)) + for _, key := range keys { + if value, ok := s.values[key]; ok { + out[key] = value + } + } + return out, nil +} + +func (s *settingOIDCRepoStub) SetMultiple(ctx context.Context, settings map[string]string) error { + panic("unexpected SetMultiple call") +} + +func (s *settingOIDCRepoStub) GetAll(ctx context.Context) (map[string]string, error) { + panic("unexpected GetAll call") +} + +func (s *settingOIDCRepoStub) Delete(ctx context.Context, key string) error { + panic("unexpected Delete call") +} + +func TestGetOIDCConnectOAuthConfig_ResolvesEndpointsFromIssuerDiscovery(t *testing.T) { + var discoveryHits int + var baseURL string + + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path != "/issuer/.well-known/openid-configuration" { + http.NotFound(w, r) + return + } + discoveryHits++ + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(fmt.Sprintf(`{ + "authorization_endpoint":"%s/issuer/protocol/openid-connect/auth", + "token_endpoint":"%s/issuer/protocol/openid-connect/token", + "userinfo_endpoint":"%s/issuer/protocol/openid-connect/userinfo", + "jwks_uri":"%s/issuer/protocol/openid-connect/certs" + }`, baseURL, baseURL, baseURL, baseURL))) + })) + defer srv.Close() + baseURL = srv.URL + + cfg := &config.Config{ + OIDC: config.OIDCConnectConfig{ + Enabled: true, + ProviderName: "OIDC", + ClientID: "oidc-client", + ClientSecret: "oidc-secret", + IssuerURL: srv.URL + "/issuer", + RedirectURL: "https://example.com/api/v1/auth/oauth/oidc/callback", + FrontendRedirectURL: "/auth/oidc/callback", + Scopes: "openid email profile", + TokenAuthMethod: "client_secret_post", + ValidateIDToken: true, + AllowedSigningAlgs: "RS256", + ClockSkewSeconds: 120, + }, + } + + repo := &settingOIDCRepoStub{values: map[string]string{}} + svc := NewSettingService(repo, cfg) + + got, err := svc.GetOIDCConnectOAuthConfig(context.Background()) + require.NoError(t, err) + require.Equal(t, 1, discoveryHits) + require.Equal(t, srv.URL+"/issuer/.well-known/openid-configuration", got.DiscoveryURL) + require.Equal(t, srv.URL+"/issuer/protocol/openid-connect/auth", got.AuthorizeURL) + require.Equal(t, srv.URL+"/issuer/protocol/openid-connect/token", got.TokenURL) + require.Equal(t, srv.URL+"/issuer/protocol/openid-connect/userinfo", got.UserInfoURL) + require.Equal(t, srv.URL+"/issuer/protocol/openid-connect/certs", got.JWKSURL) +} diff --git a/backend/internal/service/setting_service_public_test.go b/backend/internal/service/setting_service_public_test.go index b511cd29dc..6dfa627cfa 100644 --- a/backend/internal/service/setting_service_public_test.go +++ b/backend/internal/service/setting_service_public_test.go @@ -62,3 +62,18 @@ func TestSettingService_GetPublicSettings_ExposesRegistrationEmailSuffixWhitelis require.NoError(t, err) require.Equal(t, []string{"@example.com", "@foo.bar"}, settings.RegistrationEmailSuffixWhitelist) } + +func TestSettingService_GetPublicSettings_ExposesTablePreferences(t *testing.T) { + repo := &settingPublicRepoStub{ + values: map[string]string{ + SettingKeyTableDefaultPageSize: "50", + SettingKeyTablePageSizeOptions: "[20,50,100]", + }, + } + svc := NewSettingService(repo, &config.Config{}) + + settings, err := svc.GetPublicSettings(context.Background()) + require.NoError(t, err) + require.Equal(t, 50, settings.TableDefaultPageSize) + require.Equal(t, []int{20, 50, 100}, settings.TablePageSizeOptions) +} diff --git a/backend/internal/service/setting_service_update_test.go b/backend/internal/service/setting_service_update_test.go index 1de08611e2..28c7ad0220 100644 --- a/backend/internal/service/setting_service_update_test.go +++ b/backend/internal/service/setting_service_update_test.go @@ -202,3 +202,24 @@ func TestParseDefaultSubscriptions_NormalizesValues(t *testing.T) { {GroupID: 12, ValidityDays: MaxValidityDays}, }, got) } + +func TestSettingService_UpdateSettings_TablePreferences(t *testing.T) { + repo := &settingUpdateRepoStub{} + svc := NewSettingService(repo, &config.Config{}) + + err := svc.UpdateSettings(context.Background(), &SystemSettings{ + TableDefaultPageSize: 50, + TablePageSizeOptions: []int{20, 50, 100}, + }) + require.NoError(t, err) + require.Equal(t, "50", repo.updates[SettingKeyTableDefaultPageSize]) + require.Equal(t, "[20,50,100]", repo.updates[SettingKeyTablePageSizeOptions]) + + err = svc.UpdateSettings(context.Background(), &SystemSettings{ + TableDefaultPageSize: 1000, + TablePageSizeOptions: []int{20, 100}, + }) + require.NoError(t, err) + require.Equal(t, "1000", repo.updates[SettingKeyTableDefaultPageSize]) + require.Equal(t, "[20,100]", repo.updates[SettingKeyTablePageSizeOptions]) +} diff --git a/backend/internal/service/settings_view.go b/backend/internal/service/settings_view.go index 0d6cb0cd05..e3ae55491b 100644 --- a/backend/internal/service/settings_view.go +++ b/backend/internal/service/settings_view.go @@ -31,6 +31,31 @@ type SystemSettings struct { LinuxDoConnectClientSecretConfigured bool LinuxDoConnectRedirectURL string + // Generic OIDC OAuth 登录 + OIDCConnectEnabled bool + OIDCConnectProviderName string + OIDCConnectClientID string + OIDCConnectClientSecret string + OIDCConnectClientSecretConfigured bool + OIDCConnectIssuerURL string + OIDCConnectDiscoveryURL string + OIDCConnectAuthorizeURL string + OIDCConnectTokenURL string + OIDCConnectUserInfoURL string + OIDCConnectJWKSURL string + OIDCConnectScopes string + OIDCConnectRedirectURL string + OIDCConnectFrontendRedirectURL string + OIDCConnectTokenAuthMethod string + OIDCConnectUsePKCE bool + OIDCConnectValidateIDToken bool + OIDCConnectAllowedSigningAlgs string + OIDCConnectClockSkewSeconds int + OIDCConnectRequireEmailVerified bool + OIDCConnectUserInfoEmailPath string + OIDCConnectUserInfoIDPath string + OIDCConnectUserInfoUsernamePath string + SiteName string SiteLogo string SiteSubtitle string @@ -41,6 +66,8 @@ type SystemSettings struct { HideCcsImportButton bool PurchaseSubscriptionEnabled bool PurchaseSubscriptionURL string + TableDefaultPageSize int + TablePageSizeOptions []int CustomMenuItems string // JSON array of custom menu items CustomEndpoints string // JSON array of custom endpoints @@ -78,6 +105,7 @@ type SystemSettings struct { // Gateway forwarding behavior EnableFingerprintUnification bool // 是否统一 OAuth 账号的指纹头(默认 true) EnableMetadataPassthrough bool // 是否透传客户端原始 metadata(默认 false) + EnableCCHSigning bool // 是否对 billing header cch 进行签名(默认 false) } type DefaultSubscriptionSetting struct { @@ -106,13 +134,17 @@ type PublicSettings struct { PurchaseSubscriptionEnabled bool PurchaseSubscriptionURL string + TableDefaultPageSize int + TablePageSizeOptions []int CustomMenuItems string // JSON array of custom menu items CustomEndpoints string // JSON array of custom endpoints - LinuxDoOAuthEnabled bool - BackendModeEnabled bool - PaymentEnabled bool - Version string + LinuxDoOAuthEnabled bool + BackendModeEnabled bool + PaymentEnabled bool + OIDCOAuthEnabled bool + OIDCOAuthProviderName string + Version string } // StreamTimeoutSettings 流超时处理配置(仅控制超时后的处理方式,超时判定由网关配置控制) @@ -179,10 +211,13 @@ const ( // BetaPolicyRule 单条 Beta 策略规则 type BetaPolicyRule struct { - BetaToken string `json:"beta_token"` // beta token 值 - Action string `json:"action"` // "pass" | "filter" | "block" - Scope string `json:"scope"` // "all" | "oauth" | "apikey" | "bedrock" - ErrorMessage string `json:"error_message,omitempty"` // 自定义错误消息 (action=block 时生效) + BetaToken string `json:"beta_token"` // beta token 值 + Action string `json:"action"` // "pass" | "filter" | "block" + Scope string `json:"scope"` // "all" | "oauth" | "apikey" | "bedrock" + ErrorMessage string `json:"error_message,omitempty"` // 自定义错误消息 (action=block 时生效) + ModelWhitelist []string `json:"model_whitelist,omitempty"` // 模型匹配模式列表(为空=对所有模型生效) + FallbackAction string `json:"fallback_action,omitempty"` // 未匹配白名单的模型的处理方式 + FallbackErrorMessage string `json:"fallback_error_message,omitempty"` // 未匹配白名单时的自定义错误消息 (fallback_action=block 时生效) } // BetaPolicySettings Beta 策略配置 diff --git a/backend/migrations/091_add_group_messages_dispatch_model_config.sql b/backend/migrations/091_add_group_messages_dispatch_model_config.sql new file mode 100644 index 0000000000..8ddfcb0f72 --- /dev/null +++ b/backend/migrations/091_add_group_messages_dispatch_model_config.sql @@ -0,0 +1,2 @@ +ALTER TABLE groups +ADD COLUMN IF NOT EXISTS messages_dispatch_model_config JSONB NOT NULL DEFAULT '{}'::jsonb; diff --git a/deploy/Dockerfile b/deploy/Dockerfile index 7caa5ca63e..b0b6036c67 100644 --- a/deploy/Dockerfile +++ b/deploy/Dockerfile @@ -7,7 +7,7 @@ # ============================================================================= ARG NODE_IMAGE=node:24-alpine -ARG GOLANG_IMAGE=golang:1.26.1-alpine +ARG GOLANG_IMAGE=golang:1.26.2-alpine ARG ALPINE_IMAGE=alpine:3.20 ARG GOPROXY=https://goproxy.cn,direct ARG GOSUMDB=sum.golang.google.cn diff --git a/deploy/codex-instructions.md.tmpl b/deploy/codex-instructions.md.tmpl new file mode 100644 index 0000000000..87ad0a3d84 --- /dev/null +++ b/deploy/codex-instructions.md.tmpl @@ -0,0 +1,5 @@ +You are Codex, based on GPT-5. You are running as a coding agent in the Codex CLI on a user's computer. + +{{ if .ExistingInstructions }} +{{ .ExistingInstructions }} +{{ end }} diff --git a/deploy/config.example.yaml b/deploy/config.example.yaml index 8f60acd5ed..358f6a31d9 100644 --- a/deploy/config.example.yaml +++ b/deploy/config.example.yaml @@ -202,6 +202,32 @@ gateway: # # 注意:开启后会影响所有客户端的行为(不仅限于 VS Code / Codex CLI),请谨慎开启。 force_codex_cli: false + # Optional: template file used to build the final top-level Codex `instructions`. + # 可选:用于构建最终 Codex 顶层 `instructions` 的模板文件路径。 + # + # This is applied on the `/v1/messages -> Responses/Codex` conversion path, + # after Claude `system` has already been normalized into Codex `instructions`. + # 该模板作用于 `/v1/messages -> Responses/Codex` 转换链路,且发生在 Claude `system` + # 已经被归一化为 Codex `instructions` 之后。 + # + # The template can reference: + # 模板可引用: + # - {{ .ExistingInstructions }} : converted client instructions/system + # - {{ .OriginalModel }} : original requested model + # - {{ .NormalizedModel }} : normalized routing model + # - {{ .BillingModel }} : billing model + # - {{ .UpstreamModel }} : final upstream model + # + # If you want to preserve client system prompts, keep {{ .ExistingInstructions }} + # somewhere in the template. If omitted, the template output fully replaces it. + # 如需保留客户端 system 提示词,请在模板中显式包含 {{ .ExistingInstructions }}。 + # 若省略,则模板输出会完全覆盖它。 + # + # Docker users can mount a host file to /app/data/codex-instructions.md.tmpl + # and point this field there. + # Docker 用户可将宿主机文件挂载到 /app/data/codex-instructions.md.tmpl, + # 然后把本字段指向该路径。 + forced_codex_instructions_template_file: "" # OpenAI 透传模式是否放行客户端超时头(如 x-stainless-timeout) # 默认 false:过滤超时头,降低上游提前断流风险。 openai_passthrough_allow_timeout_headers: false @@ -820,6 +846,46 @@ linuxdo_connect: userinfo_id_path: "" userinfo_username_path: "" +# ============================================================================= +# Generic OIDC OAuth Login (SSO) +# 通用 OIDC OAuth 登录(用于 Sub2API 用户登录) +# ============================================================================= +oidc_connect: + enabled: false + provider_name: "OIDC" + client_id: "" + client_secret: "" + # 例如: "https://keycloak.example.com/realms/myrealm" + issuer_url: "" + # 可选: OIDC Discovery URL。为空时可手动填写 authorize/token/userinfo/jwks + discovery_url: "" + authorize_url: "" + token_url: "" + # 可选(仅补充 email/username,不用于 sub 可信绑定) + userinfo_url: "" + # validate_id_token=true 时必填 + jwks_url: "" + scopes: "openid email profile" + # 示例: "https://your-domain.com/api/v1/auth/oauth/oidc/callback" + redirect_url: "" + # 安全提示: + # - 建议使用同源相对路径(以 / 开头),避免把 token 重定向到意外的第三方域名 + # - 该地址不应包含 #fragment(本实现使用 URL fragment 传递 access_token) + frontend_redirect_url: "/auth/oidc/callback" + token_auth_method: "client_secret_post" # client_secret_post | client_secret_basic | none + # 注意:当 token_auth_method=none(public client)时,必须启用 PKCE + use_pkce: false + # 开启后强制校验 id_token 的签名和 claims(推荐) + validate_id_token: true + allowed_signing_algs: "RS256,ES256,PS256" + # 允许的时钟偏移(秒) + clock_skew_seconds: 120 + # 若 Provider 返回 email_verified=false,是否拒绝登录 + require_email_verified: false + userinfo_email_path: "" + userinfo_id_path: "" + userinfo_username_path: "" + # ============================================================================= # Default Settings # 默认设置 diff --git a/deploy/docker-compose.yml b/deploy/docker-compose.yml index e0e9e54f49..082541282f 100644 --- a/deploy/docker-compose.yml +++ b/deploy/docker-compose.yml @@ -31,6 +31,10 @@ services: # Optional: Mount custom config.yaml (uncomment and create the file first) # Copy config.example.yaml to config.yaml, modify it, then uncomment: # - ./config.yaml:/app/data/config.yaml + # Optional: Mount a custom Codex instructions template file, then point + # gateway.forced_codex_instructions_template_file at /app/data/codex-instructions.md.tmpl + # in config.yaml. + # - ./codex-instructions.md.tmpl:/app/data/codex-instructions.md.tmpl:ro environment: # ======================================================================= # Auto Setup (REQUIRED for Docker deployment) @@ -146,7 +150,17 @@ services: networks: - sub2api-network healthcheck: - test: ["CMD", "wget", "-q", "-T", "5", "-O", "/dev/null", "http://localhost:8080/health"] + test: + [ + "CMD", + "wget", + "-q", + "-T", + "5", + "-O", + "/dev/null", + "http://localhost:8080/health", + ] interval: 30s timeout: 10s retries: 3 @@ -177,11 +191,17 @@ services: networks: - sub2api-network healthcheck: - test: ["CMD-SHELL", "pg_isready -U ${POSTGRES_USER:-sub2api} -d ${POSTGRES_DB:-sub2api}"] + test: + [ + "CMD-SHELL", + "pg_isready -U ${POSTGRES_USER:-sub2api} -d ${POSTGRES_DB:-sub2api}", + ] interval: 10s timeout: 5s retries: 5 start_period: 10s + ports: + - 5432:5432 # 注意:不暴露端口到宿主机,应用通过内部网络连接 # 如需调试,可临时添加:ports: ["127.0.0.1:5433:5432"] @@ -199,12 +219,12 @@ services: volumes: - redis_data:/data command: > - sh -c ' - redis-server - --save 60 1 - --appendonly yes - --appendfsync everysec - ${REDIS_PASSWORD:+--requirepass "$REDIS_PASSWORD"}' + sh -c ' + redis-server + --save 60 1 + --appendonly yes + --appendfsync everysec + ${REDIS_PASSWORD:+--requirepass "$REDIS_PASSWORD"}' environment: - TZ=${TZ:-Asia/Shanghai} # REDISCLI_AUTH is used by redis-cli for authentication (safer than -a flag) @@ -217,7 +237,8 @@ services: timeout: 5s retries: 5 start_period: 5s - + ports: + - 6379:6379 # ============================================================================= # Volumes # ============================================================================= diff --git a/frontend/src/api/admin/accounts.ts b/frontend/src/api/admin/accounts.ts index 14ad56bee3..9f4768688d 100644 --- a/frontend/src/api/admin/accounts.ts +++ b/frontend/src/api/admin/accounts.ts @@ -38,6 +38,8 @@ export async function list( search?: string privacy_mode?: string lite?: string + sort_by?: string + sort_order?: 'asc' | 'desc' }, options?: { signal?: AbortSignal @@ -71,6 +73,8 @@ export async function listWithEtag( search?: string privacy_mode?: string lite?: string + sort_by?: string + sort_order?: 'asc' | 'desc' }, options?: { signal?: AbortSignal @@ -500,7 +504,11 @@ export async function exportData(options?: { platform?: string type?: string status?: string + group?: string + privacy_mode?: string search?: string + sort_by?: string + sort_order?: 'asc' | 'desc' } includeProxies?: boolean }): Promise { @@ -508,11 +516,15 @@ export async function exportData(options?: { if (options?.ids && options.ids.length > 0) { params.ids = options.ids.join(',') } else if (options?.filters) { - const { platform, type, status, search } = options.filters + const { platform, type, status, group, privacy_mode, search, sort_by, sort_order } = options.filters if (platform) params.platform = platform if (type) params.type = type if (status) params.status = status + if (group) params.group = group + if (privacy_mode) params.privacy_mode = privacy_mode if (search) params.search = search + if (sort_by) params.sort_by = sort_by + if (sort_order) params.sort_order = sort_order } if (options?.includeProxies === false) { params.include_proxies = 'false' diff --git a/frontend/src/api/admin/announcements.ts b/frontend/src/api/admin/announcements.ts index d02fdda772..92392a6753 100644 --- a/frontend/src/api/admin/announcements.ts +++ b/frontend/src/api/admin/announcements.ts @@ -17,10 +17,16 @@ export async function list( filters?: { status?: string search?: string + sort_by?: string + sort_order?: 'asc' | 'desc' + }, + options?: { + signal?: AbortSignal } ): Promise> { const { data } = await apiClient.get>('/admin/announcements', { - params: { page, page_size: pageSize, ...filters } + params: { page, page_size: pageSize, ...filters }, + signal: options?.signal }) return data } @@ -49,11 +55,21 @@ export async function getReadStatus( id: number, page: number = 1, pageSize: number = 20, - search: string = '' + filters?: { + search?: string + sort_by?: string + sort_order?: 'asc' | 'desc' + }, + options?: { + signal?: AbortSignal + } ): Promise> { const { data } = await apiClient.get>( `/admin/announcements/${id}/read-status`, - { params: { page, page_size: pageSize, search } } + { + params: { page, page_size: pageSize, ...filters }, + signal: options?.signal + } ) return data } @@ -68,4 +84,3 @@ const announcementsAPI = { } export default announcementsAPI - diff --git a/frontend/src/api/admin/channels.ts b/frontend/src/api/admin/channels.ts index 5334dd473d..b34550222b 100644 --- a/frontend/src/api/admin/channels.ts +++ b/frontend/src/api/admin/channels.ts @@ -83,6 +83,8 @@ export async function list( filters?: { status?: string search?: string + sort_by?: string + sort_order?: 'asc' | 'desc' }, options?: { signal?: AbortSignal } ): Promise> { diff --git a/frontend/src/api/admin/groups.ts b/frontend/src/api/admin/groups.ts index 5885dc6ad7..8739d5cbe9 100644 --- a/frontend/src/api/admin/groups.ts +++ b/frontend/src/api/admin/groups.ts @@ -27,6 +27,8 @@ export async function list( status?: 'active' | 'inactive' is_exclusive?: boolean search?: string + sort_by?: string + sort_order?: 'asc' | 'desc' }, options?: { signal?: AbortSignal diff --git a/frontend/src/api/admin/promo.ts b/frontend/src/api/admin/promo.ts index 6a8c4559e9..b24dffc200 100644 --- a/frontend/src/api/admin/promo.ts +++ b/frontend/src/api/admin/promo.ts @@ -17,10 +17,16 @@ export async function list( filters?: { status?: string search?: string + sort_by?: string + sort_order?: 'asc' | 'desc' + }, + options?: { + signal?: AbortSignal } ): Promise> { const { data } = await apiClient.get>('/admin/promo-codes', { - params: { page, page_size: pageSize, ...filters } + params: { page, page_size: pageSize, ...filters }, + signal: options?.signal }) return data } diff --git a/frontend/src/api/admin/proxies.ts b/frontend/src/api/admin/proxies.ts index 5e31ae20f7..3e041ba9a5 100644 --- a/frontend/src/api/admin/proxies.ts +++ b/frontend/src/api/admin/proxies.ts @@ -29,6 +29,8 @@ export async function list( protocol?: string status?: 'active' | 'inactive' search?: string + sort_by?: string + sort_order?: 'asc' | 'desc' }, options?: { signal?: AbortSignal @@ -227,16 +229,20 @@ export async function exportData(options?: { protocol?: string status?: 'active' | 'inactive' search?: string + sort_by?: string + sort_order?: 'asc' | 'desc' } }): Promise { const params: Record = {} if (options?.ids && options.ids.length > 0) { params.ids = options.ids.join(',') } else if (options?.filters) { - const { protocol, status, search } = options.filters + const { protocol, status, search, sort_by, sort_order } = options.filters if (protocol) params.protocol = protocol if (status) params.status = status if (search) params.search = search + if (sort_by) params.sort_by = sort_by + if (sort_order) params.sort_order = sort_order } const { data } = await apiClient.get('/admin/proxies/data', { params }) return data diff --git a/frontend/src/api/admin/redeem.ts b/frontend/src/api/admin/redeem.ts index a53c3566bb..57626b1efd 100644 --- a/frontend/src/api/admin/redeem.ts +++ b/frontend/src/api/admin/redeem.ts @@ -25,6 +25,8 @@ export async function list( type?: RedeemCodeType status?: 'active' | 'used' | 'expired' | 'unused' search?: string + sort_by?: string + sort_order?: 'asc' | 'desc' }, options?: { signal?: AbortSignal @@ -151,7 +153,10 @@ export async function getStats(): Promise<{ */ export async function exportCodes(filters?: { type?: RedeemCodeType - status?: 'active' | 'used' | 'expired' + status?: 'used' | 'expired' | 'unused' + search?: string + sort_by?: string + sort_order?: 'asc' | 'desc' }): Promise { const response = await apiClient.get('/admin/redeem-codes/export', { params: filters, diff --git a/frontend/src/api/admin/settings.ts b/frontend/src/api/admin/settings.ts index a4ee114041..504abe9c9b 100644 --- a/frontend/src/api/admin/settings.ts +++ b/frontend/src/api/admin/settings.ts @@ -38,6 +38,8 @@ export interface SystemSettings { doc_url: string home_content: string hide_ccs_import_button: boolean + table_default_page_size: number + table_page_size_options: number[] backend_mode_enabled: boolean custom_menu_items: CustomMenuItem[] custom_endpoints: CustomEndpoint[] @@ -60,6 +62,30 @@ export interface SystemSettings { linuxdo_connect_client_secret_configured: boolean linuxdo_connect_redirect_url: string + // Generic OIDC OAuth settings + oidc_connect_enabled: boolean + oidc_connect_provider_name: string + oidc_connect_client_id: string + oidc_connect_client_secret_configured: boolean + oidc_connect_issuer_url: string + oidc_connect_discovery_url: string + oidc_connect_authorize_url: string + oidc_connect_token_url: string + oidc_connect_userinfo_url: string + oidc_connect_jwks_url: string + oidc_connect_scopes: string + oidc_connect_redirect_url: string + oidc_connect_frontend_redirect_url: string + oidc_connect_token_auth_method: string + oidc_connect_use_pkce: boolean + oidc_connect_validate_id_token: boolean + oidc_connect_allowed_signing_algs: string + oidc_connect_clock_skew_seconds: number + oidc_connect_require_email_verified: boolean + oidc_connect_userinfo_email_path: string + oidc_connect_userinfo_id_path: string + oidc_connect_userinfo_username_path: string + // Model fallback configuration enable_model_fallback: boolean fallback_model_anthropic: string @@ -87,6 +113,7 @@ export interface SystemSettings { // Gateway forwarding behavior enable_fingerprint_unification: boolean enable_metadata_passthrough: boolean + enable_cch_signing: boolean // Payment configuration payment_enabled: boolean @@ -129,6 +156,8 @@ export interface UpdateSettingsRequest { doc_url?: string home_content?: string hide_ccs_import_button?: boolean + table_default_page_size?: number + table_page_size_options?: number[] backend_mode_enabled?: boolean custom_menu_items?: CustomMenuItem[] custom_endpoints?: CustomEndpoint[] @@ -146,6 +175,28 @@ export interface UpdateSettingsRequest { linuxdo_connect_client_id?: string linuxdo_connect_client_secret?: string linuxdo_connect_redirect_url?: string + oidc_connect_enabled?: boolean + oidc_connect_provider_name?: string + oidc_connect_client_id?: string + oidc_connect_client_secret?: string + oidc_connect_issuer_url?: string + oidc_connect_discovery_url?: string + oidc_connect_authorize_url?: string + oidc_connect_token_url?: string + oidc_connect_userinfo_url?: string + oidc_connect_jwks_url?: string + oidc_connect_scopes?: string + oidc_connect_redirect_url?: string + oidc_connect_frontend_redirect_url?: string + oidc_connect_token_auth_method?: string + oidc_connect_use_pkce?: boolean + oidc_connect_validate_id_token?: boolean + oidc_connect_allowed_signing_algs?: string + oidc_connect_clock_skew_seconds?: number + oidc_connect_require_email_verified?: boolean + oidc_connect_userinfo_email_path?: string + oidc_connect_userinfo_id_path?: string + oidc_connect_userinfo_username_path?: string enable_model_fallback?: boolean fallback_model_anthropic?: string fallback_model_openai?: string @@ -162,6 +213,7 @@ export interface UpdateSettingsRequest { allow_ungrouped_key_scheduling?: boolean enable_fingerprint_unification?: boolean enable_metadata_passthrough?: boolean + enable_cch_signing?: boolean // Payment configuration payment_enabled?: boolean payment_min_amount?: number @@ -394,6 +446,9 @@ export interface BetaPolicyRule { action: 'pass' | 'filter' | 'block' scope: 'all' | 'oauth' | 'apikey' | 'bedrock' error_message?: string + model_whitelist?: string[] + fallback_action?: 'pass' | 'filter' | 'block' + fallback_error_message?: string } /** diff --git a/frontend/src/api/admin/usage.ts b/frontend/src/api/admin/usage.ts index d21b28dcdb..37df7553ff 100644 --- a/frontend/src/api/admin/usage.ts +++ b/frontend/src/api/admin/usage.ts @@ -81,6 +81,8 @@ export interface AdminUsageQueryParams extends UsageQueryParams { user_id?: number exact_total?: boolean billing_mode?: string + sort_by?: string + sort_order?: 'asc' | 'desc' } // ==================== API Functions ==================== diff --git a/frontend/src/api/admin/users.ts b/frontend/src/api/admin/users.ts index bbf0ab5120..39cb1dfa69 100644 --- a/frontend/src/api/admin/users.ts +++ b/frontend/src/api/admin/users.ts @@ -24,6 +24,8 @@ export async function list( group_name?: string // fuzzy filter by allowed group name attributes?: Record // attributeId -> value include_subscriptions?: boolean + sort_by?: string + sort_order?: 'asc' | 'desc' }, options?: { signal?: AbortSignal @@ -37,7 +39,9 @@ export async function list( role: filters?.role, search: filters?.search, group_name: filters?.group_name, - include_subscriptions: filters?.include_subscriptions + include_subscriptions: filters?.include_subscriptions, + sort_by: filters?.sort_by, + sort_order: filters?.sort_order } // Add attribute filters as attr[id]=value diff --git a/frontend/src/api/auth.ts b/frontend/src/api/auth.ts index c5e1f35daf..837c4f4cf7 100644 --- a/frontend/src/api/auth.ts +++ b/frontend/src/api/auth.ts @@ -357,6 +357,28 @@ export async function completeLinuxDoOAuthRegistration( return data } +/** + * Complete OIDC OAuth registration by supplying an invitation code + * @param pendingOAuthToken - Short-lived JWT from the OAuth callback + * @param invitationCode - Invitation code entered by the user + * @returns Token pair on success + */ +export async function completeOIDCOAuthRegistration( + pendingOAuthToken: string, + invitationCode: string +): Promise<{ access_token: string; refresh_token: string; expires_in: number; token_type: string }> { + const { data } = await apiClient.post<{ + access_token: string + refresh_token: string + expires_in: number + token_type: string + }>('/auth/oauth/oidc/complete-registration', { + pending_oauth_token: pendingOAuthToken, + invitation_code: invitationCode + }) + return data +} + export const authAPI = { login, login2FA, @@ -380,7 +402,8 @@ export const authAPI = { resetPassword, refreshToken, revokeAllSessions, - completeLinuxDoOAuthRegistration + completeLinuxDoOAuthRegistration, + completeOIDCOAuthRegistration } export default authAPI diff --git a/frontend/src/api/keys.ts b/frontend/src/api/keys.ts index 137e10bafc..34dd5b4b76 100644 --- a/frontend/src/api/keys.ts +++ b/frontend/src/api/keys.ts @@ -17,7 +17,13 @@ import type { ApiKey, CreateApiKeyRequest, UpdateApiKeyRequest, PaginatedRespons export async function list( page: number = 1, pageSize: number = 10, - filters?: { search?: string; status?: string; group_id?: number | string }, + filters?: { + search?: string + status?: string + group_id?: number | string + sort_by?: string + sort_order?: 'asc' | 'desc' + }, options?: { signal?: AbortSignal } diff --git a/frontend/src/api/usage.ts b/frontend/src/api/usage.ts index 6efd7657fc..802c428f80 100644 --- a/frontend/src/api/usage.ts +++ b/frontend/src/api/usage.ts @@ -91,7 +91,7 @@ export async function list( * @returns Paginated list of usage logs */ export async function query( - params: UsageQueryParams, + params: UsageQueryParams & { sort_by?: string; sort_order?: 'asc' | 'desc' }, config: { signal?: AbortSignal } = {} ): Promise> { const { data } = await apiClient.get>('/usage', { diff --git a/frontend/src/components/admin/account/AccountTableFilters.vue b/frontend/src/components/admin/account/AccountTableFilters.vue index 6b4741830f..b33dad84e0 100644 --- a/frontend/src/components/admin/account/AccountTableFilters.vue +++ b/frontend/src/components/admin/account/AccountTableFilters.vue @@ -27,7 +27,7 @@ const updatePrivacyMode = (value: string | number | boolean | null) => { emit('u const updateGroup = (value: string | number | boolean | null) => { emit('update:filters', { ...props.filters, group: value }) } const pOpts = computed(() => [{ value: '', label: t('admin.accounts.allPlatforms') }, { value: 'anthropic', label: 'Anthropic' }, { value: 'openai', label: 'OpenAI' }, { value: 'gemini', label: 'Gemini' }, { value: 'antigravity', label: 'Antigravity' }]) const tOpts = computed(() => [{ value: '', label: t('admin.accounts.allTypes') }, { value: 'oauth', label: t('admin.accounts.oauthType') }, { value: 'setup-token', label: t('admin.accounts.setupToken') }, { value: 'apikey', label: t('admin.accounts.apiKey') }, { value: 'bedrock', label: 'AWS Bedrock' }]) -const sOpts = computed(() => [{ value: '', label: t('admin.accounts.allStatus') }, { value: 'active', label: t('admin.accounts.status.active') }, { value: 'inactive', label: t('admin.accounts.status.inactive') }, { value: 'error', label: t('admin.accounts.status.error') }, { value: 'rate_limited', label: t('admin.accounts.status.rateLimited') }, { value: 'temp_unschedulable', label: t('admin.accounts.status.tempUnschedulable') }]) +const sOpts = computed(() => [{ value: '', label: t('admin.accounts.allStatus') }, { value: 'active', label: t('admin.accounts.status.active') }, { value: 'inactive', label: t('admin.accounts.status.inactive') }, { value: 'error', label: t('admin.accounts.status.error') }, { value: 'rate_limited', label: t('admin.accounts.status.rateLimited') }, { value: 'temp_unschedulable', label: t('admin.accounts.status.tempUnschedulable') }, { value: 'unschedulable', label: t('admin.accounts.status.unschedulable') }]) const privacyOpts = computed(() => [ { value: '', label: t('admin.accounts.allPrivacyModes') }, { value: '__unset__', label: t('admin.accounts.privacyUnset') }, diff --git a/frontend/src/components/admin/account/__tests__/AccountTableFilters.spec.ts b/frontend/src/components/admin/account/__tests__/AccountTableFilters.spec.ts deleted file mode 100644 index 5a0044e5f2..0000000000 --- a/frontend/src/components/admin/account/__tests__/AccountTableFilters.spec.ts +++ /dev/null @@ -1,56 +0,0 @@ -import { describe, expect, it, vi } from 'vitest' -import { mount } from '@vue/test-utils' - -import AccountTableFilters from '../AccountTableFilters.vue' - -vi.mock('vue-i18n', async () => { - const actual = await vi.importActual('vue-i18n') - return { - ...actual, - useI18n: () => ({ - t: (key: string) => key - }) - } -}) - -describe('AccountTableFilters', () => { - it('renders privacy mode options and emits privacy_mode updates', async () => { - const wrapper = mount(AccountTableFilters, { - props: { - searchQuery: '', - filters: { - platform: '', - type: '', - status: '', - group: '', - privacy_mode: '' - }, - groups: [] - }, - global: { - stubs: { - SearchInput: { - template: '
' - }, - Select: { - props: ['modelValue', 'options'], - emits: ['update:modelValue', 'change'], - template: '
' - } - } - } - }) - - const selects = wrapper.findAll('.select-stub') - expect(selects).toHaveLength(5) - - const privacyOptions = JSON.parse(selects[3].attributes('data-options')) - expect(privacyOptions).toEqual([ - { value: '', label: 'admin.accounts.allPrivacyModes' }, - { value: '__unset__', label: 'admin.accounts.privacyUnset' }, - { value: 'training_off', label: 'Privacy' }, - { value: 'training_set_cf_blocked', label: 'CF' }, - { value: 'training_set_failed', label: 'Fail' } - ]) - }) -}) diff --git a/frontend/src/components/admin/announcements/AnnouncementReadStatusDialog.vue b/frontend/src/components/admin/announcements/AnnouncementReadStatusDialog.vue index a0d9de3ca8..60c01c6dba 100644 --- a/frontend/src/components/admin/announcements/AnnouncementReadStatusDialog.vue +++ b/frontend/src/components/admin/announcements/AnnouncementReadStatusDialog.vue @@ -21,7 +21,15 @@
- + @@ -62,7 +70,7 @@ diff --git a/frontend/src/components/admin/announcements/__tests__/AnnouncementReadStatusDialog.spec.ts b/frontend/src/components/admin/announcements/__tests__/AnnouncementReadStatusDialog.spec.ts new file mode 100644 index 0000000000..26c87d73dc --- /dev/null +++ b/frontend/src/components/admin/announcements/__tests__/AnnouncementReadStatusDialog.spec.ts @@ -0,0 +1,95 @@ +import { describe, it, expect, vi, beforeEach } from 'vitest' +import { flushPromises, mount } from '@vue/test-utils' + +import AnnouncementReadStatusDialog from '../AnnouncementReadStatusDialog.vue' + +const { getReadStatus, showError } = vi.hoisted(() => ({ + getReadStatus: vi.fn(), + showError: vi.fn(), +})) + +vi.mock('@/api/admin', () => ({ + adminAPI: { + announcements: { + getReadStatus, + }, + }, +})) + +vi.mock('@/stores/app', () => ({ + useAppStore: () => ({ + showError, + }), +})) + +vi.mock('vue-i18n', async () => { + const actual = await vi.importActual('vue-i18n') + return { + ...actual, + useI18n: () => ({ + t: (key: string) => key, + }), + } +}) + +vi.mock('@/composables/usePersistedPageSize', () => ({ + getPersistedPageSize: () => 20, +})) + +const BaseDialogStub = { + props: ['show', 'title', 'width'], + emits: ['close'], + template: '
', +} + +describe('AnnouncementReadStatusDialog', () => { + beforeEach(() => { + getReadStatus.mockReset() + showError.mockReset() + vi.useFakeTimers() + }) + + it('closes by aborting active requests and clearing debounced reloads', async () => { + let activeSignal: AbortSignal | undefined + getReadStatus.mockImplementation(async (...args: any[]) => { + activeSignal = args[4]?.signal + return new Promise(() => {}) + }) + + const wrapper = mount(AnnouncementReadStatusDialog, { + props: { + show: false, + announcementId: 1, + }, + global: { + stubs: { + BaseDialog: BaseDialogStub, + DataTable: true, + Pagination: true, + Icon: true, + }, + }, + }) + + await wrapper.setProps({ show: true }) + await flushPromises() + + expect(getReadStatus).toHaveBeenCalledTimes(1) + expect(activeSignal?.aborted).toBe(false) + + const setupState = (wrapper.vm as any).$?.setupState + setupState.search = 'alice' + setupState.handleSearch() + + setupState.handleClose() + await flushPromises() + + expect(activeSignal?.aborted).toBe(true) + expect(wrapper.emitted('close')).toHaveLength(1) + + vi.advanceTimersByTime(350) + await flushPromises() + + expect(getReadStatus).toHaveBeenCalledTimes(1) + }) +}) diff --git a/frontend/src/components/admin/group/GroupRateMultipliersModal.vue b/frontend/src/components/admin/group/GroupRateMultipliersModal.vue index cbd18af6b4..bf79bea200 100644 --- a/frontend/src/components/admin/group/GroupRateMultipliersModal.vue +++ b/frontend/src/components/admin/group/GroupRateMultipliersModal.vue @@ -196,7 +196,6 @@ :total="localEntries.length" :page="currentPage" :page-size="pageSize" - :page-size-options="[10, 20, 50]" @update:page="currentPage = $event" @update:pageSize="handlePageSizeChange" /> diff --git a/frontend/src/components/admin/usage/UsageTable.vue b/frontend/src/components/admin/usage/UsageTable.vue index 9bbdb380b6..f4494e69c2 100644 --- a/frontend/src/components/admin/usage/UsageTable.vue +++ b/frontend/src/components/admin/usage/UsageTable.vue @@ -1,7 +1,15 @@ diff --git a/frontend/src/views/admin/PromoCodesView.vue b/frontend/src/views/admin/PromoCodesView.vue index 7bf670f78f..ceca0dd4e8 100644 --- a/frontend/src/views/admin/PromoCodesView.vue +++ b/frontend/src/views/admin/PromoCodesView.vue @@ -39,7 +39,15 @@