diff --git a/.github/scripts/ui-baseline.json b/.github/scripts/ui-baseline.json index c8119337e..3e209a35e 100644 --- a/.github/scripts/ui-baseline.json +++ b/.github/scripts/ui-baseline.json @@ -1,7 +1,7 @@ { "schema_version": 1, - "baseline_ref": "ui-approved-2026-10-01-v0211-version-badge", - "baseline_commit": "be36f46044301526977db2a2800dfb6be0f01de7", + "baseline_ref": "ui-approved-2026-10-02-v0212-version-badge", + "baseline_commit": "0a614bcb02f258881e466f75e3198bfbdf6c843b", "protected_paths": [ "landing/", "frontend/", diff --git a/.github/scripts/verify-upstream-boundary.test.mjs b/.github/scripts/verify-upstream-boundary.test.mjs index f700dd9bb..2fae3dd8d 100644 --- a/.github/scripts/verify-upstream-boundary.test.mjs +++ b/.github/scripts/verify-upstream-boundary.test.mjs @@ -763,6 +763,7 @@ test('allows named immutable exceptions while adjacent seam files still fail', ( 'frontend/src/api/admin/affiliates.ts', 'frontend/src/api/admin/ops.ts', 'frontend/src/api/__tests__/admin.affiliates.spec.ts', + 'frontend/src/api/__tests__/settings.authSourceDefaults.spec.ts', 'frontend/src/api/__tests__/admin.users.spec.ts', 'frontend/src/api/admin/redeem.ts', 'frontend/src/api/__tests__/admin.redeem.spec.ts', @@ -926,6 +927,12 @@ test('allows named immutable exceptions while adjacent seam files still fail', ( path: 'frontend/src/api/admin/ops.ts', immutable_path: 'frontend/src/api/', }, + { + name: 'public-capabilities-auth-source-quota-type-regression', + owner: 'Public Capabilities', + path: 'frontend/src/api/__tests__/settings.authSourceDefaults.spec.ts', + immutable_path: 'frontend/src/api/', + }, ], ) }) diff --git a/.github/upstream-baseline.json b/.github/upstream-baseline.json index 1db056662..d9e29e944 100644 --- a/.github/upstream-baseline.json +++ b/.github/upstream-baseline.json @@ -1,13 +1,13 @@ { "schema_version": 5, "repository": "Wei-Shaw/sub2api", - "release": "v0.2.11", - "commit": "96f4c115c9749078f90cbf210a01d39baf3f53b6", + "release": "v0.2.12", + "commit": "5106065716e494204fc0e8db16f68f6e9d576be0", "upstream_sync": { - "previous_release": "v0.2.8", - "previous_commit": "fd80b08c90b55edcad5b00171b53f08721d30da1", - "product_commit": "ce899e5ae03cf4265855022a54d2d1ef5b5c5a68", - "merge_commit": "9d9aecbb860b8bceb1ae0c031a7c77f08c7b3d83" + "previous_release": "v0.2.11", + "previous_commit": "96f4c115c9749078f90cbf210a01d39baf3f53b6", + "product_commit": "6a540fca7db34220add8d4c6a2f71f1d8ea61617", + "merge_commit": "a29346e2dd0406fa2c82bc3102ba804884f88f98" }, "approved_backports": [], "preserve_bytes_on_upstream_sync": [ @@ -1315,7 +1315,10 @@ "docs/upgrades/v0.2.11.md", "docs/upgrades/v0.2.11-change-map.json", "backend/internal/service/account_stats_pricing_test.go", - "backend/internal/service/billing_inflight_reservation_test.go" + "backend/internal/service/billing_inflight_reservation_test.go", + "docs/upgrades/v0.2.12-change-map.json", + "docs/upgrades/v0.2.12.md", + "frontend/src/api/__tests__/settings.authSourceDefaults.spec.ts" ], "retired_preserved_paths": [ { @@ -1874,7 +1877,7 @@ "legacy_hotfixes": [ { "name": "approved-v0.1.179-production-correctness-and-billing-compatibility", - "valid_for_release": "v0.2.11", + "valid_for_release": "v0.2.12", "exit_condition": "Remove each path when the first stable upstream release containing its equivalent production-correctness fix becomes the baseline; remove the entire block when no listed path remains.", "paths": [ "backend/internal/handler/admin/admin_basic_handlers_test.go", @@ -1904,7 +1907,7 @@ }, { "name": "approved-v0.1.179-version-metadata-alignment", - "valid_for_release": "v0.2.11", + "valid_for_release": "v0.2.12", "exit_condition": "Remove this metadata alignment when the first stable upstream release makes backend/cmd/server/VERSION match the baseline release.", "paths": [ "backend/cmd/server/VERSION" @@ -1912,7 +1915,7 @@ }, { "name": "approved-v0.1.179-race-safety", - "valid_for_release": "v0.2.11", + "valid_for_release": "v0.2.12", "exit_condition": "Remove each path when the first stable upstream release contains the equivalent race-safe runtime cache and asynchronous test fixture handling.", "paths": [ "backend/internal/service/content_moderation.go", @@ -1927,7 +1930,7 @@ }, { "name": "approved-v0.1.179-remove-unconditional-sticky-debug-logs", - "valid_for_release": "v0.2.11", + "valid_for_release": "v0.2.12", "exit_condition": "Remove this hotfix when the first stable upstream release removes or explicitly gates the equivalent per-request sticky-session debug logs.", "paths": [ "backend/internal/service/gateway_scheduling.go" @@ -1935,7 +1938,7 @@ }, { "name": "approved-v0.1.181-grok-mapping-test-isolation", - "valid_for_release": "v0.2.11", + "valid_for_release": "v0.2.12", "exit_condition": "Remove this hotfix when the first stable upstream release contains equivalent runtime model-mapping isolation in the canonical scheduling test.", "paths": [ "backend/internal/service/openai_model_mapping_test.go" @@ -1943,7 +1946,7 @@ }, { "name": "approved-v0.1.182-go-dependency-security-updates", - "valid_for_release": "v0.2.11", + "valid_for_release": "v0.2.12", "exit_condition": "Remove this hotfix when the first stable upstream release uses golang.org/x/image v0.45.0 or later, Testcontainers v0.44.0 or later, github.com/moby/go-archive v0.3.0 or later, and no longer depends on github.com/docker/docker.", "paths": [ "backend/go.mod", @@ -1955,7 +1958,7 @@ }, { "name": "approved-v0.2.7-sse-keepalive-test-determinism", - "valid_for_release": "v0.2.11", + "valid_for_release": "v0.2.12", "exit_condition": "Remove this hotfix when the first stable upstream release makes the ordinary-client SSE keepalive test tolerate a comment payload and at least two ticker periods of runner scheduling delay.", "paths": [ "backend/internal/service/gemini_sse_comment_compat_test.go" @@ -2594,7 +2597,10 @@ "docs/upgrades/v0.2.11.md", "docs/upgrades/v0.2.11-change-map.json", "backend/internal/service/account_stats_pricing_test.go", - "backend/internal/service/billing_inflight_reservation_test.go" + "backend/internal/service/billing_inflight_reservation_test.go", + "docs/upgrades/v0.2.12-change-map.json", + "docs/upgrades/v0.2.12.md", + "frontend/src/api/__tests__/settings.authSourceDefaults.spec.ts" ] }, { @@ -2847,6 +2853,12 @@ "owner": "Public Capabilities", "path": "frontend/src/api/admin/ops.ts", "immutable_path": "frontend/src/api/" + }, + { + "name": "public-capabilities-auth-source-quota-type-regression", + "owner": "Public Capabilities", + "path": "frontend/src/api/__tests__/settings.authSourceDefaults.spec.ts", + "immutable_path": "frontend/src/api/" } ] } diff --git a/CONTEXT.md b/CONTEXT.md index e47337940..68035baec 100644 --- a/CONTEXT.md +++ b/CONTEXT.md @@ -71,7 +71,7 @@ _Avoid_: 普通用户、站长账号 _Avoid_: User、客户账号 **Provider Platform Catalog(供应商平台目录)**: -Provider Account、分组筛选、渠道定价、渠道监控、运维筛选、订阅筛选、错误透传规则与 Composite 路由目标共同使用的有序一等平台集合;新增平台必须在一个目录中同时进入这些管理入口。它不同于模型白名单,Qwen、Mistral 等模型名称不自动成为 Provider Account 平台。 +Provider Account、分组筛选、渠道定价、渠道监控、运维筛选、订阅筛选、错误透传规则与 Composite 路由目标共同使用的有序一等平台集合;新增平台必须在一个目录中同时进入这些管理入口。它不同于模型白名单,Qwen、Mistral 等模型名称不自动成为 Provider Account 平台。TypeSafe / Jev System One 是原生结构化判定平台,不是对话端点,因此不进入主动渠道监控的探测子集。 _Avoid_: 模型白名单、各页面复制的平台数组、任意 OpenAI 兼容模型列表 **API Key(密钥)**: diff --git a/README.md b/README.md index cbb60d950..de4a7ae62 100644 --- a/README.md +++ b/README.md @@ -816,6 +816,31 @@ Administrators can override automatic media eligibility through the account crea --- +## TypeSafe / Jev Support + +Sub2API supports TypeSafe API-key accounts through Jev's native, non-streaming System One protocol. + +- Platform: `typesafe`; account type: API Key +- Default upstream: `https://api.typesafe.ai` +- Public endpoint: `POST /v1/systemone` +- Model: `jev-latest`, also returned by `/v1/models` for TypeSafe groups +- Questions: `noul`, `choice`, and `score` + +Requests and successful responses retain the native System One JSON structure. This endpoint is not compatible with Chat Completions, Responses, Anthropic Messages, or streaming clients. + +Question validation follows the TypeSafe OpenAPI wire schema (also used by SDK v0.5.7). `instructions` may be omitted or `null` for all question types. Noul `criteria` may be omitted or `null`; its `true`/`false` descriptions and Choice descriptions accept strings, objects, arrays, or `null`. Score `criteria` must be a non-empty array of string, object, or array descriptions; a single level is valid. SDK integer-keyed Score maps are normalized to arrays by the SDK before sending. + +```bash +curl https://your-sub2api.example.com/v1/systemone \ + -H 'Authorization: Bearer sk-your-sub2api-key' \ + -H 'Content-Type: application/json' \ + --data '{"model":"jev-latest","state":"Text to evaluate","questions":{"safety":{"type":"noul","instructions":"Evaluate whether the text is unsafe"}}}' +``` + +The built-in `jev-latest` price is `$0.042` per million input tokens and `$0` for output tokens. Channel pricing can override both values. Credential, billing, permission, rate-limit, overload, server, and network failures (`401`, `402`, `403`, `429`, `529`, `5xx`, transport errors) use the existing account error policy (including custom error codes and temporary-unschedulable rules) and fail over to another account; request errors (`400`, `413`, and `422`) are returned without retrying another account and never change account state. TypeSafe groups (and Composite requests routed to TypeSafe) reject Messages, Chat Completions, Responses, and count_tokens requests with `404`. + +--- + ## Antigravity Support Sub2API supports [Antigravity](https://antigravity.so/) accounts. After authorization, dedicated endpoints are available for Claude and Gemini models. diff --git a/README_CN.md b/README_CN.md index 9c35d23b1..6780d4ac9 100644 --- a/README_CN.md +++ b/README_CN.md @@ -745,6 +745,31 @@ go generate ./cmd/server --- +## TypeSafe / Jev 使用说明 + +Sub2API 支持使用 TypeSafe API Key 账户,通过 Jev 原生、非流式的 System One 协议调用模型。 + +- 平台:`typesafe`;账号类型:API Key +- 默认上游:`https://api.typesafe.ai` +- 对外端点:`POST /v1/systemone` +- 模型:`jev-latest`,TypeSafe 分组的 `/v1/models` 也会返回该模型 +- 问题类型:`noul`、`choice`、`score` + +请求和成功响应保持 System One 原生 JSON 结构。该端点不兼容 Chat Completions、Responses、Anthropic Messages 或流式客户端。 + +问题校验遵循 TypeSafe OpenAPI 的线上协议 schema(SDK v0.5.7 也使用该 schema)。所有问题的 `instructions` 都可以省略或为 `null`。Noul 的 `criteria` 可以省略或为 `null`,其中 `true`/`false` 的描述和 Choice 描述支持字符串、对象、数组或 `null`。Score 的 `criteria` 必须是至少包含一档描述的数组,每档支持字符串、对象或数组;单档也合法。SDK 的整数键 Score 映射会由 SDK 在发送前转换为数组。 + +```bash +curl https://your-sub2api.example.com/v1/systemone \ + -H 'Authorization: Bearer sk-your-sub2api-key' \ + -H 'Content-Type: application/json' \ + --data '{"model":"jev-latest","state":"待评估文本","questions":{"safety":{"type":"noul","instructions":"评估文本是否不安全"}}}' +``` + +`jev-latest` 内置价格为输入 `$0.042/百万 tokens`、输出 `$0`,渠道定价可以覆盖。凭据、欠费、权限、限流、过载、服务端和网络错误(`401`、`402`、`403`、`429`、`529`、`5xx`、传输错误)沿用现有账号错误策略(含自定义错误码与临时不可调度规则)并切换账号;请求错误(`400`、`413`、`422`)不会切换账号重试,也不会改变账号状态。TypeSafe 分组(以及路由到 TypeSafe 的 Composite 请求)调用 Messages、Chat Completions、Responses、count_tokens 时返回 `404`。 + +--- + ## Antigravity 使用说明 Sub2API 支持 [Antigravity](https://antigravity.so/) 账户,授权后可通过专用端点访问 Claude 和 Gemini 模型。 diff --git a/backend/cmd/server/VERSION b/backend/cmd/server/VERSION index d3b5ba4bf..f2722b133 100644 --- a/backend/cmd/server/VERSION +++ b/backend/cmd/server/VERSION @@ -1 +1 @@ -0.2.11 +0.2.12 diff --git a/backend/ent/migrate/schema.go b/backend/ent/migrate/schema.go index b447dd5b9..e8bb12f28 100644 --- a/backend/ent/migrate/schema.go +++ b/backend/ent/migrate/schema.go @@ -1129,6 +1129,7 @@ var ( {Name: "amount", Type: field.TypeFloat64, SchemaType: map[string]string{"postgres": "decimal(20,2)"}}, {Name: "pay_amount", Type: field.TypeFloat64, SchemaType: map[string]string{"postgres": "decimal(20,2)"}}, {Name: "fee_rate", Type: field.TypeFloat64, Default: 0, SchemaType: map[string]string{"postgres": "decimal(10,4)"}}, + {Name: "bonus_amount", Type: field.TypeFloat64, Default: 0, SchemaType: map[string]string{"postgres": "decimal(20,2)"}}, {Name: "recharge_code", Type: field.TypeString, Size: 64}, {Name: "out_trade_no", Type: field.TypeString, Size: 64, Default: ""}, {Name: "payment_type", Type: field.TypeString, Size: 30}, @@ -1171,7 +1172,7 @@ var ( ForeignKeys: []*schema.ForeignKey{ { Symbol: "payment_orders_users_payment_orders", - Columns: []*schema.Column{PaymentOrdersColumns[39]}, + Columns: []*schema.Column{PaymentOrdersColumns[40]}, RefColumns: []*schema.Column{UsersColumns[0]}, OnDelete: schema.NoAction, }, @@ -1180,7 +1181,7 @@ var ( { Name: "paymentorder_out_trade_no", Unique: true, - Columns: []*schema.Column{PaymentOrdersColumns[8]}, + Columns: []*schema.Column{PaymentOrdersColumns[9]}, Annotation: &entsql.IndexAnnotation{ Where: "out_trade_no <> ''", }, @@ -1188,37 +1189,37 @@ var ( { Name: "paymentorder_user_id", Unique: false, - Columns: []*schema.Column{PaymentOrdersColumns[39]}, + Columns: []*schema.Column{PaymentOrdersColumns[40]}, }, { Name: "paymentorder_status", Unique: false, - Columns: []*schema.Column{PaymentOrdersColumns[21]}, + Columns: []*schema.Column{PaymentOrdersColumns[22]}, }, { Name: "paymentorder_expires_at", Unique: false, - Columns: []*schema.Column{PaymentOrdersColumns[29]}, + Columns: []*schema.Column{PaymentOrdersColumns[30]}, }, { Name: "paymentorder_created_at", Unique: false, - Columns: []*schema.Column{PaymentOrdersColumns[37]}, + Columns: []*schema.Column{PaymentOrdersColumns[38]}, }, { Name: "paymentorder_paid_at", Unique: false, - Columns: []*schema.Column{PaymentOrdersColumns[30]}, + Columns: []*schema.Column{PaymentOrdersColumns[31]}, }, { Name: "paymentorder_payment_type_paid_at", Unique: false, - Columns: []*schema.Column{PaymentOrdersColumns[9], PaymentOrdersColumns[30]}, + Columns: []*schema.Column{PaymentOrdersColumns[10], PaymentOrdersColumns[31]}, }, { Name: "paymentorder_order_type", Unique: false, - Columns: []*schema.Column{PaymentOrdersColumns[14]}, + Columns: []*schema.Column{PaymentOrdersColumns[15]}, }, }, } diff --git a/backend/ent/mutation.go b/backend/ent/mutation.go index 922672ae5..a9572c014 100644 --- a/backend/ent/mutation.go +++ b/backend/ent/mutation.go @@ -30167,6 +30167,8 @@ type PaymentOrderMutation struct { addpay_amount *float64 fee_rate *float64 addfee_rate *float64 + bonus_amount *float64 + addbonus_amount *float64 recharge_code *string out_trade_no *string payment_type *string @@ -30634,6 +30636,62 @@ func (m *PaymentOrderMutation) ResetFeeRate() { m.addfee_rate = nil } +// SetBonusAmount sets the "bonus_amount" field. +func (m *PaymentOrderMutation) SetBonusAmount(f float64) { + m.bonus_amount = &f + m.addbonus_amount = nil +} + +// BonusAmount returns the value of the "bonus_amount" field in the mutation. +func (m *PaymentOrderMutation) BonusAmount() (r float64, exists bool) { + v := m.bonus_amount + if v == nil { + return + } + return *v, true +} + +// OldBonusAmount returns the old "bonus_amount" field's value of the PaymentOrder entity. +// If the PaymentOrder object wasn't provided to the builder, the object is fetched from the database. +// An error is returned if the mutation operation is not UpdateOne, or the database query fails. +func (m *PaymentOrderMutation) OldBonusAmount(ctx context.Context) (v float64, err error) { + if !m.op.Is(OpUpdateOne) { + return v, errors.New("OldBonusAmount is only allowed on UpdateOne operations") + } + if m.id == nil || m.oldValue == nil { + return v, errors.New("OldBonusAmount requires an ID field in the mutation") + } + oldValue, err := m.oldValue(ctx) + if err != nil { + return v, fmt.Errorf("querying old value for OldBonusAmount: %w", err) + } + return oldValue.BonusAmount, nil +} + +// AddBonusAmount adds f to the "bonus_amount" field. +func (m *PaymentOrderMutation) AddBonusAmount(f float64) { + if m.addbonus_amount != nil { + *m.addbonus_amount += f + } else { + m.addbonus_amount = &f + } +} + +// AddedBonusAmount returns the value that was added to the "bonus_amount" field in this mutation. +func (m *PaymentOrderMutation) AddedBonusAmount() (r float64, exists bool) { + v := m.addbonus_amount + if v == nil { + return + } + return *v, true +} + +// ResetBonusAmount resets all changes to the "bonus_amount" field. +func (m *PaymentOrderMutation) ResetBonusAmount() { + m.bonus_amount = nil + m.addbonus_amount = nil +} + // SetRechargeCode sets the "recharge_code" field. func (m *PaymentOrderMutation) SetRechargeCode(s string) { m.recharge_code = &s @@ -32177,7 +32235,7 @@ func (m *PaymentOrderMutation) Type() string { // order to get all numeric fields that were incremented/decremented, call // AddedFields(). func (m *PaymentOrderMutation) Fields() []string { - fields := make([]string, 0, 39) + fields := make([]string, 0, 40) if m.user != nil { fields = append(fields, paymentorder.FieldUserID) } @@ -32199,6 +32257,9 @@ func (m *PaymentOrderMutation) Fields() []string { if m.fee_rate != nil { fields = append(fields, paymentorder.FieldFeeRate) } + if m.bonus_amount != nil { + fields = append(fields, paymentorder.FieldBonusAmount) + } if m.recharge_code != nil { fields = append(fields, paymentorder.FieldRechargeCode) } @@ -32317,6 +32378,8 @@ func (m *PaymentOrderMutation) Field(name string) (ent.Value, bool) { return m.PayAmount() case paymentorder.FieldFeeRate: return m.FeeRate() + case paymentorder.FieldBonusAmount: + return m.BonusAmount() case paymentorder.FieldRechargeCode: return m.RechargeCode() case paymentorder.FieldOutTradeNo: @@ -32404,6 +32467,8 @@ func (m *PaymentOrderMutation) OldField(ctx context.Context, name string) (ent.V return m.OldPayAmount(ctx) case paymentorder.FieldFeeRate: return m.OldFeeRate(ctx) + case paymentorder.FieldBonusAmount: + return m.OldBonusAmount(ctx) case paymentorder.FieldRechargeCode: return m.OldRechargeCode(ctx) case paymentorder.FieldOutTradeNo: @@ -32526,6 +32591,13 @@ func (m *PaymentOrderMutation) SetField(name string, value ent.Value) error { } m.SetFeeRate(v) return nil + case paymentorder.FieldBonusAmount: + v, ok := value.(float64) + if !ok { + return fmt.Errorf("unexpected type %T for field %s", value, name) + } + m.SetBonusAmount(v) + return nil case paymentorder.FieldRechargeCode: v, ok := value.(string) if !ok { @@ -32767,6 +32839,9 @@ func (m *PaymentOrderMutation) AddedFields() []string { if m.addfee_rate != nil { fields = append(fields, paymentorder.FieldFeeRate) } + if m.addbonus_amount != nil { + fields = append(fields, paymentorder.FieldBonusAmount) + } if m.addplan_id != nil { fields = append(fields, paymentorder.FieldPlanID) } @@ -32793,6 +32868,8 @@ func (m *PaymentOrderMutation) AddedField(name string) (ent.Value, bool) { return m.AddedPayAmount() case paymentorder.FieldFeeRate: return m.AddedFeeRate() + case paymentorder.FieldBonusAmount: + return m.AddedBonusAmount() case paymentorder.FieldPlanID: return m.AddedPlanID() case paymentorder.FieldSubscriptionGroupID: @@ -32831,6 +32908,13 @@ func (m *PaymentOrderMutation) AddField(name string, value ent.Value) error { } m.AddFeeRate(v) return nil + case paymentorder.FieldBonusAmount: + v, ok := value.(float64) + if !ok { + return fmt.Errorf("unexpected type %T for field %s", value, name) + } + m.AddBonusAmount(v) + return nil case paymentorder.FieldPlanID: v, ok := value.(int64) if !ok { @@ -33030,6 +33114,9 @@ func (m *PaymentOrderMutation) ResetField(name string) error { case paymentorder.FieldFeeRate: m.ResetFeeRate() return nil + case paymentorder.FieldBonusAmount: + m.ResetBonusAmount() + return nil case paymentorder.FieldRechargeCode: m.ResetRechargeCode() return nil diff --git a/backend/ent/paymentorder.go b/backend/ent/paymentorder.go index b131b8c88..db5699499 100644 --- a/backend/ent/paymentorder.go +++ b/backend/ent/paymentorder.go @@ -33,6 +33,8 @@ type PaymentOrder struct { PayAmount float64 `json:"pay_amount,omitempty"` // FeeRate holds the value of the "fee_rate" field. FeeRate float64 `json:"fee_rate,omitempty"` + // BonusAmount holds the value of the "bonus_amount" field. + BonusAmount float64 `json:"bonus_amount,omitempty"` // RechargeCode holds the value of the "recharge_code" field. RechargeCode string `json:"recharge_code,omitempty"` // OutTradeNo holds the value of the "out_trade_no" field. @@ -132,7 +134,7 @@ func (*PaymentOrder) scanValues(columns []string) ([]any, error) { values[i] = new([]byte) case paymentorder.FieldForceRefund: values[i] = new(sql.NullBool) - case paymentorder.FieldAmount, paymentorder.FieldPayAmount, paymentorder.FieldFeeRate, paymentorder.FieldRefundAmount: + case paymentorder.FieldAmount, paymentorder.FieldPayAmount, paymentorder.FieldFeeRate, paymentorder.FieldBonusAmount, paymentorder.FieldRefundAmount: values[i] = new(sql.NullFloat64) case paymentorder.FieldID, paymentorder.FieldUserID, paymentorder.FieldPlanID, paymentorder.FieldSubscriptionGroupID, paymentorder.FieldSubscriptionDays: values[i] = new(sql.NullInt64) @@ -204,6 +206,12 @@ func (_m *PaymentOrder) assignValues(columns []string, values []any) error { } else if value.Valid { _m.FeeRate = value.Float64 } + case paymentorder.FieldBonusAmount: + if value, ok := values[i].(*sql.NullFloat64); !ok { + return fmt.Errorf("unexpected type %T for field bonus_amount", values[i]) + } else if value.Valid { + _m.BonusAmount = value.Float64 + } case paymentorder.FieldRechargeCode: if value, ok := values[i].(*sql.NullString); !ok { return fmt.Errorf("unexpected type %T for field recharge_code", values[i]) @@ -480,6 +488,9 @@ func (_m *PaymentOrder) String() string { builder.WriteString("fee_rate=") builder.WriteString(fmt.Sprintf("%v", _m.FeeRate)) builder.WriteString(", ") + builder.WriteString("bonus_amount=") + builder.WriteString(fmt.Sprintf("%v", _m.BonusAmount)) + builder.WriteString(", ") builder.WriteString("recharge_code=") builder.WriteString(_m.RechargeCode) builder.WriteString(", ") diff --git a/backend/ent/paymentorder/paymentorder.go b/backend/ent/paymentorder/paymentorder.go index 628837943..391389b15 100644 --- a/backend/ent/paymentorder/paymentorder.go +++ b/backend/ent/paymentorder/paymentorder.go @@ -28,6 +28,8 @@ const ( FieldPayAmount = "pay_amount" // FieldFeeRate holds the string denoting the fee_rate field in the database. FieldFeeRate = "fee_rate" + // FieldBonusAmount holds the string denoting the bonus_amount field in the database. + FieldBonusAmount = "bonus_amount" // FieldRechargeCode holds the string denoting the recharge_code field in the database. FieldRechargeCode = "recharge_code" // FieldOutTradeNo holds the string denoting the out_trade_no field in the database. @@ -115,6 +117,7 @@ var Columns = []string{ FieldAmount, FieldPayAmount, FieldFeeRate, + FieldBonusAmount, FieldRechargeCode, FieldOutTradeNo, FieldPaymentType, @@ -166,6 +169,8 @@ var ( UserNameValidator func(string) error // DefaultFeeRate holds the default value on creation for the "fee_rate" field. DefaultFeeRate float64 + // DefaultBonusAmount holds the default value on creation for the "bonus_amount" field. + DefaultBonusAmount float64 // RechargeCodeValidator is a validator for the "recharge_code" field. It is called by the builders before save. RechargeCodeValidator func(string) error // DefaultOutTradeNo holds the default value on creation for the "out_trade_no" field. @@ -249,6 +254,11 @@ func ByFeeRate(opts ...sql.OrderTermOption) OrderOption { return sql.OrderByField(FieldFeeRate, opts...).ToFunc() } +// ByBonusAmount orders the results by the bonus_amount field. +func ByBonusAmount(opts ...sql.OrderTermOption) OrderOption { + return sql.OrderByField(FieldBonusAmount, opts...).ToFunc() +} + // ByRechargeCode orders the results by the recharge_code field. func ByRechargeCode(opts ...sql.OrderTermOption) OrderOption { return sql.OrderByField(FieldRechargeCode, opts...).ToFunc() diff --git a/backend/ent/paymentorder/where.go b/backend/ent/paymentorder/where.go index e96bf51eb..33e526779 100644 --- a/backend/ent/paymentorder/where.go +++ b/backend/ent/paymentorder/where.go @@ -90,6 +90,11 @@ func FeeRate(v float64) predicate.PaymentOrder { return predicate.PaymentOrder(sql.FieldEQ(FieldFeeRate, v)) } +// BonusAmount applies equality check predicate on the "bonus_amount" field. It's identical to BonusAmountEQ. +func BonusAmount(v float64) predicate.PaymentOrder { + return predicate.PaymentOrder(sql.FieldEQ(FieldBonusAmount, v)) +} + // RechargeCode applies equality check predicate on the "recharge_code" field. It's identical to RechargeCodeEQ. func RechargeCode(v string) predicate.PaymentOrder { return predicate.PaymentOrder(sql.FieldEQ(FieldRechargeCode, v)) @@ -590,6 +595,46 @@ func FeeRateLTE(v float64) predicate.PaymentOrder { return predicate.PaymentOrder(sql.FieldLTE(FieldFeeRate, v)) } +// BonusAmountEQ applies the EQ predicate on the "bonus_amount" field. +func BonusAmountEQ(v float64) predicate.PaymentOrder { + return predicate.PaymentOrder(sql.FieldEQ(FieldBonusAmount, v)) +} + +// BonusAmountNEQ applies the NEQ predicate on the "bonus_amount" field. +func BonusAmountNEQ(v float64) predicate.PaymentOrder { + return predicate.PaymentOrder(sql.FieldNEQ(FieldBonusAmount, v)) +} + +// BonusAmountIn applies the In predicate on the "bonus_amount" field. +func BonusAmountIn(vs ...float64) predicate.PaymentOrder { + return predicate.PaymentOrder(sql.FieldIn(FieldBonusAmount, vs...)) +} + +// BonusAmountNotIn applies the NotIn predicate on the "bonus_amount" field. +func BonusAmountNotIn(vs ...float64) predicate.PaymentOrder { + return predicate.PaymentOrder(sql.FieldNotIn(FieldBonusAmount, vs...)) +} + +// BonusAmountGT applies the GT predicate on the "bonus_amount" field. +func BonusAmountGT(v float64) predicate.PaymentOrder { + return predicate.PaymentOrder(sql.FieldGT(FieldBonusAmount, v)) +} + +// BonusAmountGTE applies the GTE predicate on the "bonus_amount" field. +func BonusAmountGTE(v float64) predicate.PaymentOrder { + return predicate.PaymentOrder(sql.FieldGTE(FieldBonusAmount, v)) +} + +// BonusAmountLT applies the LT predicate on the "bonus_amount" field. +func BonusAmountLT(v float64) predicate.PaymentOrder { + return predicate.PaymentOrder(sql.FieldLT(FieldBonusAmount, v)) +} + +// BonusAmountLTE applies the LTE predicate on the "bonus_amount" field. +func BonusAmountLTE(v float64) predicate.PaymentOrder { + return predicate.PaymentOrder(sql.FieldLTE(FieldBonusAmount, v)) +} + // RechargeCodeEQ applies the EQ predicate on the "recharge_code" field. func RechargeCodeEQ(v string) predicate.PaymentOrder { return predicate.PaymentOrder(sql.FieldEQ(FieldRechargeCode, v)) diff --git a/backend/ent/paymentorder_create.go b/backend/ent/paymentorder_create.go index 3ee24f8e9..a8693d091 100644 --- a/backend/ent/paymentorder_create.go +++ b/backend/ent/paymentorder_create.go @@ -81,6 +81,20 @@ func (_c *PaymentOrderCreate) SetNillableFeeRate(v *float64) *PaymentOrderCreate return _c } +// SetBonusAmount sets the "bonus_amount" field. +func (_c *PaymentOrderCreate) SetBonusAmount(v float64) *PaymentOrderCreate { + _c.mutation.SetBonusAmount(v) + return _c +} + +// SetNillableBonusAmount sets the "bonus_amount" field if the given value is not nil. +func (_c *PaymentOrderCreate) SetNillableBonusAmount(v *float64) *PaymentOrderCreate { + if v != nil { + _c.SetBonusAmount(*v) + } + return _c +} + // SetRechargeCode sets the "recharge_code" field. func (_c *PaymentOrderCreate) SetRechargeCode(v string) *PaymentOrderCreate { _c.mutation.SetRechargeCode(v) @@ -517,6 +531,10 @@ func (_c *PaymentOrderCreate) defaults() { v := paymentorder.DefaultFeeRate _c.mutation.SetFeeRate(v) } + if _, ok := _c.mutation.BonusAmount(); !ok { + v := paymentorder.DefaultBonusAmount + _c.mutation.SetBonusAmount(v) + } if _, ok := _c.mutation.OutTradeNo(); !ok { v := paymentorder.DefaultOutTradeNo _c.mutation.SetOutTradeNo(v) @@ -577,6 +595,9 @@ func (_c *PaymentOrderCreate) check() error { if _, ok := _c.mutation.FeeRate(); !ok { return &ValidationError{Name: "fee_rate", err: errors.New(`ent: missing required field "PaymentOrder.fee_rate"`)} } + if _, ok := _c.mutation.BonusAmount(); !ok { + return &ValidationError{Name: "bonus_amount", err: errors.New(`ent: missing required field "PaymentOrder.bonus_amount"`)} + } if _, ok := _c.mutation.RechargeCode(); !ok { return &ValidationError{Name: "recharge_code", err: errors.New(`ent: missing required field "PaymentOrder.recharge_code"`)} } @@ -725,6 +746,10 @@ func (_c *PaymentOrderCreate) createSpec() (*PaymentOrder, *sqlgraph.CreateSpec) _spec.SetField(paymentorder.FieldFeeRate, field.TypeFloat64, value) _node.FeeRate = value } + if value, ok := _c.mutation.BonusAmount(); ok { + _spec.SetField(paymentorder.FieldBonusAmount, field.TypeFloat64, value) + _node.BonusAmount = value + } if value, ok := _c.mutation.RechargeCode(); ok { _spec.SetField(paymentorder.FieldRechargeCode, field.TypeString, value) _node.RechargeCode = value @@ -1030,6 +1055,24 @@ func (u *PaymentOrderUpsert) AddFeeRate(v float64) *PaymentOrderUpsert { return u } +// SetBonusAmount sets the "bonus_amount" field. +func (u *PaymentOrderUpsert) SetBonusAmount(v float64) *PaymentOrderUpsert { + u.Set(paymentorder.FieldBonusAmount, v) + return u +} + +// UpdateBonusAmount sets the "bonus_amount" field to the value that was provided on create. +func (u *PaymentOrderUpsert) UpdateBonusAmount() *PaymentOrderUpsert { + u.SetExcluded(paymentorder.FieldBonusAmount) + return u +} + +// AddBonusAmount adds v to the "bonus_amount" field. +func (u *PaymentOrderUpsert) AddBonusAmount(v float64) *PaymentOrderUpsert { + u.Add(paymentorder.FieldBonusAmount, v) + return u +} + // SetRechargeCode sets the "recharge_code" field. func (u *PaymentOrderUpsert) SetRechargeCode(v string) *PaymentOrderUpsert { u.Set(paymentorder.FieldRechargeCode, v) @@ -1711,6 +1754,27 @@ func (u *PaymentOrderUpsertOne) UpdateFeeRate() *PaymentOrderUpsertOne { }) } +// SetBonusAmount sets the "bonus_amount" field. +func (u *PaymentOrderUpsertOne) SetBonusAmount(v float64) *PaymentOrderUpsertOne { + return u.Update(func(s *PaymentOrderUpsert) { + s.SetBonusAmount(v) + }) +} + +// AddBonusAmount adds v to the "bonus_amount" field. +func (u *PaymentOrderUpsertOne) AddBonusAmount(v float64) *PaymentOrderUpsertOne { + return u.Update(func(s *PaymentOrderUpsert) { + s.AddBonusAmount(v) + }) +} + +// UpdateBonusAmount sets the "bonus_amount" field to the value that was provided on create. +func (u *PaymentOrderUpsertOne) UpdateBonusAmount() *PaymentOrderUpsertOne { + return u.Update(func(s *PaymentOrderUpsert) { + s.UpdateBonusAmount() + }) +} + // SetRechargeCode sets the "recharge_code" field. func (u *PaymentOrderUpsertOne) SetRechargeCode(v string) *PaymentOrderUpsertOne { return u.Update(func(s *PaymentOrderUpsert) { @@ -2643,6 +2707,27 @@ func (u *PaymentOrderUpsertBulk) UpdateFeeRate() *PaymentOrderUpsertBulk { }) } +// SetBonusAmount sets the "bonus_amount" field. +func (u *PaymentOrderUpsertBulk) SetBonusAmount(v float64) *PaymentOrderUpsertBulk { + return u.Update(func(s *PaymentOrderUpsert) { + s.SetBonusAmount(v) + }) +} + +// AddBonusAmount adds v to the "bonus_amount" field. +func (u *PaymentOrderUpsertBulk) AddBonusAmount(v float64) *PaymentOrderUpsertBulk { + return u.Update(func(s *PaymentOrderUpsert) { + s.AddBonusAmount(v) + }) +} + +// UpdateBonusAmount sets the "bonus_amount" field to the value that was provided on create. +func (u *PaymentOrderUpsertBulk) UpdateBonusAmount() *PaymentOrderUpsertBulk { + return u.Update(func(s *PaymentOrderUpsert) { + s.UpdateBonusAmount() + }) +} + // SetRechargeCode sets the "recharge_code" field. func (u *PaymentOrderUpsertBulk) SetRechargeCode(v string) *PaymentOrderUpsertBulk { return u.Update(func(s *PaymentOrderUpsert) { diff --git a/backend/ent/paymentorder_update.go b/backend/ent/paymentorder_update.go index 378e0dad2..3ab893cbe 100644 --- a/backend/ent/paymentorder_update.go +++ b/backend/ent/paymentorder_update.go @@ -154,6 +154,27 @@ func (_u *PaymentOrderUpdate) AddFeeRate(v float64) *PaymentOrderUpdate { return _u } +// SetBonusAmount sets the "bonus_amount" field. +func (_u *PaymentOrderUpdate) SetBonusAmount(v float64) *PaymentOrderUpdate { + _u.mutation.ResetBonusAmount() + _u.mutation.SetBonusAmount(v) + return _u +} + +// SetNillableBonusAmount sets the "bonus_amount" field if the given value is not nil. +func (_u *PaymentOrderUpdate) SetNillableBonusAmount(v *float64) *PaymentOrderUpdate { + if v != nil { + _u.SetBonusAmount(*v) + } + return _u +} + +// AddBonusAmount adds value to the "bonus_amount" field. +func (_u *PaymentOrderUpdate) AddBonusAmount(v float64) *PaymentOrderUpdate { + _u.mutation.AddBonusAmount(v) + return _u +} + // SetRechargeCode sets the "recharge_code" field. func (_u *PaymentOrderUpdate) SetRechargeCode(v string) *PaymentOrderUpdate { _u.mutation.SetRechargeCode(v) @@ -881,6 +902,12 @@ func (_u *PaymentOrderUpdate) sqlSave(ctx context.Context) (_node int, err error if value, ok := _u.mutation.AddedFeeRate(); ok { _spec.AddField(paymentorder.FieldFeeRate, field.TypeFloat64, value) } + if value, ok := _u.mutation.BonusAmount(); ok { + _spec.SetField(paymentorder.FieldBonusAmount, field.TypeFloat64, value) + } + if value, ok := _u.mutation.AddedBonusAmount(); ok { + _spec.AddField(paymentorder.FieldBonusAmount, field.TypeFloat64, value) + } if value, ok := _u.mutation.RechargeCode(); ok { _spec.SetField(paymentorder.FieldRechargeCode, field.TypeString, value) } @@ -1217,6 +1244,27 @@ func (_u *PaymentOrderUpdateOne) AddFeeRate(v float64) *PaymentOrderUpdateOne { return _u } +// SetBonusAmount sets the "bonus_amount" field. +func (_u *PaymentOrderUpdateOne) SetBonusAmount(v float64) *PaymentOrderUpdateOne { + _u.mutation.ResetBonusAmount() + _u.mutation.SetBonusAmount(v) + return _u +} + +// SetNillableBonusAmount sets the "bonus_amount" field if the given value is not nil. +func (_u *PaymentOrderUpdateOne) SetNillableBonusAmount(v *float64) *PaymentOrderUpdateOne { + if v != nil { + _u.SetBonusAmount(*v) + } + return _u +} + +// AddBonusAmount adds value to the "bonus_amount" field. +func (_u *PaymentOrderUpdateOne) AddBonusAmount(v float64) *PaymentOrderUpdateOne { + _u.mutation.AddBonusAmount(v) + return _u +} + // SetRechargeCode sets the "recharge_code" field. func (_u *PaymentOrderUpdateOne) SetRechargeCode(v string) *PaymentOrderUpdateOne { _u.mutation.SetRechargeCode(v) @@ -1974,6 +2022,12 @@ func (_u *PaymentOrderUpdateOne) sqlSave(ctx context.Context) (_node *PaymentOrd if value, ok := _u.mutation.AddedFeeRate(); ok { _spec.AddField(paymentorder.FieldFeeRate, field.TypeFloat64, value) } + if value, ok := _u.mutation.BonusAmount(); ok { + _spec.SetField(paymentorder.FieldBonusAmount, field.TypeFloat64, value) + } + if value, ok := _u.mutation.AddedBonusAmount(); ok { + _spec.AddField(paymentorder.FieldBonusAmount, field.TypeFloat64, value) + } if value, ok := _u.mutation.RechargeCode(); ok { _spec.SetField(paymentorder.FieldRechargeCode, field.TypeString, value) } diff --git a/backend/ent/runtime/runtime.go b/backend/ent/runtime/runtime.go index 4db7e1d69..ffff74abb 100644 --- a/backend/ent/runtime/runtime.go +++ b/backend/ent/runtime/runtime.go @@ -1327,70 +1327,74 @@ func init() { paymentorderDescFeeRate := paymentorderFields[6].Descriptor() // paymentorder.DefaultFeeRate holds the default value on creation for the fee_rate field. paymentorder.DefaultFeeRate = paymentorderDescFeeRate.Default.(float64) + // paymentorderDescBonusAmount is the schema descriptor for bonus_amount field. + paymentorderDescBonusAmount := paymentorderFields[7].Descriptor() + // paymentorder.DefaultBonusAmount holds the default value on creation for the bonus_amount field. + paymentorder.DefaultBonusAmount = paymentorderDescBonusAmount.Default.(float64) // paymentorderDescRechargeCode is the schema descriptor for recharge_code field. - paymentorderDescRechargeCode := paymentorderFields[7].Descriptor() + paymentorderDescRechargeCode := paymentorderFields[8].Descriptor() // paymentorder.RechargeCodeValidator is a validator for the "recharge_code" field. It is called by the builders before save. paymentorder.RechargeCodeValidator = paymentorderDescRechargeCode.Validators[0].(func(string) error) // paymentorderDescOutTradeNo is the schema descriptor for out_trade_no field. - paymentorderDescOutTradeNo := paymentorderFields[8].Descriptor() + paymentorderDescOutTradeNo := paymentorderFields[9].Descriptor() // paymentorder.DefaultOutTradeNo holds the default value on creation for the out_trade_no field. paymentorder.DefaultOutTradeNo = paymentorderDescOutTradeNo.Default.(string) // paymentorder.OutTradeNoValidator is a validator for the "out_trade_no" field. It is called by the builders before save. paymentorder.OutTradeNoValidator = paymentorderDescOutTradeNo.Validators[0].(func(string) error) // paymentorderDescPaymentType is the schema descriptor for payment_type field. - paymentorderDescPaymentType := paymentorderFields[9].Descriptor() + paymentorderDescPaymentType := paymentorderFields[10].Descriptor() // paymentorder.PaymentTypeValidator is a validator for the "payment_type" field. It is called by the builders before save. paymentorder.PaymentTypeValidator = paymentorderDescPaymentType.Validators[0].(func(string) error) // paymentorderDescPaymentTradeNo is the schema descriptor for payment_trade_no field. - paymentorderDescPaymentTradeNo := paymentorderFields[10].Descriptor() + paymentorderDescPaymentTradeNo := paymentorderFields[11].Descriptor() // paymentorder.PaymentTradeNoValidator is a validator for the "payment_trade_no" field. It is called by the builders before save. paymentorder.PaymentTradeNoValidator = paymentorderDescPaymentTradeNo.Validators[0].(func(string) error) // paymentorderDescOrderType is the schema descriptor for order_type field. - paymentorderDescOrderType := paymentorderFields[14].Descriptor() + paymentorderDescOrderType := paymentorderFields[15].Descriptor() // paymentorder.DefaultOrderType holds the default value on creation for the order_type field. paymentorder.DefaultOrderType = paymentorderDescOrderType.Default.(string) // paymentorder.OrderTypeValidator is a validator for the "order_type" field. It is called by the builders before save. paymentorder.OrderTypeValidator = paymentorderDescOrderType.Validators[0].(func(string) error) // paymentorderDescProviderInstanceID is the schema descriptor for provider_instance_id field. - paymentorderDescProviderInstanceID := paymentorderFields[18].Descriptor() + paymentorderDescProviderInstanceID := paymentorderFields[19].Descriptor() // paymentorder.ProviderInstanceIDValidator is a validator for the "provider_instance_id" field. It is called by the builders before save. paymentorder.ProviderInstanceIDValidator = paymentorderDescProviderInstanceID.Validators[0].(func(string) error) // paymentorderDescProviderKey is the schema descriptor for provider_key field. - paymentorderDescProviderKey := paymentorderFields[19].Descriptor() + paymentorderDescProviderKey := paymentorderFields[20].Descriptor() // paymentorder.ProviderKeyValidator is a validator for the "provider_key" field. It is called by the builders before save. paymentorder.ProviderKeyValidator = paymentorderDescProviderKey.Validators[0].(func(string) error) // paymentorderDescStatus is the schema descriptor for status field. - paymentorderDescStatus := paymentorderFields[21].Descriptor() + paymentorderDescStatus := paymentorderFields[22].Descriptor() // paymentorder.DefaultStatus holds the default value on creation for the status field. paymentorder.DefaultStatus = paymentorderDescStatus.Default.(string) // paymentorder.StatusValidator is a validator for the "status" field. It is called by the builders before save. paymentorder.StatusValidator = paymentorderDescStatus.Validators[0].(func(string) error) // paymentorderDescRefundAmount is the schema descriptor for refund_amount field. - paymentorderDescRefundAmount := paymentorderFields[22].Descriptor() + paymentorderDescRefundAmount := paymentorderFields[23].Descriptor() // paymentorder.DefaultRefundAmount holds the default value on creation for the refund_amount field. paymentorder.DefaultRefundAmount = paymentorderDescRefundAmount.Default.(float64) // paymentorderDescForceRefund is the schema descriptor for force_refund field. - paymentorderDescForceRefund := paymentorderFields[25].Descriptor() + paymentorderDescForceRefund := paymentorderFields[26].Descriptor() // paymentorder.DefaultForceRefund holds the default value on creation for the force_refund field. paymentorder.DefaultForceRefund = paymentorderDescForceRefund.Default.(bool) // paymentorderDescRefundRequestedBy is the schema descriptor for refund_requested_by field. - paymentorderDescRefundRequestedBy := paymentorderFields[28].Descriptor() + paymentorderDescRefundRequestedBy := paymentorderFields[29].Descriptor() // paymentorder.RefundRequestedByValidator is a validator for the "refund_requested_by" field. It is called by the builders before save. paymentorder.RefundRequestedByValidator = paymentorderDescRefundRequestedBy.Validators[0].(func(string) error) // paymentorderDescClientIP is the schema descriptor for client_ip field. - paymentorderDescClientIP := paymentorderFields[34].Descriptor() + paymentorderDescClientIP := paymentorderFields[35].Descriptor() // paymentorder.ClientIPValidator is a validator for the "client_ip" field. It is called by the builders before save. paymentorder.ClientIPValidator = paymentorderDescClientIP.Validators[0].(func(string) error) // paymentorderDescSrcHost is the schema descriptor for src_host field. - paymentorderDescSrcHost := paymentorderFields[35].Descriptor() + paymentorderDescSrcHost := paymentorderFields[36].Descriptor() // paymentorder.SrcHostValidator is a validator for the "src_host" field. It is called by the builders before save. paymentorder.SrcHostValidator = paymentorderDescSrcHost.Validators[0].(func(string) error) // paymentorderDescCreatedAt is the schema descriptor for created_at field. - paymentorderDescCreatedAt := paymentorderFields[37].Descriptor() + paymentorderDescCreatedAt := paymentorderFields[38].Descriptor() // paymentorder.DefaultCreatedAt holds the default value on creation for the created_at field. paymentorder.DefaultCreatedAt = paymentorderDescCreatedAt.Default.(func() time.Time) // paymentorderDescUpdatedAt is the schema descriptor for updated_at field. - paymentorderDescUpdatedAt := paymentorderFields[38].Descriptor() + paymentorderDescUpdatedAt := paymentorderFields[39].Descriptor() // paymentorder.DefaultUpdatedAt holds the default value on creation for the updated_at field. paymentorder.DefaultUpdatedAt = paymentorderDescUpdatedAt.Default.(func() time.Time) // paymentorder.UpdateDefaultUpdatedAt holds the default value on update for the updated_at field. diff --git a/backend/ent/schema/payment_order.go b/backend/ent/schema/payment_order.go index d25d1e5e1..c8ce5f1ed 100644 --- a/backend/ent/schema/payment_order.go +++ b/backend/ent/schema/payment_order.go @@ -50,6 +50,10 @@ func (PaymentOrder) Fields() []ent.Field { field.Float("fee_rate"). SchemaType(map[string]string{dialect.Postgres: "decimal(10,4)"}). Default(0), + // 充值赠送额度(USD)。已计入 amount;单独记录用于订单展示与推广返利基数剔除。 + field.Float("bonus_amount"). + SchemaType(map[string]string{dialect.Postgres: "decimal(20,2)"}). + Default(0), field.String("recharge_code"). MaxLen(64), diff --git a/backend/ent/schema/user_platform_quota.go b/backend/ent/schema/user_platform_quota.go index 20c0b3aa1..663be0a6e 100644 --- a/backend/ent/schema/user_platform_quota.go +++ b/backend/ent/schema/user_platform_quota.go @@ -42,7 +42,7 @@ func (UserPlatformQuota) Fields() []ent.Field { // 此处为 ent 构建期约束,需与 service.AllowedQuotaPlatforms 保持同步。 switch s { case "anthropic", "openai", "gemini", "antigravity", "grok", - "kimi", "zhipu", "deepseek", "minimax", "opencode_go": + "kimi", "zhipu", "deepseek", "minimax", "opencode_go", "typesafe": return nil default: return fmt.Errorf("platform %q is not allowed", s) diff --git a/backend/internal/domain/constants.go b/backend/internal/domain/constants.go index f6b74fec3..75dd47496 100644 --- a/backend/internal/domain/constants.go +++ b/backend/internal/domain/constants.go @@ -28,6 +28,7 @@ const ( PlatformZhipu = "zhipu" // 智谱 GLM (bigmodel) PlatformDeepseek = "deepseek" // DeepSeek PlatformMiniMax = "minimax" // MiniMax (M 系列) + PlatformTypeSafe = "typesafe" // TypeSafe AI System One (Jev) // PlatformOpenCodeGo 是 OpenCode 平台(账号类型 Zen 按量 / Go 订阅)。 // 值保持 opencode_go 以兼容已落库的分组、配额与 Composite 路由 CHECK。 PlatformOpenCodeGo = "opencode_go" diff --git a/backend/internal/handler/admin/account_handler.go b/backend/internal/handler/admin/account_handler.go index e43e3bf1f..84ee4371c 100644 --- a/backend/internal/handler/admin/account_handler.go +++ b/backend/internal/handler/admin/account_handler.go @@ -27,6 +27,7 @@ import ( "github.com/Wei-Shaw/sub2api/internal/pkg/openai" "github.com/Wei-Shaw/sub2api/internal/pkg/response" "github.com/Wei-Shaw/sub2api/internal/pkg/timezone" + "github.com/Wei-Shaw/sub2api/internal/pkg/typesafe" "github.com/Wei-Shaw/sub2api/internal/pkg/xai" "github.com/Wei-Shaw/sub2api/internal/service" @@ -2954,6 +2955,12 @@ func (h *AccountHandler) GetAvailableModels(c *gin.Context) { return } + // TypeSafe accounts serve only the native System One model. + if account.IsTypeSafe() { + response.Success(c, []claude.Model{{ID: typesafe.JevLatestModel, Type: "model", DisplayName: typesafe.JevLatestModel}}) + return + } + // Handle Claude/Anthropic accounts // For OAuth and Setup-Token accounts: return default models if account.IsOAuth() { diff --git a/backend/internal/handler/admin/account_handler_available_models_test.go b/backend/internal/handler/admin/account_handler_available_models_test.go index 28dc6d149..3bcfce11d 100644 --- a/backend/internal/handler/admin/account_handler_available_models_test.go +++ b/backend/internal/handler/admin/account_handler_available_models_test.go @@ -526,3 +526,36 @@ func TestAccountHandlerSyncUpstreamModels_MetadataEnrichmentFailureReturnsWarnin require.Len(t, resp.Data.Warnings, 1) require.Equal(t, "upstream_model_metadata_incomplete", resp.Data.Warnings[0].Code) } + +func TestAccountHandlerGetAvailableModels_TypeSafeOnlyReturnsJev(t *testing.T) { + for _, credentials := range []map[string]any{ + {"api_key": "ts-secret"}, + {"api_key": "ts-secret", "model_mapping": map[string]any{"jev-latest": "jev-latest"}}, + } { + svc := &availableModelsAdminService{ + stubAdminService: newStubAdminService(), + account: service.Account{ + ID: 46, + Name: "typesafe", + Platform: service.PlatformTypeSafe, + Type: service.AccountTypeAPIKey, + Status: service.StatusActive, + Credentials: credentials, + }, + } + router := setupAvailableModelsRouter(svc) + + rec := httptest.NewRecorder() + router.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/api/v1/admin/accounts/46/models", nil)) + require.Equal(t, http.StatusOK, rec.Code) + + var resp struct { + Data []struct { + ID string `json:"id"` + } `json:"data"` + } + require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &resp)) + require.Len(t, resp.Data, 1) + require.Equal(t, "jev-latest", resp.Data[0].ID) + } +} diff --git a/backend/internal/handler/admin/channel_handler.go b/backend/internal/handler/admin/channel_handler.go index 376e08f88..065aedb2f 100644 --- a/backend/internal/handler/admin/channel_handler.go +++ b/backend/internal/handler/admin/channel_handler.go @@ -644,6 +644,7 @@ var platformToLiteLLMProvider = map[string]string{ service.PlatformDeepseek: "deepseek", service.PlatformMiniMax: "minimax", service.PlatformOpenCodeGo: "opencode-go", + service.PlatformTypeSafe: "typesafe", } // SyncPricingModels 返回 LiteLLM 定价目录中指定平台的最新模型列表 diff --git a/backend/internal/handler/admin/channel_handler_test.go b/backend/internal/handler/admin/channel_handler_test.go index 5678a9499..0c01631ec 100644 --- a/backend/internal/handler/admin/channel_handler_test.go +++ b/backend/internal/handler/admin/channel_handler_test.go @@ -565,7 +565,7 @@ func TestSyncPricingModels_ValidPlatform_EmptyService(t *testing.T) { svc := service.NewPricingService(nil, nil) router := setupSyncPricingModelsRouter(svc) - for _, platform := range []string{"anthropic", "openai", "gemini", "antigravity", "grok", "kimi", "zhipu", "deepseek", "minimax"} { + for _, platform := range []string{"anthropic", "openai", "gemini", "antigravity", "grok", "kimi", "zhipu", "deepseek", "minimax", "typesafe"} { req := httptest.NewRequest(http.MethodGet, "/channels/pricing/sync-models?platform="+platform, nil) w := httptest.NewRecorder() router.ServeHTTP(w, req) diff --git a/backend/internal/handler/admin/group_handler.go b/backend/internal/handler/admin/group_handler.go index 6148408e6..ccb05eed5 100644 --- a/backend/internal/handler/admin/group_handler.go +++ b/backend/internal/handler/admin/group_handler.go @@ -187,7 +187,7 @@ func sanitizeUpdateGroupRequestForSimpleMode(req *UpdateGroupRequest) { type CreateGroupRequest struct { Name string `json:"name" binding:"required"` Description string `json:"description"` - Platform string `json:"platform" binding:"omitempty,oneof=anthropic openai gemini antigravity grok kimi zhipu deepseek minimax opencode_go composite"` + Platform string `json:"platform" binding:"omitempty,oneof=anthropic openai gemini antigravity grok kimi zhipu deepseek minimax opencode_go typesafe composite"` RateMultiplier float64 `json:"rate_multiplier"` IsExclusive bool `json:"is_exclusive"` SubscriptionType string `json:"subscription_type" binding:"omitempty,oneof=standard subscription"` @@ -262,7 +262,7 @@ type CreateGroupRequest struct { type UpdateGroupRequest struct { Name string `json:"name"` Description *string `json:"description"` - Platform string `json:"platform" binding:"omitempty,oneof=anthropic openai gemini antigravity grok kimi zhipu deepseek minimax opencode_go composite"` + Platform string `json:"platform" binding:"omitempty,oneof=anthropic openai gemini antigravity grok kimi zhipu deepseek minimax opencode_go typesafe composite"` RateMultiplier *float64 `json:"rate_multiplier"` IsExclusive *bool `json:"is_exclusive"` Status string `json:"status" binding:"omitempty,oneof=active inactive"` @@ -368,7 +368,7 @@ func resolveUpdateGroupModelAllowlist( type CompositeRouteRequest struct { PublicModel string `json:"public_model" binding:"required"` MatchType string `json:"match_type" binding:"omitempty,oneof=exact prefix"` - TargetPlatform string `json:"target_platform" binding:"required,oneof=anthropic openai gemini antigravity grok kimi zhipu deepseek minimax opencode_go"` + TargetPlatform string `json:"target_platform" binding:"required,oneof=anthropic openai gemini antigravity grok kimi zhipu deepseek minimax opencode_go typesafe"` UpstreamModel string `json:"upstream_model"` Endpoint string `json:"endpoint" binding:"omitempty,oneof=any messages count_tokens responses chat_completions embeddings images gemini"` Priority int `json:"priority"` diff --git a/backend/internal/handler/admin/group_handler_platform_test.go b/backend/internal/handler/admin/group_handler_platform_test.go index 0c820868b..d5d48845e 100644 --- a/backend/internal/handler/admin/group_handler_platform_test.go +++ b/backend/internal/handler/admin/group_handler_platform_test.go @@ -28,6 +28,7 @@ func TestGroupPlatformBinding_AllowedPlatforms(t *testing.T) { allowed := []string{ "anthropic", "openai", "gemini", "antigravity", "grok", "kimi", "zhipu", "deepseek", "minimax", "opencode_go", "composite", + "typesafe", } for _, platform := range allowed { t.Run("create_"+platform, func(t *testing.T) { @@ -71,8 +72,8 @@ func TestGroupPlatformBinding_RejectsInvalidPlatforms(t *testing.T) { } } -func TestCompositeRouteTargetPlatform_AllowsCNProviders(t *testing.T) { - for _, platform := range []string{"kimi", "zhipu", "deepseek", "minimax", "opencode_go"} { +func TestCompositeRouteTargetPlatform_AllowsConcreteProviders(t *testing.T) { + for _, platform := range []string{"kimi", "zhipu", "deepseek", "minimax", "opencode_go", "typesafe"} { var req CompositeRouteRequest body := fmt.Sprintf(`{"public_model":"m","target_platform":%q}`, platform) require.NoError(t, bindGroupPlatformJSON(t, &req, body)) diff --git a/backend/internal/handler/admin/payment_handler.go b/backend/internal/handler/admin/payment_handler.go index 1749d1967..3598635cb 100644 --- a/backend/internal/handler/admin/payment_handler.go +++ b/backend/internal/handler/admin/payment_handler.go @@ -125,6 +125,7 @@ type AdminPaymentOrderResult struct { Amount float64 `json:"amount"` PayAmount float64 `json:"pay_amount"` FeeRate float64 `json:"fee_rate"` + BonusAmount float64 `json:"bonus_amount"` Currency string `json:"currency"` RechargeCode string `json:"recharge_code,omitempty"` OutTradeNo string `json:"out_trade_no"` @@ -182,6 +183,7 @@ func sanitizeAdminPaymentOrderForResponse(order *dbent.PaymentOrder) *AdminPayme Amount: order.Amount, PayAmount: order.PayAmount, FeeRate: order.FeeRate, + BonusAmount: order.BonusAmount, Currency: service.PaymentOrderCurrency(order), RechargeCode: order.RechargeCode, OutTradeNo: order.OutTradeNo, diff --git a/backend/internal/handler/admin/setting_handler.go b/backend/internal/handler/admin/setting_handler.go index 6971c2620..ab8bbe023 100644 --- a/backend/internal/handler/admin/setting_handler.go +++ b/backend/internal/handler/admin/setting_handler.go @@ -400,6 +400,9 @@ func (h *SettingHandler) GetSettings(c *gin.Context) { PaymentBalanceRechargeMultiplier: paymentCfg.BalanceRechargeMultiplier, PaymentSubscriptionUSDToCNYRate: paymentCfg.SubscriptionUSDToCNYRate, PaymentRechargeFeeRate: paymentCfg.RechargeFeeRate, + PaymentRechargeBonusTiers: rechargeBonusTiersToDTO(paymentCfg.RechargeBonusTiers), + PaymentRechargeBonusMode: rechargeBonusModeToDTO(paymentCfg.RechargeBonusMode), + PaymentRechargeBonusNotice: paymentCfg.RechargeBonusNotice, PaymentLoadBalanceStrat: paymentCfg.LoadBalanceStrategy, PaymentProductNamePrefix: paymentCfg.ProductNamePrefix, PaymentProductNameSuffix: paymentCfg.ProductNameSuffix, diff --git a/backend/internal/handler/admin/setting_handler_recharge_bonus.go b/backend/internal/handler/admin/setting_handler_recharge_bonus.go new file mode 100644 index 000000000..334aedcef --- /dev/null +++ b/backend/internal/handler/admin/setting_handler_recharge_bonus.go @@ -0,0 +1,33 @@ +package admin + +import ( + "github.com/Wei-Shaw/sub2api/internal/handler/dto" + "github.com/Wei-Shaw/sub2api/internal/service" +) + +// rechargeBonusTiersFromDTO 请求 nil 表示未携带该字段(保持现值);空数组表示清空阶梯。 +func rechargeBonusTiersFromDTO(items *[]dto.RechargeBonusTier) *[]service.RechargeBonusTier { + if items == nil { + return nil + } + out := make([]service.RechargeBonusTier, 0, len(*items)) + for _, item := range *items { + out = append(out, service.RechargeBonusTier{MinAmount: item.MinAmount, BonusPercent: item.BonusPercent}) + } + return &out +} + +// rechargeBonusModeToDTO 输出已归一化的模式(空/非法按 bonus)。 +func rechargeBonusModeToDTO(mode string) string { + normalized, _ := service.NormalizeRechargeBonusMode(mode) + return normalized +} + +// rechargeBonusTiersToDTO 始终返回非 nil 切片,空配置输出 []。 +func rechargeBonusTiersToDTO(items []service.RechargeBonusTier) []dto.RechargeBonusTier { + out := make([]dto.RechargeBonusTier, 0, len(items)) + for _, item := range items { + out = append(out, dto.RechargeBonusTier{MinAmount: item.MinAmount, BonusPercent: item.BonusPercent}) + } + return out +} diff --git a/backend/internal/handler/admin/setting_handler_update.go b/backend/internal/handler/admin/setting_handler_update.go index accf16b38..f717e962c 100644 --- a/backend/internal/handler/admin/setting_handler_update.go +++ b/backend/internal/handler/admin/setting_handler_update.go @@ -326,11 +326,15 @@ type UpdateSettingsRequest struct { PaymentBalanceRechargeMultiplier *float64 `json:"payment_balance_recharge_multiplier"` PaymentSubscriptionUSDToCNYRate *float64 `json:"payment_subscription_usd_to_cny_rate"` PaymentRechargeFeeRate *float64 `json:"payment_recharge_fee_rate"` - PaymentLoadBalanceStrat *string `json:"payment_load_balance_strategy"` - PaymentProductNamePrefix *string `json:"payment_product_name_prefix"` - PaymentProductNameSuffix *string `json:"payment_product_name_suffix"` - PaymentHelpImageURL *string `json:"payment_help_image_url"` - PaymentHelpText *string `json:"payment_help_text"` + // nil 表示不更新;空数组表示清空阶梯 + PaymentRechargeBonusTiers *[]dto.RechargeBonusTier `json:"payment_recharge_bonus_tiers"` + PaymentRechargeBonusMode *string `json:"payment_recharge_bonus_mode"` + PaymentRechargeBonusNotice *string `json:"payment_recharge_bonus_notice"` + PaymentLoadBalanceStrat *string `json:"payment_load_balance_strategy"` + PaymentProductNamePrefix *string `json:"payment_product_name_prefix"` + PaymentProductNameSuffix *string `json:"payment_product_name_suffix"` + PaymentHelpImageURL *string `json:"payment_help_image_url"` + PaymentHelpText *string `json:"payment_help_text"` // Cancel rate limit PaymentCancelRateLimitEnabled *bool `json:"payment_cancel_rate_limit_enabled"` @@ -2252,6 +2256,9 @@ func (h *SettingHandler) UpdateSettings(c *gin.Context) { BalanceRechargeMultiplier: req.PaymentBalanceRechargeMultiplier, SubscriptionUSDToCNYRate: req.PaymentSubscriptionUSDToCNYRate, RechargeFeeRate: req.PaymentRechargeFeeRate, + RechargeBonusTiers: rechargeBonusTiersFromDTO(req.PaymentRechargeBonusTiers), + RechargeBonusMode: req.PaymentRechargeBonusMode, + RechargeBonusNotice: req.PaymentRechargeBonusNotice, LoadBalanceStrategy: req.PaymentLoadBalanceStrat, ProductNamePrefix: req.PaymentProductNamePrefix, ProductNameSuffix: req.PaymentProductNameSuffix, @@ -2555,6 +2562,9 @@ func (h *SettingHandler) UpdateSettings(c *gin.Context) { PaymentBalanceRechargeMultiplier: updatedPaymentCfg.BalanceRechargeMultiplier, PaymentSubscriptionUSDToCNYRate: updatedPaymentCfg.SubscriptionUSDToCNYRate, PaymentRechargeFeeRate: updatedPaymentCfg.RechargeFeeRate, + PaymentRechargeBonusTiers: rechargeBonusTiersToDTO(updatedPaymentCfg.RechargeBonusTiers), + PaymentRechargeBonusMode: rechargeBonusModeToDTO(updatedPaymentCfg.RechargeBonusMode), + PaymentRechargeBonusNotice: updatedPaymentCfg.RechargeBonusNotice, PaymentLoadBalanceStrat: updatedPaymentCfg.LoadBalanceStrategy, PaymentProductNamePrefix: updatedPaymentCfg.ProductNamePrefix, PaymentProductNameSuffix: updatedPaymentCfg.ProductNameSuffix, @@ -2628,6 +2638,7 @@ func hasPaymentFields(req UpdateSettingsRequest) bool { req.PaymentEnabledTypes != nil || req.PaymentBalanceDisabled != nil || req.PaymentBalanceRechargeMultiplier != nil || req.PaymentSubscriptionUSDToCNYRate != nil || req.PaymentRechargeFeeRate != nil || + req.PaymentRechargeBonusTiers != nil || req.PaymentRechargeBonusMode != nil || req.PaymentRechargeBonusNotice != nil || req.PaymentLoadBalanceStrat != nil || req.PaymentProductNamePrefix != nil || req.PaymentProductNameSuffix != nil || req.PaymentHelpImageURL != nil || req.PaymentHelpText != nil || req.PaymentCancelRateLimitEnabled != nil || diff --git a/backend/internal/handler/auth_oauth_pending_flow_test.go b/backend/internal/handler/auth_oauth_pending_flow_test.go index 6d12a9ad1..24be95a43 100644 --- a/backend/internal/handler/auth_oauth_pending_flow_test.go +++ b/backend/internal/handler/auth_oauth_pending_flow_test.go @@ -3606,6 +3606,19 @@ func (oauthPendingFlowTotpEncryptorStub) Decrypt(ciphertext string) (string, err return ciphertext, nil } +func (s *oauthPendingFlowEmailCacheStub) IncrVerificationCodeAttempts(_ context.Context, email string) (int, error) { + data := s.verificationCodes[email] + if data == nil { + return 0, errors.New("verification code not found") + } + data.Attempts++ + return data.Attempts, nil +} + +func (s *oauthPendingFlowEmailCacheStub) IncrNotifyVerifyCodeAttempts(context.Context, string) (int, error) { + return 0, errors.New("notify verification code not found") +} + func (s *oauthPendingFlowEmailCacheStub) ConsumePasswordResetToken(context.Context, string, string) (bool, error) { return false, nil } diff --git a/backend/internal/handler/dto/recharge_bonus_tiers.go b/backend/internal/handler/dto/recharge_bonus_tiers.go new file mode 100644 index 000000000..21731566a --- /dev/null +++ b/backend/internal/handler/dto/recharge_bonus_tiers.go @@ -0,0 +1,7 @@ +package dto + +// RechargeBonusTier 充值赠送档位:余额充值支付金额 ≥ MinAmount 时,在到账基数上赠送 BonusPercent%。 +type RechargeBonusTier struct { + MinAmount float64 `json:"min_amount"` + BonusPercent float64 `json:"bonus_percent"` +} diff --git a/backend/internal/handler/dto/settings.go b/backend/internal/handler/dto/settings.go index 5a4a82d0a..a4657332e 100644 --- a/backend/internal/handler/dto/settings.go +++ b/backend/internal/handler/dto/settings.go @@ -299,11 +299,15 @@ type SystemSettings struct { PaymentBalanceRechargeMultiplier float64 `json:"payment_balance_recharge_multiplier"` PaymentSubscriptionUSDToCNYRate float64 `json:"payment_subscription_usd_to_cny_rate"` PaymentRechargeFeeRate float64 `json:"payment_recharge_fee_rate"` - PaymentLoadBalanceStrat string `json:"payment_load_balance_strategy"` - PaymentProductNamePrefix string `json:"payment_product_name_prefix"` - PaymentProductNameSuffix string `json:"payment_product_name_suffix"` - PaymentHelpImageURL string `json:"payment_help_image_url"` - PaymentHelpText string `json:"payment_help_text"` + // 充值赠送阶梯与活动文案 + PaymentRechargeBonusTiers []RechargeBonusTier `json:"payment_recharge_bonus_tiers"` + PaymentRechargeBonusMode string `json:"payment_recharge_bonus_mode"` + PaymentRechargeBonusNotice string `json:"payment_recharge_bonus_notice"` + PaymentLoadBalanceStrat string `json:"payment_load_balance_strategy"` + PaymentProductNamePrefix string `json:"payment_product_name_prefix"` + PaymentProductNameSuffix string `json:"payment_product_name_suffix"` + PaymentHelpImageURL string `json:"payment_help_image_url"` + PaymentHelpText string `json:"payment_help_text"` // Cancel rate limit PaymentCancelRateLimitEnabled bool `json:"payment_cancel_rate_limit_enabled"` diff --git a/backend/internal/handler/endpoint.go b/backend/internal/handler/endpoint.go index 2115c2d7c..82d8279e9 100644 --- a/backend/internal/handler/endpoint.go +++ b/backend/internal/handler/endpoint.go @@ -16,6 +16,7 @@ import ( const ( EndpointMessages = "/v1/messages" + EndpointSystemOne = "/v1/systemone" EndpointChatCompletions = "/v1/chat/completions" EndpointEmbeddings = "/v1/embeddings" EndpointAlphaSearch = "/v1/alpha/search" @@ -94,6 +95,8 @@ func NormalizeInboundEndpoint(path string) string { return EndpointChatCompletions case strings.Contains(path, EndpointMessages): return EndpointMessages + case strings.Contains(path, EndpointSystemOne): + return EndpointSystemOne case strings.Contains(path, EndpointImagesGenerations) || strings.Contains(path, "/images/generations"): return EndpointImagesGenerations case strings.Contains(path, EndpointImagesEdits) || strings.Contains(path, "/images/edits"): @@ -223,6 +226,9 @@ func DeriveUpstreamEndpoint(inbound, rawRequestPath, platform string) string { case service.PlatformAnthropic: return EndpointMessages + case service.PlatformTypeSafe: + return EndpointSystemOne + case service.PlatformGemini: return EndpointGeminiModels diff --git a/backend/internal/handler/endpoint_test.go b/backend/internal/handler/endpoint_test.go index aaa7e9f8f..6941e5637 100644 --- a/backend/internal/handler/endpoint_test.go +++ b/backend/internal/handler/endpoint_test.go @@ -23,6 +23,7 @@ func TestNormalizeInboundEndpoint(t *testing.T) { }{ // Direct canonical paths. {"/v1/messages", EndpointMessages}, + {"/v1/systemone", EndpointSystemOne}, {"/v1/chat/completions", EndpointChatCompletions}, {"/v1/embeddings", EndpointEmbeddings}, {"/v1/alpha/search", EndpointAlphaSearch}, @@ -97,6 +98,7 @@ func TestDeriveUpstreamEndpoint(t *testing.T) { }{ // Anthropic. {"anthropic messages", EndpointMessages, "/v1/messages", service.PlatformAnthropic, EndpointMessages}, + {"typesafe system one", EndpointSystemOne, "/v1/systemone", service.PlatformTypeSafe, EndpointSystemOne}, // Gemini. {"gemini models", EndpointGeminiModels, "/v1beta/models/gemini:gen", service.PlatformGemini, EndpointGeminiModels}, diff --git a/backend/internal/handler/gateway_handler.go b/backend/internal/handler/gateway_handler.go index 3a5eb5f51..219d80e4d 100644 --- a/backend/internal/handler/gateway_handler.go +++ b/backend/internal/handler/gateway_handler.go @@ -24,6 +24,7 @@ import ( "github.com/Wei-Shaw/sub2api/internal/pkg/logger" "github.com/Wei-Shaw/sub2api/internal/pkg/openai" "github.com/Wei-Shaw/sub2api/internal/pkg/timezone" + "github.com/Wei-Shaw/sub2api/internal/pkg/typesafe" "github.com/Wei-Shaw/sub2api/internal/pkg/xai" "github.com/Wei-Shaw/sub2api/internal/securityaudit" middleware2 "github.com/Wei-Shaw/sub2api/internal/server/middleware" @@ -222,6 +223,9 @@ func (h *GatewayHandler) Messages(c *gin.Context) { h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "Model is not supported by composite groups") return } + if rejectSystemOneOnlyPlatform(c, apiKey, h.errorResponse) { + return + } if decision := h.checkSecurityAudit(c, reqLog, apiKey, subject, service.ContentModerationProtocolAnthropicMessages, reqModel, body); decision != nil && !decision.AllowNextStage { h.anthropicSecurityAuditError(c, decision) @@ -1145,7 +1149,7 @@ func (h *GatewayHandler) Models(c *gin.Context) { } if platform == service.PlatformComposite { - availableModels := h.compositeAvailableModels(c.Request.Context(), groupID) + availableModels := h.compositeAvailableModels(c.Request.Context(), groupID, true) if apiKey != nil && apiKey.Group != nil && apiKey.Group.ModelAllowlistEnabled() { source := availableModels if len(source) == 0 { @@ -1189,6 +1193,10 @@ func (h *GatewayHandler) Models(c *gin.Context) { writeGrokModelsList(c, xai.DefaultModelIDs()) return } + if platform == service.PlatformTypeSafe { + writeModelsList(c, platform, []string{typesafe.JevLatestModel}) + return + } writeModelsListResponse(c, claude.DefaultModels) } @@ -1240,7 +1248,7 @@ func (h *GatewayHandler) codexModelIDsForGroup(ctx context.Context, group *servi platform = group.Platform } if platform == service.PlatformComposite { - availableModels := h.compositeAvailableModels(ctx, groupID) + availableModels := h.compositeAvailableModels(ctx, groupID, false) fallbackModels := defaultCodexModelIDsForPlatform(service.PlatformComposite) if group.ModelAllowlistEnabled() { source := availableModels @@ -1266,14 +1274,20 @@ func (h *GatewayHandler) codexModelIDsForGroup(ctx context.Context, group *servi return fallbackModels } -func (h *GatewayHandler) compositeAvailableModels(ctx context.Context, groupID *int64) []string { +// compositeAvailableModels lists the models the composite group can serve. +// includeSystemOne adds TypeSafe models, which only work through /v1/systemone; +// LLM client catalogs (Codex) must exclude them. +func (h *GatewayHandler) compositeAvailableModels(ctx context.Context, groupID *int64, includeSystemOne bool) []string { if h == nil || h.gatewayService == nil { return nil } seen := make(map[string]struct{}) models := make([]string, 0) schedulablePlatforms := h.gatewayService.GetSchedulablePlatforms(ctx, groupID) - for _, platform := range []string{service.PlatformAnthropic, service.PlatformGemini, service.PlatformOpenAI, service.PlatformAntigravity, service.PlatformGrok, service.PlatformKimi, service.PlatformZhipu, service.PlatformDeepseek, service.PlatformMiniMax, service.PlatformOpenCodeGo} { + for _, platform := range []string{service.PlatformAnthropic, service.PlatformGemini, service.PlatformOpenAI, service.PlatformAntigravity, service.PlatformGrok, service.PlatformKimi, service.PlatformZhipu, service.PlatformDeepseek, service.PlatformMiniMax, service.PlatformOpenCodeGo, service.PlatformTypeSafe} { + if platform == service.PlatformTypeSafe && !includeSystemOne { + continue + } platformModels := h.gatewayService.GetAvailableModels(ctx, groupID, platform) if len(platformModels) == 0 { // CN 供应商没有静态默认模型列表(defaultModelIDsForPlatform 的 @@ -1457,9 +1471,14 @@ func defaultModelIDsForPlatform(platform string) []string { return xai.DefaultModelIDs() case service.PlatformOpenCodeGo: return service.DefaultOpenCodeGoModelIDs() + case service.PlatformTypeSafe: + return []string{"jev-latest"} case service.PlatformComposite: ids := make([]string, 0) seen := make(map[string]struct{}) + // TypeSafe is deliberately absent: jev-latest only works through + // /v1/systemone, so the static fallback never advertises it to LLM + // clients. compositeAvailableModels lists it when the group can serve it. for _, concretePlatform := range []string{service.PlatformAnthropic, service.PlatformGemini, service.PlatformOpenAI, service.PlatformAntigravity, service.PlatformGrok, service.PlatformKimi, service.PlatformZhipu, service.PlatformDeepseek, service.PlatformMiniMax, service.PlatformOpenCodeGo} { for _, id := range defaultModelIDsForPlatform(concretePlatform) { if _, ok := seen[id]; ok { @@ -2145,6 +2164,9 @@ func (h *GatewayHandler) CountTokens(c *gin.Context) { h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "Model is not supported by composite groups") return } + if rejectSystemOneOnlyPlatform(c, apiKey, h.errorResponse) { + return + } setOpsRequestContext(c, parsedReq.Model, parsedReq.Stream) setOpsEndpointContext(c, "", int16(service.RequestTypeFromLegacy(parsedReq.Stream, false))) diff --git a/backend/internal/handler/gateway_handler_chat_completions.go b/backend/internal/handler/gateway_handler_chat_completions.go index dc037da56..017316c20 100644 --- a/backend/internal/handler/gateway_handler_chat_completions.go +++ b/backend/internal/handler/gateway_handler_chat_completions.go @@ -82,6 +82,9 @@ func (h *GatewayHandler) ChatCompletions(c *gin.Context) { h.chatCompletionsErrorResponse(c, http.StatusBadRequest, "invalid_request_error", "Model is not supported by composite groups") return } + if rejectSystemOneOnlyPlatform(c, apiKey, h.chatCompletionsErrorResponse) { + return + } reqStream, ok := parseOpenAICompatibleStream(body) if !ok { h.chatCompletionsErrorResponse(c, http.StatusBadRequest, "invalid_request_error", invalidStreamFieldTypeMessage) diff --git a/backend/internal/handler/gateway_handler_responses.go b/backend/internal/handler/gateway_handler_responses.go index c7e97fe21..91986468e 100644 --- a/backend/internal/handler/gateway_handler_responses.go +++ b/backend/internal/handler/gateway_handler_responses.go @@ -82,6 +82,9 @@ func (h *GatewayHandler) Responses(c *gin.Context) { h.responsesErrorResponse(c, http.StatusBadRequest, "invalid_request_error", "Model is not supported by composite groups") return } + if rejectSystemOneOnlyPlatform(c, apiKey, h.responsesErrorResponse) { + return + } reqStream, ok := parseOpenAICompatibleStream(body) if !ok { h.responsesErrorResponse(c, http.StatusBadRequest, "invalid_request_error", invalidStreamFieldTypeMessage) diff --git a/backend/internal/handler/gateway_models_retrieve_test.go b/backend/internal/handler/gateway_models_retrieve_test.go index 26a5b6107..fc590a35d 100644 --- a/backend/internal/handler/gateway_models_retrieve_test.go +++ b/backend/internal/handler/gateway_models_retrieve_test.go @@ -27,7 +27,7 @@ func requestModelForTest(h *GatewayHandler, group *service.Group, modelID, etag func TestRetrieveModelMatchesVisibleCatalogue(t *testing.T) { gin.SetMode(gin.TestMode) - for _, platform := range []string{service.PlatformOpenAI, service.PlatformAnthropic, service.PlatformGemini, service.PlatformGrok, service.PlatformComposite} { + for _, platform := range []string{service.PlatformOpenAI, service.PlatformAnthropic, service.PlatformGemini, service.PlatformGrok, service.PlatformTypeSafe, service.PlatformComposite} { for _, mapped := range []bool{false, true} { name := platform + "/fallback" if mapped { diff --git a/backend/internal/handler/gateway_models_test.go b/backend/internal/handler/gateway_models_test.go index 0b49ca5f3..d2fdb5cfa 100644 --- a/backend/internal/handler/gateway_models_test.go +++ b/backend/internal/handler/gateway_models_test.go @@ -932,6 +932,10 @@ func TestDefaultCodexModelIDsForPlatform_DeepSeekUsesDeepSeekModels(t *testing.T require.Equal(t, defaultModelIDsForPlatform(service.PlatformAnthropic), defaultCodexModelIDsForPlatform(service.PlatformAnthropic)) } +func TestDefaultModelIDsForPlatform_TypeSafeUsesJev(t *testing.T) { + require.Equal(t, []string{"jev-latest"}, defaultModelIDsForPlatform(service.PlatformTypeSafe)) +} + func TestGatewayCodexModels_DeepSeekWithoutMappingUsesDeepSeekDefaults(t *testing.T) { gin.SetMode(gin.TestMode) const groupID int64 = 130 @@ -1522,3 +1526,47 @@ func TestGatewayModels_GPT6SolLunaDiscoveryRespectsGroupAndAccountRestrictions(t }) } } + +// Scenario: jev-latest only works through /v1/systemone, so Composite groups list +// it in /v1/models only when they can serve it, and never in the Codex manifest. +func TestGatewayModels_CompositeTypeSafeListingScope(t *testing.T) { + gin.SetMode(gin.TestMode) + + groupID := int64(66) + h := newGatewayModelsHandlerForTest(&gatewayModelsAccountRepoStub{ + byGroup: map[int64][]service.Account{ + groupID: {{ID: 1, Platform: service.PlatformAnthropic}, {ID: 2, Platform: service.PlatformTypeSafe, Type: service.AccountTypeAPIKey}}, + }, + }) + newContext := func(path string) (*gin.Context, *httptest.ResponseRecorder) { + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodGet, path, nil) + c.Set(string(middleware2.ContextKeyAPIKey), &service.APIKey{ + Group: &service.Group{ID: groupID, Platform: service.PlatformComposite}, + }) + return c, rec + } + + c, rec := newContext("/v1/models") + h.Models(c) + require.Equal(t, http.StatusOK, rec.Code) + var models gatewayModelsResponseForTest + require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &models)) + require.Contains(t, modelIDsForTest(models.Data), "jev-latest") + require.Contains(t, modelIDsForTest(models.Data), "claude-opus-4-6") + + c, rec = newContext("/models?client_version=0.147.0") + h.CodexModels(c) + require.Equal(t, http.StatusOK, rec.Code) + var manifest codexModelsResponseForTest + require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &manifest)) + slugs := codexModelSlugsForTest(manifest.Models) + require.Contains(t, slugs, "claude-opus-4-6") + require.NotContains(t, slugs, "jev-latest") +} + +func TestDefaultModelIDsForPlatform_CompositeFallbackExcludesTypeSafe(t *testing.T) { + require.NotContains(t, defaultModelIDsForPlatform(service.PlatformComposite), "jev-latest") + require.NotContains(t, defaultCodexModelIDsForPlatform(service.PlatformComposite), "jev-latest") +} diff --git a/backend/internal/handler/gateway_systemone.go b/backend/internal/handler/gateway_systemone.go new file mode 100644 index 000000000..f9cc753c7 --- /dev/null +++ b/backend/internal/handler/gateway_systemone.go @@ -0,0 +1,287 @@ +package handler + +import ( + "context" + "errors" + "net/http" + "strconv" + "time" + + "github.com/Wei-Shaw/sub2api/internal/pkg/ip" + "github.com/Wei-Shaw/sub2api/internal/pkg/logger" + "github.com/Wei-Shaw/sub2api/internal/pkg/typesafe" + middleware2 "github.com/Wei-Shaw/sub2api/internal/server/middleware" + "github.com/Wei-Shaw/sub2api/internal/service" + "github.com/gin-gonic/gin" + "go.uber.org/zap" +) + +const systemOneOnlyPlatformMessage = "TypeSafe models are only available through POST /v1/systemone" + +// rejectSystemOneOnlyPlatform stops TypeSafe traffic from entering a non System +// One protocol chain. TypeSafe accounts only speak the native System One +// protocol; letting them reach the Anthropic/OpenAI converters would send +// foreign payloads (and the account key) to the wrong upstream path and feed +// the resulting auth failures back into account state. +func rejectSystemOneOnlyPlatform(c *gin.Context, apiKey *service.APIKey, writeError func(*gin.Context, int, string, string)) bool { + if _, forced := middleware2.GetForcePlatformFromContext(c); forced { + return false + } + if effectiveAPIKeyPlatform(c, apiKey) != service.PlatformTypeSafe { + return false + } + service.MarkOpsClientBusinessLimited(c, service.OpsClientBusinessLimitedReasonLocalFeatureGate) + writeError(c, http.StatusNotFound, "not_found_error", systemOneOnlyPlatformMessage) + return true +} + +// SystemOne proxies TypeSafe's native, non-streaming System One protocol. +func (h *GatewayHandler) SystemOne(c *gin.Context) { + requestStart := time.Now() + apiKey, ok := middleware2.GetAPIKeyFromContext(c) + if !ok || apiKey == nil || apiKey.Group == nil { + h.errorResponse(c, http.StatusUnauthorized, "authentication_error", "Invalid API key") + return + } + subject, ok := middleware2.GetAuthSubjectFromContext(c) + if !ok { + h.errorResponse(c, http.StatusInternalServerError, "api_error", "User context not found") + return + } + reqLog := requestLogger(c, "handler.gateway.systemone", + zap.Int64("user_id", subject.UserID), + zap.Int64("api_key_id", apiKey.ID), + zap.Any("group_id", apiKey.GroupID), + ) + + body, err := readLenientJSONRequestBodyWithPrealloc(c.Request, h.cfg) + if err != nil { + if maxErr, ok := extractMaxBytesError(err); ok { + h.errorResponse(c, http.StatusRequestEntityTooLarge, "invalid_request_error", buildBodyTooLargeMessage(maxErr.Limit)) + return + } + h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "Failed to read request body") + return + } + model, err := typesafe.ValidateSystemOneRequest(body) + if err != nil { + h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", err.Error()) + return + } + ensureCompositeTargetPlatform(c, apiKey, model) + if apiKey.Group.Platform != service.PlatformTypeSafe && + (apiKey.Group.Platform != service.PlatformComposite || !compositeTargetPlatformAllowed(c, apiKey, model, service.PlatformTypeSafe)) { + h.errorResponse(c, http.StatusNotFound, "not_found_error", "System One is only available for TypeSafe and compatible Composite groups") + return + } + if decision := h.checkSecurityAudit(c, reqLog, apiKey, subject, service.ContentModerationProtocolTypeSafeSystemOne, model, body); decision != nil && !decision.AllowNextStage { + h.anthropicSecurityAuditError(c, decision) + return + } + + setOpsRequestContext(c, model, false) + setOpsEndpointContext(c, "", int16(service.RequestTypeSync)) + service.SetOpsLatencyMs(c, service.OpsAuthLatencyMsKey, time.Since(requestStart).Milliseconds()) + pricingCtx, pricingAt := service.WithGatewayTokenRequestPricing(c.Request.Context()) + c.Request = c.Request.WithContext(pricingCtx) + channelMapping, _ := h.gatewayService.ResolveChannelMappingAndRestrict(c.Request.Context(), apiKey.GroupID, model) + subscription, _ := middleware2.GetSubscriptionFromContext(c) + + streamStarted := false + userRelease, err := h.concurrencyHelper.AcquireUserSlotWithWait(c, subject.UserID, subject.Concurrency, false, &streamStarted) + if err != nil { + reqLog.Warn("systemone.user_slot_acquire_failed", zap.Error(err)) + h.handleConcurrencyError(c, err, "user", false) + return + } + // 在请求结束或 Context 取消时确保释放槽位,避免客户端断开造成泄漏。 + userRelease = wrapReleaseOnDone(c.Request.Context(), userRelease) + if userRelease != nil { + defer userRelease() + } + if err := h.billingCacheService.CheckBillingEligibility(c.Request.Context(), apiKey.User, apiKey, apiKey.Group, subscription, service.QuotaPlatform(c.Request.Context(), apiKey)); err != nil { + reqLog.Info("systemone.billing_eligibility_check_failed", zap.Error(err)) + status, code, message, retryAfter := billingErrorDetails(err) + if retryAfter > 0 { + c.Header("Retry-After", strconv.Itoa(retryAfter)) + } + h.errorResponse(c, status, code, message) + return + } + // 余额模式在途预留:防止并发请求在预检时看到同一份余额而集体透支。 + inflightRelease, err := reserveInflightBalance(c, h.billingCacheService, h.gatewayService, apiKey, subscription, tokenInflightEstimate(model, body)) + if err != nil { + reqLog.Info("systemone.inflight_reservation_rejected", zap.Error(err)) + status, code, message, retryAfter := billingErrorDetails(err) + if retryAfter > 0 { + c.Header("Retry-After", strconv.Itoa(retryAfter)) + } + h.errorResponse(c, status, code, message) + return + } + defer inflightRelease() + + fs := NewFailoverState(h.maxAccountSwitches, false) + for { + if failoverClientGone(c) { + return + } + selection, err := h.gatewayService.SelectAccountWithLoadAwareness(c.Request.Context(), apiKey.GroupID, "", model, fs.FailedAccountIDs, "", subject.UserID) + if err == nil && (selection == nil || selection.Account == nil) { + err = service.ErrNoAvailableAccounts + } + if err != nil { + if failoverClientGone(c) { + reqLog.Info("systemone.account_select_aborted_client_disconnected", zap.Error(err)) + return + } + if len(fs.FailedAccountIDs) == 0 { + cls := classifyNoAccountErrorFromGin(c, h.gatewayService, apiKey, model, model, service.PlatformTypeSafe) + cls = classifySelectionFailureError(err, cls) + if !cls.ModelNotFound { + markOpsRoutingCapacityLimitedIfNoAvailable(c, err) + } + reqLog.Warn("systemone.select_account_no_available", zap.Bool("model_not_found", cls.ModelNotFound), zap.Error(err)) + h.errorResponse(c, cls.Status, cls.ErrType, cls.Message) + return + } + switch fs.HandleSelectionExhausted(c.Request.Context()) { + case FailoverContinue: + continue + case FailoverCanceled: + failoverClientGone(c) + return + default: + if fs.LastFailoverErr != nil { + h.handleFailoverExhausted(c, fs.LastFailoverErr, service.PlatformTypeSafe, false) + } else { + h.handleFailoverExhaustedSimple(c, http.StatusBadGateway, false) + } + return + } + } + account := selection.Account + + accountRelease := selection.ReleaseFunc + if !selection.Acquired { + if selection.WaitPlan == nil { + markOpsRoutingCapacityLimited(c) + h.errorResponse(c, http.StatusServiceUnavailable, "api_error", "No available accounts") + return + } + accountRelease, err = h.concurrencyHelper.AcquireAccountSlotWithWaitTimeout(c, account.ID, selection.WaitPlan.MaxConcurrency, selection.WaitPlan.Timeout, false, &streamStarted) + if err != nil { + reqLog.Warn("systemone.account_slot_acquire_failed", zap.Int64("account_id", account.ID), zap.Error(err)) + h.handleConcurrencyError(c, err, "account", false) + return + } + } + // 准入终检:与其他网关入口一致,利润控制否决的账号不得承接本次请求。 + admissionCtx := service.ContextWithSelectionProfitGate(c.Request.Context(), selection) + latest, vetoed, reason := h.gatewayService.GatewayProfitControlVetoLatest(admissionCtx, account) + if vetoed { + if accountRelease != nil { + accountRelease() + } + reqLog.Debug("systemone.account_slot_profit_vetoed", zap.Int64("account_id", account.ID), zap.String("reason", reason)) + if fs.RecordProfitVeto(account.ID) == FailoverExhausted { + reqLog.Warn("systemone.profit_veto_attempts_exhausted", zap.Int("profit_veto_count", fs.ProfitVetoCount())) + h.errorResponse(c, http.StatusServiceUnavailable, "api_error", profitVetoExhaustedMessage) + return + } + continue + } + account = latest + accountRelease = wrapReleaseOnDone(c.Request.Context(), accountRelease) + setOpsSelectedAccount(c, account.ID, account.Platform) + service.SetOpsUpstreamModel(c, model) + + forwardStart := time.Now() + result, forwardErr := h.gatewayService.ForwardSystemOne(c.Request.Context(), c, account, body) + if accountRelease != nil { + accountRelease() + } + service.SetOpsLatencyMs(c, service.OpsResponseLatencyMsKey, time.Since(forwardStart).Milliseconds()) + + if forwardErr != nil { + var failoverErr *service.UpstreamFailoverError + if errors.As(forwardErr, &failoverErr) { + switch fs.HandleFailoverError(c.Request.Context(), h.gatewayService, account.ID, account.Platform, account.GetPoolModeRetryCount(), failoverErr) { + case FailoverContinue: + reqLog.Warn("systemone.upstream_failover_switching", + zap.Int64("account_id", account.ID), + zap.Int("upstream_status", failoverErr.StatusCode), + zap.Int("switch_count", fs.SwitchCount), + ) + continue + case FailoverExhausted: + h.handleFailoverExhausted(c, fs.LastFailoverErr, service.PlatformTypeSafe, false) + return + case FailoverCanceled: + failoverClientGone(c) + return + } + } + if failoverClientGone(c) { + return + } + var upstreamErr *service.SystemOneUpstreamError + if errors.As(forwardErr, &upstreamErr) { + status := upstreamErr.StatusCode + if !service.IsSystemOneRequestErrorStatus(status) { + status = http.StatusBadGateway + } + h.errorResponse(c, status, "upstream_error", "TypeSafe rejected the request") + return + } + reqLog.Warn("systemone.forward_failed", zap.Int64("account_id", account.ID), zap.Error(forwardErr)) + if errors.Is(forwardErr, typesafe.ErrSystemOneResponseTooLarge) { + h.errorResponse(c, http.StatusBadGateway, "upstream_error", "TypeSafe response exceeds the gateway size limit") + return + } + h.errorResponse(c, http.StatusBadGateway, "upstream_error", "TypeSafe upstream request failed") + return + } + + c.Data(result.StatusCode, result.ContentType, result.Body) + h.recordSystemOneUsage(c, apiKey, account, subscription, channelMapping, model, body, result, subject.UserID, pricingAt) + return + } +} + +func (h *GatewayHandler) recordSystemOneUsage(c *gin.Context, apiKey *service.APIKey, account *service.Account, subscription *service.UserSubscription, mapping service.ChannelMappingResult, model string, body []byte, result *service.SystemOneForwardResult, userID int64, pricingAt time.Time) { + userAgent := c.GetHeader("User-Agent") + clientIP := ip.GetClientIP(c) + inboundEndpoint := GetInboundEndpoint(c) + upstreamEndpoint := GetUpstreamEndpoint(c, account.Platform) + quotaPlatform := service.QuotaPlatform(c.Request.Context(), apiKey) + sessionID := service.ExtractClientSessionID(c) + requestPayloadHash := service.HashUsageRequestPayload(body) + + h.submitMandatoryUsageRecordTask(c.Request.Context(), func(ctx context.Context) { + if err := h.gatewayService.RecordUsage(ctx, &service.RecordUsageInput{ + Result: &result.ForwardResult, + APIKey: apiKey, + User: apiKey.User, + Account: account, + Subscription: subscription, + PricingAt: pricingAt, + InboundEndpoint: inboundEndpoint, + UpstreamEndpoint: upstreamEndpoint, + UserAgent: userAgent, + IPAddress: clientIP, + SessionID: sessionID, + RequestPayloadHash: requestPayloadHash, + APIKeyService: h.apiKeyService, + QuotaPlatform: quotaPlatform, + ChannelUsageFields: clientRequestedUsageFields(c, mapping, model, result.UpstreamModel), + }); err != nil { + logger.L().With( + zap.String("component", "handler.gateway.systemone"), + zap.Int64("user_id", userID), + zap.Int64("api_key_id", apiKey.ID), + zap.Int64("account_id", account.ID), + ).Error("systemone.record_usage_failed", zap.Error(err)) + } + }) +} diff --git a/backend/internal/handler/gateway_systemone_test.go b/backend/internal/handler/gateway_systemone_test.go new file mode 100644 index 000000000..084b37ed9 --- /dev/null +++ b/backend/internal/handler/gateway_systemone_test.go @@ -0,0 +1,109 @@ +package handler + +import ( + "net/http" + "net/http/httptest" + "strings" + "testing" + + "github.com/Wei-Shaw/sub2api/internal/config" + middleware2 "github.com/Wei-Shaw/sub2api/internal/server/middleware" + "github.com/Wei-Shaw/sub2api/internal/service" + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" +) + +const validSystemOneHandlerBody = `{"model":"jev-latest","state":"sample","questions":{"q":{"type":"noul","instructions":"Evaluate"}}}` + +func newSystemOneHandlerContext(body string) (*gin.Context, *httptest.ResponseRecorder) { + gin.SetMode(gin.TestMode) + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/systemone", strings.NewReader(body)) + c.Request.Header.Set("Content-Type", "application/json") + return c, recorder +} + +func TestSystemOneRequiresAuthentication(t *testing.T) { + c, recorder := newSystemOneHandlerContext(validSystemOneHandlerBody) + (&GatewayHandler{}).SystemOne(c) + require.Equal(t, http.StatusUnauthorized, recorder.Code) + require.Contains(t, recorder.Body.String(), "authentication_error") +} + +func TestSystemOneRejectsNonTypeSafeGroupBeforeScheduling(t *testing.T) { + c, recorder := newSystemOneHandlerContext(validSystemOneHandlerBody) + groupID := int64(3) + c.Set(string(middleware2.ContextKeyAPIKey), &service.APIKey{ID: 4, UserID: 5, GroupID: &groupID, Group: &service.Group{ID: groupID, Platform: service.PlatformOpenAI}}) + c.Set(string(middleware2.ContextKeyUser), middleware2.AuthSubject{UserID: 5, Concurrency: 1}) + + (&GatewayHandler{cfg: &config.Config{Gateway: config.GatewayConfig{MaxBodySize: 1 << 20}}}).SystemOne(c) + require.Equal(t, http.StatusNotFound, recorder.Code) + require.Contains(t, recorder.Body.String(), "only available for TypeSafe") +} + +func newTypeSafeGroupContext(t *testing.T, path, body, groupPlatform string) (*gin.Context, *httptest.ResponseRecorder) { + t.Helper() + c, recorder := newSystemOneHandlerContext(body) + c.Request = httptest.NewRequest(http.MethodPost, path, strings.NewReader(body)) + c.Request.Header.Set("Content-Type", "application/json") + groupID := int64(9) + c.Set(string(middleware2.ContextKeyAPIKey), &service.APIKey{ID: 4, UserID: 5, GroupID: &groupID, Group: &service.Group{ID: groupID, Platform: groupPlatform}}) + c.Set(string(middleware2.ContextKeyUser), middleware2.AuthSubject{UserID: 5, Concurrency: 1}) + return c, recorder +} + +func TestRejectSystemOneOnlyPlatform(t *testing.T) { + write := func(c *gin.Context, status int, errType, message string) { + c.JSON(status, gin.H{"type": errType, "message": message}) + } + + c, recorder := newTypeSafeGroupContext(t, "/v1/messages", `{}`, service.PlatformTypeSafe) + require.True(t, rejectSystemOneOnlyPlatform(c, mustAPIKey(t, c), write)) + require.Equal(t, http.StatusNotFound, recorder.Code) + require.Contains(t, recorder.Body.String(), systemOneOnlyPlatformMessage) + + c, _ = newTypeSafeGroupContext(t, "/v1/messages", `{}`, service.PlatformComposite) + require.False(t, rejectSystemOneOnlyPlatform(c, mustAPIKey(t, c), write)) + c.Request = c.Request.WithContext(service.WithResolvedTargetPlatform(c.Request.Context(), service.PlatformTypeSafe)) + require.True(t, rejectSystemOneOnlyPlatform(c, mustAPIKey(t, c), write)) + + c, _ = newTypeSafeGroupContext(t, "/v1/messages", `{}`, service.PlatformAnthropic) + require.False(t, rejectSystemOneOnlyPlatform(c, mustAPIKey(t, c), write)) + + c, _ = newTypeSafeGroupContext(t, "/antigravity/v1/messages", `{}`, service.PlatformTypeSafe) + c.Set(string(middleware2.ContextKeyForcePlatform), service.PlatformAntigravity) + require.False(t, rejectSystemOneOnlyPlatform(c, mustAPIKey(t, c), write)) +} + +func mustAPIKey(t *testing.T, c *gin.Context) *service.APIKey { + t.Helper() + apiKey, ok := middleware2.GetAPIKeyFromContext(c) + require.True(t, ok) + return apiKey +} + +func TestTypeSafeGroupsRejectNonSystemOneProtocolsBeforeScheduling(t *testing.T) { + h := &GatewayHandler{ + cfg: &config.Config{Gateway: config.GatewayConfig{MaxBodySize: 1 << 20}}, + gatewayService: &service.GatewayService{}, + } + for _, tc := range []struct { + name string + path string + body string + handler func(*gin.Context) + }{ + {"messages", "/v1/messages", `{"model":"jev-latest","max_tokens":1,"messages":[{"role":"user","content":"hi"}]}`, h.Messages}, + {"count_tokens", "/v1/messages/count_tokens", `{"model":"jev-latest","messages":[{"role":"user","content":"hi"}]}`, h.CountTokens}, + {"chat_completions", "/v1/chat/completions", `{"model":"jev-latest","messages":[{"role":"user","content":"hi"}]}`, h.ChatCompletions}, + {"responses", "/v1/responses", `{"model":"jev-latest","input":"hi"}`, h.Responses}, + } { + t.Run(tc.name, func(t *testing.T) { + c, recorder := newTypeSafeGroupContext(t, tc.path, tc.body, service.PlatformTypeSafe) + tc.handler(c) + require.Equal(t, http.StatusNotFound, recorder.Code, recorder.Body.String()) + require.Contains(t, recorder.Body.String(), systemOneOnlyPlatformMessage) + }) + } +} diff --git a/backend/internal/handler/payment_handler.go b/backend/internal/handler/payment_handler.go index 9aab32550..24e318cc5 100644 --- a/backend/internal/handler/payment_handler.go +++ b/backend/internal/handler/payment_handler.go @@ -149,6 +149,9 @@ func (h *PaymentHandler) GetCheckoutInfo(c *gin.Context) { BalanceRechargeMultiplier: cfg.BalanceRechargeMultiplier, SubscriptionUSDToCNYRate: cfg.SubscriptionUSDToCNYRate, RechargeFeeRate: cfg.RechargeFeeRate, + RechargeBonusTiers: cfg.RechargeBonusTiers, + RechargeBonusMode: cfg.RechargeBonusMode, + RechargeBonusNotice: cfg.RechargeBonusNotice, HelpText: cfg.HelpText, HelpImageURL: cfg.HelpImageURL, StripePublishableKey: cfg.StripePublishableKey, @@ -166,6 +169,9 @@ type checkoutInfoResponse struct { BalanceRechargeMultiplier float64 `json:"balance_recharge_multiplier"` SubscriptionUSDToCNYRate float64 `json:"subscription_usd_to_cny_rate"` RechargeFeeRate float64 `json:"recharge_fee_rate"` + RechargeBonusTiers []service.RechargeBonusTier `json:"recharge_bonus_tiers"` + RechargeBonusMode string `json:"recharge_bonus_mode"` + RechargeBonusNotice string `json:"recharge_bonus_notice"` HelpText string `json:"help_text"` HelpImageURL string `json:"help_image_url"` StripePublishableKey string `json:"stripe_publishable_key"` @@ -482,6 +488,7 @@ type PublicOrderResult struct { Amount float64 `json:"amount"` PayAmount float64 `json:"pay_amount"` FeeRate float64 `json:"fee_rate"` + BonusAmount float64 `json:"bonus_amount"` Currency string `json:"currency"` PaymentType string `json:"payment_type"` OrderType string `json:"order_type"` @@ -517,6 +524,7 @@ func buildPublicOrderResult(order *dbent.PaymentOrder) PublicOrderResult { Amount: order.Amount, PayAmount: order.PayAmount, FeeRate: order.FeeRate, + BonusAmount: order.BonusAmount, Currency: service.PaymentOrderCurrency(order), PaymentType: order.PaymentType, OrderType: order.OrderType, @@ -625,6 +633,7 @@ type PaymentOrderResult struct { Amount float64 `json:"amount"` PayAmount float64 `json:"pay_amount"` FeeRate float64 `json:"fee_rate"` + BonusAmount float64 `json:"bonus_amount"` Currency string `json:"currency"` PaymentType string `json:"payment_type"` OutTradeNo string `json:"out_trade_no"` @@ -663,6 +672,7 @@ func sanitizePaymentOrderForResponse(order *dbent.PaymentOrder) *PaymentOrderRes Amount: order.Amount, PayAmount: order.PayAmount, FeeRate: order.FeeRate, + BonusAmount: order.BonusAmount, Currency: service.PaymentOrderCurrency(order), PaymentType: order.PaymentType, OutTradeNo: order.OutTradeNo, diff --git a/backend/internal/handler/user_handler_test.go b/backend/internal/handler/user_handler_test.go index c30de147f..796c64f33 100644 --- a/backend/internal/handler/user_handler_test.go +++ b/backend/internal/handler/user_handler_test.go @@ -6,6 +6,7 @@ import ( "bytes" "context" "encoding/json" + "errors" "net/http" "net/http/httptest" "testing" @@ -812,6 +813,18 @@ func TestUserHandlerStartIdentityBindingReturnsAuthorizeURL(t *testing.T) { require.Contains(t, resp.Data.AuthorizeURL, "redirect=%2Fsettings%2Fprofile") } +func (s *userHandlerEmailCacheStub) IncrVerificationCodeAttempts(context.Context, string) (int, error) { + if s.data == nil { + return 0, errors.New("verification code not found") + } + s.data.Attempts++ + return s.data.Attempts, nil +} + +func (s *userHandlerEmailCacheStub) IncrNotifyVerifyCodeAttempts(context.Context, string) (int, error) { + return 0, errors.New("notify verification code not found") +} + func (s *userHandlerEmailCacheStub) ConsumePasswordResetToken(context.Context, string, string) (bool, error) { return false, nil } diff --git a/backend/internal/model/error_passthrough_rule.go b/backend/internal/model/error_passthrough_rule.go index a0c81934d..da928125f 100644 --- a/backend/internal/model/error_passthrough_rule.go +++ b/backend/internal/model/error_passthrough_rule.go @@ -46,6 +46,7 @@ const ( PlatformDeepseek = domain.PlatformDeepseek PlatformMiniMax = domain.PlatformMiniMax PlatformOpenCodeGo = domain.PlatformOpenCodeGo + PlatformTypeSafe = domain.PlatformTypeSafe ) // AllPlatforms 返回所有支持的平台列表 @@ -61,6 +62,7 @@ func AllPlatforms() []string { PlatformDeepseek, PlatformMiniMax, PlatformOpenCodeGo, + PlatformTypeSafe, } } diff --git a/backend/internal/model/error_passthrough_rule_test.go b/backend/internal/model/error_passthrough_rule_test.go index 34ceb414e..17e6fdcbf 100644 --- a/backend/internal/model/error_passthrough_rule_test.go +++ b/backend/internal/model/error_passthrough_rule_test.go @@ -18,5 +18,6 @@ func TestAllPlatformsIncludesEveryConcretePlatform(t *testing.T) { "deepseek", "minimax", "opencode_go", + "typesafe", }, AllPlatforms()) } diff --git a/backend/internal/pkg/typesafe/client.go b/backend/internal/pkg/typesafe/client.go index 35783249e..0974204e2 100644 --- a/backend/internal/pkg/typesafe/client.go +++ b/backend/internal/pkg/typesafe/client.go @@ -13,6 +13,12 @@ import ( "strings" ) +const ( + DefaultBaseURL = "https://api.typesafe.ai" + SystemOnePath = "/v1/systemone" + JevLatestModel = "jev-latest" +) + type Question struct { Type string `json:"type"` Instructions string `json:"instructions"` @@ -35,22 +41,103 @@ type Usage struct { OutputTokens int `json:"output_tokens"` } -// Evaluate performs one attempt. The caller owns timeouts, retries and key rotation. -func Evaluate(ctx context.Context, client *http.Client, baseURL, key string, input Request) (*Result, int, error) { - endpoint, err := url.JoinPath(strings.TrimRight(baseURL, "/"), "/v1/systemone") +type SystemOneResponse struct { + Body []byte + Model string + Usage Usage +} + +func NewSystemOneRequest(ctx context.Context, baseURL, key string, body []byte) (*http.Request, error) { + endpoint, err := url.JoinPath(strings.TrimRight(baseURL, "/"), SystemOnePath) if err != nil { - return nil, 0, errors.New("typesafe invalid endpoint") + return nil, errors.New("typesafe invalid endpoint") + } + req, err := http.NewRequestWithContext(ctx, http.MethodPost, endpoint, bytes.NewReader(body)) + if err != nil { + return nil, errors.New("typesafe invalid request") + } + req.Header.Set("Authorization", "Bearer "+key) + req.Header.Set("Content-Type", "application/json") + return req, nil +} + +// MaxSystemOneResponseBytes bounds a buffered System One response body. +const MaxSystemOneResponseBytes = 4 << 20 + +var ErrSystemOneResponseTooLarge = errors.New("typesafe response exceeds size limit") + +func DecodeSystemOneResponse(r io.Reader) (*SystemOneResponse, error) { + body, err := io.ReadAll(io.LimitReader(r, MaxSystemOneResponseBytes+1)) + if err != nil { + return nil, errors.New("typesafe invalid response") + } + // A truncated body would otherwise surface as a misleading "invalid JSON". + if len(body) > MaxSystemOneResponseBytes { + return nil, ErrSystemOneResponseTooLarge + } + if !json.Valid(body) { + return nil, errors.New("typesafe invalid response") } + var envelope map[string]json.RawMessage + if err := json.Unmarshal(body, &envelope); err != nil || envelope == nil { + return nil, errors.New("typesafe invalid response") + } + // The upstream already answered (and charged); an unexpected model or usage + // shape must not discard the answer, so both are decoded leniently. + var model string + _ = json.Unmarshal(envelope["model"], &model) + var usage map[string]json.RawMessage + _ = json.Unmarshal(envelope["usage"], &usage) + return &SystemOneResponse{ + Body: body, + Model: model, + Usage: Usage{ + InputTokens: systemOneTokenCount(usage["input_tokens"]), + OutputTokens: systemOneTokenCount(usage["output_tokens"]), + }, + }, nil +} + +// maxSystemOneTokenCount bounds a reported token count before int conversion. +const maxSystemOneTokenCount = 1 << 40 + +// systemOneTokenCount accepts integer, float, or numeric-string token counts. +func systemOneTokenCount(raw json.RawMessage) int { + raw = bytes.TrimSpace(raw) + if len(raw) == 0 { + return 0 + } + if raw[0] == '"' { + var text string + if json.Unmarshal(raw, &text) != nil { + return 0 + } + raw = json.RawMessage(strings.TrimSpace(text)) + } + var number json.Number + if json.Unmarshal(raw, &number) != nil { + return 0 + } + value, err := number.Float64() + if err != nil || math.IsNaN(value) || value <= 0 { + return 0 + } + if value > maxSystemOneTokenCount { + value = maxSystemOneTokenCount + } + return int(math.Round(value)) +} + +// Evaluate performs one attempt. The caller owns timeouts, retries and key rotation. +func Evaluate(ctx context.Context, client *http.Client, baseURL, key string, input Request) (*Result, int, error) { body, err := json.Marshal(input) if err != nil { return nil, 0, errors.New("typesafe invalid request") } - req, err := http.NewRequestWithContext(ctx, http.MethodPost, endpoint, bytes.NewReader(body)) + req, err := NewSystemOneRequest(ctx, baseURL, key, body) if err != nil { - return nil, 0, errors.New("typesafe invalid request") + return nil, 0, err } - req.Header.Set("Authorization", "Bearer "+key) - req.Header.Set("Content-Type", "application/json") resp, err := client.Do(req) if err != nil { if ctx.Err() != nil { diff --git a/backend/internal/pkg/typesafe/systemone.go b/backend/internal/pkg/typesafe/systemone.go new file mode 100644 index 000000000..4fab1f1d6 --- /dev/null +++ b/backend/internal/pkg/typesafe/systemone.go @@ -0,0 +1,202 @@ +package typesafe + +import ( + "bytes" + "encoding/json" + "errors" + "fmt" + "strings" +) + +var ErrStreamingUnsupported = errors.New("typesafe system one does not support streaming") + +type systemOneEnvelope struct { + Model json.RawMessage `json:"model"` + State json.RawMessage `json:"state"` + Questions json.RawMessage `json:"questions"` + Stream json.RawMessage `json:"stream"` +} + +var ( + systemOneRequestFields = []string{"model", "state", "questions", "stream"} + systemOneQuestionFields = []string{"type", "instructions", "criteria"} +) + +func ValidateSystemOneRequest(body []byte) (string, error) { + var envelope systemOneEnvelope + if err := json.Unmarshal(body, &envelope); err != nil { + return "", errors.New("invalid JSON request") + } + // encoding/json matches struct fields case-insensitively and keeps the last + // duplicate, while the raw body is forwarded upstream unchanged. Reject any + // spelling the upstream could read differently from this validator. + if err := checkSystemOneObjectKeys(body, "request", systemOneRequestFields); err != nil { + return "", err + } + + model, err := requiredString(envelope.Model, "model") + if err != nil { + return "", err + } + if model != JevLatestModel { + return "", fmt.Errorf("model must be %s", JevLatestModel) + } + if err := validateStringObjectOrArray(envelope.State, "state"); err != nil { + return "", err + } + questionsRaw := bytes.TrimSpace(envelope.Questions) + var questions map[string]json.RawMessage + if len(questionsRaw) == 0 || questionsRaw[0] != '{' || json.Unmarshal(questionsRaw, &questions) != nil || len(questions) == 0 { + return "", errors.New("questions must be a non-empty object") + } + if err := checkSystemOneObjectKeys(questionsRaw, "questions", nil); err != nil { + return "", err + } + if len(envelope.Stream) > 0 && string(envelope.Stream) != "null" { + var stream bool + if err := json.Unmarshal(envelope.Stream, &stream); err != nil { + return "", errors.New("stream must be a boolean") + } + if stream { + return "", ErrStreamingUnsupported + } + } + for id, raw := range questions { + if err := validateQuestion(id, raw); err != nil { + return "", err + } + } + return model, nil +} + +func validateQuestion(id string, raw json.RawMessage) error { + raw = bytes.TrimSpace(raw) + var question struct { + Type json.RawMessage `json:"type"` + Instructions json.RawMessage `json:"instructions"` + Criteria json.RawMessage `json:"criteria"` + } + if len(raw) == 0 || raw[0] != '{' || json.Unmarshal(raw, &question) != nil { + return fmt.Errorf("question %q must be an object", id) + } + if err := checkSystemOneObjectKeys(raw, fmt.Sprintf("question %q", id), systemOneQuestionFields); err != nil { + return err + } + typ, err := requiredString(question.Type, "question type") + if err != nil { + return fmt.Errorf("question %q: %w", id, err) + } + // Mirror the wire schema generated from https://api.typesafe.ai/openapi.json + // in TypeSafe SDK v0.5.7: instructions are optional and nullable for all types. + if err := validateOptionalDescription(question.Instructions, "instructions"); err != nil { + return fmt.Errorf("question %q: %w", id, err) + } + + switch typ { + case "noul": + criteria := bytes.TrimSpace(question.Criteria) + if len(criteria) == 0 || bytes.Equal(criteria, []byte("null")) { + return nil + } + var descriptions map[string]json.RawMessage + if criteria[0] != '{' || json.Unmarshal(criteria, &descriptions) != nil { + return fmt.Errorf("question %q: noul criteria must be an object", id) + } + for _, outcome := range []string{"true", "false"} { + if err := validateOptionalDescription(descriptions[outcome], "noul criteria "+outcome); err != nil { + return fmt.Errorf("question %q: %w", id, err) + } + } + case "choice": + var criteria map[string]json.RawMessage + if json.Unmarshal(question.Criteria, &criteria) != nil || criteria == nil { + return fmt.Errorf("question %q: choice criteria must be an object", id) + } + for _, value := range criteria { + if err := validateOptionalDescription(value, "choice criteria value"); err != nil { + return fmt.Errorf("question %q: %w", id, err) + } + } + case "score": + var criteria []json.RawMessage + if json.Unmarshal(question.Criteria, &criteria) != nil || len(criteria) == 0 { + return fmt.Errorf("question %q: score criteria must contain at least one level", id) + } + for _, value := range criteria { + if err := validateStringObjectOrArray(value, "score criteria value"); err != nil { + return fmt.Errorf("question %q: %w", id, err) + } + } + default: + return fmt.Errorf("question %q: unsupported type %q", id, typ) + } + return nil +} + +func validateOptionalDescription(raw json.RawMessage, name string) error { + raw = bytes.TrimSpace(raw) + if len(raw) == 0 || bytes.Equal(raw, []byte("null")) { + return nil + } + return validateStringObjectOrArray(raw, name) +} + +func requiredString(raw json.RawMessage, name string) (string, error) { + var value string + if len(raw) == 0 || json.Unmarshal(raw, &value) != nil || strings.TrimSpace(value) == "" { + return "", fmt.Errorf("%s must be a non-empty string", name) + } + return value, nil +} + +func validateStringObjectOrArray(raw json.RawMessage, name string) error { + raw = bytes.TrimSpace(raw) + if len(raw) == 0 || string(raw) == "null" { + return fmt.Errorf("%s is required", name) + } + switch raw[0] { + case '{', '[': + return nil + case '"': + if rawString(raw) { + return nil + } + } + return fmt.Errorf("%s must be a string, object, or array", name) +} + +func rawString(raw json.RawMessage) bool { + var value string + return json.Unmarshal(raw, &value) == nil +} + +// checkSystemOneObjectKeys rejects duplicate keys and non-canonical spellings +// (any case variant) of the known fields in one JSON object level. +func checkSystemOneObjectKeys(raw []byte, scope string, canonical []string) error { + decoder := json.NewDecoder(bytes.NewReader(raw)) + if token, err := decoder.Token(); err != nil || token != json.Delim('{') { + return fmt.Errorf("%s must be an object", scope) + } + seen := make(map[string]struct{}) + for decoder.More() { + token, err := decoder.Token() + if err != nil { + return errors.New("invalid JSON request") + } + key, _ := token.(string) + if _, duplicate := seen[key]; duplicate { + return fmt.Errorf("%s contains duplicate field %q", scope, key) + } + seen[key] = struct{}{} + for _, field := range canonical { + if key != field && strings.EqualFold(key, field) { + return fmt.Errorf("%s field %q must be written as %q", scope, key, field) + } + } + var value json.RawMessage + if err := decoder.Decode(&value); err != nil { + return errors.New("invalid JSON request") + } + } + return nil +} diff --git a/backend/internal/pkg/typesafe/systemone_test.go b/backend/internal/pkg/typesafe/systemone_test.go new file mode 100644 index 000000000..65991cf14 --- /dev/null +++ b/backend/internal/pkg/typesafe/systemone_test.go @@ -0,0 +1,124 @@ +package typesafe + +import ( + "errors" + "strings" + "testing" + + "github.com/stretchr/testify/require" +) + +func TestValidateSystemOneRequestValidQuestionTypes(t *testing.T) { + for _, tc := range []struct { + name string + body string + }{ + {"noul string state", `{"model":"jev-latest","state":"sample","questions":{"safety":{"type":"noul","instructions":"Evaluate safety","criteria":{"safe":"No harm"}}}}`}, + {"choice object state", `{"model":"jev-latest","state":{"text":"sample"},"questions":{"label":{"type":"choice","instructions":{"task":"Classify"},"criteria":{"safe":"Allowed","unsafe":null}}},"stream":false}`}, + {"score array state", `{"model":"jev-latest","state":["sample"],"questions":{"quality":{"type":"score","instructions":["Rate quality"],"criteria":["poor","good"]}}}`}, + {"noul omitted instructions", `{"model":"jev-latest","state":"sample","questions":{"q":{"type":"noul"}}}`}, + {"noul nullable instructions and criteria", `{"model":"jev-latest","state":"sample","questions":{"q":{"type":"noul","instructions":null,"criteria":null}}}`}, + {"noul structured descriptions", `{"model":"jev-latest","state":"sample","questions":{"q":{"type":"noul","criteria":{"true":{"examples":["yes",null]},"false":["no",null],"extension":42}}}}`}, + {"choice structured descriptions", `{"model":"jev-latest","state":"sample","questions":{"q":{"type":"choice","instructions":null,"criteria":{"a":{"description":"A","extra":null},"b":["B",null],"c":null}}}}`}, + {"choice empty criteria", `{"model":"jev-latest","state":"sample","questions":{"q":{"type":"choice","criteria":{}}}}`}, + {"score one level", `{"model":"jev-latest","state":"sample","questions":{"q":{"type":"score","criteria":["only"]}}}`}, + {"score object level", `{"model":"jev-latest","state":"sample","questions":{"q":{"type":"score","instructions":null,"criteria":[{"description":"only","extra":null}]}}}`}, + {"score array level", `{"model":"jev-latest","state":"sample","questions":{"q":{"type":"score","criteria":[["only",null]]}}}`}, + {"native extensions preserved", `{"model":"jev-latest","state":"sample","questions":{"q":{"type":"noul","extension":{"kept":true}}},"extension":[1,null]}`}, + } { + t.Run(tc.name, func(t *testing.T) { + model, err := ValidateSystemOneRequest([]byte(tc.body)) + require.NoError(t, err) + require.Equal(t, JevLatestModel, model) + }) + } +} + +func TestValidateSystemOneRequestRejectsInvalidRequests(t *testing.T) { + for _, tc := range []struct { + name string + body string + want string + }{ + {"invalid json", `{`, "invalid JSON"}, + {"missing model", `{"state":"x","questions":{"q":{"type":"noul","instructions":"x"}}}`, "model"}, + {"illegal model", `{"model":"jev-old","state":"x","questions":{"q":{"type":"noul","instructions":"x"}}}`, "jev-latest"}, + {"model with whitespace", `{"model":" jev-latest ","state":"x","questions":{"q":{"type":"noul","instructions":"x"}}}`, "jev-latest"}, + {"missing state", `{"model":"jev-latest","questions":{"q":{"type":"noul","instructions":"x"}}}`, "state"}, + {"scalar state", `{"model":"jev-latest","state":42,"questions":{"q":{"type":"noul","instructions":"x"}}}`, "state"}, + {"empty questions", `{"model":"jev-latest","state":"x","questions":{}}`, "questions"}, + {"array questions", `{"model":"jev-latest","state":"x","questions":[{"type":"noul"}]}`, "questions must be a non-empty object"}, + {"null questions", `{"model":"jev-latest","state":"x","questions":null}`, "questions must be a non-empty object"}, + {"unknown question type", `{"model":"jev-latest","state":"x","questions":{"q":{"type":"boolean","instructions":"x"}}}`, "unsupported type"}, + {"question type with whitespace", `{"model":"jev-latest","state":"x","questions":{"q":{"type":" noul ","instructions":"x"}}}`, "unsupported type"}, + {"noul criteria array", `{"model":"jev-latest","state":"x","questions":{"q":{"type":"noul","instructions":"x","criteria":[]}}}`, "noul criteria"}, + {"noul numeric description", `{"model":"jev-latest","state":"x","questions":{"q":{"type":"noul","criteria":{"true":1}}}}`, "noul criteria"}, + {"noul boolean description", `{"model":"jev-latest","state":"x","questions":{"q":{"type":"noul","criteria":{"false":false}}}}`, "noul criteria"}, + {"missing choice criteria", `{"model":"jev-latest","state":"x","questions":{"q":{"type":"choice"}}}`, "choice criteria"}, + {"null choice criteria", `{"model":"jev-latest","state":"x","questions":{"q":{"type":"choice","criteria":null}}}`, "choice criteria"}, + {"choice criteria array", `{"model":"jev-latest","state":"x","questions":{"q":{"type":"choice","criteria":[]}}}`, "choice criteria"}, + {"choice numeric value", `{"model":"jev-latest","state":"x","questions":{"q":{"type":"choice","instructions":"x","criteria":{"one":1}}}}`, "choice criteria"}, + {"missing score criteria", `{"model":"jev-latest","state":"x","questions":{"q":{"type":"score"}}}`, "score criteria"}, + {"empty score criteria", `{"model":"jev-latest","state":"x","questions":{"q":{"type":"score","criteria":[]}}}`, "score criteria"}, + {"null score criteria", `{"model":"jev-latest","state":"x","questions":{"q":{"type":"score","criteria":null}}}`, "score criteria"}, + {"score map is not wire protocol", `{"model":"jev-latest","state":"x","questions":{"q":{"type":"score","criteria":{"0":"low"}}}}`, "score criteria"}, + {"numeric score level", `{"model":"jev-latest","state":"x","questions":{"q":{"type":"score","criteria":["low",2]}}}`, "score criteria"}, + {"null score level", `{"model":"jev-latest","state":"x","questions":{"q":{"type":"score","criteria":[null]}}}`, "score criteria"}, + {"numeric instructions", `{"model":"jev-latest","state":"x","questions":{"q":{"type":"noul","instructions":1}}}`, "instructions"}, + {"boolean instructions", `{"model":"jev-latest","state":"x","questions":{"q":{"type":"choice","instructions":true,"criteria":{}}}}`, "instructions"}, + {"stream true", `{"model":"jev-latest","state":"x","questions":{"q":{"type":"noul","instructions":"x"}},"stream":true}`, "streaming"}, + {"case variant model smuggles upstream model", `{"model":"jev-pro","MODEL":"jev-latest","state":"x","questions":{"q":{"type":"noul"}}}`, `must be written as "model"`}, + {"case variant stream hides streaming", `{"model":"jev-latest","state":"x","questions":{"q":{"type":"noul"}},"stream":true,"Stream":false}`, `must be written as "stream"`}, + {"unicode fold variant state", `{"model":"jev-latest","state":"x","ſtate":"y","questions":{"q":{"type":"noul"}}}`, `must be written as "state"`}, + {"duplicate model", `{"model":"jev-pro","model":"jev-latest","state":"x","questions":{"q":{"type":"noul"}}}`, `duplicate field "model"`}, + {"escaped duplicate model", `{"\u006dodel":"jev-pro","model":"jev-latest","state":"x","questions":{"q":{"type":"noul"}}}`, `duplicate field "model"`}, + {"duplicate state", `{"model":"jev-latest","state":"benign","state":"payload","questions":{"q":{"type":"noul"}}}`, `duplicate field "state"`}, + {"duplicate question id", `{"model":"jev-latest","state":"x","questions":{"q":{"type":"noul"},"q":{"type":"noul"}}}`, `duplicate field "q"`}, + {"duplicate question instructions", `{"model":"jev-latest","state":"x","questions":{"q":{"type":"noul","instructions":"a","instructions":"b"}}}`, `duplicate field "instructions"`}, + {"case variant question type", `{"model":"jev-latest","state":"x","questions":{"q":{"type":"noul","Type":"choice"}}}`, `must be written as "type"`}, + } { + t.Run(tc.name, func(t *testing.T) { + _, err := ValidateSystemOneRequest([]byte(tc.body)) + require.Error(t, err) + require.Contains(t, err.Error(), tc.want) + if tc.name == "stream true" { + require.True(t, errors.Is(err, ErrStreamingUnsupported)) + } + }) + } +} + +func TestDecodeSystemOneResponseToleratesUsageShapes(t *testing.T) { + for _, tc := range []struct { + name string + body string + model string + input, output int + }{ + {"integers", `{"model":"jev-1","usage":{"input_tokens":12,"output_tokens":3}}`, "jev-1", 12, 3}, + {"floats", `{"model":"jev-1","usage":{"input_tokens":12.0,"output_tokens":2.6}}`, "jev-1", 12, 3}, + {"numeric strings", `{"usage":{"input_tokens":"15","output_tokens":" 4 "}}`, "", 15, 4}, + {"non numeric usage", `{"model":7,"usage":{"input_tokens":"many","output_tokens":null}}`, "", 0, 0}, + {"negative usage", `{"usage":{"input_tokens":-5}}`, "", 0, 0}, + {"usage not object", `{"model":"jev-1","usage":"none","answers":{}}`, "jev-1", 0, 0}, + {"missing usage", `{"answers":{}}`, "", 0, 0}, + } { + t.Run(tc.name, func(t *testing.T) { + decoded, err := DecodeSystemOneResponse(strings.NewReader(tc.body)) + require.NoError(t, err) + require.Equal(t, []byte(tc.body), decoded.Body) + require.Equal(t, tc.model, decoded.Model) + require.Equal(t, tc.input, decoded.Usage.InputTokens) + require.Equal(t, tc.output, decoded.Usage.OutputTokens) + }) + } +} + +func TestDecodeSystemOneResponseRejectsNonObjects(t *testing.T) { + for _, body := range []string{`[]`, `"text"`, `null`, `{`, `data: {}`} { + _, err := DecodeSystemOneResponse(strings.NewReader(body)) + require.Error(t, err, body) + } + _, err := DecodeSystemOneResponse(strings.NewReader(`{"pad":"` + strings.Repeat("a", MaxSystemOneResponseBytes) + `"}`)) + require.ErrorIs(t, err, ErrSystemOneResponseTooLarge) +} diff --git a/backend/internal/pkg/xai/billing.go b/backend/internal/pkg/xai/billing.go index 39e1a2c8b..735baa86d 100644 --- a/backend/internal/pkg/xai/billing.go +++ b/backend/internal/pkg/xai/billing.go @@ -19,10 +19,7 @@ const ( // repository and service layers build their own client identity from it, so // one bump here covers OAuth traffic and billing probes together. // Keep in sync with https://x.ai/cli/stable. - CLIClientVersion = "0.2.120" - // billingCLIUserAgent is the legacy pager/shell UA used by billing probes. - // Distinct from CLIUserAgent() in cli_identity.go (workspace-style UA). - billingCLIUserAgent = "grok-pager/" + CLIClientVersion + " grok-shell/" + CLIClientVersion + " (macos; aarch64)" + CLIClientVersion = "1.0.46" BillingWeeklyPath = "/billing?format=credits" BillingMonthlyPath = "/billing" @@ -148,7 +145,8 @@ func ApplyCLIBillingHeaders(req *http.Request, accessToken string) { req.Header.Set("Content-Type", "application/json") req.Header.Set(CLITokenAuthHeader, CLITokenAuthValue) req.Header.Set(CLIClientVersionHeader, CLIClientVersion) - req.Header.Set("User-Agent", billingCLIUserAgent) + req.Header.Set("User-Agent", CLIUserAgent(CLIClientVersion)) + req.Header.Set("x-grok-client-mode", CLIClientMode) } // ParseBillingPayload unmarshals a billing API response body. diff --git a/backend/internal/pkg/xai/billing_test.go b/backend/internal/pkg/xai/billing_test.go index 3dbfc89bf..51e3f2b1b 100644 --- a/backend/internal/pkg/xai/billing_test.go +++ b/backend/internal/pkg/xai/billing_test.go @@ -38,7 +38,8 @@ func TestApplyCLIBillingHeaders(t *testing.T) { require.Equal(t, "Bearer token", req.Header.Get("Authorization")) require.Equal(t, CLITokenAuthValue, req.Header.Get(CLITokenAuthHeader)) require.Equal(t, CLIClientVersion, req.Header.Get(CLIClientVersionHeader)) - require.Equal(t, "grok-pager/"+CLIClientVersion+" grok-shell/"+CLIClientVersion+" (macos; aarch64)", req.UserAgent()) + require.Equal(t, CLIUserAgent(CLIClientVersion), req.UserAgent()) + require.Equal(t, "interactive", req.Header.Get("x-grok-client-mode")) } func TestBuildBillingSummaryWeeklyAndMonthly(t *testing.T) { diff --git a/backend/internal/pkg/xai/cli_identity.go b/backend/internal/pkg/xai/cli_identity.go index 480e2e40a..e930867f3 100644 --- a/backend/internal/pkg/xai/cli_identity.go +++ b/backend/internal/pkg/xai/cli_identity.go @@ -3,6 +3,7 @@ package xai import ( "net/http" "os" + "runtime" "strings" "golang.org/x/mod/semver" @@ -16,7 +17,7 @@ const ( CLIProxyHost = "cli-chat-proxy.grok.com" // CLIStableVersion is the known-good minimum client version accepted by cli-chat-proxy. - CLIStableVersion = "0.2.93" + CLIStableVersion = "1.0.13" // CLIVersionEnv is the optional operator override for CLIStableVersion. CLIVersionEnv = "XAI_GROK_CLI_VERSION" @@ -24,11 +25,11 @@ const ( // CLITokenAuth is required by cli-chat-proxy for Grok Build OAuth tokens. CLITokenAuth = "xai-grok-cli" - // CLIClientIdentifier is the x-grok-client-identifier value used by Grok shell/CLI. - CLIClientIdentifier = "grok-shell" + // CLIClientIdentifier 对齐官方交互式 CLI 主请求的客户端标识。 + CLIClientIdentifier = "grok-pager" - // CLIClientMode is used by billing / quota probes on the CLI surface. - CLIClientMode = "cli" + // CLIClientMode 对齐官方 CLI 正常交互模式。 + CLIClientMode = "interactive" ) // ResolveCLIVersion returns a supported CLI client version. @@ -54,12 +55,24 @@ func IsSupportedCLIVersion(version string) bool { semver.Compare(canonical, minimum) >= 0 } -// CLIUserAgent builds the workspace-style User-Agent for a CLI client version. +// CLIUserAgent 对齐官方交互式 CLI 的 UA,平台名称使用 Rust 的格式。 func CLIUserAgent(version string) string { if strings.TrimSpace(version) == "" { version = CLIClientVersion } - return "xai-grok-workspace/" + version + platform, arch := runtime.GOOS, runtime.GOARCH + if platform == "darwin" { + platform = "macos" + } + switch arch { + case "amd64": + arch = "x86_64" + case "arm64": + arch = "aarch64" + case "386": + arch = "x86" + } + return "grok-pager/" + version + " grok-shell/" + version + " (" + platform + "; " + arch + ")" } // ApplyCLIProxyHeaders stamps the fixed Grok CLI identity when the request @@ -75,5 +88,7 @@ func ApplyCLIProxyHeaders(req *http.Request) { req.Header.Set("X-XAI-Token-Auth", CLITokenAuth) req.Header.Set("x-grok-client-version", version) req.Header.Set("x-grok-client-identifier", CLIClientIdentifier) + req.Header.Set("x-grok-client-mode", CLIClientMode) + req.Header.Set("x-authenticateresponse", "authenticate-response") req.Header.Set("User-Agent", CLIUserAgent(version)) } diff --git a/backend/internal/pkg/xai/cli_identity_test.go b/backend/internal/pkg/xai/cli_identity_test.go index 2057b9eff..015114919 100644 --- a/backend/internal/pkg/xai/cli_identity_test.go +++ b/backend/internal/pkg/xai/cli_identity_test.go @@ -2,11 +2,40 @@ package xai import ( "net/http" + "runtime" "testing" "github.com/stretchr/testify/require" + "golang.org/x/mod/semver" ) +func TestApplyCLIProxyHeadersMeetsUpstreamMinimumVersion(t *testing.T) { + // 上游 426 明确要求至少 1.0.13;断言独立于生产版本常量,防止旧覆盖值绕过下限。 + for _, override := range []string{"", "0.2.93", "0.2.120", "1.0.12", "1.0.13-beta.1", "1.0.13", "1.0.14-alpha.1"} { + t.Run("override="+override, func(t *testing.T) { + t.Setenv(CLIVersionEnv, override) + req, err := http.NewRequest(http.MethodPost, "https://cli-chat-proxy.grok.com/v1/responses", nil) + require.NoError(t, err) + + ApplyCLIProxyHeaders(req) + + version := req.Header.Get("x-grok-client-version") + require.True(t, semver.IsValid("v"+version)) + require.GreaterOrEqual(t, semver.Compare("v"+version, "v1.0.13"), 0, + "Grok Responses 会以 426 拒绝版本 %s,最低要求为 1.0.13", version) + require.Contains(t, req.Header.Get("User-Agent"), "grok-shell/"+version+" (") + }) + } +} + +func TestCLIUserAgentMatchesOfficialInteractiveCapture(t *testing.T) { + // 来源:官方 CLI 1.0.46 交互式界面在 Linux x86_64 的主 Responses 请求抓包。 + if runtime.GOOS != "linux" || runtime.GOARCH != "amd64" { + t.Skip("该抓包样本来自 Linux x86_64") + } + require.Equal(t, "grok-pager/1.0.46 grok-shell/1.0.46 (linux; x86_64)", CLIUserAgent("1.0.46")) +} + func TestResolveCLIVersionDefaultsToPinnedClientVersion(t *testing.T) { t.Setenv(CLIVersionEnv, "") // Default advertise pin is CLIClientVersion; CLIStableVersion is only the floor. @@ -16,18 +45,18 @@ func TestResolveCLIVersionDefaultsToPinnedClientVersion(t *testing.T) { } func TestResolveCLIVersionAcceptsValidOverride(t *testing.T) { - t.Setenv(CLIVersionEnv, "0.2.95-alpha.1") - require.Equal(t, "0.2.95-alpha.1", ResolveCLIVersion()) + t.Setenv(CLIVersionEnv, "1.0.14-alpha.1") + require.Equal(t, "1.0.14-alpha.1", ResolveCLIVersion()) } func TestResolveCLIVersionRejectsUnsafeOrTooOld(t *testing.T) { for _, version := range []string{ - "0.2.92", - "0.2.93-beta.1", - "0.2.95\r\nX-Injected: true", - "0.2.093", - "0.3", - "1", + "1.0.12", + "1.0.13-beta.1", + "1.0.14\r\nX-Injected: true", + "1.0.014", + "1.1", + "2", } { t.Run(version, func(t *testing.T) { t.Setenv(CLIVersionEnv, version) @@ -48,11 +77,13 @@ func TestApplyCLIProxyHeaders(t *testing.T) { require.Equal(t, CLIClientVersion, req.Header.Get("x-grok-client-version")) require.Equal(t, CLIClientIdentifier, req.Header.Get("x-grok-client-identifier")) require.Equal(t, CLITokenAuth, req.Header.Get("X-XAI-Token-Auth")) + require.Equal(t, "interactive", req.Header.Get("x-grok-client-mode")) + require.Equal(t, "authenticate-response", req.Header.Get("x-authenticateresponse")) require.Equal(t, CLIUserAgent(CLIClientVersion), req.Header.Get("User-Agent")) } func TestApplyCLIProxyHeadersLeavesAPIHostUnchanged(t *testing.T) { - t.Setenv(CLIVersionEnv, "0.2.95") + t.Setenv(CLIVersionEnv, "1.0.14") req, err := http.NewRequest(http.MethodPost, "https://api.x.ai/v1/responses", nil) require.NoError(t, err) @@ -63,5 +94,7 @@ func TestApplyCLIProxyHeadersLeavesAPIHostUnchanged(t *testing.T) { require.Empty(t, req.Header.Get("x-grok-client-version")) require.Empty(t, req.Header.Get("x-grok-client-identifier")) require.Empty(t, req.Header.Get("X-XAI-Token-Auth")) + require.Empty(t, req.Header.Get("x-grok-client-mode")) + require.Empty(t, req.Header.Get("x-authenticateresponse")) require.Equal(t, "direct-api-client/1.0", req.Header.Get("User-Agent")) } diff --git a/backend/internal/repository/api_key_repo.go b/backend/internal/repository/api_key_repo.go index 189851621..4c55761ff 100644 --- a/backend/internal/repository/api_key_repo.go +++ b/backend/internal/repository/api_key_repo.go @@ -652,6 +652,17 @@ func apiKeyListOrder(params pagination.PaginationParams) []func(*entsql.Selector sortBy := strings.ToLower(strings.TrimSpace(params.SortBy)) sortOrder := params.NormalizedSortOrder(pagination.SortOrderDesc) + if sortBy == "group" { + // Sort before pagination, keeping ungrouped keys last in either direction. + opts := []entsql.OrderTermOption{entsql.OrderNullsLast()} + tieOrder := dbent.Asc(apikey.FieldID) + if sortOrder == pagination.SortOrderDesc { + opts = append(opts, entsql.OrderDesc()) + tieOrder = dbent.Desc(apikey.FieldID) + } + return []func(*entsql.Selector){apikey.ByGroupField(group.FieldName, opts...), tieOrder} + } + var field string switch sortBy { case "name": diff --git a/backend/internal/repository/api_key_repo_sort_test.go b/backend/internal/repository/api_key_repo_sort_test.go new file mode 100644 index 000000000..369b6a590 --- /dev/null +++ b/backend/internal/repository/api_key_repo_sort_test.go @@ -0,0 +1,78 @@ +package repository + +import ( + "context" + "testing" + + dbent "github.com/Wei-Shaw/sub2api/ent" + "github.com/Wei-Shaw/sub2api/internal/pkg/pagination" + "github.com/Wei-Shaw/sub2api/internal/service" + "github.com/stretchr/testify/require" +) + +func TestAPIKeyRepositoryListByUserIDSortByGroup(t *testing.T) { + repo, client := newAPIKeyRepoSQLite(t) + ctx := context.Background() + user := mustCreateAPIKeyRepoUser(t, ctx, client, "group-sort@test.com") + otherUser := mustCreateAPIKeyRepoUser(t, ctx, client, "other-group-sort@test.com") + createGroup := func(name string) *dbent.Group { + g, err := client.Group.Create().SetName(name).Save(ctx) + require.NoError(t, err) + return g + } + // Create groups in reverse name order so sorting by group ID cannot pass. + zulu := createGroup("Zulu") + alpha := createGroup("Alpha") + createKey := func(userID int64, name string, groupID *int64, status string) int64 { + key := &service.APIKey{ + UserID: userID, Key: "sk-" + name, Name: name, GroupID: groupID, Status: status, + } + require.NoError(t, repo.Create(ctx, key)) + return key.ID + } + zuluFirst := createKey(user.ID, "match-zulu-first", &zulu.ID, service.StatusActive) + ungroupedFirst := createKey(user.ID, "match-ungrouped-first", nil, service.StatusActive) + alphaFirst := createKey(user.ID, "match-alpha-first", &alpha.ID, service.StatusActive) + zuluSecond := createKey(user.ID, "match-zulu-second", &zulu.ID, service.StatusActive) + alphaSecond := createKey(user.ID, "match-alpha-second", &alpha.ID, service.StatusActive) + ungroupedSecond := createKey(user.ID, "match-ungrouped-second", nil, service.StatusActive) + createKey(otherUser.ID, "match-other-user", &alpha.ID, service.StatusActive) + createKey(user.ID, "excluded-search", &alpha.ID, service.StatusActive) + createKey(user.ID, "match-inactive", &alpha.ID, service.StatusDisabled) + deleted := createKey(user.ID, "match-deleted", &alpha.ID, service.StatusActive) + require.NoError(t, repo.Delete(ctx, deleted)) + + ungroupedID := int64(0) + for _, tc := range []struct { + name string + order string + groupID *int64 + want []int64 + }{ + {"ascending", "asc", nil, []int64{alphaFirst, alphaSecond, zuluFirst, zuluSecond, ungroupedFirst, ungroupedSecond}}, + {"descending", "desc", nil, []int64{zuluSecond, zuluFirst, alphaSecond, alphaFirst, ungroupedSecond, ungroupedFirst}}, + {"group filter", "asc", &alpha.ID, []int64{alphaFirst, alphaSecond}}, + {"ungrouped filter", "desc", &ungroupedID, []int64{ungroupedSecond, ungroupedFirst}}, + } { + t.Run(tc.name, func(t *testing.T) { + var got []int64 + const pageSize = 3 + pages := (len(tc.want) + pageSize - 1) / pageSize + for page := 1; page <= pages; page++ { + keys, result, err := repo.ListByUserID(ctx, user.ID, pagination.PaginationParams{ + Page: page, PageSize: pageSize, SortBy: "group", SortOrder: tc.order, + }, service.APIKeyListFilters{Search: "match", Status: service.StatusActive, GroupID: tc.groupID}) + require.NoError(t, err) + require.EqualValues(t, len(tc.want), result.Total) + require.Equal(t, pages, result.Pages) + for _, key := range keys { + got = append(got, key.ID) + if key.GroupID != nil { + require.NotNil(t, key.Group, "sorting must preserve group preloading") + } + } + } + require.Equal(t, tc.want, got) + }) + } +} diff --git a/backend/internal/repository/email_cache.go b/backend/internal/repository/email_cache.go index b917824b6..9f59ff05b 100644 --- a/backend/internal/repository/email_cache.go +++ b/backend/internal/repository/email_cache.go @@ -17,8 +17,42 @@ const ( passwordResetKeyPrefix = "password_reset:" passwordResetSentAtKeyPrefix = "password_reset_sent:" notifyCodeUserRateKeyPrefix = "notify_code_user_rate:" + + // attemptsKeySuffix stores the failed-attempt counter next to a verification code. + // Kept in a separate key so it can be incremented atomically with INCR. + attemptsKeySuffix = ":attempts" ) +// incrAttemptsScript atomically increments the attempt counter for an existing +// verification code and aligns the counter TTL with the code TTL. +// KEYS[1] = code key, KEYS[2] = attempts key. Returns -1 when the code is missing. +var incrAttemptsScript = redis.NewScript(` +if redis.call('EXISTS', KEYS[1]) == 0 then + return -1 +end +local n = redis.call('INCR', KEYS[2]) +local ttl = redis.call('PTTL', KEYS[1]) +if ttl > 0 then + redis.call('PEXPIRE', KEYS[2], ttl) +end +return n +`) + +// consumeResetTokenScript atomically compares the stored token hash and deletes it. +// KEYS[1] = reset key, ARGV[1] = expected token hash. Returns 1 on success, 0 otherwise. +var consumeResetTokenScript = redis.NewScript(` +local v = redis.call('GET', KEYS[1]) +if not v then + return 0 +end +local ok, d = pcall(cjson.decode, v) +if not ok or type(d) ~= 'table' or d['Token'] ~= ARGV[1] then + return 0 +end +redis.call('DEL', KEYS[1]) +return 1 +`) + // verifyCodeKey generates the Redis key for email verification code. // Email is lowercased for case-insensitive consistency. func verifyCodeKey(email string) string { @@ -50,8 +84,7 @@ func NewEmailCache(rdb *redis.Client) service.EmailCache { return &emailCache{rdb: rdb} } -func (c *emailCache) GetVerificationCode(ctx context.Context, email string) (*service.VerificationCodeData, error) { - key := verifyCodeKey(email) +func (c *emailCache) getCode(ctx context.Context, key string) (*service.VerificationCodeData, error) { val, err := c.rdb.Get(ctx, key).Result() if err != nil { return nil, err @@ -60,21 +93,56 @@ func (c *emailCache) GetVerificationCode(ctx context.Context, email string) (*se if err := json.Unmarshal([]byte(val), &data); err != nil { return nil, err } + if n, err := c.rdb.Get(ctx, key+attemptsKeySuffix).Int(); err == nil && n > data.Attempts { + data.Attempts = n + } return &data, nil } -func (c *emailCache) SetVerificationCode(ctx context.Context, email string, data *service.VerificationCodeData, ttl time.Duration) error { - key := verifyCodeKey(email) +func (c *emailCache) setCode(ctx context.Context, key string, data *service.VerificationCodeData, ttl time.Duration) error { val, err := json.Marshal(data) if err != nil { return err } - return c.rdb.Set(ctx, key, val, ttl).Err() + pipe := c.rdb.TxPipeline() + pipe.Set(ctx, key, val, ttl) + pipe.Del(ctx, key+attemptsKeySuffix) + if data.Attempts > 0 { + pipe.Set(ctx, key+attemptsKeySuffix, data.Attempts, ttl) + } + _, err = pipe.Exec(ctx) + return err +} + +func (c *emailCache) incrCodeAttempts(ctx context.Context, key string) (int, error) { + n, err := incrAttemptsScript.Run(ctx, c.rdb, []string{key, key + attemptsKeySuffix}).Int() + if err != nil { + return 0, err + } + if n < 0 { + return 0, redis.Nil + } + return n, nil +} + +func (c *emailCache) deleteCode(ctx context.Context, key string) error { + return c.rdb.Del(ctx, key, key+attemptsKeySuffix).Err() +} + +func (c *emailCache) GetVerificationCode(ctx context.Context, email string) (*service.VerificationCodeData, error) { + return c.getCode(ctx, verifyCodeKey(email)) +} + +func (c *emailCache) SetVerificationCode(ctx context.Context, email string, data *service.VerificationCodeData, ttl time.Duration) error { + return c.setCode(ctx, verifyCodeKey(email), data, ttl) +} + +func (c *emailCache) IncrVerificationCodeAttempts(ctx context.Context, email string) (int, error) { + return c.incrCodeAttempts(ctx, verifyCodeKey(email)) } func (c *emailCache) DeleteVerificationCode(ctx context.Context, email string) error { - key := verifyCodeKey(email) - return c.rdb.Del(ctx, key).Err() + return c.deleteCode(ctx, verifyCodeKey(email)) } // Password reset token methods @@ -101,25 +169,21 @@ func (c *emailCache) SetPasswordResetToken(ctx context.Context, email string, da return c.rdb.Set(ctx, key, val, ttl).Err() } +// ConsumePasswordResetToken atomically deletes the stored reset token when its +// stored hash equals tokenHash. Returns true only for the single winning caller. +func (c *emailCache) ConsumePasswordResetToken(ctx context.Context, email, tokenHash string) (bool, error) { + n, err := consumeResetTokenScript.Run(ctx, c.rdb, []string{passwordResetKey(email)}, tokenHash).Int() + if err != nil { + return false, err + } + return n == 1, nil +} + func (c *emailCache) DeletePasswordResetToken(ctx context.Context, email string) error { key := passwordResetKey(email) return c.rdb.Del(ctx, key).Err() } -// 比较与删除在同一 Redis 脚本中执行:并发请求只能成功一次,旧请求不能删除新邮件的令牌。 -var consumePasswordResetTokenScript = redis.NewScript(` -local raw = redis.call('GET', KEYS[1]) -if not raw then return 0 end -local ok, data = pcall(cjson.decode, raw) -if not ok or type(data) ~= 'table' or data.Token ~= ARGV[1] then return 0 end -return redis.call('DEL', KEYS[1]) -`) - -func (c *emailCache) ConsumePasswordResetToken(ctx context.Context, email, token string) (bool, error) { - consumed, err := consumePasswordResetTokenScript.Run(ctx, c.rdb, []string{passwordResetKey(email)}, token).Int() - return consumed == 1, err -} - // Password reset email cooldown methods func (c *emailCache) IsPasswordResetEmailInCooldown(ctx context.Context, email string) bool { @@ -136,30 +200,19 @@ func (c *emailCache) SetPasswordResetEmailCooldown(ctx context.Context, email st // Notify email verification code methods func (c *emailCache) GetNotifyVerifyCode(ctx context.Context, email string) (*service.VerificationCodeData, error) { - key := notifyVerifyKey(email) - val, err := c.rdb.Get(ctx, key).Result() - if err != nil { - return nil, err - } - var data service.VerificationCodeData - if err := json.Unmarshal([]byte(val), &data); err != nil { - return nil, err - } - return &data, nil + return c.getCode(ctx, notifyVerifyKey(email)) } func (c *emailCache) SetNotifyVerifyCode(ctx context.Context, email string, data *service.VerificationCodeData, ttl time.Duration) error { - key := notifyVerifyKey(email) - val, err := json.Marshal(data) - if err != nil { - return err - } - return c.rdb.Set(ctx, key, val, ttl).Err() + return c.setCode(ctx, notifyVerifyKey(email), data, ttl) +} + +func (c *emailCache) IncrNotifyVerifyCodeAttempts(ctx context.Context, email string) (int, error) { + return c.incrCodeAttempts(ctx, notifyVerifyKey(email)) } func (c *emailCache) DeleteNotifyVerifyCode(ctx context.Context, email string) error { - key := notifyVerifyKey(email) - return c.rdb.Del(ctx, key).Err() + return c.deleteCode(ctx, notifyVerifyKey(email)) } // User-level rate limiting for notify email verification codes diff --git a/backend/internal/repository/email_cache_atomic_test.go b/backend/internal/repository/email_cache_atomic_test.go new file mode 100644 index 000000000..0086bbe46 --- /dev/null +++ b/backend/internal/repository/email_cache_atomic_test.go @@ -0,0 +1,150 @@ +package repository + +import ( + "context" + "crypto/sha256" + "encoding/hex" + "errors" + "strings" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/Wei-Shaw/sub2api/internal/service" + "github.com/alicebob/miniredis/v2" + "github.com/redis/go-redis/v9" + "github.com/stretchr/testify/require" +) + +func newMiniredisEmailCache(t *testing.T) (service.EmailCache, *miniredis.Miniredis, *redis.Client) { + t.Helper() + mr := miniredis.RunT(t) + rdb := redis.NewClient(&redis.Options{Addr: mr.Addr()}) + t.Cleanup(func() { _ = rdb.Close() }) + return NewEmailCache(rdb), mr, rdb +} + +func TestEmailCache_ConcurrentWrongCodesCannotExceedAttemptCap(t *testing.T) { + cache, _, _ := newMiniredisEmailCache(t) + ctx := context.Background() + email := "user@example.com" + + svc := service.NewEmailService(nil, cache) + require.NoError(t, cache.SetVerificationCode(ctx, email, &service.VerificationCodeData{ + Code: "123456", + CreatedAt: time.Now(), + ExpiresAt: time.Now().Add(15 * time.Minute), + }, 15*time.Minute)) + + const workers = 50 + var invalid, maxed atomic.Int32 + var wg sync.WaitGroup + for i := 0; i < workers; i++ { + wg.Add(1) + go func() { + defer wg.Done() + err := svc.VerifyCode(ctx, email, "000000") + switch { + case errors.Is(err, service.ErrInvalidVerifyCode): + invalid.Add(1) + case errors.Is(err, service.ErrVerifyCodeMaxAttempts): + maxed.Add(1) + default: + t.Errorf("unexpected result: %v", err) + } + }() + } + wg.Wait() + + // Only attempts 1..4 may return "invalid"; every other guess is rejected by the cap. + require.LessOrEqual(t, int(invalid.Load()), 4) + require.Equal(t, workers, int(invalid.Load()+maxed.Load())) + + // Even the correct code is now rejected. + require.ErrorIs(t, svc.VerifyCode(ctx, email, "123456"), service.ErrVerifyCodeMaxAttempts) + + data, err := cache.GetVerificationCode(ctx, email) + require.NoError(t, err) + require.GreaterOrEqual(t, data.Attempts, 5) +} + +func TestEmailCache_AttemptsResetOnNewCodeAndTTLFollowsCode(t *testing.T) { + cache, mr, _ := newMiniredisEmailCache(t) + ctx := context.Background() + email := "User@Example.com" + + require.NoError(t, cache.SetVerificationCode(ctx, email, &service.VerificationCodeData{Code: "1"}, time.Minute)) + n, err := cache.IncrVerificationCodeAttempts(ctx, email) + require.NoError(t, err) + require.Equal(t, 1, n) + require.Greater(t, mr.TTL(verifyCodeKey(email)+attemptsKeySuffix), time.Duration(0)) + + require.NoError(t, cache.SetVerificationCode(ctx, email, &service.VerificationCodeData{Code: "2"}, time.Minute)) + data, err := cache.GetVerificationCode(ctx, email) + require.NoError(t, err) + require.Equal(t, 0, data.Attempts) + + require.NoError(t, cache.DeleteVerificationCode(ctx, email)) + _, err = cache.IncrVerificationCodeAttempts(ctx, email) + require.Error(t, err) + require.False(t, mr.Exists(verifyCodeKey(email)+attemptsKeySuffix)) +} + +func TestEmailCache_PasswordResetTokenHashedAndSingleUse(t *testing.T) { + cache, mr, _ := newMiniredisEmailCache(t) + ctx := context.Background() + email := "reset@example.com" + + svc := service.NewEmailService(nil, cache) + + // Seed the token the same way SendPasswordResetEmail does (hash only). + token, err := svc.GeneratePasswordResetToken() + require.NoError(t, err) + sum := sha256.Sum256([]byte(token)) + require.NoError(t, cache.SetPasswordResetToken(ctx, email, &service.PasswordResetTokenData{ + Token: hex.EncodeToString(sum[:]), CreatedAt: time.Now(), + }, 30*time.Minute)) + + raw, err := mr.Get(passwordResetKey(email)) + require.NoError(t, err) + require.False(t, strings.Contains(raw, token), "plaintext token must not be stored") + + require.NoError(t, svc.VerifyPasswordResetToken(ctx, email, token)) + require.ErrorIs(t, svc.ConsumePasswordResetToken(ctx, email, "wrong"), service.ErrInvalidResetToken) + + const workers = 30 + var ok atomic.Int32 + var wg sync.WaitGroup + for i := 0; i < workers; i++ { + wg.Add(1) + go func() { + defer wg.Done() + if svc.ConsumePasswordResetToken(ctx, email, token) == nil { + ok.Add(1) + } + }() + } + wg.Wait() + require.Equal(t, int32(1), ok.Load()) + require.False(t, mr.Exists(passwordResetKey(email))) +} + +func TestEmailCache_ConsumePasswordResetTokenMismatchKeepsToken(t *testing.T) { + cache, mr, _ := newMiniredisEmailCache(t) + ctx := context.Background() + email := "keep@example.com" + require.NoError(t, cache.SetPasswordResetToken(ctx, email, &service.PasswordResetTokenData{Token: "abc"}, time.Minute)) + + ok, err := cache.ConsumePasswordResetToken(ctx, email, "xyz") + require.NoError(t, err) + require.False(t, ok) + require.True(t, mr.Exists(passwordResetKey(email))) + + ok, err = cache.ConsumePasswordResetToken(ctx, email, "abc") + require.NoError(t, err) + require.True(t, ok) + ok, err = cache.ConsumePasswordResetToken(ctx, email, "abc") + require.NoError(t, err) + require.False(t, ok) +} diff --git a/backend/internal/repository/email_cache_password_reset_test.go b/backend/internal/repository/email_cache_password_reset_test.go index 1408a82b4..9607f968a 100644 --- a/backend/internal/repository/email_cache_password_reset_test.go +++ b/backend/internal/repository/email_cache_password_reset_test.go @@ -2,6 +2,8 @@ package repository import ( "context" + "crypto/sha256" + "encoding/hex" "errors" "sync" "sync/atomic" @@ -31,7 +33,7 @@ func TestPasswordResetTokenConsumptionFailureDoesNotSucceed(t *testing.T) { t.Cleanup(func() { _ = client.Close() }) cache := &failedResetConsumeCache{EmailCache: NewEmailCache(client)} ctx := context.Background() - require.NoError(t, cache.SetPasswordResetToken(ctx, "test@example.com", &service.PasswordResetTokenData{Token: "valid"}, time.Minute)) + require.NoError(t, cache.SetPasswordResetToken(ctx, "test@example.com", &service.PasswordResetTokenData{Token: resetTokenHash("valid")}, time.Minute)) svc := service.NewEmailService(nil, cache) require.ErrorIs(t, svc.ConsumePasswordResetToken(ctx, "test@example.com", "valid"), service.ErrServiceUnavailable) } @@ -80,7 +82,7 @@ func TestPasswordResetTokenCanOnlyBeConsumedOnceConcurrently(t *testing.T) { cache := &simultaneousResetReadCache{EmailCache: NewEmailCache(client)} ctx := context.Background() require.NoError(t, cache.SetPasswordResetToken(ctx, "test@example.com", &service.PasswordResetTokenData{ - Token: "one-time-token", CreatedAt: time.Now(), + Token: resetTokenHash("one-time-token"), CreatedAt: time.Now(), }, time.Minute)) svc := service.NewEmailService(nil, cache) const callers = 8 @@ -107,3 +109,8 @@ func TestPasswordResetTokenCanOnlyBeConsumedOnceConcurrently(t *testing.T) { require.ErrorIs(t, err, service.ErrInvalidResetToken) } } + +func resetTokenHash(token string) string { + sum := sha256.Sum256([]byte(token)) + return hex.EncodeToString(sum[:]) +} diff --git a/backend/internal/repository/http_upstream.go b/backend/internal/repository/http_upstream.go index b27f20ece..ff19b03d8 100644 --- a/backend/internal/repository/http_upstream.go +++ b/backend/internal/repository/http_upstream.go @@ -466,6 +466,8 @@ func newGrokOfficialAPIFallbackRequest(req *http.Request) (*http.Request, error) for _, header := range []string{ "X-XAI-Token-Auth", "X-Grok-Client-Version", + "X-Grok-Client-Mode", + "X-Authenticateresponse", "X-Grok-Client-Surface", "X-UserID", "X-Email", @@ -531,6 +533,8 @@ func applyGrokCLIProxyHeaders(req *http.Request) { req.Header.Set("X-XAI-Token-Auth", xai.CLITokenAuth) req.Header.Set("x-grok-client-version", version) req.Header.Set("x-grok-client-identifier", xai.CLIClientIdentifier) + req.Header.Set("x-grok-client-mode", xai.CLIClientMode) + req.Header.Set("x-authenticateresponse", "authenticate-response") req.Header.Set("User-Agent", xai.CLIUserAgent(version)) } diff --git a/backend/internal/repository/http_upstream_test.go b/backend/internal/repository/http_upstream_test.go index 49881ae51..dba89e7d1 100644 --- a/backend/internal/repository/http_upstream_test.go +++ b/backend/internal/repository/http_upstream_test.go @@ -237,6 +237,10 @@ func TestHTTPUpstreamDoAppliesGrokCLIIdentityBeforeOAuthRoundTrip(t *testing.T) require.NoError(t, resp.Body.Close()) require.Equal(t, xai.CLIClientVersion, capturedHeaders.Get("x-grok-client-version")) + require.Equal(t, "1.0.46", capturedHeaders.Get("x-grok-client-version")) + require.Equal(t, "grok-pager", capturedHeaders.Get("x-grok-client-identifier")) + require.Equal(t, "interactive", capturedHeaders.Get("x-grok-client-mode")) + require.Equal(t, "authenticate-response", capturedHeaders.Get("x-authenticateresponse")) require.Equal(t, "xai-grok-cli", capturedHeaders.Get("X-XAI-Token-Auth")) require.Equal(t, xai.CLIUserAgent(xai.CLIClientVersion), capturedHeaders.Get("User-Agent")) }) @@ -309,6 +313,8 @@ func TestHTTPUpstreamDoFallsBackToOfficialGrokAPIOnCLIAccessDenied(t *testing.T) require.Equal(t, "Bearer oauth-token", fallbackHeaders.Get("Authorization")) require.Empty(t, fallbackHeaders.Get("X-XAI-Token-Auth")) require.Empty(t, fallbackHeaders.Get("x-grok-client-version")) + require.Empty(t, fallbackHeaders.Get("x-grok-client-mode")) + require.Empty(t, fallbackHeaders.Get("x-authenticateresponse")) require.Empty(t, fallbackHeaders.Get("User-Agent")) } @@ -465,18 +471,18 @@ func TestApplyGrokCLIProxyHeaders(t *testing.T) { }) t.Run("accepts a valid operator override", func(t *testing.T) { - t.Setenv("XAI_GROK_CLI_VERSION", "0.2.121-alpha.1") + t.Setenv("XAI_GROK_CLI_VERSION", "1.0.47-alpha.1") req, err := http.NewRequest(http.MethodPost, "https://cli-chat-proxy.grok.com/v1/chat/completions", nil) require.NoError(t, err) applyGrokCLIProxyHeaders(req) - require.Equal(t, "0.2.121-alpha.1", req.Header.Get("x-grok-client-version")) - require.Equal(t, xai.CLIUserAgent("0.2.121-alpha.1"), req.Header.Get("User-Agent")) + require.Equal(t, "1.0.47-alpha.1", req.Header.Get("x-grok-client-version")) + require.Equal(t, xai.CLIUserAgent("1.0.47-alpha.1"), req.Header.Get("User-Agent")) }) t.Run("rejects an unsafe override", func(t *testing.T) { - t.Setenv("XAI_GROK_CLI_VERSION", "0.2.121\r\nX-Injected: true") + t.Setenv("XAI_GROK_CLI_VERSION", "1.0.47\r\nX-Injected: true") req, err := http.NewRequest(http.MethodPost, "https://cli-chat-proxy.grok.com/v1/responses", nil) require.NoError(t, err) @@ -487,7 +493,7 @@ func TestApplyGrokCLIProxyHeaders(t *testing.T) { }) t.Run("rejects an override below the supported minimum", func(t *testing.T) { - t.Setenv("XAI_GROK_CLI_VERSION", "0.2.119") + t.Setenv("XAI_GROK_CLI_VERSION", "1.0.45") req, err := http.NewRequest(http.MethodPost, "https://cli-chat-proxy.grok.com/v1/responses", nil) require.NoError(t, err) @@ -498,7 +504,7 @@ func TestApplyGrokCLIProxyHeaders(t *testing.T) { }) t.Run("rejects a prerelease override at the minimum version", func(t *testing.T) { - t.Setenv("XAI_GROK_CLI_VERSION", "0.2.120-beta.1") + t.Setenv("XAI_GROK_CLI_VERSION", "1.0.46-beta.1") req, err := http.NewRequest(http.MethodPost, "https://cli-chat-proxy.grok.com/v1/responses", nil) require.NoError(t, err) @@ -511,11 +517,11 @@ func TestApplyGrokCLIProxyHeaders(t *testing.T) { // Every entry sits above the pinned minimum, so a rejection here can only be // caused by the malformed semver and never by the version being too old. for _, version := range []string{ - "0.2.0121", - "0.2.121-alpha..1", - "0.3", - "1", - "0.2.121+build.1", + "1.0.047", + "1.0.47-alpha..1", + "1.1", + "2", + "1.0.47+build.1", } { t.Run("rejects invalid semver "+version, func(t *testing.T) { t.Setenv("XAI_GROK_CLI_VERSION", version) @@ -530,7 +536,7 @@ func TestApplyGrokCLIProxyHeaders(t *testing.T) { } t.Run("leaves direct xAI API requests unchanged", func(t *testing.T) { - t.Setenv("XAI_GROK_CLI_VERSION", "0.2.95") + t.Setenv("XAI_GROK_CLI_VERSION", "1.0.47") req, err := http.NewRequest(http.MethodPost, "https://api.x.ai/v1/responses", nil) require.NoError(t, err) req.Header.Set("User-Agent", "direct-api-client/1.0") diff --git a/backend/internal/securityaudit/prompt_snapshot.go b/backend/internal/securityaudit/prompt_snapshot.go index b47e84674..74a0d2597 100644 --- a/backend/internal/securityaudit/prompt_snapshot.go +++ b/backend/internal/securityaudit/prompt_snapshot.go @@ -106,6 +106,8 @@ func extractProtocolSegments(protocol string, document any) []promptSegment { return append(extractInstructions(root["instructions"]), extractResponses(root["input"])...) case "openai_images", "grok_media", "media", "images": return userPromptSegments(extractMediaPrompts(root)) + case "typesafe_systemone": + return extractSystemOneSegments(root) default: if segments := extractChatLikeSegments(root); len(segments) > 0 { return segments @@ -120,6 +122,79 @@ func extractProtocolSegments(protocol string, document any) []promptSegment { } } +// extractSystemOneSegments collects every client-controlled text of a TypeSafe +// System One request: question IDs, every question field except the validated +// type, unknown top-level extension fields, and the evaluated state. Object keys +// are text too (Jev reads the whole JSON), so they are collected with values. +// The state comes last so it is the prioritized segment, and keys are visited +// in sorted order to keep the prompt hash stable across requests. +func extractSystemOneSegments(root map[string]any) []promptSegment { + if root == nil { + return nil + } + texts := make([]string, 0, 4) + questions, isObject := root["questions"].(map[string]any) + if !isObject { + texts = appendJSONStringLeaves(texts, root["questions"]) + } + for _, id := range sortedJSONKeys(questions) { + texts = appendJSONStringLeaves(texts, id) + question, ok := questions[id].(map[string]any) + if !ok { + texts = appendJSONStringLeaves(texts, questions[id]) + continue + } + for _, field := range sortedJSONKeys(question) { + switch field { + case "type": + continue + case "instructions", "criteria": + default: + texts = appendJSONStringLeaves(texts, field) + } + texts = appendJSONStringLeaves(texts, question[field]) + } + } + for _, field := range sortedJSONKeys(root) { + switch field { + case "model", "stream", "state", "questions": + continue + } + texts = appendJSONStringLeaves(texts, field) + texts = appendJSONStringLeaves(texts, root[field]) + } + texts = appendJSONStringLeaves(texts, root["state"]) + return userPromptSegments(texts) +} + +func appendJSONStringLeaves(texts []string, value any) []string { + switch typed := value.(type) { + case string: + if text := strings.TrimSpace(typed); text != "" { + texts = append(texts, text) + } + case []any: + for _, item := range typed { + texts = appendJSONStringLeaves(texts, item) + } + case map[string]any: + for _, key := range sortedJSONKeys(typed) { + texts = appendJSONStringLeaves(texts, key) + texts = appendJSONStringLeaves(texts, typed[key]) + } + } + return texts +} + +func sortedJSONKeys(values map[string]any) []string { + keys := make([]string, 0, len(values)) + for key := range values { + keys = append(keys, key) + } + sort.Strings(keys) + return keys +} + // clientInstructionRoles are roles a client may freely populate. Attackers can // place jailbreak/PII text in assistant/tool turns, so blocking audit must scan // them too—not only user/system/developer instructions. diff --git a/backend/internal/securityaudit/prompt_snapshot_test.go b/backend/internal/securityaudit/prompt_snapshot_test.go index 70d26f925..751043f55 100644 --- a/backend/internal/securityaudit/prompt_snapshot_test.go +++ b/backend/internal/securityaudit/prompt_snapshot_test.go @@ -412,3 +412,53 @@ func mustJSON(t *testing.T, value string) []byte { func metadataTextForTest(scanText string) string { return strings.Replace(scanText, promptAuditPrioritySeparator, "\n\n", 1) } + +func TestPromptSnapshotTypeSafeSystemOneCollectsStateAndQuestions(t *testing.T) { + body := `{"model":"jev-latest","state":{"STATE_KEY":"STATE_TITLE","items":["STATE_ITEM",{"text":"STATE_NESTED"},3]},` + + `"questions":{"b":{"type":"choice","instructions":"CHOICE_INSTRUCTIONS","criteria":{"OPTION_LABEL":{"description":"OPTION_DESC"},"empty":null}},` + + `"a":{"type":"noul","instructions":["NOUL_INSTRUCTIONS"],"criteria":{"true":"NOUL_TRUE","false":"NOUL_FALSE"},"EXTENSION_KEY":"EXTENSION_VALUE"},` + + `"QUESTION_ID":{"type":"score","criteria":["SCORE_LOW",{"description":"SCORE_HIGH"}]}},"TOP_EXTENSION":{"x":"TOP_VALUE"}}` + + snapshot, err := ExtractPromptSnapshot(Request{Protocol: "typesafe_systemone", Body: []byte(body)}) + require.NoError(t, err) + for _, text := range []string{"STATE_KEY", "STATE_TITLE", "STATE_ITEM", "STATE_NESTED", "CHOICE_INSTRUCTIONS", "OPTION_LABEL", "OPTION_DESC", + "NOUL_INSTRUCTIONS", "NOUL_TRUE", "NOUL_FALSE", "EXTENSION_KEY", "EXTENSION_VALUE", "QUESTION_ID", "SCORE_LOW", "SCORE_HIGH", + "TOP_EXTENSION", "TOP_VALUE"} { + require.Contains(t, snapshot.ScanText, text) + } + // Canonical field names and validated enum values are not client text. + for _, text := range []string{"jev-latest", "instructions", "criteria", "questions", "choice", "score"} { + require.NotContains(t, snapshot.ScanText, text) + } + + // Map iteration order must not change the audited text or its hash. + for range 20 { + again, err := ExtractPromptSnapshot(Request{Protocol: "typesafe_systemone", Body: []byte(body)}) + require.NoError(t, err) + require.Equal(t, snapshot.PromptHash, again.PromptHash) + require.Equal(t, snapshot.ScanText, again.ScanText) + } + + blocking, err := ExtractBlockingPromptSnapshot(Request{Protocol: "typesafe_systemone", Body: []byte(body)}, true) + require.NoError(t, err) + require.Contains(t, blocking.ScanText, "STATE_TITLE") + require.Contains(t, blocking.ScanText, "CHOICE_INSTRUCTIONS") + require.Contains(t, blocking.ScanText, "QUESTION_ID") +} + +func TestPromptSnapshotTypeSafeSystemOneStringStateIsPrioritized(t *testing.T) { + snapshot, err := ExtractPromptSnapshot(Request{Protocol: "typesafe_systemone", Body: []byte(`{"model":"jev-latest","state":"plain state","questions":{"q":{"type":"noul"}}}`)}) + require.NoError(t, err) + require.True(t, strings.HasPrefix(snapshot.ScanText, "plain state")) + require.Equal(t, 2, snapshot.MessageCount) +} + +func TestPromptSnapshotTypeSafeSystemOneAuditsKeyOnlyPayloads(t *testing.T) { + body := `{"model":"jev-latest","state":{"HIDDEN_STATE_KEY":1},"questions":{"HIDDEN_QUESTION_ID":{"type":"noul"}}}` + for _, latestTurnOnly := range []bool{false, true} { + snapshot, err := ExtractBlockingPromptSnapshot(Request{Protocol: "typesafe_systemone", Body: []byte(body)}, latestTurnOnly) + require.NoError(t, err) + require.Contains(t, snapshot.ScanText, "HIDDEN_STATE_KEY") + require.Contains(t, snapshot.ScanText, "HIDDEN_QUESTION_ID") + } +} diff --git a/backend/internal/server/api_contract_test.go b/backend/internal/server/api_contract_test.go index acaadcd94..68d98a951 100644 --- a/backend/internal/server/api_contract_test.go +++ b/backend/internal/server/api_contract_test.go @@ -873,7 +873,7 @@ func TestAPIContracts(t *testing.T) { "force_email_on_third_party_signup": false, "default_concurrency": 5, "default_balance": 1.25, - "default_platform_quotas": {"anthropic":{"daily":null,"weekly":null,"monthly":null},"antigravity":{"daily":null,"weekly":null,"monthly":null},"deepseek":{"daily":null,"weekly":null,"monthly":null},"gemini":{"daily":null,"weekly":null,"monthly":null},"grok":{"daily":null,"weekly":null,"monthly":null},"kimi":{"daily":null,"weekly":null,"monthly":null},"minimax":{"daily":null,"weekly":null,"monthly":null},"openai":{"daily":null,"weekly":null,"monthly":null},"opencode_go":{"daily":null,"weekly":null,"monthly":null},"zhipu":{"daily":null,"weekly":null,"monthly":null}}, + "default_platform_quotas": {"anthropic":{"daily":null,"weekly":null,"monthly":null},"antigravity":{"daily":null,"weekly":null,"monthly":null},"deepseek":{"daily":null,"weekly":null,"monthly":null},"gemini":{"daily":null,"weekly":null,"monthly":null},"grok":{"daily":null,"weekly":null,"monthly":null},"kimi":{"daily":null,"weekly":null,"monthly":null},"minimax":{"daily":null,"weekly":null,"monthly":null},"openai":{"daily":null,"weekly":null,"monthly":null},"opencode_go":{"daily":null,"weekly":null,"monthly":null},"typesafe":{"daily":null,"weekly":null,"monthly":null},"zhipu":{"daily":null,"weekly":null,"monthly":null}}, "auth_source_default_email_platform_quotas": null, "auth_source_default_github_platform_quotas": null, "auth_source_default_google_platform_quotas": null, @@ -982,6 +982,9 @@ func TestAPIContracts(t *testing.T) { "payment_balance_recharge_multiplier": 0, "payment_subscription_usd_to_cny_rate": 0, "payment_recharge_fee_rate": 0, + "payment_recharge_bonus_tiers": [], + "payment_recharge_bonus_mode": "bonus", + "payment_recharge_bonus_notice": "", "payment_load_balance_strategy": "", "payment_product_name_prefix": "", "payment_product_name_suffix": "", @@ -1207,7 +1210,7 @@ func TestAPIContracts(t *testing.T) { "purchase_subscription_url": "", "table_default_page_size": 20, "table_page_size_options": [10, 20, 50], - "default_platform_quotas": {"anthropic":{"daily":null,"weekly":null,"monthly":null},"antigravity":{"daily":null,"weekly":null,"monthly":null},"deepseek":{"daily":null,"weekly":null,"monthly":null},"gemini":{"daily":null,"weekly":null,"monthly":null},"grok":{"daily":null,"weekly":null,"monthly":null},"kimi":{"daily":null,"weekly":null,"monthly":null},"minimax":{"daily":null,"weekly":null,"monthly":null},"openai":{"daily":null,"weekly":null,"monthly":null},"opencode_go":{"daily":null,"weekly":null,"monthly":null},"zhipu":{"daily":null,"weekly":null,"monthly":null}}, + "default_platform_quotas": {"anthropic":{"daily":null,"weekly":null,"monthly":null},"antigravity":{"daily":null,"weekly":null,"monthly":null},"deepseek":{"daily":null,"weekly":null,"monthly":null},"gemini":{"daily":null,"weekly":null,"monthly":null},"grok":{"daily":null,"weekly":null,"monthly":null},"kimi":{"daily":null,"weekly":null,"monthly":null},"minimax":{"daily":null,"weekly":null,"monthly":null},"openai":{"daily":null,"weekly":null,"monthly":null},"opencode_go":{"daily":null,"weekly":null,"monthly":null},"typesafe":{"daily":null,"weekly":null,"monthly":null},"zhipu":{"daily":null,"weekly":null,"monthly":null}}, "auth_source_default_email_platform_quotas": null, "auth_source_default_github_platform_quotas": null, "auth_source_default_google_platform_quotas": null, @@ -1313,6 +1316,9 @@ func TestAPIContracts(t *testing.T) { "payment_balance_recharge_multiplier": 0, "payment_subscription_usd_to_cny_rate": 0, "payment_recharge_fee_rate": 0, + "payment_recharge_bonus_tiers": [], + "payment_recharge_bonus_mode": "bonus", + "payment_recharge_bonus_notice": "", "payment_load_balance_strategy": "", "payment_product_name_prefix": "", "payment_product_name_suffix": "", diff --git a/backend/internal/server/router.go b/backend/internal/server/router.go index 52333e33c..0e81c8c95 100644 --- a/backend/internal/server/router.go +++ b/backend/internal/server/router.go @@ -131,7 +131,7 @@ func registerRoutes( routes.RegisterModelPlazaRoutes(v1, h, optionalJWTAuth, settingService, panelRateLimiter) routes.RegisterAdminRoutes(v1, h, adminAuth, auditLog, stepUpAuth, settingService, panelRateLimiter) routes.RegisterGatewayRoutes(r, h, apiKeyAuth, apiKeyService, subscriptionService, opsService, settingService, compositeResolver, cfg) - routes.RegisterPaymentRoutes(v1, h.Payment, h.PaymentWebhook, h.Admin.Payment, jwtAuth, adminAuth, auditLog, settingService, panelRateLimiter) + routes.RegisterPaymentRoutes(v1, h.Payment, h.PaymentWebhook, h.Admin.Payment, jwtAuth, adminAuth, auditLog, settingService, panelRateLimiter, redisClient) handler.RegisterPageRoutes(v1, cfg.Pricing.DataDir, gin.HandlerFunc(jwtAuth), gin.HandlerFunc(adminAuth), settingService) } diff --git a/backend/internal/server/routes/gateway.go b/backend/internal/server/routes/gateway.go index 83782d60c..eb6ff6188 100644 --- a/backend/internal/server/routes/gateway.go +++ b/backend/internal/server/routes/gateway.go @@ -204,6 +204,8 @@ func RegisterGatewayRoutes( } h.Gateway.Messages(c) }) + // System One carries only JSON text, so it uses the text body limit. + gateway.POST("/systemone", textBodyLimit, h.Gateway.SystemOne) // /v1/messages/count_tokens: OpenAI bridges upstream, Grok estimates // locally, and Anthropic-compatible platforms retain their existing path. gateway.POST("/messages/count_tokens", countTokensHandler) diff --git a/backend/internal/server/routes/gateway_model_allowlist_test.go b/backend/internal/server/routes/gateway_model_allowlist_test.go index e1585dbd2..a75aa944d 100644 --- a/backend/internal/server/routes/gateway_model_allowlist_test.go +++ b/backend/internal/server/routes/gateway_model_allowlist_test.go @@ -152,6 +152,7 @@ func TestGatewayRoutesGroupModelAllowlistCoversRootAliasRoutes(t *testing.T) { {http.MethodGet, "/realtime?model=gpt-4.1", ""}, {http.MethodPost, "/v1/responses", `{"model":"gpt-4.1"}`}, {http.MethodPost, "/v1/messages", `{"model":"gpt-4.1"}`}, + {http.MethodPost, "/v1/systemone", `{"model":"gpt-4.1","state":"x","questions":{"q":{"type":"noul","instructions":"x"}}}`}, {http.MethodPost, "/v1/messages/count_tokens", `{"model":"gpt-4.1","messages":[]}`}, {http.MethodPost, "/v1/chat/completions", `{"model":"gpt-4.1"}`}, {http.MethodPost, "/v1/embeddings", `{"model":"gpt-4.1","input":"hi"}`}, diff --git a/backend/internal/server/routes/payment.go b/backend/internal/server/routes/payment.go index ecda25f53..26d0fcd66 100644 --- a/backend/internal/server/routes/payment.go +++ b/backend/internal/server/routes/payment.go @@ -1,12 +1,25 @@ package routes import ( + "time" + "github.com/Wei-Shaw/sub2api/internal/handler" "github.com/Wei-Shaw/sub2api/internal/handler/admin" + ratelimit "github.com/Wei-Shaw/sub2api/internal/middleware" "github.com/Wei-Shaw/sub2api/internal/server/middleware" "github.com/Wei-Shaw/sub2api/internal/service" "github.com/gin-gonic/gin" + "github.com/redis/go-redis/v9" +) + +// publicOrderVerifyRateLimit caps anonymous legacy out_trade_no lookups per +// client IP. The payment result page polls at most a handful of times per +// order, so this leaves ample headroom for real users while making +// out_trade_no enumeration impractical. +const ( + publicOrderVerifyRateLimit = 20 + publicOrderVerifyRateLimitWindow = time.Minute ) // RegisterPaymentRoutes registers all payment-related routes: @@ -21,6 +34,7 @@ func RegisterPaymentRoutes( auditLog middleware.AuditLogMiddleware, settingService *service.SettingService, panelRateLimiter *middleware.PanelRateLimiter, + redisClient *redis.Client, ) { // --- User-facing payment endpoints (authenticated) --- authenticated := v1.Group("/payment") @@ -50,9 +64,15 @@ func RegisterPaymentRoutes( // Signed resume-token recovery is the preferred public lookup path. // The legacy anonymous out_trade_no verify endpoint remains available as a // persisted-state compatibility path for staggered upgrades. + // The anonymous verify endpoint is IP rate-limited to prevent out_trade_no + // enumeration. It fails open on Redis errors so an outage never blocks a + // user who is mid-payment from seeing their result. + publicRateLimiter := ratelimit.NewRateLimiter(redisClient) public := v1.Group("/payment/public") { - public.POST("/orders/verify", paymentHandler.VerifyOrderPublic) + public.POST("/orders/verify", + publicRateLimiter.Limit("payment-public-order-verify", publicOrderVerifyRateLimit, publicOrderVerifyRateLimitWindow), + paymentHandler.VerifyOrderPublic) public.POST("/orders/resolve", paymentHandler.ResolveOrderPublicByResumeToken) } diff --git a/backend/internal/server/routes/payment_public_rate_limit_test.go b/backend/internal/server/routes/payment_public_rate_limit_test.go new file mode 100644 index 000000000..3f8c86eab --- /dev/null +++ b/backend/internal/server/routes/payment_public_rate_limit_test.go @@ -0,0 +1,84 @@ +package routes + +import ( + "net/http" + "net/http/httptest" + "strings" + "testing" + "time" + + "github.com/Wei-Shaw/sub2api/internal/handler" + "github.com/Wei-Shaw/sub2api/internal/handler/admin" + servermiddleware "github.com/Wei-Shaw/sub2api/internal/server/middleware" + "github.com/alicebob/miniredis/v2" + "github.com/gin-gonic/gin" + "github.com/redis/go-redis/v9" + "github.com/stretchr/testify/require" +) + +func newPaymentRoutesTestRouter(redisClient *redis.Client) *gin.Engine { + gin.SetMode(gin.TestMode) + router := gin.New() + v1 := router.Group("/api/v1") + noop := func(c *gin.Context) { c.Next() } + + RegisterPaymentRoutes( + v1, + &handler.PaymentHandler{}, + &handler.PaymentWebhookHandler{}, + &admin.PaymentHandler{}, + servermiddleware.JWTAuthMiddleware(noop), + servermiddleware.AdminAuthMiddleware(noop), + servermiddleware.AuditLogMiddleware(noop), + nil, + nil, + redisClient, + ) + return router +} + +func postPublicOrderVerify(router *gin.Engine, remoteAddr string) *httptest.ResponseRecorder { + // Empty body fails binding in the handler, so no service is touched and a + // non-429 response proves the request passed the limiter. + req := httptest.NewRequest(http.MethodPost, "/api/v1/payment/public/orders/verify", strings.NewReader(`{}`)) + req.Header.Set("Content-Type", "application/json") + req.RemoteAddr = remoteAddr + w := httptest.NewRecorder() + router.ServeHTTP(w, req) + return w +} + +func TestPublicOrderVerifyRateLimitedPerIP(t *testing.T) { + mr := miniredis.RunT(t) + rdb := redis.NewClient(&redis.Options{Addr: mr.Addr()}) + t.Cleanup(func() { _ = rdb.Close() }) + + router := newPaymentRoutesTestRouter(rdb) + + for i := 1; i <= publicOrderVerifyRateLimit; i++ { + w := postPublicOrderVerify(router, "198.51.100.20:1234") + require.Equal(t, http.StatusBadRequest, w.Code, "request %d should reach the handler", i) + } + + w := postPublicOrderVerify(router, "198.51.100.20:1234") + require.Equal(t, http.StatusTooManyRequests, w.Code) + require.Contains(t, w.Body.String(), "rate limit exceeded") + + // A different client IP has its own budget. + w = postPublicOrderVerify(router, "198.51.100.21:1234") + require.Equal(t, http.StatusBadRequest, w.Code) +} + +func TestPublicOrderVerifyRateLimitFailsOpenWhenRedisUnavailable(t *testing.T) { + rdb := redis.NewClient(&redis.Options{ + Addr: "127.0.0.1:1", + DialTimeout: 50 * time.Millisecond, + ReadTimeout: 50 * time.Millisecond, + WriteTimeout: 50 * time.Millisecond, + }) + t.Cleanup(func() { _ = rdb.Close() }) + + router := newPaymentRoutesTestRouter(rdb) + w := postPublicOrderVerify(router, "203.0.113.30:1234") + require.Equal(t, http.StatusBadRequest, w.Code, "users mid-payment must not be blocked by a Redis outage") +} diff --git a/backend/internal/server/routes/prompt_audit_route_coverage_test.go b/backend/internal/server/routes/prompt_audit_route_coverage_test.go index af4272cb8..2d82f7c1d 100644 --- a/backend/internal/server/routes/prompt_audit_route_coverage_test.go +++ b/backend/internal/server/routes/prompt_audit_route_coverage_test.go @@ -29,6 +29,7 @@ func TestEveryGatewayPOSTRouteIsClassifiedForPromptAuditCoverage(t *testing.T) { audited := map[string][]string{ "/messages": {"gateway_handler.go", "openai_gateway_handler.go"}, + "/systemone": {"gateway_systemone.go"}, "/responses": {"gateway_handler_responses.go", "openai_gateway_handler.go"}, "/responses/*subpath": {"gateway_handler_responses.go", "openai_gateway_handler.go"}, "/chat/completions": {"gateway_handler_chat_completions.go", "openai_chat_completions.go"}, diff --git a/backend/internal/service/account.go b/backend/internal/service/account.go index 28dc28795..5262134a2 100644 --- a/backend/internal/service/account.go +++ b/backend/internal/service/account.go @@ -17,6 +17,7 @@ import ( "github.com/Wei-Shaw/sub2api/internal/domain" "github.com/Wei-Shaw/sub2api/internal/pkg/geminicli" "github.com/Wei-Shaw/sub2api/internal/pkg/openai_compat" + "github.com/Wei-Shaw/sub2api/internal/pkg/typesafe" "github.com/Wei-Shaw/sub2api/internal/pkg/xai" ) @@ -268,6 +269,10 @@ func (a *Account) IsGrok() bool { return a.Platform == PlatformGrok } +func (a *Account) IsTypeSafe() bool { + return a != nil && a.Platform == PlatformTypeSafe +} + func (a *Account) IsGrokOAuth() bool { return a.IsGrok() && a.Type == AccountTypeOAuth } @@ -980,6 +985,10 @@ func (a *Account) GetBaseURL() string { } baseURL := a.GetCredential("base_url") if baseURL == "" { + // TypeSafe keys must never fall back to the Anthropic host. + if a.Platform == PlatformTypeSafe { + return typesafe.DefaultBaseURL + } return "https://api.anthropic.com" } if a.Platform == PlatformAntigravity { @@ -1001,6 +1010,28 @@ func (a *Account) GetGeminiBaseURL(defaultBaseURL string) string { return baseURL } +func (a *Account) GetTypeSafeBaseURL() string { + if a == nil || !a.IsTypeSafe() || a.Type != AccountTypeAPIKey { + return "" + } + baseURL := strings.TrimRight(strings.TrimSpace(a.GetCredential("base_url")), "/") + // The System One path already carries /v1; accept a base URL pasted with it. + if len(baseURL) >= 3 && strings.EqualFold(baseURL[len(baseURL)-3:], "/v1") { + baseURL = strings.TrimRight(baseURL[:len(baseURL)-3], "/") + } + if baseURL == "" { + return typesafe.DefaultBaseURL + } + return baseURL +} + +func (a *Account) GetTypeSafeAPIKey() string { + if a == nil || !a.IsTypeSafe() || a.Type != AccountTypeAPIKey { + return "" + } + return strings.TrimSpace(a.GetCredential("api_key")) +} + func (a *Account) GetExtraString(key string) string { if a.Extra == nil { return "" diff --git a/backend/internal/service/account_service.go b/backend/internal/service/account_service.go index 205907abc..bbf668eaf 100644 --- a/backend/internal/service/account_service.go +++ b/backend/internal/service/account_service.go @@ -2,6 +2,7 @@ package service import ( "context" + "errors" "fmt" "time" @@ -223,6 +224,9 @@ func NewAccountService(accountRepo AccountRepository, groupRepo GroupRepository) // Create 创建账号 func (s *AccountService) Create(ctx context.Context, req CreateAccountRequest) (*Account, error) { + if req.Platform == PlatformTypeSafe && req.Type != AccountTypeAPIKey { + return nil, errors.New("typesafe accounts only support apikey credentials") + } // 验证分组是否存在(如果指定了分组) if len(req.GroupIDs) > 0 { if err := s.validateGroupIDsExist(ctx, req.GroupIDs); err != nil { @@ -515,6 +519,9 @@ func (s *AccountService) TestCredentials(ctx context.Context, id int64) error { case PlatformGrok: // Grok OAuth credentials are validated via token exchange/refresh and request-path probes. return nil + case PlatformTypeSafe: + // TypeSafe credentials are API keys; inference failures drive health and cooldown state. + return nil case PlatformKimi, PlatformZhipu, PlatformDeepseek, PlatformMiniMax, PlatformOpenCodeGo: // 国产 OpenAI 兼容供应商与 OpenCode:凭证为 API Key,实际可用性经余额/额度探测与转发路径验证。 return nil diff --git a/backend/internal/service/account_test_service.go b/backend/internal/service/account_test_service.go index 7d16eeaae..4da30d610 100644 --- a/backend/internal/service/account_test_service.go +++ b/backend/internal/service/account_test_service.go @@ -416,6 +416,10 @@ func (s *AccountTestService) TestAccountConnection(c *gin.Context, accountID int return s.testOpenCodeGoAccountConnection(c, account, modelID, prompt) } + if account.IsTypeSafe() { + return s.testTypeSafeAccountConnection(c, account, prompt) + } + return s.testClaudeAccountConnection(c, account, modelID) } diff --git a/backend/internal/service/account_test_service_typesafe.go b/backend/internal/service/account_test_service_typesafe.go new file mode 100644 index 000000000..6c9501d89 --- /dev/null +++ b/backend/internal/service/account_test_service_typesafe.go @@ -0,0 +1,95 @@ +package service + +import ( + "encoding/json" + "fmt" + "io" + "net/http" + "strings" + + "github.com/Wei-Shaw/sub2api/internal/pkg/typesafe" + "github.com/gin-gonic/gin" +) + +const ( + typeSafeTestDefaultState = "Sub2API connection test" + typeSafeTestQuestionID = "connection_test" + typeSafeTestMaxPreviewBytes = 2000 +) + +// testTypeSafeAccountConnection probes a TypeSafe account with a minimal native +// System One request. TypeSafe accounts never speak the Claude protocol, so +// they must not fall through to testClaudeAccountConnection (which would send +// the key to /v1/messages and could misclassify the account). +func (s *AccountTestService) testTypeSafeAccountConnection(c *gin.Context, account *Account, prompt string) error { + ctx := c.Request.Context() + if account.Type != AccountTypeAPIKey { + return s.sendErrorAndEnd(c, fmt.Sprintf("Unsupported account type: %s", account.Type)) + } + apiKey := account.GetTypeSafeAPIKey() + if apiKey == "" { + return s.sendErrorAndEnd(c, "No API key available") + } + baseURL, err := s.validateUpstreamBaseURL(account.GetTypeSafeBaseURL()) + if err != nil { + return s.sendErrorAndEnd(c, fmt.Sprintf("Invalid base URL: %s", err.Error())) + } + + state := strings.TrimSpace(prompt) + if state == "" { + state = typeSafeTestDefaultState + } + payload, err := json.Marshal(map[string]any{ + "model": typesafe.JevLatestModel, + "state": state, + "questions": map[string]any{ + typeSafeTestQuestionID: map[string]any{ + "type": "noul", + "instructions": "Is this text a connection test?", + }, + }, + }) + if err != nil { + return s.sendErrorAndEnd(c, "Failed to create test payload") + } + + c.Writer.Header().Set("Content-Type", "text/event-stream") + c.Writer.Header().Set("Cache-Control", "no-cache") + c.Writer.Header().Set("Connection", "keep-alive") + c.Writer.Header().Set("X-Accel-Buffering", "no") + c.Writer.Flush() + + s.sendEvent(c, TestEvent{Type: "test_start", Model: typesafe.JevLatestModel}) + + req, err := typesafe.NewSystemOneRequest(ctx, baseURL, apiKey, payload) + if err != nil { + return s.sendErrorAndEnd(c, "Failed to create request") + } + proxyURL := "" + if account.ProxyID != nil && account.Proxy != nil { + proxyURL = account.Proxy.URL() + } + resp, err := s.httpUpstream.Do(req, proxyURL, account.ID, account.Concurrency) + if err != nil { + return s.sendErrorAndEnd(c, fmt.Sprintf("Request failed: %s", sanitizeUpstreamErrorMessage(err.Error()))) + } + defer func() { _ = resp.Body.Close() }() + + if resp.StatusCode != http.StatusOK { + body, _ := io.ReadAll(io.LimitReader(resp.Body, 64<<10)) + errMsg := fmt.Sprintf("API returned %d: %s", resp.StatusCode, truncateString(string(body), typeSafeTestMaxPreviewBytes)) + // 401/403 表示 API Key 无效或被上游拒绝,标记为 error 状态。 + if (resp.StatusCode == http.StatusUnauthorized || resp.StatusCode == http.StatusForbidden) && s.accountRepo != nil { + _ = s.accountRepo.SetError(ctx, account.ID, errMsg) + } + return s.sendErrorAndEnd(c, errMsg) + } + + decoded, err := typesafe.DecodeSystemOneResponse(resp.Body) + if err != nil { + return s.sendErrorAndEnd(c, fmt.Sprintf("Invalid System One response: %s", err.Error())) + } + s.sendEvent(c, TestEvent{Type: "content", Text: truncateString(string(decoded.Body), typeSafeTestMaxPreviewBytes)}) + s.sendEvent(c, TestEvent{Type: "test_complete", Success: true}) + return nil +} diff --git a/backend/internal/service/account_test_service_typesafe_test.go b/backend/internal/service/account_test_service_typesafe_test.go new file mode 100644 index 000000000..c60abf42c --- /dev/null +++ b/backend/internal/service/account_test_service_typesafe_test.go @@ -0,0 +1,92 @@ +package service + +import ( + "encoding/json" + "io" + "net/http" + "net/http/httptest" + "strings" + "testing" + + "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/Wei-Shaw/sub2api/internal/pkg/typesafe" + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" +) + +func newTypeSafeAccountTestFixture(t *testing.T, status int, responseBody string) (*AccountTestService, *systemOnePolicyAccountRepo, *[]*http.Request, *gin.Context, *httptest.ResponseRecorder) { + t.Helper() + account := &Account{ID: 31, Platform: PlatformTypeSafe, Type: AccountTypeAPIKey, Credentials: map[string]any{"api_key": "ts-secret"}} + repo := &systemOnePolicyAccountRepo{account: account} + var requests []*http.Request + upstream := &systemOneHTTPUpstream{do: func(req *http.Request) (*http.Response, error) { + requests = append(requests, req) + return &http.Response{StatusCode: status, Header: make(http.Header), Body: io.NopCloser(strings.NewReader(responseBody))}, nil + }} + svc := &AccountTestService{ + accountRepo: repo, + httpUpstream: upstream, + cfg: &config.Config{}, + } + gin.SetMode(gin.TestMode) + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/api/v1/admin/accounts/31/test", nil) + return svc, repo, &requests, c, rec +} + +func TestTypeSafeAccountTestSendsNativeSystemOneRequest(t *testing.T) { + svc, repo, requests, c, rec := newTypeSafeAccountTestFixture(t, http.StatusOK, `{"model":"jev-1.13.0","answers":{"connection_test":{"type":"noul"}},"usage":{"input_tokens":9}}`) + + // The Claude default model must be ignored: TypeSafe only serves jev-latest. + require.NoError(t, svc.TestAccountConnection(c, 31, "claude-sonnet-4-5", "", AccountTestModeDefault)) + + require.Len(t, *requests, 1) + req := (*requests)[0] + require.Equal(t, typesafe.DefaultBaseURL+typesafe.SystemOnePath, req.URL.String()) + require.Equal(t, "Bearer ts-secret", req.Header.Get("Authorization")) + require.Empty(t, req.Header.Get("x-api-key")) + body, err := io.ReadAll(req.Body) + require.NoError(t, err) + model, err := typesafe.ValidateSystemOneRequest(body) + require.NoError(t, err) + require.Equal(t, typesafe.JevLatestModel, model) + + output := rec.Body.String() + require.Contains(t, output, `"type":"test_start"`) + require.Contains(t, output, `"model":"jev-latest"`) + require.Contains(t, output, `"type":"test_complete"`) + require.Contains(t, output, "jev-1.13.0") + require.Zero(t, repo.errorCalls) +} + +func TestTypeSafeAccountTestMarksRejectedKey(t *testing.T) { + for _, status := range []int{http.StatusUnauthorized, http.StatusForbidden} { + t.Run(http.StatusText(status), func(t *testing.T) { + svc, repo, _, c, rec := newTypeSafeAccountTestFixture(t, status, `{"detail":"invalid key"}`) + require.Error(t, svc.TestAccountConnection(c, 31, "", "", AccountTestModeDefault)) + require.Equal(t, 1, repo.errorCalls) + responseText, errMsg := parseTestSSEOutput(rec.Body.String()) + require.Empty(t, responseText) + require.Contains(t, errMsg, "API returned") + }) + } +} + +func TestTypeSafeAccountTestTransientFailureKeepsAccountState(t *testing.T) { + svc, repo, _, c, _ := newTypeSafeAccountTestFixture(t, http.StatusServiceUnavailable, `{"detail":"busy"}`) + require.Error(t, svc.TestAccountConnection(c, 31, "", "", AccountTestModeDefault)) + require.Zero(t, repo.errorCalls) +} + +func TestTypeSafeAccountTestUsesPromptAsState(t *testing.T) { + svc, _, requests, c, _ := newTypeSafeAccountTestFixture(t, http.StatusOK, `{"answers":{}}`) + require.NoError(t, svc.TestAccountConnection(c, 31, "", "custom state", AccountTestModeDefault)) + body, err := io.ReadAll((*requests)[0].Body) + require.NoError(t, err) + var payload struct { + State string `json:"state"` + } + require.NoError(t, json.Unmarshal(body, &payload)) + require.Equal(t, "custom state", payload.State) +} diff --git a/backend/internal/service/admin_account.go b/backend/internal/service/admin_account.go index 0cf93355f..ef4a47fa1 100644 --- a/backend/internal/service/admin_account.go +++ b/backend/internal/service/admin_account.go @@ -411,6 +411,9 @@ func normalizeOpenAILongContextBillingUpdateExtra(account *Account, input *Updat // Grok media eligibility helpers live in account_grok_media_eligibility.go. func buildAccountForCreate(input *CreateAccountInput, accountExtra map[string]any) (*Account, error) { + if input.Platform == PlatformTypeSafe && input.Type != AccountTypeAPIKey { + return nil, errors.New("typesafe accounts only support apikey credentials") + } // Probe/session state is system-managed. New accounts always start with automatic refresh disabled. delete(accountExtra, UpstreamBillingProbeEnabledExtraKey) delete(accountExtra, UpstreamBillingRateSyncEnabledExtraKey) @@ -576,6 +579,9 @@ func (s *adminServiceImpl) UpdateAccount(ctx context.Context, id int64, input *U if err != nil { return nil, err } + if account.Platform == PlatformTypeSafe && input.Type != "" && input.Type != AccountTypeAPIKey { + return nil, errors.New("typesafe accounts only support apikey credentials") + } var normalizedExtra map[string]any if input.Extra != nil { normalizedExtra, err = normalizeOpenAILongContextBillingUpdateExtra(account, input) diff --git a/backend/internal/service/admin_group.go b/backend/internal/service/admin_group.go index 3068344f7..b0f8b56a1 100644 --- a/backend/internal/service/admin_group.go +++ b/backend/internal/service/admin_group.go @@ -17,6 +17,7 @@ import ( "github.com/Wei-Shaw/sub2api/internal/pkg/logger" "github.com/Wei-Shaw/sub2api/internal/pkg/openai" "github.com/Wei-Shaw/sub2api/internal/pkg/pagination" + "github.com/Wei-Shaw/sub2api/internal/pkg/typesafe" "github.com/Wei-Shaw/sub2api/internal/pkg/xai" ) @@ -296,6 +297,8 @@ func defaultModelsListCandidateIDs(platform string) []string { return xai.DefaultModelIDs() case PlatformOpenCodeGo: return DefaultOpenCodeGoModelIDs() + case PlatformTypeSafe: + return []string{typesafe.JevLatestModel} case PlatformComposite: return compositeDefaultModelsListCandidateIDs() default: @@ -316,6 +319,9 @@ func defaultAllowImageGenerationForPlatform(platform string) bool { func compositeDefaultModelsListCandidateIDs() []string { seen := make(map[string]struct{}) ids := make([]string, 0) + // TypeSafe stays out of the static composite candidates (jev-latest only works + // through /v1/systemone); groups with TypeSafe accounts still get it from the + // account model mappings collected by GetGroupModelsListCandidates. for _, platform := range []string{PlatformAnthropic, PlatformGemini, PlatformOpenAI, PlatformAntigravity, PlatformGrok, PlatformKimi, PlatformZhipu, PlatformDeepseek, PlatformMiniMax, PlatformOpenCodeGo} { for _, id := range defaultModelsListCandidateIDs(platform) { if _, ok := seen[id]; ok { diff --git a/backend/internal/service/antigravity_gateway_gemini.go b/backend/internal/service/antigravity_gateway_gemini.go index 3ce13e6de..4a8c70dcf 100644 --- a/backend/internal/service/antigravity_gateway_gemini.go +++ b/backend/internal/service/antigravity_gateway_gemini.go @@ -194,7 +194,6 @@ func (s *AntigravityGatewayService) ForwardGemini(ctx context.Context, c *gin.Co // 处理错误响应 if resp.StatusCode >= 400 { respBody := s.readUpstreamErrorBody(resp) - contentType := resp.Header.Get("Content-Type") // 尽早关闭原始响应体,释放连接;后续逻辑仍可能需要读取 body,因此用内存副本重新包装。 _ = resp.Body.Close() resp.Body = io.NopCloser(bytes.NewReader(respBody)) @@ -301,7 +300,6 @@ func (s *AntigravityGatewayService) ForwardGemini(ctx context.Context, c *gin.Co Header: retryResp.Header.Clone(), Body: io.NopCloser(bytes.NewReader(retryRespBody)), } - contentType = resp.Header.Get("Content-Type") } } else { if switchErr, ok := IsAntigravityAccountSwitchError(retryErr); ok { @@ -393,9 +391,6 @@ func (s *AntigravityGatewayService) ForwardGemini(ctx context.Context, c *gin.Co }) return nil, &UpstreamFailoverError{StatusCode: resp.StatusCode, ResponseBody: unwrappedForOps} } - if contentType == "" { - contentType = "application/json" - } appendOpsUpstreamError(c, OpsUpstreamErrorEvent{ ProxyID: opsUpstreamProxyID(account), ProxyName: opsUpstreamProxyName(account), @@ -410,7 +405,8 @@ func (s *AntigravityGatewayService) ForwardGemini(ctx context.Context, c *gin.Co }) logger.LegacyPrintf("service.antigravity_gateway", "[antigravity-Forward] upstream error status=%d body=%s", resp.StatusCode, truncateForLog(unwrappedForOps, 500)) MarkResponseCommitted(c) - c.Data(resp.StatusCode, contentType, unwrappedForOps) + // 原始上游错误体仅保留在 ops 日志中;返回客户端的错误体需脱敏,避免泄露账号池身份(项目号/服务账号等) + c.Data(resp.StatusCode, "application/json", buildAntigravityClientErrorBody(resp.StatusCode, unwrappedForOps)) return nil, fmt.Errorf("antigravity upstream error: %d", resp.StatusCode) } diff --git a/backend/internal/service/antigravity_upstream_error_sanitize.go b/backend/internal/service/antigravity_upstream_error_sanitize.go new file mode 100644 index 000000000..b937970d6 --- /dev/null +++ b/backend/internal/service/antigravity_upstream_error_sanitize.go @@ -0,0 +1,97 @@ +package service + +import ( + "encoding/json" + "net/http" + "regexp" + "strings" +) + +var ( + antigravityProjectRefRegex = regexp.MustCompile(`(?i)\bprojects/[a-z0-9][a-z0-9._:-]*`) + antigravityEmailRegex = regexp.MustCompile(`[A-Za-z0-9._%+\-]+@[A-Za-z0-9.\-]+\.[A-Za-z]{2,}`) + antigravityConsumerRegex = regexp.MustCompile(`(?i)\b(consumer|project(?:[ _-]?(?:id|number))?)(\s*[:=]?\s*['"]?)[0-9]{6,}`) +) + +// sanitizeAntigravityErrorText 清除上游错误文本中的 GCP 项目号/项目 ID、服务账号邮箱及敏感查询参数。 +func sanitizeAntigravityErrorText(msg string) string { + if msg == "" { + return msg + } + msg = sanitizeUpstreamErrorMessage(msg) + msg = antigravityProjectRefRegex.ReplaceAllString(msg, "projects/***") + msg = antigravityEmailRegex.ReplaceAllString(msg, "***") + msg = antigravityConsumerRegex.ReplaceAllString(msg, "$1$2***") + return msg +} + +// buildAntigravityClientErrorBody 为客户端构造 Gemini 风格的错误体: +// 仅保留 code/status/message(message 已脱敏),丢弃 details 等可能含账号身份的字段。 +func buildAntigravityClientErrorBody(statusCode int, body []byte) []byte { + code := statusCode + status := "" + message := "" + + var parsed struct { + Error struct { + Code int `json:"code"` + Message string `json:"message"` + Status string `json:"status"` + } `json:"error"` + } + if err := json.Unmarshal(body, &parsed); err == nil { + if parsed.Error.Code != 0 { + code = parsed.Error.Code + } + status = parsed.Error.Status + message = parsed.Error.Message + } + if strings.TrimSpace(message) == "" { + message = strings.TrimSpace(extractUpstreamErrorMessage(body)) + } + if strings.TrimSpace(message) == "" { + message = http.StatusText(statusCode) + if message == "" { + message = "Upstream request failed" + } + } + if status == "" { + status = antigravityGeminiStatusFromHTTP(statusCode) + } + + out, err := json.Marshal(map[string]any{ + "error": map[string]any{ + "code": code, + "message": sanitizeAntigravityErrorText(message), + "status": status, + }, + }) + if err != nil { + return []byte(`{"error":{"code":500,"message":"Upstream request failed","status":"INTERNAL"}}`) + } + return out +} + +func antigravityGeminiStatusFromHTTP(statusCode int) string { + switch statusCode { + case http.StatusBadRequest: + return "INVALID_ARGUMENT" + case http.StatusUnauthorized: + return "UNAUTHENTICATED" + case http.StatusForbidden: + return "PERMISSION_DENIED" + case http.StatusNotFound: + return "NOT_FOUND" + case http.StatusTooManyRequests: + return "RESOURCE_EXHAUSTED" + case http.StatusServiceUnavailable: + return "UNAVAILABLE" + case http.StatusGatewayTimeout: + return "DEADLINE_EXCEEDED" + default: + if statusCode >= 500 { + return "INTERNAL" + } + return "UNKNOWN" + } +} diff --git a/backend/internal/service/antigravity_upstream_error_sanitize_test.go b/backend/internal/service/antigravity_upstream_error_sanitize_test.go new file mode 100644 index 000000000..3f4abf04a --- /dev/null +++ b/backend/internal/service/antigravity_upstream_error_sanitize_test.go @@ -0,0 +1,50 @@ +//go:build unit + +package service + +import ( + "encoding/json" + "net/http" + "net/http/httptest" + "strings" + "testing" + + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" +) + +func TestBuildAntigravityClientErrorBody_ScrubsPoolIdentity(t *testing.T) { + gin.SetMode(gin.TestMode) + upstream := []byte(`{"error":{"code":403,"message":"Permission denied on resource project projects/123456789 for consumer: projects/123456789; caller pool-sa@my-gcp-proj.iam.gserviceaccount.com","status":"PERMISSION_DENIED","details":[{"@type":"type.googleapis.com/google.rpc.ErrorInfo","metadata":{"consumer":"projects/123456789","service":"cloudcode-pa.googleapis.com"}}]}}`) + + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Data(http.StatusForbidden, "application/json", buildAntigravityClientErrorBody(http.StatusForbidden, upstream)) + + require.Equal(t, http.StatusForbidden, rec.Code) + out := rec.Body.String() + require.NotContains(t, out, "123456789") + require.NotContains(t, out, "pool-sa@") + require.NotContains(t, out, "gserviceaccount.com") + require.NotContains(t, out, "details") + + var parsed struct { + Error struct { + Code int `json:"code"` + Message string `json:"message"` + Status string `json:"status"` + } `json:"error"` + } + require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &parsed)) + require.Equal(t, 403, parsed.Error.Code) + require.Equal(t, "PERMISSION_DENIED", parsed.Error.Status) + require.True(t, strings.Contains(parsed.Error.Message, "Permission denied")) +} + +func TestBuildAntigravityClientErrorBody_NonJSONBody(t *testing.T) { + out := string(buildAntigravityClientErrorBody(http.StatusTooManyRequests, []byte("quota exceeded for consumer 987654321 sa@x.iam.gserviceaccount.com"))) + require.NotContains(t, out, "987654321") + require.NotContains(t, out, "gserviceaccount") + require.Contains(t, out, `"status":"RESOURCE_EXHAUSTED"`) + require.Contains(t, out, `"code":429`) +} diff --git a/backend/internal/service/auth_service_email_bind_test.go b/backend/internal/service/auth_service_email_bind_test.go index 10d459af6..c41734c40 100644 --- a/backend/internal/service/auth_service_email_bind_test.go +++ b/backend/internal/service/auth_service_email_bind_test.go @@ -1166,6 +1166,18 @@ func cloneEmailBindUser(user *service.User) *service.User { return &cloned } +func (s *emailBindCacheStub) IncrVerificationCodeAttempts(context.Context, string) (int, error) { + if s.data == nil { + return 0, errors.New("verification code not found") + } + s.data.Attempts++ + return s.data.Attempts, nil +} + +func (s *emailBindCacheStub) IncrNotifyVerifyCodeAttempts(context.Context, string) (int, error) { + return 0, errors.New("notify verification code not found") +} + func (s *emailBindCacheStub) ConsumePasswordResetToken(context.Context, string, string) (bool, error) { return false, nil } diff --git a/backend/internal/service/auth_service_register_test.go b/backend/internal/service/auth_service_register_test.go index e96488958..e5c0922e3 100644 --- a/backend/internal/service/auth_service_register_test.go +++ b/backend/internal/service/auth_service_register_test.go @@ -1053,6 +1053,18 @@ func TestCanBypassRegistrationDisabledForOAuth(t *testing.T) { } } +func (s *emailCacheStub) IncrVerificationCodeAttempts(context.Context, string) (int, error) { + if s.data == nil { + return 0, errors.New("verification code not found") + } + s.data.Attempts++ + return s.data.Attempts, nil +} + +func (s *emailCacheStub) IncrNotifyVerifyCodeAttempts(context.Context, string) (int, error) { + return 0, errors.New("notify verification code not found") +} + func (s *emailCacheStub) ConsumePasswordResetToken(context.Context, string, string) (bool, error) { return false, nil } diff --git a/backend/internal/service/billing_service.go b/backend/internal/service/billing_service.go index b64a11429..6d95656b1 100644 --- a/backend/internal/service/billing_service.go +++ b/backend/internal/service/billing_service.go @@ -698,6 +698,12 @@ func (s *BillingService) initFallbackPricing() { SupportsCacheBreakdown: false, } + // TypeSafe Jev bills input tokens only: $0.042 per million tokens. + s.fallbackPrices["jev-latest"] = &ModelPricing{ + InputPricePerToken: 0.042 / 1_000_000, + OutputPricePerToken: 0, + } + // ---- 智谱 GLM(Z.AI)---- // Source: https://docs.z.ai/guides/overview/pricing (USD per 1M tokens) // 注意:CacheReadPricePerToken 即"缓存命中"价格,CacheCreationPricePerToken 留空(智谱未公开写入价,按 0 处理)。 @@ -994,6 +1000,9 @@ func (s *BillingService) initFallbackPricing() { // getFallbackPricing 根据模型系列获取回退价格 func (s *BillingService) getFallbackPricing(model string) *ModelPricing { modelLower := strings.ToLower(model) + if modelLower == "jev-latest" { + return s.fallbackPrices["jev-latest"] + } // 按模型系列匹配 if isClaudeFable51Model(modelLower) { diff --git a/backend/internal/service/billing_service_test.go b/backend/internal/service/billing_service_test.go index f2776870b..b8ae7b349 100644 --- a/backend/internal/service/billing_service_test.go +++ b/backend/internal/service/billing_service_test.go @@ -130,6 +130,20 @@ func TestGetModelPricing_CaseInsensitive(t *testing.T) { require.Equal(t, p1.InputPricePerToken, p2.InputPricePerToken) } +func TestGetModelPricing_JevLatestInputOnlyAndChannelOverride(t *testing.T) { + svc := newTestBillingService() + pricing, err := svc.GetModelPricing("jev-latest") + require.NoError(t, err) + require.InDelta(t, 0.042/1_000_000, pricing.InputPricePerToken, 1e-15) + require.Zero(t, pricing.OutputPricePerToken) + + input, output := 0.25/1_000_000, 0.5/1_000_000 + pricing, err = svc.GetModelPricingWithChannel("jev-latest", &ChannelModelPricing{InputPrice: &input, OutputPrice: &output}) + require.NoError(t, err) + require.Equal(t, input, pricing.InputPricePerToken) + require.Equal(t, output, pricing.OutputPricePerToken) +} + // issue #3394: fallback warn 应按模型名去重,每个模型每进程最多打一条, // 避免热路径每请求刷屏 ops_system_logs。 func TestGetModelPricing_FallbackWarnLoggedOncePerModel(t *testing.T) { diff --git a/backend/internal/service/channel_service.go b/backend/internal/service/channel_service.go index fb59d7f7b..b2c605cd8 100644 --- a/backend/internal/service/channel_service.go +++ b/backend/internal/service/channel_service.go @@ -368,7 +368,7 @@ func isPlatformPricingMatch(groupPlatform, pricingPlatform string) bool { // fallback used before a request target has been resolved. func matchingPlatforms(groupPlatform string) []string { if groupPlatform == PlatformComposite { - return []string{PlatformAnthropic, PlatformGemini, PlatformOpenAI, PlatformAntigravity, PlatformGrok, PlatformKimi, PlatformZhipu, PlatformDeepseek, PlatformMiniMax, PlatformOpenCodeGo} + return []string{PlatformAnthropic, PlatformGemini, PlatformOpenAI, PlatformAntigravity, PlatformGrok, PlatformKimi, PlatformZhipu, PlatformDeepseek, PlatformMiniMax, PlatformOpenCodeGo, PlatformTypeSafe} } return []string{groupPlatform} } diff --git a/backend/internal/service/channel_service_test.go b/backend/internal/service/channel_service_test.go index 9b85aa3e0..62ba8e0b9 100644 --- a/backend/internal/service/channel_service_test.go +++ b/backend/internal/service/channel_service_test.go @@ -2127,7 +2127,7 @@ func TestMatchingPlatforms(t *testing.T) { {"anthropic returns itself", PlatformAnthropic, []string{PlatformAnthropic}}, {"gemini returns itself", PlatformGemini, []string{PlatformGemini}}, {"openai returns itself", PlatformOpenAI, []string{PlatformOpenAI}}, - {"composite returns concrete platforms", PlatformComposite, []string{PlatformAnthropic, PlatformGemini, PlatformOpenAI, PlatformAntigravity, PlatformGrok, PlatformKimi, PlatformZhipu, PlatformDeepseek, PlatformMiniMax, PlatformOpenCodeGo}}, + {"composite returns concrete platforms", PlatformComposite, []string{PlatformAnthropic, PlatformGemini, PlatformOpenAI, PlatformAntigravity, PlatformGrok, PlatformKimi, PlatformZhipu, PlatformDeepseek, PlatformMiniMax, PlatformOpenCodeGo, PlatformTypeSafe}}, } for _, tt := range tests { diff --git a/backend/internal/service/composite_platform.go b/backend/internal/service/composite_platform.go index 6ddb9cfa9..c10e335cb 100644 --- a/backend/internal/service/composite_platform.go +++ b/backend/internal/service/composite_platform.go @@ -114,6 +114,8 @@ func DetectModelPlatform(model string) (string, bool) { return PlatformDeepseek, true case "minimax": return PlatformMiniMax, true + case "typesafe", "jev": + return PlatformTypeSafe, true } if rest != "" { normalized = strings.TrimPrefix(rest, "models/") @@ -155,6 +157,8 @@ func DetectModelPlatform(model string) (string, bool) { strings.HasPrefix(normalized, "abab6"), strings.HasPrefix(normalized, "abab7"): return PlatformMiniMax, true + case normalized == "jev-latest" || strings.HasPrefix(normalized, "jev-"): + return PlatformTypeSafe, true default: return "", false } @@ -202,7 +206,7 @@ func (s *GatewayService) resolveCompositeRouteDecision(ctx context.Context, grou func isConcreteRequestPlatform(platform string) bool { switch platform { case PlatformAnthropic, PlatformOpenAI, PlatformGemini, PlatformAntigravity, PlatformGrok, - PlatformKimi, PlatformZhipu, PlatformDeepseek, PlatformMiniMax, PlatformOpenCodeGo: + PlatformKimi, PlatformZhipu, PlatformDeepseek, PlatformMiniMax, PlatformOpenCodeGo, PlatformTypeSafe: return true default: return false diff --git a/backend/internal/service/composite_platform_test.go b/backend/internal/service/composite_platform_test.go index 20505a398..4b11366a2 100644 --- a/backend/internal/service/composite_platform_test.go +++ b/backend/internal/service/composite_platform_test.go @@ -180,6 +180,8 @@ func TestDetectModelPlatform(t *testing.T) { {name: "minimax prefix", model: "minimax/MiniMax-M2.5", platform: PlatformMiniMax, ok: true}, {name: "abab legacy", model: "abab6.5-chat", platform: PlatformMiniMax, ok: true}, {name: "abab7 legacy", model: "abab7-chat-preview", platform: PlatformMiniMax, ok: true}, + {name: "jev", model: "jev-latest", platform: PlatformTypeSafe, ok: true}, + {name: "typesafe prefix", model: "typesafe/jev-latest", platform: PlatformTypeSafe, ok: true}, {name: "abab unrelated namespace", model: "abab-other", ok: false}, {name: "unknown k3 alias", model: "k3-preview", ok: false}, {name: "unknown", model: "llama-4-maverick", ok: false}, @@ -216,13 +218,13 @@ func TestCompositeGroupSchedulerHasAllCanonicalPlatformBuckets(t *testing.T) { platforms = append(platforms, platform) } require.ElementsMatch(t, - []string{PlatformAnthropic, PlatformGemini, PlatformOpenAI, PlatformAntigravity, PlatformGrok, PlatformKimi, PlatformZhipu, PlatformDeepseek, PlatformMiniMax, PlatformOpenCodeGo}, + []string{PlatformAnthropic, PlatformGemini, PlatformOpenAI, PlatformAntigravity, PlatformGrok, PlatformKimi, PlatformZhipu, PlatformDeepseek, PlatformMiniMax, PlatformOpenCodeGo, PlatformTypeSafe}, platforms, ) } func TestCompositeConcretePlatformsIncludeCNProviders(t *testing.T) { - for _, platform := range []string{PlatformKimi, PlatformZhipu, PlatformDeepseek, PlatformMiniMax, PlatformOpenCodeGo} { + for _, platform := range []string{PlatformKimi, PlatformZhipu, PlatformDeepseek, PlatformMiniMax, PlatformOpenCodeGo, PlatformTypeSafe} { require.True(t, isConcreteRequestPlatform(platform)) require.True(t, canCopyAccountsFromGroupPlatform(PlatformComposite, platform)) } diff --git a/backend/internal/service/content_moderation.go b/backend/internal/service/content_moderation.go index 9c083a516..6e696d2fe 100644 --- a/backend/internal/service/content_moderation.go +++ b/backend/internal/service/content_moderation.go @@ -56,6 +56,7 @@ const ( ContentModerationProtocolOpenAIChat = "openai_chat_completions" ContentModerationProtocolGemini = "gemini" ContentModerationProtocolOpenAIImages = "openai_images" + ContentModerationProtocolTypeSafeSystemOne = "typesafe_systemone" defaultContentModerationBaseURL = "https://api.openai.com" defaultContentModerationModel = "omni-moderation-latest" diff --git a/backend/internal/service/content_moderation_input.go b/backend/internal/service/content_moderation_input.go index 3886bfff8..b52adfaa2 100644 --- a/backend/internal/service/content_moderation_input.go +++ b/backend/internal/service/content_moderation_input.go @@ -46,6 +46,10 @@ func extractContentModerationInput(protocol string, body []byte, filterReminders case ContentModerationProtocolOpenAIImages: collector.addModerationText(&parts, gjson.GetBytes(body, "prompt").String()) collector.collectContentValue(gjson.GetBytes(body, "images"), &parts, &images) + case ContentModerationProtocolTypeSafeSystemOne: + // System One carries no client-harness reminder blocks, so a literal + // is ordinary user text and must never be skipped. + moderationTextCollector{}.collectSystemOneInput(body, &parts) default: collector.collectLastResponsesInput(gjson.GetBytes(body, "input"), &parts, &images) collector.collectLastRoleMessage(gjson.GetBytes(body, "messages"), "user", &parts, &images) @@ -61,6 +65,67 @@ func extractContentModerationInput(protocol string, body []byte, filterReminders return out } +// collectSystemOneInput moderates every client-controlled text of a System One +// request: question IDs, every question field except the validated type, +// unknown top-level extension fields, and the evaluated state. Object keys are +// sent to Jev as part of the JSON, so they are moderated like values. +func (collector moderationTextCollector) collectSystemOneInput(body []byte, parts *[]string) { + root := gjson.ParseBytes(body) + questions := root.Get("questions") + if !questions.IsObject() { + collector.collectSystemOneText(questions, parts) + } + questions.ForEach(func(id, question gjson.Result) bool { + collector.addModerationText(parts, id.String()) + if !question.IsObject() { + collector.collectSystemOneText(question, parts) + return true + } + question.ForEach(func(field, value gjson.Result) bool { + switch field.String() { + case "type": + return true + case "instructions", "criteria": + default: + collector.addModerationText(parts, field.String()) + } + collector.collectSystemOneText(value, parts) + return true + }) + return true + }) + root.ForEach(func(field, value gjson.Result) bool { + switch field.String() { + case "model", "stream", "state", "questions": + return true + } + collector.addModerationText(parts, field.String()) + collector.collectSystemOneText(value, parts) + return true + }) + collector.collectSystemOneText(root.Get("state"), parts) +} + +func (collector moderationTextCollector) collectSystemOneText(value gjson.Result, parts *[]string) { + switch { + case !value.Exists(): + return + case value.Type == gjson.String: + collector.addModerationText(parts, value.String()) + case value.IsArray(): + value.ForEach(func(_, child gjson.Result) bool { + collector.collectSystemOneText(child, parts) + return true + }) + case value.IsObject(): + value.ForEach(func(key, child gjson.Result) bool { + collector.addModerationText(parts, key.String()) + collector.collectSystemOneText(child, parts) + return true + }) + } +} + func (collector moderationTextCollector) collectLastRoleMessage(messages gjson.Result, role string, parts *[]string, images *[]string) { if !messages.IsArray() { return diff --git a/backend/internal/service/content_moderation_systemone_test.go b/backend/internal/service/content_moderation_systemone_test.go new file mode 100644 index 000000000..8f2eba816 --- /dev/null +++ b/backend/internal/service/content_moderation_systemone_test.go @@ -0,0 +1,34 @@ +package service + +import ( + "testing" + + "github.com/stretchr/testify/require" +) + +func TestExtractContentModerationInputTypeSafeSystemOne(t *testing.T) { + for _, tc := range []struct { + name string + body string + want string + }{ + {"string", `{"state":"plain text"}`, "plain text"}, + {"object", `{"state":{"title":"hello","nested":{"body":"world"}}}`, "title hello nested body world"}, + {"array", `{"state":["first",{"text":"second"},3]}`, "first text second"}, + {"questions", `{"state":"state text","questions":{"c":{"type":"choice","instructions":"pick one","criteria":{"label a":{"description":"desc a"},"label b":null}},"s":{"type":"score","criteria":["low",{"text":"high"}]},"n":{"type":"noul","instructions":["judge"],"criteria":{"true":"yes"},"ext":"extra"}},"top":{"k":"v"}}`, "c pick one label a description desc a label b s low text high n judge true yes ext extra top k v state text"}, + {"key only payload", `{"state":{"hidden state key":1},"questions":{"hidden question id":{"type":"noul"}}}`, "hidden question id hidden state key"}, + } { + t.Run(tc.name, func(t *testing.T) { + input := ExtractContentModerationInput(ContentModerationProtocolTypeSafeSystemOne, []byte(tc.body)) + require.Equal(t, tc.want, input.Text) + require.Empty(t, input.Images) + }) + } +} + +func TestExtractContentModerationInputTypeSafeSystemOneKeepsReminderText(t *testing.T) { + body := `{"state":"hidden payload","questions":{"q":{"type":"noul","instructions":"hidden instructions"}}}` + input := ExtractContentModerationInput(ContentModerationProtocolTypeSafeSystemOne, []byte(body)) + require.Contains(t, input.Text, "hidden payload") + require.Contains(t, input.Text, "hidden instructions") +} diff --git a/backend/internal/service/domain_constants.go b/backend/internal/service/domain_constants.go index 79e5379e6..4dad31988 100644 --- a/backend/internal/service/domain_constants.go +++ b/backend/internal/service/domain_constants.go @@ -48,6 +48,7 @@ const ( PlatformZhipu = domain.PlatformZhipu PlatformDeepseek = domain.PlatformDeepseek PlatformMiniMax = domain.PlatformMiniMax + PlatformTypeSafe = domain.PlatformTypeSafe PlatformOpenCodeGo = domain.PlatformOpenCodeGo PlatformComposite = domain.PlatformComposite // PlatformKiro is retained for unsupported-platform threshold tests and legacy @@ -135,6 +136,7 @@ var AllowedQuotaPlatforms = []string{ PlatformDeepseek, PlatformMiniMax, PlatformOpenCodeGo, + PlatformTypeSafe, } // AllowedSchedulingThresholdPlatforms 是允许设置账号自动停调阈值的平台列表。 diff --git a/backend/internal/service/email_service.go b/backend/internal/service/email_service.go index 6a4004a83..b7088b5d7 100644 --- a/backend/internal/service/email_service.go +++ b/backend/internal/service/email_service.go @@ -3,6 +3,7 @@ package service import ( "context" "crypto/rand" + "crypto/sha256" "crypto/subtle" "crypto/tls" "crypto/x509" @@ -37,17 +38,23 @@ type EmailCache interface { GetVerificationCode(ctx context.Context, email string) (*VerificationCodeData, error) SetVerificationCode(ctx context.Context, email string, data *VerificationCodeData, ttl time.Duration) error DeleteVerificationCode(ctx context.Context, email string) error + // IncrVerificationCodeAttempts atomically increments the attempt counter of an + // existing code and returns the new count. Returns an error if the code is missing. + IncrVerificationCodeAttempts(ctx context.Context, email string) (int, error) // Notify email verification code methods GetNotifyVerifyCode(ctx context.Context, email string) (*VerificationCodeData, error) SetNotifyVerifyCode(ctx context.Context, email string, data *VerificationCodeData, ttl time.Duration) error DeleteNotifyVerifyCode(ctx context.Context, email string) error + IncrNotifyVerifyCodeAttempts(ctx context.Context, email string) (int, error) // Password reset token methods GetPasswordResetToken(ctx context.Context, email string) (*PasswordResetTokenData, error) SetPasswordResetToken(ctx context.Context, email string, data *PasswordResetTokenData, ttl time.Duration) error DeletePasswordResetToken(ctx context.Context, email string) error - ConsumePasswordResetToken(ctx context.Context, email, token string) (bool, error) + // ConsumePasswordResetToken atomically compares the stored token hash with + // tokenHash and deletes it on match. Only one concurrent caller can succeed. + ConsumePasswordResetToken(ctx context.Context, email, tokenHash string) (bool, error) // Password reset email cooldown methods // Returns true if in cooldown period (email was sent recently) @@ -69,6 +76,7 @@ type VerificationCodeData struct { // PasswordResetTokenData represents password reset token data type PasswordResetTokenData struct { + // Token holds the hex-encoded SHA-256 hash of the reset token (never the plaintext). Token string CreatedAt time.Time } @@ -377,35 +385,49 @@ func (s *EmailService) SendVerifyCode(ctx context.Context, email, siteName strin // VerifyCode 验证验证码 func (s *EmailService) VerifyCode(ctx context.Context, email, code string) error { - data, err := s.cache.GetVerificationCode(ctx, email) + return verifyCodeWithAttempts(ctx, email, code, + s.cache.GetVerificationCode, s.cache.IncrVerificationCodeAttempts, + func() { + if err := s.cache.DeleteVerificationCode(ctx, email); err != nil { + slog.Error("failed to delete verification code after success", "email", email, "error", err) + } + }) +} + +// verifyCodeWithAttempts checks a verification code while enforcing the attempt cap +// atomically: each check first reserves an attempt via an atomic increment, so +// concurrent guesses can never evaluate more than maxVerifyCodeAttempts codes. +func verifyCodeWithAttempts( + ctx context.Context, email, code string, + get func(context.Context, string) (*VerificationCodeData, error), + incr func(context.Context, string) (int, error), + onSuccess func(), +) error { + data, err := get(ctx, email) if err != nil || data == nil { return ErrInvalidVerifyCode } - - // 检查是否已达到最大尝试次数 if data.Attempts >= maxVerifyCodeAttempts { return ErrVerifyCodeMaxAttempts } + attempts, err := incr(ctx, email) + if err != nil { + return ErrInvalidVerifyCode + } + if attempts > maxVerifyCodeAttempts { + return ErrVerifyCodeMaxAttempts + } - // 验证码不匹配 (constant-time comparison to prevent timing attacks) + // constant-time comparison to prevent timing attacks if subtle.ConstantTimeCompare([]byte(data.Code), []byte(code)) != 1 { - data.Attempts++ - remaining := time.Until(data.ExpiresAt) - if remaining <= 0 { - return ErrInvalidVerifyCode - } - if err := s.cache.SetVerificationCode(ctx, email, data, remaining); err != nil { - slog.Error("failed to update verification attempt count", "email", email, "error", err) - } - if data.Attempts >= maxVerifyCodeAttempts { + if attempts >= maxVerifyCodeAttempts { return ErrVerifyCodeMaxAttempts } return ErrInvalidVerifyCode } - // 验证成功,删除验证码 - if err := s.cache.DeleteVerificationCode(ctx, email); err != nil { - slog.Error("failed to delete verification code after success", "email", email, "error", err) + if onSuccess != nil { + onSuccess() } return nil } @@ -481,33 +503,18 @@ func (s *EmailService) GeneratePasswordResetToken() (string, error) { // SendPasswordResetEmail sends a password reset email with a reset link func (s *EmailService) SendPasswordResetEmail(ctx context.Context, email, siteName, resetURL string, locale ...string) error { - var token string - var needSaveToken bool - - // Check if token already exists - existing, err := s.cache.GetPasswordResetToken(ctx, email) - if err == nil && existing != nil { - // Token exists, reuse it (allows resending email without generating new token) - token = existing.Token - needSaveToken = false - } else { - // Generate new token - token, err = s.GeneratePasswordResetToken() - if err != nil { - return fmt.Errorf("generate token: %w", err) - } - needSaveToken = true + // Only the SHA-256 hash of the token is stored, so an existing token cannot be + // re-sent; always issue a fresh token (the email cooldown bounds resend frequency). + token, err := s.GeneratePasswordResetToken() + if err != nil { + return fmt.Errorf("generate token: %w", err) } - - // Save token to Redis (only if new token generated) - if needSaveToken { - data := &PasswordResetTokenData{ - Token: token, - CreatedAt: time.Now(), - } - if err := s.cache.SetPasswordResetToken(ctx, email, data, passwordResetTokenTTL); err != nil { - return fmt.Errorf("save reset token: %w", err) - } + data := &PasswordResetTokenData{ + Token: hashPasswordResetToken(token), + CreatedAt: time.Now(), + } + if err := s.cache.SetPasswordResetToken(ctx, email, data, passwordResetTokenTTL); err != nil { + return fmt.Errorf("save reset token: %w", err) } // Build full reset URL with URL-encoded token and email @@ -567,29 +574,35 @@ func (s *EmailService) SendPasswordResetEmailWithCooldown(ctx context.Context, e return nil } +// hashPasswordResetToken returns the hex-encoded SHA-256 of a reset token. +func hashPasswordResetToken(token string) string { + sum := sha256.Sum256([]byte(token)) + return hex.EncodeToString(sum[:]) +} + // VerifyPasswordResetToken verifies the password reset token without consuming it func (s *EmailService) VerifyPasswordResetToken(ctx context.Context, email, token string) error { data, err := s.cache.GetPasswordResetToken(ctx, email) - if err != nil || data == nil { + if err != nil || data == nil || token == "" { return ErrInvalidResetToken } // Use constant-time comparison to prevent timing attacks - if subtle.ConstantTimeCompare([]byte(data.Token), []byte(token)) != 1 { + if subtle.ConstantTimeCompare([]byte(data.Token), []byte(hashPasswordResetToken(token))) != 1 { return ErrInvalidResetToken } return nil } -// ConsumePasswordResetToken 保留恒定时间校验,再原子核销匹配的令牌。 +// ConsumePasswordResetToken verifies and deletes the token atomically (one-time use). func (s *EmailService) ConsumePasswordResetToken(ctx context.Context, email, token string) error { - // Verify first + // Constant-time pre-check in Go; the atomic compare-and-delete below is authoritative. if err := s.VerifyPasswordResetToken(ctx, email, token); err != nil { return err } - consumed, err := s.cache.ConsumePasswordResetToken(ctx, email, token) + consumed, err := s.cache.ConsumePasswordResetToken(ctx, email, hashPasswordResetToken(token)) if err != nil { slog.Error("failed to consume password reset token", "error", err) return ErrServiceUnavailable diff --git a/backend/internal/service/email_service_reset_token_test.go b/backend/internal/service/email_service_reset_token_test.go new file mode 100644 index 000000000..15efd50d6 --- /dev/null +++ b/backend/internal/service/email_service_reset_token_test.go @@ -0,0 +1,50 @@ +//go:build unit + +package service + +import ( + "context" + "crypto/sha256" + "encoding/hex" + "testing" + + "github.com/stretchr/testify/require" +) + +type resetTokenCacheStub struct { + emailCacheStub + stored *PasswordResetTokenData + consumedHash string +} + +func (s *resetTokenCacheStub) GetPasswordResetToken(context.Context, string) (*PasswordResetTokenData, error) { + return s.stored, nil +} + +func (s *resetTokenCacheStub) ConsumePasswordResetToken(_ context.Context, _ string, tokenHash string) (bool, error) { + s.consumedHash = tokenHash + if s.stored == nil || s.stored.Token != tokenHash { + return false, nil + } + s.stored = nil + return true, nil +} + +func TestConsumePasswordResetToken_ComparesHashNotPlaintext(t *testing.T) { + token := "deadbeef" + sum := sha256.Sum256([]byte(token)) + hash := hex.EncodeToString(sum[:]) + require.Equal(t, hash, hashPasswordResetToken(token)) + require.NotEqual(t, token, hashPasswordResetToken(token)) + + cache := &resetTokenCacheStub{stored: &PasswordResetTokenData{Token: hash}} + svc := NewEmailService(nil, cache) + + require.NoError(t, svc.ConsumePasswordResetToken(context.Background(), "a@b.c", token)) + require.Equal(t, hash, cache.consumedHash) + require.ErrorIs(t, svc.ConsumePasswordResetToken(context.Background(), "a@b.c", token), ErrInvalidResetToken) + + // A legacy plaintext value (issued before upgrade) no longer validates. + cache.stored = &PasswordResetTokenData{Token: token} + require.ErrorIs(t, svc.ConsumePasswordResetToken(context.Background(), "a@b.c", token), ErrInvalidResetToken) +} diff --git a/backend/internal/service/gateway_systemone.go b/backend/internal/service/gateway_systemone.go new file mode 100644 index 000000000..0add1999d --- /dev/null +++ b/backend/internal/service/gateway_systemone.go @@ -0,0 +1,182 @@ +package service + +import ( + "context" + "errors" + "fmt" + "mime" + "net/http" + "strings" + "time" + + "github.com/Wei-Shaw/sub2api/internal/pkg/typesafe" + "github.com/gin-gonic/gin" +) + +type SystemOneForwardResult struct { + ForwardResult + StatusCode int + Body []byte + ContentType string +} + +type SystemOneUpstreamError struct { + StatusCode int +} + +const TypeSafeCredentialRejectedReason GatewayFailureReason = "typesafe_api_key_rejected" + +func (e *SystemOneUpstreamError) Error() string { + return fmt.Sprintf("typesafe upstream rejected request with status %d", e.StatusCode) +} + +func (s *GatewayService) ForwardSystemOne(ctx context.Context, c *gin.Context, account *Account, body []byte) (*SystemOneForwardResult, error) { + started := time.Now() + if account == nil || !account.IsTypeSafe() || account.Type != AccountTypeAPIKey { + return nil, errors.New("invalid typesafe account") + } + key := account.GetTypeSafeAPIKey() + if key == "" { + return nil, errors.New("typesafe api key is missing") + } + baseURL, err := s.validateUpstreamBaseURL(account.GetTypeSafeBaseURL()) + if err != nil { + return nil, err + } + req, err := typesafe.NewSystemOneRequest(ctx, baseURL, key, body) + if err != nil { + return nil, err + } + upstreamURL := req.URL.Scheme + "://" + req.URL.Host + req.URL.Path + proxyURL := "" + if account.ProxyID != nil && account.Proxy != nil { + proxyURL = account.Proxy.URL() + } + resp, err := s.httpUpstream.Do(req, proxyURL, account.ID, account.Concurrency) + if err != nil { + return nil, s.handleUpstreamTransportError(ctx, c, account, err, OpsUpstreamErrorEvent{ + Passthrough: true, + UpstreamURL: upstreamURL, + }) + } + defer func() { _ = resp.Body.Close() }() + + if resp.StatusCode < http.StatusOK || resp.StatusCode >= http.StatusMultipleChoices { + return nil, s.handleSystemOneErrorResponse(ctx, c, account, resp, upstreamURL) + } + + decoded, err := typesafe.DecodeSystemOneResponse(resp.Body) + if err != nil { + // The upstream accepted (and may have charged) this request but the + // gateway cannot relay it; keep an ops trail for reconciliation. + setOpsUpstreamError(c, resp.StatusCode, err.Error(), "") + appendOpsUpstreamError(c, OpsUpstreamErrorEvent{ + Passthrough: true, + ProxyID: opsUpstreamProxyID(account), + ProxyName: opsUpstreamProxyName(account), + Platform: account.Platform, + AccountID: account.ID, + AccountName: account.Name, + UpstreamStatusCode: resp.StatusCode, + UpstreamRequestID: resp.Header.Get("x-request-id"), + UpstreamURL: upstreamURL, + Kind: "response_error", + Message: err.Error(), + }) + return nil, err + } + return &SystemOneForwardResult{ + ForwardResult: ForwardResult{ + RequestID: resp.Header.Get("x-request-id"), + UpstreamHeaders: resp.Header.Clone(), + Usage: ClaudeUsage{InputTokens: decoded.Usage.InputTokens, OutputTokens: decoded.Usage.OutputTokens}, + Model: typesafe.JevLatestModel, + UpstreamResponseModel: decoded.Model, + Duration: time.Since(started), + }, + StatusCode: resp.StatusCode, + Body: decoded.Body, + ContentType: systemOneResponseContentType(resp.Header.Get("Content-Type")), + }, nil +} + +// IsSystemOneRequestErrorStatus reports upstream statuses that describe the +// caller's own payload (malformed, unprocessable, or too large). +func IsSystemOneRequestErrorStatus(status int) bool { + switch status { + case http.StatusBadRequest, http.StatusRequestEntityTooLarge, http.StatusUnprocessableEntity: + return true + default: + return false + } +} + +// handleSystemOneErrorResponse applies the shared account error policy to a +// non-2xx System One response. 400/413/422 describe the caller's own payload, so +// they never touch account state (a tenant must not be able to disable an +// account with bad input) and are not retried elsewhere. Every other status +// goes through the account error policy (custom error codes, temporary +// unschedulable rules, pool mode) and fails over when the status is retryable +// or the policy took the account out of rotation. +func (s *GatewayService) handleSystemOneErrorResponse(ctx context.Context, c *gin.Context, account *Account, resp *http.Response, upstreamURL string) error { + respBody, _ := s.readUpstreamErrorBody(resp) + upstreamMsg := sanitizeUpstreamErrorMessage(strings.TrimSpace(extractUpstreamErrorMessage(respBody))) + setOpsUpstreamError(c, resp.StatusCode, upstreamMsg, "") + event := OpsUpstreamErrorEvent{ + Passthrough: true, + ProxyID: opsUpstreamProxyID(account), + ProxyName: opsUpstreamProxyName(account), + Platform: account.Platform, + AccountID: account.ID, + AccountName: account.Name, + UpstreamStatusCode: resp.StatusCode, + UpstreamRequestID: resp.Header.Get("x-request-id"), + UpstreamURL: upstreamURL, + Kind: "http_error", + Message: upstreamMsg, + } + + if IsSystemOneRequestErrorStatus(resp.StatusCode) { + appendOpsUpstreamError(c, event) + return &SystemOneUpstreamError{StatusCode: resp.StatusCode} + } + + shouldDisable := false + if s.rateLimitService != nil { + shouldDisable = s.rateLimitService.HandleUpstreamError(ctx, account, resp.StatusCode, resp.Header, respBody, typesafe.JevLatestModel) + } + if !shouldDisable && !s.shouldFailoverUpstreamError(resp.StatusCode) { + appendOpsUpstreamError(c, event) + return &SystemOneUpstreamError{StatusCode: resp.StatusCode} + } + + event.Kind = "failover" + appendOpsUpstreamError(c, event) + failoverErr := &UpstreamFailoverError{ + StatusCode: resp.StatusCode, + ResponseBody: respBody, + ResponseHeaders: resp.Header.Clone(), + RetryableOnSameAccount: !shouldDisable && account.IsPoolMode() && account.IsPoolModeRetryableStatus(resp.StatusCode), + } + if resp.StatusCode == http.StatusUnauthorized { + failoverErr.Stage = GatewayFailureStageAccountAuth + failoverErr.Scope = GatewayFailureScopeAccount + failoverErr.Reason = TypeSafeCredentialRejectedReason + failoverErr.NextAccountAction = NextAccountRetry + } + return failoverErr +} + +// systemOneResponseContentType keeps the upstream JSON media type (and its +// charset) but never relays a non-JSON type for a body already validated as +// JSON, so the gateway origin cannot be made to serve it as HTML. +func systemOneResponseContentType(raw string) string { + mediaType, _, err := mime.ParseMediaType(strings.TrimSpace(raw)) + if err != nil { + return "application/json" + } + if mediaType == "application/json" || (strings.HasPrefix(mediaType, "application/") && strings.HasSuffix(mediaType, "+json")) { + return strings.TrimSpace(raw) + } + return "application/json" +} diff --git a/backend/internal/service/gateway_systemone_test.go b/backend/internal/service/gateway_systemone_test.go new file mode 100644 index 000000000..d3cb13141 --- /dev/null +++ b/backend/internal/service/gateway_systemone_test.go @@ -0,0 +1,402 @@ +package service + +import ( + "context" + "errors" + "io" + "net/http" + "net/http/httptest" + "strings" + "testing" + "time" + + "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/Wei-Shaw/sub2api/internal/pkg/tlsfingerprint" + "github.com/Wei-Shaw/sub2api/internal/pkg/typesafe" + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" +) + +type systemOneHTTPUpstream struct { + do func(*http.Request) (*http.Response, error) +} + +type systemOnePolicyAccountRepo struct { + AccountRepository + account *Account + rateLimitedCalls int + overloadedCalls int + errorCalls int +} + +func (r *systemOnePolicyAccountRepo) GetByID(context.Context, int64) (*Account, error) { + return r.account, nil +} + +func (r *systemOnePolicyAccountRepo) SetRateLimited(context.Context, int64, time.Time) error { + r.rateLimitedCalls++ + return nil +} + +func (r *systemOnePolicyAccountRepo) SetOverloaded(context.Context, int64, time.Time) error { + r.overloadedCalls++ + return nil +} + +func (r *systemOnePolicyAccountRepo) SetError(context.Context, int64, string) error { + r.errorCalls++ + return nil +} + +func (u *systemOneHTTPUpstream) Do(req *http.Request, _ string, _ int64, _ int) (*http.Response, error) { + return u.do(req) +} + +func (u *systemOneHTTPUpstream) DoWithTLS(req *http.Request, _ string, _ int64, _ int, _ *tlsfingerprint.Profile) (*http.Response, error) { + return u.do(req) +} + +func newSystemOneTestService(upstream HTTPUpstream) *GatewayService { + return &GatewayService{ + cfg: &config.Config{Security: config.SecurityConfig{URLAllowlist: config.URLAllowlistConfig{AllowInsecureHTTP: true}}}, + httpUpstream: upstream, + } +} + +func newSystemOneTestContext() *gin.Context { + gin.SetMode(gin.TestMode) + c, _ := gin.CreateTestContext(httptest.NewRecorder()) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/systemone", nil) + return c +} + +func TestForwardSystemOneForwardsNativeProtocolAndUsage(t *testing.T) { + requestBody := []byte(`{"model":"jev-latest","state":{"text":"sample"},"questions":{"q":{"type":"choice","instructions":"Pick","criteria":{"a":"A","b":"B"}}}}`) + responseBody := []byte(`{"model":"jev-1.13.0","answers":{"q":{"type":"choice","choice":"a"}},"usage":{"input_tokens":123,"output_tokens":7},"provider_extension":{"kept":true}}`) + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + require.Equal(t, typesafeSystemOnePathForTest, r.URL.Path) + require.Equal(t, "Bearer ts-secret", r.Header.Get("Authorization")) + require.Equal(t, "application/json", r.Header.Get("Content-Type")) + got, err := io.ReadAll(r.Body) + require.NoError(t, err) + require.Equal(t, requestBody, got) + w.Header().Set("Content-Type", "application/json; charset=utf-8") + w.Header().Set("x-request-id", "req-jev") + _, err = w.Write(responseBody) + require.NoError(t, err) + })) + defer server.Close() + + svc := newSystemOneTestService(&systemOneHTTPUpstream{do: server.Client().Do}) + account := &Account{ID: 7, Platform: PlatformTypeSafe, Type: AccountTypeAPIKey, Credentials: map[string]any{"base_url": server.URL, "api_key": "ts-secret"}} + result, err := svc.ForwardSystemOne(context.Background(), newSystemOneTestContext(), account, requestBody) + require.NoError(t, err) + require.Equal(t, responseBody, result.Body) + require.Equal(t, http.StatusOK, result.StatusCode) + require.Equal(t, "application/json; charset=utf-8", result.ContentType) + require.Equal(t, "req-jev", result.RequestID) + require.Equal(t, "jev-latest", result.Model) + require.Equal(t, "jev-1.13.0", result.UpstreamResponseModel) + require.Equal(t, 123, result.Usage.InputTokens) + require.Equal(t, 7, result.Usage.OutputTokens) +} + +func TestForwardSystemOneAllowsSuccessfulResponseWithoutModel(t *testing.T) { + responseBody := []byte(`{"answers":{"q":{"type":"noul","answer":"ok"}},"usage":{"input_tokens":12}}`) + upstream := &systemOneHTTPUpstream{do: func(*http.Request) (*http.Response, error) { + return &http.Response{StatusCode: http.StatusOK, Header: make(http.Header), Body: io.NopCloser(strings.NewReader(string(responseBody)))}, nil + }} + svc := newSystemOneTestService(upstream) + account := &Account{ID: 11, Platform: PlatformTypeSafe, Type: AccountTypeAPIKey, Credentials: map[string]any{"base_url": "http://typesafe.test", "api_key": "ts-secret"}} + + result, err := svc.ForwardSystemOne(context.Background(), newSystemOneTestContext(), account, []byte(`{}`)) + require.NoError(t, err) + require.Equal(t, responseBody, result.Body) + require.Empty(t, result.UpstreamResponseModel) + require.Equal(t, 12, result.Usage.InputTokens) +} + +const typesafeSystemOnePathForTest = "/v1/systemone" + +func TestForwardSystemOneSchemaConformancePassthrough(t *testing.T) { + for _, tc := range []struct { + name string + body string + }{ + {"omitted noul instructions", `{"model":"jev-latest","state":"sample","questions":{"q":{"type":"noul","extension":{"kept":true}}}}`}, + {"nullable noul fields", `{"model":"jev-latest","state":"sample","questions":{"q":{"type":"noul","instructions":null,"criteria":null}}}`}, + {"structured noul descriptions", `{"model":"jev-latest","state":"sample","questions":{"q":{"type":"noul","criteria":{"true":{"reason":"Yes"},"false":["No",null]}}}}`}, + {"object choice description", `{"model":"jev-latest","state":"sample","questions":{"q":{"type":"choice","criteria":{"a":{"description":"A","extra":null}}}}}`}, + {"array choice description", `{"model":"jev-latest","state":"sample","questions":{"q":{"type":"choice","instructions":null,"criteria":{"a":["A",null],"b":null}}}}`}, + {"empty choice criteria", `{"model":"jev-latest","state":"sample","questions":{"q":{"type":"choice","criteria":{}}}}`}, + {"one-level score", `{"model":"jev-latest","state":"sample","questions":{"q":{"type":"score","criteria":["only"]}}}`}, + {"object score level", `{"model":"jev-latest","state":"sample","questions":{"q":{"type":"score","instructions":null,"criteria":[{"description":"only","extra":null}]}}}`}, + {"array score level", `{"model":"jev-latest","state":"sample","questions":{"q":{"type":"score","criteria":[["only",null]]}}}`}, + } { + t.Run(tc.name, func(t *testing.T) { + requestBody := []byte(tc.body) + _, err := typesafe.ValidateSystemOneRequest(requestBody) + require.NoError(t, err) + responseBody := []byte(`{"answers":{},"usage":{"input_tokens":12}}`) + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + require.Equal(t, "/v1/systemone", r.URL.Path) + require.Equal(t, "Bearer ts-mock-key", r.Header.Get("Authorization")) + got, err := io.ReadAll(r.Body) + require.NoError(t, err) + require.Equal(t, requestBody, got) + w.Header().Set("Content-Type", "application/json") + _, err = w.Write(responseBody) + require.NoError(t, err) + })) + defer server.Close() + + svc := newSystemOneTestService(&systemOneHTTPUpstream{do: server.Client().Do}) + account := &Account{ID: 12, Platform: PlatformTypeSafe, Type: AccountTypeAPIKey, Credentials: map[string]any{"base_url": server.URL, "api_key": "ts-mock-key"}} + result, err := svc.ForwardSystemOne(context.Background(), newSystemOneTestContext(), account, requestBody) + require.NoError(t, err) + require.Equal(t, responseBody, result.Body) + require.Equal(t, 12, result.Usage.InputTokens) + }) + } +} + +func TestForwardSystemOneErrorPolicy(t *testing.T) { + for _, status := range []int{400, 422, 401, 429, 529, 500, 503} { + t.Run(http.StatusText(status), func(t *testing.T) { + upstream := &systemOneHTTPUpstream{do: func(*http.Request) (*http.Response, error) { + return &http.Response{StatusCode: status, Header: make(http.Header), Body: io.NopCloser(strings.NewReader("private upstream body ts-secret"))}, nil + }} + svc := newSystemOneTestService(upstream) + account := &Account{ID: 8, Platform: PlatformTypeSafe, Type: AccountTypeAPIKey, Credentials: map[string]any{"base_url": "http://typesafe.test", "api_key": "ts-secret"}} + result, err := svc.ForwardSystemOne(context.Background(), newSystemOneTestContext(), account, []byte(`{"private":"request"}`)) + require.Nil(t, result) + require.Error(t, err) + require.NotContains(t, err.Error(), "private") + require.NotContains(t, err.Error(), "ts-secret") + if status == 400 || status == 422 { + var upstreamErr *SystemOneUpstreamError + require.ErrorAs(t, err, &upstreamErr) + require.Equal(t, status, upstreamErr.StatusCode) + return + } + var failoverErr *UpstreamFailoverError + require.ErrorAs(t, err, &failoverErr) + require.True(t, failoverErr.ShouldRetryNextAccount()) + if status == http.StatusUnauthorized { + require.Equal(t, GatewayFailureStageAccountAuth, failoverErr.Stage) + require.Equal(t, GatewayFailureScopeAccount, failoverErr.Scope) + require.Equal(t, TypeSafeCredentialRejectedReason, failoverErr.Reason) + } + }) + } +} + +func TestForwardSystemOneAppliesExistingAccountStatePolicy(t *testing.T) { + for _, tc := range []struct { + name string + status int + wantRateLimited int + wantOverloaded int + wantError int + }{ + {name: "unauthorized", status: http.StatusUnauthorized, wantError: 1}, + {name: "rate limited", status: http.StatusTooManyRequests, wantRateLimited: 1}, + {name: "overloaded", status: 529, wantOverloaded: 1}, + } { + t.Run(tc.name, func(t *testing.T) { + account := &Account{ID: 10, Platform: PlatformTypeSafe, Type: AccountTypeAPIKey, Credentials: map[string]any{"base_url": "http://typesafe.test", "api_key": "ts-secret"}} + repo := &systemOnePolicyAccountRepo{account: account} + upstream := &systemOneHTTPUpstream{do: func(*http.Request) (*http.Response, error) { + return &http.Response{StatusCode: tc.status, Header: make(http.Header), Body: io.NopCloser(strings.NewReader(`{"error":"private"}`))}, nil + }} + svc := newSystemOneTestService(upstream) + svc.rateLimitService = NewRateLimitService(repo, nil, &config.Config{}, nil, nil) + + _, err := svc.ForwardSystemOne(context.Background(), newSystemOneTestContext(), account, []byte(`{}`)) + require.Error(t, err) + require.Equal(t, tc.wantRateLimited, repo.rateLimitedCalls) + require.Equal(t, tc.wantOverloaded, repo.overloadedCalls) + require.Equal(t, tc.wantError, repo.errorCalls) + }) + } +} + +func TestForwardSystemOneTransportAndTimeoutFailOver(t *testing.T) { + for _, transportErr := range []error{errors.New("network unavailable"), context.DeadlineExceeded} { + upstream := &systemOneHTTPUpstream{do: func(*http.Request) (*http.Response, error) { return nil, transportErr }} + svc := newSystemOneTestService(upstream) + account := &Account{ID: 9, Name: "jev", Platform: PlatformTypeSafe, Type: AccountTypeAPIKey, Credentials: map[string]any{"base_url": "http://typesafe.test", "api_key": "ts-secret"}} + _, err := svc.ForwardSystemOne(context.Background(), newSystemOneTestContext(), account, []byte(`{}`)) + var failoverErr *UpstreamFailoverError + require.ErrorAs(t, err, &failoverErr) + } +} + +func newSystemOneStatusUpstream(status int, body string) *systemOneHTTPUpstream { + return &systemOneHTTPUpstream{do: func(*http.Request) (*http.Response, error) { + return &http.Response{StatusCode: status, Header: http.Header{"X-Request-Id": []string{"req-err"}}, Body: io.NopCloser(strings.NewReader(body))}, nil + }} +} + +func systemOneOpsEvents(t *testing.T, c *gin.Context) []*OpsUpstreamErrorEvent { + t.Helper() + raw, ok := c.Get(OpsUpstreamErrorsKey) + require.True(t, ok) + events, ok := raw.([]*OpsUpstreamErrorEvent) + require.True(t, ok) + return events +} + +func TestForwardSystemOneRequestErrorsNeverTouchAccountState(t *testing.T) { + for _, status := range []int{http.StatusBadRequest, http.StatusRequestEntityTooLarge, http.StatusUnprocessableEntity} { + t.Run(http.StatusText(status), func(t *testing.T) { + // Even a custom error-code rule must not let client input disable the account. + account := &Account{ID: 21, Platform: PlatformTypeSafe, Type: AccountTypeAPIKey, Credentials: map[string]any{ + "base_url": "http://typesafe.test", "api_key": "ts-secret", + "custom_error_codes_enabled": true, "custom_error_codes": []any{float64(status)}, + }} + repo := &systemOnePolicyAccountRepo{account: account} + svc := newSystemOneTestService(newSystemOneStatusUpstream(status, `{"detail":"bad question"}`)) + svc.rateLimitService = NewRateLimitService(repo, nil, &config.Config{}, nil, nil) + c := newSystemOneTestContext() + + _, err := svc.ForwardSystemOne(context.Background(), c, account, []byte(`{}`)) + var upstreamErr *SystemOneUpstreamError + require.ErrorAs(t, err, &upstreamErr) + require.Equal(t, status, upstreamErr.StatusCode) + require.Zero(t, repo.errorCalls+repo.rateLimitedCalls+repo.overloadedCalls) + events := systemOneOpsEvents(t, c) + require.Len(t, events, 1) + require.Equal(t, "http_error", events[0].Kind) + require.Equal(t, status, events[0].UpstreamStatusCode) + require.Equal(t, "req-err", events[0].UpstreamRequestID) + }) + } +} + +func TestForwardSystemOneAccountLevelFailuresFailOver(t *testing.T) { + for _, tc := range []struct { + name string + status int + wantError int + }{ + {name: "payment required", status: http.StatusPaymentRequired, wantError: 1}, + {name: "forbidden", status: http.StatusForbidden, wantError: 1}, + } { + t.Run(tc.name, func(t *testing.T) { + account := &Account{ID: 22, Platform: PlatformTypeSafe, Type: AccountTypeAPIKey, Credentials: map[string]any{"base_url": "http://typesafe.test", "api_key": "ts-secret"}} + repo := &systemOnePolicyAccountRepo{account: account} + svc := newSystemOneTestService(newSystemOneStatusUpstream(tc.status, `{"detail":"account problem"}`)) + svc.rateLimitService = NewRateLimitService(repo, nil, &config.Config{}, nil, nil) + c := newSystemOneTestContext() + + _, err := svc.ForwardSystemOne(context.Background(), c, account, []byte(`{}`)) + var failoverErr *UpstreamFailoverError + require.ErrorAs(t, err, &failoverErr) + require.True(t, failoverErr.ShouldRetryNextAccount()) + require.Equal(t, tc.status, failoverErr.StatusCode) + require.Equal(t, tc.wantError, repo.errorCalls) + require.Equal(t, "failover", systemOneOpsEvents(t, c)[0].Kind) + }) + } +} + +func TestForwardSystemOneHonorsCustomErrorCodes(t *testing.T) { + account := &Account{ID: 23, Platform: PlatformTypeSafe, Type: AccountTypeAPIKey, Credentials: map[string]any{ + "base_url": "http://typesafe.test", "api_key": "ts-secret", + "custom_error_codes_enabled": true, "custom_error_codes": []any{float64(http.StatusNotFound)}, + }} + repo := &systemOnePolicyAccountRepo{account: account} + svc := newSystemOneTestService(newSystemOneStatusUpstream(http.StatusNotFound, `{"detail":"missing"}`)) + svc.rateLimitService = NewRateLimitService(repo, nil, &config.Config{}, nil, nil) + + _, err := svc.ForwardSystemOne(context.Background(), newSystemOneTestContext(), account, []byte(`{}`)) + var failoverErr *UpstreamFailoverError + require.ErrorAs(t, err, &failoverErr) + require.Equal(t, 1, repo.errorCalls) +} + +func TestForwardSystemOneUnhandledStatusDoesNotFailOver(t *testing.T) { + account := &Account{ID: 24, Platform: PlatformTypeSafe, Type: AccountTypeAPIKey, Credentials: map[string]any{"base_url": "http://typesafe.test", "api_key": "ts-secret"}} + repo := &systemOnePolicyAccountRepo{account: account} + svc := newSystemOneTestService(newSystemOneStatusUpstream(http.StatusConflict, `{"detail":"conflict"}`)) + svc.rateLimitService = NewRateLimitService(repo, nil, &config.Config{}, nil, nil) + + _, err := svc.ForwardSystemOne(context.Background(), newSystemOneTestContext(), account, []byte(`{}`)) + var upstreamErr *SystemOneUpstreamError + require.ErrorAs(t, err, &upstreamErr) + require.Equal(t, http.StatusConflict, upstreamErr.StatusCode) + require.Zero(t, repo.errorCalls) +} + +func TestForwardSystemOneNormalizesResponseContentType(t *testing.T) { + for _, tc := range []struct{ upstream, want string }{ + {"application/json; charset=utf-8", "application/json; charset=utf-8"}, + {"application/problem+json", "application/problem+json"}, + {"text/html; charset=utf-8", "application/json"}, + {"", "application/json"}, + {"not a media type;;", "application/json"}, + } { + upstream := &systemOneHTTPUpstream{do: func(*http.Request) (*http.Response, error) { + header := make(http.Header) + if tc.upstream != "" { + header.Set("Content-Type", tc.upstream) + } + return &http.Response{StatusCode: http.StatusOK, Header: header, Body: io.NopCloser(strings.NewReader(`{"answers":{}}`))}, nil + }} + account := &Account{ID: 25, Platform: PlatformTypeSafe, Type: AccountTypeAPIKey, Credentials: map[string]any{"base_url": "http://typesafe.test", "api_key": "ts-secret"}} + result, err := newSystemOneTestService(upstream).ForwardSystemOne(context.Background(), newSystemOneTestContext(), account, []byte(`{}`)) + require.NoError(t, err) + require.Equal(t, tc.want, result.ContentType, tc.upstream) + } +} + +func TestForwardSystemOneRejectsOversizedResponse(t *testing.T) { + oversized := `{"pad":"` + strings.Repeat("a", typesafe.MaxSystemOneResponseBytes) + `"}` + upstream := &systemOneHTTPUpstream{do: func(*http.Request) (*http.Response, error) { + return &http.Response{StatusCode: http.StatusOK, Header: make(http.Header), Body: io.NopCloser(strings.NewReader(oversized))}, nil + }} + account := &Account{ID: 26, Platform: PlatformTypeSafe, Type: AccountTypeAPIKey, Credentials: map[string]any{"base_url": "http://typesafe.test", "api_key": "ts-secret"}} + c := newSystemOneTestContext() + _, err := newSystemOneTestService(upstream).ForwardSystemOne(context.Background(), c, account, []byte(`{}`)) + require.ErrorIs(t, err, typesafe.ErrSystemOneResponseTooLarge) + events := systemOneOpsEvents(t, c) + require.Len(t, events, 1) + require.Equal(t, "response_error", events[0].Kind) + require.Equal(t, http.StatusOK, events[0].UpstreamStatusCode) +} + +func TestForwardSystemOneBillsLenientUsageShapes(t *testing.T) { + upstream := &systemOneHTTPUpstream{do: func(*http.Request) (*http.Response, error) { + return &http.Response{StatusCode: http.StatusOK, Header: make(http.Header), Body: io.NopCloser(strings.NewReader(`{"answers":{},"usage":{"input_tokens":"21","output_tokens":1.0}}`))}, nil + }} + account := &Account{ID: 27, Platform: PlatformTypeSafe, Type: AccountTypeAPIKey, Credentials: map[string]any{"base_url": "http://typesafe.test", "api_key": "ts-secret"}} + result, err := newSystemOneTestService(upstream).ForwardSystemOne(context.Background(), newSystemOneTestContext(), account, []byte(`{}`)) + require.NoError(t, err) + require.Equal(t, 21, result.Usage.InputTokens) + require.Equal(t, 1, result.Usage.OutputTokens) +} + +func TestTypeSafeAccountBaseURLNeverFallsBackToAnthropic(t *testing.T) { + account := &Account{Platform: PlatformTypeSafe, Type: AccountTypeAPIKey, Credentials: map[string]any{"api_key": "ts-secret"}} + require.Equal(t, typesafe.DefaultBaseURL, account.GetBaseURL()) + require.Equal(t, typesafe.DefaultBaseURL, account.GetTypeSafeBaseURL()) + for raw, want := range map[string]string{ + "https://api.typesafe.ai/v1": "https://api.typesafe.ai", + "https://api.typesafe.ai/V1/": "https://api.typesafe.ai", + "https://proxy.example/typesafe/": "https://proxy.example/typesafe", + "https://proxy.example/apiv1": "https://proxy.example/apiv1", + " https://proxy.example/x/v1/ ": "https://proxy.example/x", + "/v1": typesafe.DefaultBaseURL, + } { + withBase := &Account{Platform: PlatformTypeSafe, Type: AccountTypeAPIKey, Credentials: map[string]any{"base_url": raw}} + require.Equal(t, want, withBase.GetTypeSafeBaseURL(), raw) + } + anthropic := &Account{Platform: PlatformAnthropic, Type: AccountTypeAPIKey, Credentials: map[string]any{"api_key": "sk"}} + require.Equal(t, "https://api.anthropic.com", anthropic.GetBaseURL()) +} + +func TestTypeSafeModelsListCandidates(t *testing.T) { + require.Equal(t, []string{typesafe.JevLatestModel}, defaultModelsListCandidateIDs(PlatformTypeSafe)) + require.NotContains(t, compositeDefaultModelsListCandidateIDs(), typesafe.JevLatestModel) +} diff --git a/backend/internal/service/grok_upstream_headers.go b/backend/internal/service/grok_upstream_headers.go index 4b24e4c13..75e5be5c7 100644 --- a/backend/internal/service/grok_upstream_headers.go +++ b/backend/internal/service/grok_upstream_headers.go @@ -17,7 +17,7 @@ const ( grokClientModeHeader = xai.CLIClientMode ) -// defaultGrokUpstreamUserAgent is the pinned Grok CLI / workspace UA. +// defaultGrokUpstreamUserAgent 使用固定版本的官方交互式 CLI UA。 // Grok upstream must not forward Claude Code / Codex / browser client UAs. func defaultGrokUpstreamUserAgent() string { return xai.CLIUserAgent(xai.ResolveCLIVersion()) diff --git a/backend/internal/service/grok_upstream_headers_test.go b/backend/internal/service/grok_upstream_headers_test.go index 18972f480..31b51c638 100644 --- a/backend/internal/service/grok_upstream_headers_test.go +++ b/backend/internal/service/grok_upstream_headers_test.go @@ -28,7 +28,7 @@ func TestApplyDefaultGrokUpstreamHeadersUsesCLIUserAgent(t *testing.T) { } func TestApplyDefaultGrokUpstreamHeadersHonorsCLIVersionOverride(t *testing.T) { - t.Setenv(xai.CLIVersionEnv, "0.2.95") + t.Setenv(xai.CLIVersionEnv, "1.0.14") req, err := http.NewRequest(http.MethodGet, "https://api.x.ai/v1/responses", nil) require.NoError(t, err) @@ -36,9 +36,9 @@ func TestApplyDefaultGrokUpstreamHeadersHonorsCLIVersionOverride(t *testing.T) { applyDefaultGrokUpstreamHeaders(req) - require.Equal(t, "0.2.95", req.Header.Get("x-grok-client-version")) - require.Equal(t, xai.CLIUserAgent("0.2.95"), req.Header.Get("User-Agent")) - require.Equal(t, "grok-shell", req.Header.Get("x-grok-client-identifier")) + require.Equal(t, "1.0.14", req.Header.Get("x-grok-client-version")) + require.Equal(t, xai.CLIUserAgent("1.0.14"), req.Header.Get("User-Agent")) + require.Equal(t, "grok-pager", req.Header.Get("x-grok-client-identifier")) } func TestResolveGrokUpstreamUserAgentNeverPassthrough(t *testing.T) { diff --git a/backend/internal/service/openai_gateway_grok.go b/backend/internal/service/openai_gateway_grok.go index 06a17a961..7aaa649a7 100644 --- a/backend/internal/service/openai_gateway_grok.go +++ b/backend/internal/service/openai_gateway_grok.go @@ -1613,8 +1613,8 @@ func applyGrokCLIHeaders(headers http.Header) { headers.Set("X-Grok-Client-Version", version) headers.Set("x-grok-client-version", version) headers.Set("x-grok-client-identifier", xai.CLIClientIdentifier) - // Historical mode value expected by some unit tests / older CLI probes. - headers.Set("X-Grok-Client-Mode", "interactive") + // 对齐官方 CLI 交互模式,网关请求与额度探测共用身份。 + headers.Set("X-Grok-Client-Mode", xai.CLIClientMode) } func (s *OpenAIGatewayService) updateGrokUsageSnapshot(ctx context.Context, account *Account, snapshot *xai.QuotaSnapshot) { diff --git a/backend/internal/service/payment_config_service.go b/backend/internal/service/payment_config_service.go index aac1cc8e2..ad5583da2 100644 --- a/backend/internal/service/payment_config_service.go +++ b/backend/internal/service/payment_config_service.go @@ -62,12 +62,18 @@ type PaymentConfig struct { // SubscriptionUSDToCNYRate 为 0 时订阅换算关闭(兼容存量行为)。 SubscriptionUSDToCNYRate float64 `json:"subscription_usd_to_cny_rate"` RechargeFeeRate float64 `json:"recharge_fee_rate"` - LoadBalanceStrategy string `json:"load_balance_strategy"` - ProductNamePrefix string `json:"product_name_prefix"` - ProductNameSuffix string `json:"product_name_suffix"` - HelpImageURL string `json:"help_image_url"` - HelpText string `json:"help_text"` - StripePublishableKey string `json:"stripe_publishable_key,omitempty"` + // RechargeBonusTiers 余额充值优惠阶梯(按 MinAmount 升序);空表示无优惠。 + RechargeBonusTiers []RechargeBonusTier `json:"recharge_bonus_tiers"` + // RechargeBonusMode 阶梯模式:bonus(赠金)/ discount(折扣),已归一化。 + RechargeBonusMode string `json:"recharge_bonus_mode"` + // RechargeBonusNotice 充值页展示的 Markdown 活动文案;空表示不展示。 + RechargeBonusNotice string `json:"recharge_bonus_notice"` + LoadBalanceStrategy string `json:"load_balance_strategy"` + ProductNamePrefix string `json:"product_name_prefix"` + ProductNameSuffix string `json:"product_name_suffix"` + HelpImageURL string `json:"help_image_url"` + HelpText string `json:"help_text"` + StripePublishableKey string `json:"stripe_publishable_key,omitempty"` // Cancel rate limit settings CancelRateLimitEnabled bool `json:"cancel_rate_limit_enabled"` @@ -95,11 +101,15 @@ type UpdatePaymentConfigRequest struct { BalanceRechargeMultiplier *float64 `json:"balance_recharge_multiplier"` SubscriptionUSDToCNYRate *float64 `json:"subscription_usd_to_cny_rate"` RechargeFeeRate *float64 `json:"recharge_fee_rate"` - LoadBalanceStrategy *string `json:"load_balance_strategy"` - ProductNamePrefix *string `json:"product_name_prefix"` - ProductNameSuffix *string `json:"product_name_suffix"` - HelpImageURL *string `json:"help_image_url"` - HelpText *string `json:"help_text"` + // RechargeBonusTiers nil 表示不更新;空切片表示清空阶梯。 + RechargeBonusTiers *[]RechargeBonusTier `json:"recharge_bonus_tiers"` + RechargeBonusMode *string `json:"recharge_bonus_mode"` + RechargeBonusNotice *string `json:"recharge_bonus_notice"` + LoadBalanceStrategy *string `json:"load_balance_strategy"` + ProductNamePrefix *string `json:"product_name_prefix"` + ProductNameSuffix *string `json:"product_name_suffix"` + HelpImageURL *string `json:"help_image_url"` + HelpText *string `json:"help_text"` // Cancel rate limit settings CancelRateLimitEnabled *bool `json:"cancel_rate_limit_enabled"` @@ -220,6 +230,7 @@ func (s *PaymentConfigService) GetPaymentConfig(ctx context.Context) (*PaymentCo SettingPaymentEnabled, SettingMinRechargeAmount, SettingMaxRechargeAmount, SettingDailyRechargeLimit, SettingOrderTimeoutMinutes, SettingMaxPendingOrders, SettingEnabledPaymentTypes, SettingBalancePayDisabled, SettingBalanceRechargeMult, SettingSubscriptionUSDToCNYRate, SettingRechargeFeeRate, SettingLoadBalanceStrategy, + SettingRechargeBonusTiers, SettingRechargeBonusMode, SettingRechargeBonusNotice, SettingProductNamePrefix, SettingProductNameSuffix, SettingHelpImageURL, SettingHelpText, SettingCancelRateLimitOn, SettingCancelRateLimitMax, @@ -250,6 +261,8 @@ func (s *PaymentConfigService) parsePaymentConfig(vals map[string]string) *Payme BalanceRechargeMultiplier: normalizeBalanceRechargeMultiplier(pcParseFloat(vals[SettingBalanceRechargeMult], defaultBalanceRechargeMultiplier)), SubscriptionUSDToCNYRate: normalizeSubscriptionUSDToCNYRate(pcParseFloat(vals[SettingSubscriptionUSDToCNYRate], 0)), RechargeFeeRate: pcParseFloat(vals[SettingRechargeFeeRate], 0), + RechargeBonusTiers: parseRechargeBonusTiers(vals[SettingRechargeBonusTiers]), + RechargeBonusNotice: vals[SettingRechargeBonusNotice], LoadBalanceStrategy: vals[SettingLoadBalanceStrategy], ProductNamePrefix: vals[SettingProductNamePrefix], ProductNameSuffix: vals[SettingProductNameSuffix], @@ -265,6 +278,7 @@ func (s *PaymentConfigService) parsePaymentConfig(vals map[string]string) *Payme AlipayForceQRCode: vals[SettingAlipayForceQRCode] == "true", AlipayMobilePrecreateDeepLink: vals[SettingAlipayMobilePrecreateDeepLink] == "true", } + cfg.RechargeBonusMode, _ = NormalizeRechargeBonusMode(vals[SettingRechargeBonusMode]) cfg.AlipayMobilePrecreateDeepLink = pcEnvBoolOverride( SettingAlipayMobilePrecreateDeepLink, cfg.AlipayMobilePrecreateDeepLink, @@ -343,6 +357,15 @@ func (s *PaymentConfigService) UpdatePaymentConfig(ctx context.Context, req Upda return infraerrors.BadRequest("INVALID_RECHARGE_FEE_RATE", "recharge fee rate allows at most 2 decimal places") } } + rechargeBonusTiersValue, rechargeBonusModeValue, err := s.resolveRechargeBonusUpdate(ctx, req) + if err != nil { + return err + } + if req.RechargeBonusNotice != nil { + if err := validateRechargeBonusNotice(*req.RechargeBonusNotice); err != nil { + return infraerrors.BadRequest("INVALID_RECHARGE_BONUS_NOTICE", err.Error()) + } + } m := make(map[string]string) if req.Enabled != nil { m[SettingPaymentEnabled] = formatBoolOrEmpty(req.Enabled) @@ -377,6 +400,15 @@ func (s *PaymentConfigService) UpdatePaymentConfig(ctx context.Context, req Upda if req.RechargeFeeRate != nil { m[SettingRechargeFeeRate] = formatNonNegativeFloat(req.RechargeFeeRate) } + if req.RechargeBonusTiers != nil { + m[SettingRechargeBonusTiers] = rechargeBonusTiersValue + } + if req.RechargeBonusMode != nil { + m[SettingRechargeBonusMode] = rechargeBonusModeValue + } + if req.RechargeBonusNotice != nil { + m[SettingRechargeBonusNotice] = strings.TrimSpace(*req.RechargeBonusNotice) + } if req.LoadBalanceStrategy != nil { m[SettingLoadBalanceStrategy] = derefStr(req.LoadBalanceStrategy) } diff --git a/backend/internal/service/payment_fulfillment.go b/backend/internal/service/payment_fulfillment.go index 10412260f..9cd3e2859 100644 --- a/backend/internal/service/payment_fulfillment.go +++ b/backend/internal/service/payment_fulfillment.go @@ -741,7 +741,10 @@ func affiliateRebateBaseAmount(o *dbent.PaymentOrder) float64 { return 0 } switch o.OrderType { - case payment.OrderTypeBalance, payment.OrderTypeSubscription: + case payment.OrderTypeBalance: + // 返利只按实充部分计算,赠送额度不参与 + return paymentOrderAmountWithoutBonus(o) + case payment.OrderTypeSubscription: return o.Amount default: return 0 diff --git a/backend/internal/service/payment_order.go b/backend/internal/service/payment_order.go index da7178dd7..becefc6a1 100644 --- a/backend/internal/service/payment_order.go +++ b/backend/internal/service/payment_order.go @@ -53,14 +53,6 @@ func (s *PaymentService) CreateOrder(ctx context.Context, req CreateOrderRequest if s.notificationEmailService != nil { s.notificationEmailService.RememberRecipientLocale(ctx, req.UserID, user.Email, req.Locale) } - orderAmount := req.Amount - limitAmount := req.Amount - if plan != nil { - orderAmount = plan.Price - limitAmount = plan.Price - } else if req.OrderType == payment.OrderTypeBalance { - orderAmount = calculateCreditedBalance(req.Amount, cfg.BalanceRechargeMultiplier) - } feeRate := cfg.RechargeFeeRate methodCurrency := payment.DefaultPaymentCurrency if s.configService != nil { @@ -69,6 +61,19 @@ func (s *PaymentService) CreateOrder(ctx context.Context, req CreateOrderRequest return nil, err } } + orderAmount := req.Amount + limitAmount := req.Amount + bonusAmount := 0.0 + if plan != nil { + orderAmount = plan.Price + limitAmount = plan.Price + } else if req.OrderType == payment.OrderTypeBalance { + // 阈值按支付金额命中。赠金模式:到账 = 基数 + 赠送;折扣模式:到账 = 基数,实付基数按折扣减少。 + quote := quoteRechargeBonus(cfg, req.Amount, methodCurrency) + limitAmount = quote.PayBase + bonusAmount = quote.Bonus + orderAmount = quote.Credited + } payAmountStr, payAmount, err := calculateCreateOrderPayAmountForOrderType(limitAmount, feeRate, methodCurrency, req.OrderType, cfg.SubscriptionUSDToCNYRate) if err != nil { return nil, err @@ -100,7 +105,7 @@ func (s *PaymentService) CreateOrder(ctx context.Context, req CreateOrderRequest if oauthResp != nil { return oauthResp, nil } - order, err := s.createOrderInTx(ctx, req, user, plan, cfg, orderAmount, limitAmount, feeRate, payAmount, sel) + order, err := s.createOrderInTx(ctx, req, user, plan, cfg, orderAmount, limitAmount, feeRate, payAmount, bonusAmount, sel) if err != nil { return nil, err } @@ -149,7 +154,7 @@ func (s *PaymentService) validateSubOrder(ctx context.Context, req CreateOrderRe return plan, nil } -func (s *PaymentService) createOrderInTx(ctx context.Context, req CreateOrderRequest, user *User, plan *dbent.SubscriptionPlan, cfg *PaymentConfig, orderAmount, limitAmount, feeRate, payAmount float64, sel *payment.InstanceSelection) (*dbent.PaymentOrder, error) { +func (s *PaymentService) createOrderInTx(ctx context.Context, req CreateOrderRequest, user *User, plan *dbent.SubscriptionPlan, cfg *PaymentConfig, orderAmount, limitAmount, feeRate, payAmount, bonusAmount float64, sel *payment.InstanceSelection) (*dbent.PaymentOrder, error) { tx, err := s.entClient.Tx(ctx) if err != nil { return nil, fmt.Errorf("begin transaction: %w", err) @@ -185,6 +190,7 @@ func (s *PaymentService) createOrderInTx(ctx context.Context, req CreateOrderReq SetAmount(orderAmount). SetPayAmount(payAmount). SetFeeRate(feeRate). + SetBonusAmount(bonusAmount). SetRechargeCode(""). SetOutTradeNo(outTradeNo). SetPaymentType(req.PaymentType). @@ -471,6 +477,7 @@ func (s *PaymentService) invokeProvider(ctx context.Context, order *dbent.Paymen s.writeAuditLog(ctx, order.ID, "ORDER_CREATED", fmt.Sprintf("user:%d", req.UserID), map[string]any{ "paymentAmount": req.Amount, "creditedAmount": order.Amount, + "bonusAmount": order.BonusAmount, "payAmount": order.PayAmount, "paymentType": req.PaymentType, "orderType": req.OrderType, @@ -735,6 +742,7 @@ func buildCreateOrderResponse(order *dbent.PaymentOrder, req CreateOrderRequest, Amount: order.Amount, PayAmount: payAmount, FeeRate: order.FeeRate, + BonusAmount: order.BonusAmount, Status: OrderStatusPending, ResultType: resultType, PaymentType: req.PaymentType, diff --git a/backend/internal/service/payment_order_provider_snapshot_test.go b/backend/internal/service/payment_order_provider_snapshot_test.go index 127202bc2..19047959e 100644 --- a/backend/internal/service/payment_order_provider_snapshot_test.go +++ b/backend/internal/service/payment_order_provider_snapshot_test.go @@ -88,6 +88,7 @@ func TestCreateOrderInTx_WritesProviderSnapshot(t *testing.T) { 88, 0, 88, + 0, &payment.InstanceSelection{ InstanceID: strconv.FormatInt(instance.ID, 10), ProviderKey: payment.TypeAlipay, diff --git a/backend/internal/service/payment_recharge_bonus.go b/backend/internal/service/payment_recharge_bonus.go new file mode 100644 index 000000000..8d6c1525a --- /dev/null +++ b/backend/internal/service/payment_recharge_bonus.go @@ -0,0 +1,337 @@ +package service + +import ( + "context" + "encoding/json" + "fmt" + "log/slog" + "math" + "sort" + "strings" + "unicode/utf8" + + dbent "github.com/Wei-Shaw/sub2api/ent" + "github.com/Wei-Shaw/sub2api/internal/payment" + infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" + "github.com/shopspring/decimal" +) + +// 充值优惠阶梯:余额充值订单按用户输入的支付金额命中阶梯(取不超过该金额的最大 MinAmount), +// 整条阶梯只有一种模式(RECHARGE_BONUS_MODE): +// - bonus(赠金):实付不变,在到账基数(输入 × 充值倍率)之上额外赠送 BonusPercent% 的 USD 余额; +// - discount(折扣):到账不变(输入 × 倍率),实付基数按 BonusPercent% 打折。 +// +// 两种模式落库形态相同:amount 为到账总额,bonus_amount 为其中「免费」的 USD 部分,pay_amount 为实收。 +// 订阅订单不参与。优惠在下单时按当时配置计算并落库,后续改配置不影响已建订单。 +const ( + // SettingRechargeBonusTiers 存 JSON 数组(RechargeBonusTier 列表),空/缺失表示未启用优惠。 + SettingRechargeBonusTiers = "RECHARGE_BONUS_TIERS" + // SettingRechargeBonusNotice 充值页金额卡顶部展示的 Markdown 活动文案,空表示不展示。 + SettingRechargeBonusNotice = "RECHARGE_BONUS_NOTICE" + // SettingRechargeBonusMode 阶梯模式:bonus / discount;空/非法按 bonus 解析(兼容早期配置)。 + SettingRechargeBonusMode = "RECHARGE_BONUS_MODE" +) + +const ( + RechargeBonusModeBonus = "bonus" + RechargeBonusModeDiscount = "discount" +) + +const ( + maxRechargeBonusTiers = 20 + maxRechargeBonusPercent = 1000 + maxRechargeBonusNoticeRunes = 10000 + rechargeBonusAmountEpsilon = 1e-9 +) + +// RechargeBonusTier 一个优惠档位:支付金额 ≥ MinAmount 时按 BonusPercent% 赠送(bonus)或打折(discount)。 +type RechargeBonusTier struct { + MinAmount float64 `json:"min_amount"` + BonusPercent float64 `json:"bonus_percent"` +} + +// NormalizeRechargeBonusMode 归一化模式;空按 bonus。第二个返回值表示输入是否合法。 +func NormalizeRechargeBonusMode(raw string) (string, bool) { + switch strings.ToLower(strings.TrimSpace(raw)) { + case "", RechargeBonusModeBonus: + return RechargeBonusModeBonus, true + case RechargeBonusModeDiscount: + return RechargeBonusModeDiscount, true + default: + return RechargeBonusModeBonus, false + } +} + +// ValidateRechargeBonusTiersForMode 折扣模式下百分比必须 < 100,否则实付为 0 或负数。 +func ValidateRechargeBonusTiersForMode(mode string, tiers []RechargeBonusTier) error { + if mode != RechargeBonusModeDiscount { + return nil + } + for _, tier := range tiers { + if tier.BonusPercent >= 100 { + return fmt.Errorf("discount percent must be less than 100 (tier with min amount %s)", + decimal.NewFromFloat(tier.MinAmount).Round(2).String()) + } + } + return nil +} + +func rechargeBonusValueValid(v float64, max float64) bool { + if math.IsNaN(v) || math.IsInf(v, 0) || v < 0 || v > max { + return false + } + d := decimal.NewFromFloat(v) + return d.Equal(d.Round(2)) +} + +// NormalizeRechargeBonusTiers 严格归一化(写路径):任何非法项直接报错; +// 成功时返回按 MinAmount 升序排序的副本。 +func NormalizeRechargeBonusTiers(raw []RechargeBonusTier) ([]RechargeBonusTier, error) { + if len(raw) == 0 { + return []RechargeBonusTier{}, nil + } + if len(raw) > maxRechargeBonusTiers { + return nil, fmt.Errorf("recharge bonus tiers exceed limit of %d", maxRechargeBonusTiers) + } + out := make([]RechargeBonusTier, 0, len(raw)) + seen := make(map[string]struct{}, len(raw)) + for _, tier := range raw { + if !rechargeBonusValueValid(tier.MinAmount, math.MaxFloat64) { + return nil, fmt.Errorf("recharge bonus tier min amount must be a non-negative number with at most 2 decimals") + } + if !rechargeBonusValueValid(tier.BonusPercent, maxRechargeBonusPercent) { + return nil, fmt.Errorf("recharge bonus tier percent must be between 0 and %d with at most 2 decimals", maxRechargeBonusPercent) + } + key := decimal.NewFromFloat(tier.MinAmount).Round(2).String() + if _, dup := seen[key]; dup { + return nil, fmt.Errorf("duplicate recharge bonus tier min amount: %s", key) + } + seen[key] = struct{}{} + out = append(out, RechargeBonusTier{MinAmount: tier.MinAmount, BonusPercent: tier.BonusPercent}) + } + sortRechargeBonusTiers(out) + return out, nil +} + +func sortRechargeBonusTiers(tiers []RechargeBonusTier) { + sort.SliceStable(tiers, func(i, j int) bool { + return tiers[i].MinAmount < tiers[j].MinAmount + }) +} + +// encodeRechargeBonusTiers 序列化为设置值;空列表存空串,与「未配置」保持同一形态。 +func encodeRechargeBonusTiers(tiers []RechargeBonusTier) (string, error) { + if len(tiers) == 0 { + return "", nil + } + raw, err := json.Marshal(tiers) + if err != nil { + return "", fmt.Errorf("marshal recharge bonus tiers: %w", err) + } + return string(raw), nil +} + +// parseRechargeBonusTiers 宽松解析(读路径):非法条目丢弃而非报错,避免历史错配置阻断下单。 +// 同一 MinAmount 重复时保留先出现的档位。始终返回非 nil 切片,便于 JSON 输出为 []。 +func parseRechargeBonusTiers(raw string) []RechargeBonusTier { + out := make([]RechargeBonusTier, 0) + raw = strings.TrimSpace(raw) + if raw == "" { + return out + } + var items []RechargeBonusTier + if err := json.Unmarshal([]byte(raw), &items); err != nil { + slog.Warn("[Payment] parseRechargeBonusTiers: unmarshal failed", "error", err) + return out + } + seen := make(map[string]struct{}, len(items)) + for _, tier := range items { + if !rechargeBonusValueValid(tier.MinAmount, math.MaxFloat64) || !rechargeBonusValueValid(tier.BonusPercent, maxRechargeBonusPercent) { + continue + } + key := decimal.NewFromFloat(tier.MinAmount).Round(2).String() + if _, dup := seen[key]; dup { + continue + } + seen[key] = struct{}{} + out = append(out, tier) + } + sortRechargeBonusTiers(out) + return out +} + +func validateRechargeBonusNotice(notice string) error { + if utf8.RuneCountInString(notice) > maxRechargeBonusNoticeRunes { + return fmt.Errorf("recharge bonus notice exceeds %d characters", maxRechargeBonusNoticeRunes) + } + return nil +} + +// resolveRechargeBonusUpdate 校验并归一化阶梯/模式更新。任一字段缺省时读取现值做交叉校验 +// (折扣模式下所有档位百分比必须 < 100)。返回值仅在对应请求字段非 nil 时有意义。 +func (s *PaymentConfigService) resolveRechargeBonusUpdate(ctx context.Context, req UpdatePaymentConfigRequest) (tiersValue string, modeValue string, err error) { + if req.RechargeBonusTiers == nil && req.RechargeBonusMode == nil { + return "", "", nil + } + stored := map[string]string{} + if (req.RechargeBonusTiers == nil || req.RechargeBonusMode == nil) && s != nil && s.settingRepo != nil { + stored, err = s.settingRepo.GetMultiple(ctx, []string{SettingRechargeBonusTiers, SettingRechargeBonusMode}) + if err != nil { + return "", "", fmt.Errorf("get recharge bonus settings: %w", err) + } + } + + var tiers []RechargeBonusTier + if req.RechargeBonusTiers != nil { + tiers, err = NormalizeRechargeBonusTiers(*req.RechargeBonusTiers) + if err != nil { + return "", "", infraerrors.BadRequest("INVALID_RECHARGE_BONUS_TIERS", err.Error()) + } + } else { + tiers = parseRechargeBonusTiers(stored[SettingRechargeBonusTiers]) + } + + var mode string + if req.RechargeBonusMode != nil { + normalized, ok := NormalizeRechargeBonusMode(*req.RechargeBonusMode) + if !ok { + return "", "", infraerrors.BadRequest("INVALID_RECHARGE_BONUS_MODE", "recharge bonus mode must be bonus or discount") + } + mode = normalized + } else { + mode, _ = NormalizeRechargeBonusMode(stored[SettingRechargeBonusMode]) + } + + if err := ValidateRechargeBonusTiersForMode(mode, tiers); err != nil { + return "", "", infraerrors.BadRequest("INVALID_RECHARGE_BONUS_TIERS", err.Error()) + } + tiersValue, err = encodeRechargeBonusTiers(tiers) + if err != nil { + return "", "", err + } + return tiersValue, mode, nil +} + +// matchRechargeBonusTier 返回不超过 paymentAmount 的最大档位;tiers 需已按 MinAmount 升序。 +func matchRechargeBonusTier(tiers []RechargeBonusTier, paymentAmount float64) (RechargeBonusTier, bool) { + if math.IsNaN(paymentAmount) || math.IsInf(paymentAmount, 0) || paymentAmount <= 0 { + return RechargeBonusTier{}, false + } + var matched RechargeBonusTier + found := false + for _, tier := range tiers { + if paymentAmount+rechargeBonusAmountEpsilon < tier.MinAmount { + break + } + matched = tier + found = true + } + return matched, found +} + +// calculateRechargeBonus 赠送金额 = 到账基数 × 百分比,保留两位小数(四舍五入)。 +func calculateRechargeBonus(baseCredited, bonusPercent float64) float64 { + if baseCredited <= 0 || bonusPercent <= 0 || math.IsNaN(baseCredited) || math.IsNaN(bonusPercent) { + return 0 + } + return decimal.NewFromFloat(baseCredited). + Mul(decimal.NewFromFloat(bonusPercent)). + Div(decimal.NewFromInt(100)). + Round(2). + InexactFloat64() +} + +// addRechargeBonus 到账总额 = 基数 + 赠送,两位小数。 +func addRechargeBonus(baseCredited, bonus float64) float64 { + return decimal.NewFromFloat(baseCredited). + Add(decimal.NewFromFloat(bonus)). + Round(2). + InexactFloat64() +} + +// calculateDiscountedPayBase 折扣模式实付基数 = 支付金额 × (1 − 百分比),按币种精度四舍五入。 +func calculateDiscountedPayBase(paymentAmount, discountPercent float64, currency string) float64 { + digits := int32(payment.CurrencyMaxFractionDigits(currency)) + return decimal.NewFromFloat(paymentAmount). + Mul(decimal.NewFromInt(100).Sub(decimal.NewFromFloat(discountPercent))). + Div(decimal.NewFromInt(100)). + Round(digits). + InexactFloat64() +} + +// rechargeBonusQuote 一笔余额充值的报价结果。 +type rechargeBonusQuote struct { + // PayBase 网关收款基数(支付币种,不含手续费);赠金模式等于支付金额,折扣模式为折后金额。 + PayBase float64 + // Credited 到账总额(USD),含 Bonus。 + Credited float64 + // Bonus 免费额度(USD):赠金模式为额外赠送,折扣模式为未付费却到账的部分。 + Bonus float64 + // Percent 命中档位的百分比;未命中或未产生优惠时为 0。 + Percent float64 +} + +// quoteRechargeBonus 按配置模式报价。currency 用于折扣模式实付基数的精度。 +// 未配置阶梯、未命中、或折扣百分比 ≥ 100(非法历史数据,fail-safe)时按无优惠处理。 +func quoteRechargeBonus(cfg *PaymentConfig, paymentAmount float64, currency string) rechargeBonusQuote { + multiplier := defaultBalanceRechargeMultiplier + var tiers []RechargeBonusTier + mode := RechargeBonusModeBonus + if cfg != nil { + multiplier = cfg.BalanceRechargeMultiplier + tiers = cfg.RechargeBonusTiers + mode, _ = NormalizeRechargeBonusMode(cfg.RechargeBonusMode) + } + base := calculateCreditedBalance(paymentAmount, multiplier) + quote := rechargeBonusQuote{PayBase: paymentAmount, Credited: base} + + tier, ok := matchRechargeBonusTier(tiers, paymentAmount) + if !ok || tier.BonusPercent <= 0 { + return quote + } + switch mode { + case RechargeBonusModeDiscount: + if tier.BonusPercent >= 100 { + return quote + } + payBase := calculateDiscountedPayBase(paymentAmount, tier.BonusPercent, currency) + if payBase <= 0 || payBase >= paymentAmount { + return quote + } + paidCredit := calculateCreditedBalance(payBase, multiplier) + bonus := decimal.NewFromFloat(base).Sub(decimal.NewFromFloat(paidCredit)).Round(2).InexactFloat64() + if bonus < 0 { + bonus = 0 + } + quote.PayBase = payBase + quote.Bonus = bonus + quote.Percent = tier.BonusPercent + default: + bonus := calculateRechargeBonus(base, tier.BonusPercent) + if bonus <= 0 { + return quote + } + quote.Bonus = bonus + quote.Credited = addRechargeBonus(base, bonus) + quote.Percent = tier.BonusPercent + } + return quote +} + +// paymentOrderAmountWithoutBonus 订单到账金额剔除免费额度后的实付部分(USD),用于推广返利基数。 +func paymentOrderAmountWithoutBonus(o *dbent.PaymentOrder) float64 { + if o == nil { + return 0 + } + if o.OrderType != payment.OrderTypeBalance || o.BonusAmount <= 0 { + return o.Amount + } + base := decimal.NewFromFloat(o.Amount). + Sub(decimal.NewFromFloat(o.BonusAmount)). + Round(2). + InexactFloat64() + if base < 0 { + return 0 + } + return base +} diff --git a/backend/internal/service/payment_recharge_bonus_test.go b/backend/internal/service/payment_recharge_bonus_test.go new file mode 100644 index 000000000..00af53bfe --- /dev/null +++ b/backend/internal/service/payment_recharge_bonus_test.go @@ -0,0 +1,349 @@ +//go:build unit + +package service + +import ( + "context" + "testing" + + dbent "github.com/Wei-Shaw/sub2api/ent" + "github.com/Wei-Shaw/sub2api/internal/payment" + "github.com/stretchr/testify/require" +) + +func TestNormalizeRechargeBonusTiers(t *testing.T) { + t.Run("sorts ascending by min amount", func(t *testing.T) { + out, err := NormalizeRechargeBonusTiers([]RechargeBonusTier{ + {MinAmount: 1000, BonusPercent: 35}, + {MinAmount: 100, BonusPercent: 20}, + {MinAmount: 500, BonusPercent: 30}, + }) + require.NoError(t, err) + require.Equal(t, []RechargeBonusTier{ + {MinAmount: 100, BonusPercent: 20}, + {MinAmount: 500, BonusPercent: 30}, + {MinAmount: 1000, BonusPercent: 35}, + }, out) + }) + + t.Run("empty input yields empty non-nil slice", func(t *testing.T) { + out, err := NormalizeRechargeBonusTiers(nil) + require.NoError(t, err) + require.NotNil(t, out) + require.Len(t, out, 0) + }) + + t.Run("allows zero threshold and zero percent", func(t *testing.T) { + out, err := NormalizeRechargeBonusTiers([]RechargeBonusTier{{MinAmount: 0, BonusPercent: 0}}) + require.NoError(t, err) + require.Len(t, out, 1) + }) + + t.Run("rejects invalid values", func(t *testing.T) { + cases := map[string][]RechargeBonusTier{ + "negative min": {{MinAmount: -1, BonusPercent: 10}}, + "min three decimals": {{MinAmount: 100.123, BonusPercent: 10}}, + "negative percent": {{MinAmount: 100, BonusPercent: -5}}, + "percent over limit": {{MinAmount: 100, BonusPercent: 1000.01}}, + "percent 3 decimals": {{MinAmount: 100, BonusPercent: 12.345}}, + "duplicate min": {{MinAmount: 100, BonusPercent: 10}, {MinAmount: 100, BonusPercent: 20}}, + "duplicate min 2 dec": {{MinAmount: 100, BonusPercent: 10}, {MinAmount: 100.00, BonusPercent: 20}}, + } + for name, tiers := range cases { + _, err := NormalizeRechargeBonusTiers(tiers) + require.Error(t, err, name) + } + }) + + t.Run("rejects too many tiers", func(t *testing.T) { + tiers := make([]RechargeBonusTier, 0, maxRechargeBonusTiers+1) + for i := 0; i <= maxRechargeBonusTiers; i++ { + tiers = append(tiers, RechargeBonusTier{MinAmount: float64(i + 1), BonusPercent: 1}) + } + _, err := NormalizeRechargeBonusTiers(tiers) + require.Error(t, err) + }) +} + +func TestParseRechargeBonusTiers(t *testing.T) { + t.Run("empty or invalid json yields empty slice", func(t *testing.T) { + require.NotNil(t, parseRechargeBonusTiers("")) + require.Len(t, parseRechargeBonusTiers(""), 0) + require.Len(t, parseRechargeBonusTiers("not json"), 0) + }) + + t.Run("drops invalid entries keeps first duplicate and sorts", func(t *testing.T) { + raw := `[{"min_amount":500,"bonus_percent":30},{"min_amount":-1,"bonus_percent":5},` + + `{"min_amount":100,"bonus_percent":20},{"min_amount":100,"bonus_percent":99},` + + `{"min_amount":50,"bonus_percent":5000}]` + out := parseRechargeBonusTiers(raw) + require.Equal(t, []RechargeBonusTier{ + {MinAmount: 100, BonusPercent: 20}, + {MinAmount: 500, BonusPercent: 30}, + }, out) + }) + + t.Run("round trips encode", func(t *testing.T) { + tiers := []RechargeBonusTier{{MinAmount: 100, BonusPercent: 20}, {MinAmount: 500, BonusPercent: 30}} + encoded, err := encodeRechargeBonusTiers(tiers) + require.NoError(t, err) + require.Equal(t, tiers, parseRechargeBonusTiers(encoded)) + + empty, err := encodeRechargeBonusTiers(nil) + require.NoError(t, err) + require.Equal(t, "", empty) + }) +} + +func TestMatchRechargeBonusTier(t *testing.T) { + tiers := []RechargeBonusTier{ + {MinAmount: 100, BonusPercent: 20}, + {MinAmount: 500, BonusPercent: 30}, + {MinAmount: 1000, BonusPercent: 35}, + } + cases := []struct { + amount float64 + percent float64 + ok bool + }{ + {amount: 0, ok: false}, + {amount: 99.99, ok: false}, + {amount: 100, percent: 20, ok: true}, + {amount: 499.99, percent: 20, ok: true}, + {amount: 500, percent: 30, ok: true}, + {amount: 1000, percent: 35, ok: true}, + {amount: 1000000, percent: 35, ok: true}, + } + for _, tc := range cases { + tier, ok := matchRechargeBonusTier(tiers, tc.amount) + require.Equal(t, tc.ok, ok, "amount %v", tc.amount) + if ok { + require.Equal(t, tc.percent, tier.BonusPercent, "amount %v", tc.amount) + } + } + + t.Run("float boundary 0.1+0.2 still matches 0.3 threshold", func(t *testing.T) { + _, ok := matchRechargeBonusTier([]RechargeBonusTier{{MinAmount: 0.3, BonusPercent: 1}}, 0.1+0.2) + require.True(t, ok) + }) + + t.Run("no tiers never matches", func(t *testing.T) { + _, ok := matchRechargeBonusTier(nil, 100) + require.False(t, ok) + }) +} + +func TestCalculateRechargeBonusRounding(t *testing.T) { + require.Equal(t, 20.0, calculateRechargeBonus(100, 20)) + require.Equal(t, 120.0, addRechargeBonus(100, 20)) + // 33.33 * 15% = 4.9995 → 5.00 + require.Equal(t, 5.0, calculateRechargeBonus(33.33, 15)) + require.Zero(t, calculateRechargeBonus(100, 0)) + require.Zero(t, calculateRechargeBonus(0, 20)) + + // 阈值按支付金额命中,赠送按到账基数计算:1000 CNY × 0.14 = 140 USD,命中 1000 档 30% → 42 + cfg := &PaymentConfig{ + BalanceRechargeMultiplier: 0.14, + RechargeBonusTiers: []RechargeBonusTier{{MinAmount: 100, BonusPercent: 20}, {MinAmount: 500, BonusPercent: 30}}, + } + require.Equal(t, rechargeBonusQuote{PayBase: 1000, Credited: 182, Bonus: 42, Percent: 30}, quoteRechargeBonus(cfg, 1000, "CNY")) + // 命中 0% 档位视为无优惠 + cfg.RechargeBonusTiers = []RechargeBonusTier{{MinAmount: 10, BonusPercent: 0}} + require.Equal(t, rechargeBonusQuote{PayBase: 50, Credited: 7}, quoteRechargeBonus(cfg, 50, "CNY")) +} + +func TestParsePaymentConfigRechargeBonus(t *testing.T) { + svc := &PaymentConfigService{} + + t.Run("defaults", func(t *testing.T) { + cfg := svc.parsePaymentConfig(map[string]string{}) + require.NotNil(t, cfg.RechargeBonusTiers) + require.Len(t, cfg.RechargeBonusTiers, 0) + require.Equal(t, "", cfg.RechargeBonusNotice) + }) + + t.Run("reads tiers and notice", func(t *testing.T) { + cfg := svc.parsePaymentConfig(map[string]string{ + SettingRechargeBonusTiers: `[{"min_amount":500,"bonus_percent":30},{"min_amount":100,"bonus_percent":20}]`, + SettingRechargeBonusNotice: "**满 100 送 20%**", + }) + require.Equal(t, []RechargeBonusTier{{MinAmount: 100, BonusPercent: 20}, {MinAmount: 500, BonusPercent: 30}}, cfg.RechargeBonusTiers) + require.Equal(t, "**满 100 送 20%**", cfg.RechargeBonusNotice) + }) +} + +func TestUpdatePaymentConfigRechargeBonus(t *testing.T) { + ctx := context.Background() + + t.Run("persists normalized tiers and trimmed notice", func(t *testing.T) { + repo := &paymentConfigSettingRepoStub{values: map[string]string{}} + svc := &PaymentConfigService{settingRepo: repo} + tiers := []RechargeBonusTier{{MinAmount: 500, BonusPercent: 30}, {MinAmount: 100, BonusPercent: 20}} + notice := " 活动文案 " + require.NoError(t, svc.UpdatePaymentConfig(ctx, UpdatePaymentConfigRequest{ + RechargeBonusTiers: &tiers, + RechargeBonusNotice: ¬ice, + })) + require.Equal(t, `[{"min_amount":100,"bonus_percent":20},{"min_amount":500,"bonus_percent":30}]`, repo.updates[SettingRechargeBonusTiers]) + require.Equal(t, "活动文案", repo.updates[SettingRechargeBonusNotice]) + + cfg, err := svc.GetPaymentConfig(ctx) + require.NoError(t, err) + require.Equal(t, []RechargeBonusTier{{MinAmount: 100, BonusPercent: 20}, {MinAmount: 500, BonusPercent: 30}}, cfg.RechargeBonusTiers) + }) + + t.Run("empty tiers clears setting and omitted fields are untouched", func(t *testing.T) { + repo := &paymentConfigSettingRepoStub{values: map[string]string{ + SettingRechargeBonusTiers: `[{"min_amount":100,"bonus_percent":20}]`, + SettingRechargeBonusNotice: "keep me", + }} + svc := &PaymentConfigService{settingRepo: repo} + empty := []RechargeBonusTier{} + require.NoError(t, svc.UpdatePaymentConfig(ctx, UpdatePaymentConfigRequest{RechargeBonusTiers: &empty})) + value, ok := repo.updates[SettingRechargeBonusTiers] + require.True(t, ok) + require.Equal(t, "", value) + _, touched := repo.updates[SettingRechargeBonusNotice] + require.False(t, touched) + require.Equal(t, "keep me", repo.values[SettingRechargeBonusNotice]) + }) + + t.Run("rejects invalid tiers", func(t *testing.T) { + repo := &paymentConfigSettingRepoStub{values: map[string]string{}} + svc := &PaymentConfigService{settingRepo: repo} + bad := []RechargeBonusTier{{MinAmount: 100, BonusPercent: 20}, {MinAmount: 100, BonusPercent: 30}} + err := svc.UpdatePaymentConfig(ctx, UpdatePaymentConfigRequest{RechargeBonusTiers: &bad}) + require.Error(t, err) + require.Nil(t, repo.updates) + }) +} + +func TestAffiliateRebateBaseAmountExcludesRechargeBonus(t *testing.T) { + require.Equal(t, 100.0, affiliateRebateBaseAmount(&dbent.PaymentOrder{ + OrderType: payment.OrderTypeBalance, Amount: 130, BonusAmount: 30, + })) + require.Equal(t, 130.0, affiliateRebateBaseAmount(&dbent.PaymentOrder{ + OrderType: payment.OrderTypeBalance, Amount: 130, + })) + // 订阅订单不受 bonus 字段影响 + require.Equal(t, 50.0, affiliateRebateBaseAmount(&dbent.PaymentOrder{ + OrderType: payment.OrderTypeSubscription, Amount: 50, BonusAmount: 30, + })) + // 异常数据:赠送大于总额时钳到 0 + require.Equal(t, 0.0, affiliateRebateBaseAmount(&dbent.PaymentOrder{ + OrderType: payment.OrderTypeBalance, Amount: 10, BonusAmount: 30, + })) +} + +func TestNormalizeRechargeBonusMode(t *testing.T) { + for raw, want := range map[string]string{"": RechargeBonusModeBonus, "bonus": RechargeBonusModeBonus, " Discount ": RechargeBonusModeDiscount} { + mode, ok := NormalizeRechargeBonusMode(raw) + require.True(t, ok, raw) + require.Equal(t, want, mode, raw) + } + mode, ok := NormalizeRechargeBonusMode("cashback") + require.False(t, ok) + require.Equal(t, RechargeBonusModeBonus, mode) +} + +func TestValidateRechargeBonusTiersForMode(t *testing.T) { + tiers := []RechargeBonusTier{{MinAmount: 100, BonusPercent: 20}, {MinAmount: 500, BonusPercent: 100}} + require.NoError(t, ValidateRechargeBonusTiersForMode(RechargeBonusModeBonus, tiers)) + require.Error(t, ValidateRechargeBonusTiersForMode(RechargeBonusModeDiscount, tiers)) + require.NoError(t, ValidateRechargeBonusTiersForMode(RechargeBonusModeDiscount, tiers[:1])) +} + +func TestQuoteRechargeBonus(t *testing.T) { + tiers := []RechargeBonusTier{{MinAmount: 100, BonusPercent: 20}, {MinAmount: 500, BonusPercent: 50}} + + t.Run("bonus mode keeps pay base and inflates credit", func(t *testing.T) { + cfg := &PaymentConfig{BalanceRechargeMultiplier: 1, RechargeBonusTiers: tiers, RechargeBonusMode: RechargeBonusModeBonus} + q := quoteRechargeBonus(cfg, 100, "USD") + require.Equal(t, rechargeBonusQuote{PayBase: 100, Credited: 120, Bonus: 20, Percent: 20}, q) + + // 未命中:无优惠 + q = quoteRechargeBonus(cfg, 50, "USD") + require.Equal(t, rechargeBonusQuote{PayBase: 50, Credited: 50}, q) + }) + + t.Run("discount mode keeps credit and reduces pay base", func(t *testing.T) { + cfg := &PaymentConfig{BalanceRechargeMultiplier: 1, RechargeBonusTiers: tiers, RechargeBonusMode: RechargeBonusModeDiscount} + q := quoteRechargeBonus(cfg, 500, "USD") + require.Equal(t, rechargeBonusQuote{PayBase: 250, Credited: 500, Bonus: 250, Percent: 50}, q) + + // 倍率 0.14:1000 CNY 到账 140 USD;20% off 实付 800 CNY,免费部分 = 140 − 112 = 28 USD + cfg.BalanceRechargeMultiplier = 0.14 + q = quoteRechargeBonus(cfg, 1000, "CNY") + require.Equal(t, rechargeBonusQuote{PayBase: 500, Credited: 140, Bonus: 70, Percent: 50}, q) + q = quoteRechargeBonus(cfg, 200, "CNY") + require.Equal(t, rechargeBonusQuote{PayBase: 160, Credited: 28, Bonus: 5.6, Percent: 20}, q) + }) + + t.Run("discount rounds pay base to currency precision", func(t *testing.T) { + cfg := &PaymentConfig{BalanceRechargeMultiplier: 1, RechargeBonusTiers: []RechargeBonusTier{{MinAmount: 1, BonusPercent: 15}}, RechargeBonusMode: RechargeBonusModeDiscount} + require.Equal(t, 85.85, quoteRechargeBonus(cfg, 101, "USD").PayBase) + require.Equal(t, 86.0, quoteRechargeBonus(cfg, 101, "JPY").PayBase) + }) + + t.Run("discount percent at or above 100 is ignored fail-safe", func(t *testing.T) { + cfg := &PaymentConfig{BalanceRechargeMultiplier: 1, RechargeBonusTiers: []RechargeBonusTier{{MinAmount: 1, BonusPercent: 100}}, RechargeBonusMode: RechargeBonusModeDiscount} + require.Equal(t, rechargeBonusQuote{PayBase: 100, Credited: 100}, quoteRechargeBonus(cfg, 100, "USD")) + }) + + t.Run("nil config and empty tiers yield plain conversion", func(t *testing.T) { + require.Equal(t, rechargeBonusQuote{PayBase: 100, Credited: 100}, quoteRechargeBonus(nil, 100, "USD")) + require.Equal(t, rechargeBonusQuote{PayBase: 100, Credited: 14}, quoteRechargeBonus(&PaymentConfig{BalanceRechargeMultiplier: 0.14}, 100, "CNY")) + }) +} + +func TestParsePaymentConfigRechargeBonusMode(t *testing.T) { + svc := &PaymentConfigService{} + require.Equal(t, RechargeBonusModeBonus, svc.parsePaymentConfig(map[string]string{}).RechargeBonusMode) + require.Equal(t, RechargeBonusModeDiscount, svc.parsePaymentConfig(map[string]string{SettingRechargeBonusMode: "discount"}).RechargeBonusMode) + require.Equal(t, RechargeBonusModeBonus, svc.parsePaymentConfig(map[string]string{SettingRechargeBonusMode: "junk"}).RechargeBonusMode) +} + +func TestUpdatePaymentConfigRechargeBonusMode(t *testing.T) { + ctx := context.Background() + + t.Run("persists discount mode with valid tiers", func(t *testing.T) { + repo := &paymentConfigSettingRepoStub{values: map[string]string{}} + svc := &PaymentConfigService{settingRepo: repo} + mode := "discount" + tiers := []RechargeBonusTier{{MinAmount: 100, BonusPercent: 20}} + require.NoError(t, svc.UpdatePaymentConfig(ctx, UpdatePaymentConfigRequest{RechargeBonusTiers: &tiers, RechargeBonusMode: &mode})) + require.Equal(t, "discount", repo.updates[SettingRechargeBonusMode]) + cfg, err := svc.GetPaymentConfig(ctx) + require.NoError(t, err) + require.Equal(t, RechargeBonusModeDiscount, cfg.RechargeBonusMode) + }) + + t.Run("rejects unknown mode", func(t *testing.T) { + repo := &paymentConfigSettingRepoStub{values: map[string]string{}} + svc := &PaymentConfigService{settingRepo: repo} + mode := "cashback" + require.Error(t, svc.UpdatePaymentConfig(ctx, UpdatePaymentConfigRequest{RechargeBonusMode: &mode})) + require.Nil(t, repo.updates) + }) + + t.Run("switching to discount with stored tiers at 100 percent is rejected", func(t *testing.T) { + repo := &paymentConfigSettingRepoStub{values: map[string]string{ + SettingRechargeBonusTiers: `[{"min_amount":100,"bonus_percent":100}]`, + }} + svc := &PaymentConfigService{settingRepo: repo} + mode := "discount" + require.Error(t, svc.UpdatePaymentConfig(ctx, UpdatePaymentConfigRequest{RechargeBonusMode: &mode})) + require.Nil(t, repo.updates) + }) + + t.Run("saving tiers at 100 percent while stored mode is discount is rejected", func(t *testing.T) { + repo := &paymentConfigSettingRepoStub{values: map[string]string{SettingRechargeBonusMode: "discount"}} + svc := &PaymentConfigService{settingRepo: repo} + tiers := []RechargeBonusTier{{MinAmount: 100, BonusPercent: 100}} + require.Error(t, svc.UpdatePaymentConfig(ctx, UpdatePaymentConfigRequest{RechargeBonusTiers: &tiers})) + require.Nil(t, repo.updates) + // 同样的档位在赠金模式下合法 + repo.values[SettingRechargeBonusMode] = "bonus" + require.NoError(t, svc.UpdatePaymentConfig(ctx, UpdatePaymentConfigRequest{RechargeBonusTiers: &tiers})) + }) +} diff --git a/backend/internal/service/payment_service.go b/backend/internal/service/payment_service.go index 792a842a2..dae1a7d93 100644 --- a/backend/internal/service/payment_service.go +++ b/backend/internal/service/payment_service.go @@ -92,6 +92,7 @@ type CreateOrderResponse struct { Amount float64 `json:"amount"` PayAmount float64 `json:"pay_amount"` FeeRate float64 `json:"fee_rate"` + BonusAmount float64 `json:"bonus_amount"` Status string `json:"status"` ResultType payment.CreatePaymentResultType `json:"result_type,omitempty"` PaymentType string `json:"payment_type"` diff --git a/backend/internal/service/scheduler_snapshot_service.go b/backend/internal/service/scheduler_snapshot_service.go index 14d43f0c8..6f999c189 100644 --- a/backend/internal/service/scheduler_snapshot_service.go +++ b/backend/internal/service/scheduler_snapshot_service.go @@ -609,7 +609,7 @@ func (s *SchedulerSnapshotService) handleBulkAccountEvent(ctx context.Context, p } accountGroupIDs := s.normalizeGroupIDs(account.GroupIDs) switch account.Platform { - case PlatformAnthropic, PlatformGemini, PlatformOpenAI, PlatformGrok, PlatformKimi, PlatformZhipu, PlatformDeepseek, PlatformMiniMax, PlatformOpenCodeGo: + case PlatformAnthropic, PlatformGemini, PlatformOpenAI, PlatformGrok, PlatformKimi, PlatformZhipu, PlatformDeepseek, PlatformMiniMax, PlatformOpenCodeGo, PlatformTypeSafe: addPlatformGroups(account.Platform, accountGroupIDs) case PlatformAntigravity: // 批量更新可能刚关闭 mixed_scheduling,仍需清理两个兼容平台的旧快照。 @@ -824,8 +824,8 @@ func (s *SchedulerSnapshotService) rebuildByAccount(ctx context.Context, account return s.rebuildBuckets(ctx, buckets, reason) } -func schedulerSnapshotPlatforms() [10]string { - return [10]string{PlatformAnthropic, PlatformGemini, PlatformOpenAI, PlatformAntigravity, PlatformGrok, PlatformKimi, PlatformZhipu, PlatformDeepseek, PlatformMiniMax, PlatformOpenCodeGo} +func schedulerSnapshotPlatforms() [11]string { + return [11]string{PlatformAnthropic, PlatformGemini, PlatformOpenAI, PlatformAntigravity, PlatformGrok, PlatformKimi, PlatformZhipu, PlatformDeepseek, PlatformMiniMax, PlatformOpenCodeGo, PlatformTypeSafe} } // 生命周期辅助函数有意排除 group0;full rebuild 构造 group0 canonical 集时必须显式调用 canonical helper。 diff --git a/backend/internal/service/user_service.go b/backend/internal/service/user_service.go index 8b6833a2b..ac9388637 100644 --- a/backend/internal/service/user_service.go +++ b/backend/internal/service/user_service.go @@ -4,7 +4,6 @@ import ( "bytes" "context" "crypto/sha256" - "crypto/subtle" "encoding/base64" "encoding/hex" "fmt" @@ -1326,28 +1325,7 @@ func (s *UserService) VerifyAndAddNotifyEmail(ctx context.Context, userID int64, // verifyNotifyCode validates the verification code against the cached data. func verifyNotifyCode(ctx context.Context, cache EmailCache, email, code string) error { - data, err := cache.GetNotifyVerifyCode(ctx, email) - if err != nil || data == nil { - return ErrInvalidVerifyCode - } - if data.Attempts >= maxVerifyCodeAttempts { - return ErrVerifyCodeMaxAttempts - } - if subtle.ConstantTimeCompare([]byte(data.Code), []byte(code)) != 1 { - data.Attempts++ - remaining := time.Until(data.ExpiresAt) - if remaining <= 0 { - return ErrInvalidVerifyCode - } - if err := cache.SetNotifyVerifyCode(ctx, email, data, remaining); err != nil { - slog.Error("failed to update notify verify code attempts", "email", email, "error", err) - } - if data.Attempts >= maxVerifyCodeAttempts { - return ErrVerifyCodeMaxAttempts - } - return ErrInvalidVerifyCode - } - return nil + return verifyCodeWithAttempts(ctx, email, code, cache.GetNotifyVerifyCode, cache.IncrNotifyVerifyCodeAttempts, nil) } // addOrVerifyNotifyEmail adds the email to user's extra notification emails or marks it as verified. diff --git a/backend/migrations/241_add_payment_order_bonus_amount.sql b/backend/migrations/241_add_payment_order_bonus_amount.sql new file mode 100644 index 000000000..1eeddfadb --- /dev/null +++ b/backend/migrations/241_add_payment_order_bonus_amount.sql @@ -0,0 +1,3 @@ +-- 充值赠送额度:余额充值订单命中赠送阶梯时的赠送 USD 金额。 +-- 已计入 payment_orders.amount(到账总额),单独落列用于订单展示与推广返利基数剔除。 +ALTER TABLE payment_orders ADD COLUMN IF NOT EXISTS bonus_amount DECIMAL(20,2) NOT NULL DEFAULT 0; diff --git a/backend/migrations/241_add_typesafe_platform.sql b/backend/migrations/241_add_typesafe_platform.sql new file mode 100644 index 000000000..82a81de70 --- /dev/null +++ b/backend/migrations/241_add_typesafe_platform.sql @@ -0,0 +1,26 @@ +-- Add TypeSafe (Jev System One) as a first-class platform. +-- +-- 1. user_platform_quotas.platform CHECK +-- 2. composite_model_routes.target_platform CHECK +-- +-- TypeSafe 不是对话模型,不进入渠道监控 provider,因此 channel_monitors / +-- channel_monitor_request_templates 的约束保持不变。 +-- +-- Runs after 238_opencode_go_platform.sql. DROP ... IF EXISTS 保证可重入; +-- 新约束是 238 的超集,存量行瞬时校验通过。 + +ALTER TABLE user_platform_quotas + DROP CONSTRAINT IF EXISTS user_platform_quotas_platform_check; + +ALTER TABLE user_platform_quotas + ADD CONSTRAINT user_platform_quotas_platform_check + CHECK (platform IN ('anthropic', 'openai', 'gemini', 'antigravity', 'grok', + 'kimi', 'zhipu', 'deepseek', 'minimax', 'opencode_go', 'typesafe')); + +ALTER TABLE composite_model_routes + DROP CONSTRAINT IF EXISTS composite_model_routes_target_platform_check; + +ALTER TABLE composite_model_routes + ADD CONSTRAINT composite_model_routes_target_platform_check + CHECK (target_platform IN ('anthropic', 'openai', 'gemini', 'antigravity', 'grok', + 'kimi', 'zhipu', 'deepseek', 'minimax', 'opencode_go', 'typesafe')); diff --git a/backend/migrations/typesafe_platform_migration_test.go b/backend/migrations/typesafe_platform_migration_test.go new file mode 100644 index 000000000..7eb238091 --- /dev/null +++ b/backend/migrations/typesafe_platform_migration_test.go @@ -0,0 +1,21 @@ +package migrations + +import ( + "strings" + "testing" + + "github.com/stretchr/testify/require" +) + +func TestTypeSafePlatformMigration(t *testing.T) { + content, err := FS.ReadFile("241_add_typesafe_platform.sql") + require.NoError(t, err) + + sql := strings.Join(strings.Fields(string(content)), " ") + require.Contains(t, sql, "DROP CONSTRAINT IF EXISTS user_platform_quotas_platform_check") + require.Contains(t, sql, "DROP CONSTRAINT IF EXISTS composite_model_routes_target_platform_check") + require.Contains(t, sql, + "CHECK (platform IN ('anthropic', 'openai', 'gemini', 'antigravity', 'grok', 'kimi', 'zhipu', 'deepseek', 'minimax', 'opencode_go', 'typesafe'))") + require.Contains(t, sql, + "CHECK (target_platform IN ('anthropic', 'openai', 'gemini', 'antigravity', 'grok', 'kimi', 'zhipu', 'deepseek', 'minimax', 'opencode_go', 'typesafe'))") +} diff --git a/docs/TECHNICAL-PLAN.md b/docs/TECHNICAL-PLAN.md index 91d4198d7..62ad3aed5 100644 --- a/docs/TECHNICAL-PLAN.md +++ b/docs/TECHNICAL-PLAN.md @@ -3,7 +3,7 @@ ## Baseline And Change Boundary 项目的稳定技术基线为 -[`Wei-Shaw/sub2api v0.2.4@5de5e2bed035d43591a2e10e51f420ef6a84eb98`](https://github.com/Wei-Shaw/sub2api/tree/5de5e2bed035d43591a2e10e51f420ef6a84eb98)。 +[`Wei-Shaw/sub2api v0.2.12@5106065716e494204fc0e8db16f68f6e9d576be0`](https://github.com/Wei-Shaw/sub2api/tree/5106065716e494204fc0e8db16f68f6e9d576be0)。 产品仓库 [`XiaoSiKe/zero-one-api`](https://github.com/XiaoSiKe/zero-one-api) 配置为 `origin`,官方仓库 `Wei-Shaw/sub2api` 配置为只读 `upstream`。`main` 是零一 API 唯一产品、CI 和发布分支;不保留第二产品分支。 @@ -13,7 +13,7 @@ 新仓库权限、Secrets、分支保护和包访问权限必须分别核验,不能视为已从旧仓库继承; 导入及远端验收记录见[运维清单](PRODUCTION_SERVER_CN.md#历史发布索引)。 -`.github/upstream-baseline.json` 是 schema v4 Overlay Registry,所有常规回放路径必须唯一归属 `Console Skin`、`Public Capabilities`、`Supported Preview`、`Visual Regression` 或 `Marketing Source Assets`。Registry 只接受精确文件或精确目录,不接受 glob 或未命名的顺带改动。Product Change Protection 会对固定 Upstream Baseline 后的全部产品差异做闭包检查:每个差异必须由 `preserve_on_upstream_sync`、Approved UI Snapshot、带退出条件的 legacy hotfix 或精确 backport 之一保留,仅有 Overlay 归属会使 readiness 失败。`preserve_on_upstream_sync` 禁止目录和 glob,每一项必须唯一归属一个 Overlay;上游同步不能删除保护。后续产品决策明确取代旧行为时,允许用 `retired_preserved_paths` 逐文件记录 owner、受保护 ADR 和非空原因,形成不可静默撤销的退役墓碑;同一路径不能同时处于保留与退役状态。`upstream_sync` 持久绑定旧 upstream Tag/commit、合并前 product tip 和真实双亲 merge commit,CI 与 publish 会重放 product-to-merge 差异并拒绝任一未退役的受保护文件被覆盖。临时生产正确性修补保留在带退出条件的独立 legacy hotfix 区块;安全 backport 继续锁定逐文件 SHA-256 与 Git mode。`frontend/src/api/` 与 `frontend/src/types/` 默认不可变,只有 Registry 中精确命名并绑定 owner 且逐文件进入 `preserve_on_upstream_sync` 的兼容文件例外可以通过,相邻文件仍被拒绝。v179 的渠道定价 API 例外只把数据库可空的 multiplier 字段表达为可选且可空,使已批准 Console 能继续构建,不改变请求或响应字段。 +`.github/upstream-baseline.json` 是 schema v5 Overlay Registry,所有常规回放路径必须唯一归属 `Console Skin`、`Public Capabilities`、`Supported Preview`、`Visual Regression` 或 `Marketing Source Assets`。Registry 只接受精确文件或精确目录,不接受 glob 或未命名的顺带改动。Product Change Protection 会对固定 Upstream Baseline 后的全部产品差异做闭包检查:每个差异必须由 `preserve_on_upstream_sync`、Approved UI Snapshot、带退出条件的 legacy hotfix 或精确 backport 之一保留,仅有 Overlay 归属会使 readiness 失败。`preserve_on_upstream_sync` 禁止目录和 glob,每一项必须唯一归属一个 Overlay;上游同步不能删除保护。后续产品决策明确取代旧行为时,允许用 `retired_preserved_paths` 逐文件记录 owner、受保护 ADR 和非空原因,形成不可静默撤销的退役墓碑;同一路径不能同时处于保留与退役状态。`upstream_sync` 持久绑定旧 upstream Tag/commit、合并前 product tip 和真实双亲 merge commit,CI 与 publish 会重放 product-to-merge 差异,核对契约登记连续性,并拒绝 `preserve_bytes_on_upstream_sync` 中已发布资产被覆盖;普通实现允许在契约和回归成立时演进,见 ADR 0019。临时生产正确性修补保留在带退出条件的独立 legacy hotfix 区块;安全 backport 继续锁定逐文件 SHA-256 与 Git mode。`frontend/src/api/` 与 `frontend/src/types/` 默认不可变,只有 Registry 中精确命名并绑定 owner 且逐文件进入 `preserve_on_upstream_sync` 的兼容文件例外可以通过,相邻文件仍被拒绝。v179 的渠道定价 API 例外只把数据库可空的 multiplier 字段表达为可选且可空,使已批准 Console 能继续构建,不改变请求或响应字段。 | Overlay owner | Interface and seam | | --- | --- | @@ -59,7 +59,7 @@ The React app lives in `landing/`, uses Vite with base `/_landing/`, and is buil | Authentication, billing, redeem and affiliate data | Existing route/service/repository contracts and integration tests own the invariants; migrations and original business records remain immutable. | | Console and Landing | Source, generated adapters and Approved UI Snapshot have separate roles. Versioned asset URLs and byte content remain available for old pages and rollback; identical byte storage already shares the immutable pool. | | Generated code and dependencies | Ent/Wire output follows its generator. Frontend API/type exports and optional provider/plugin integrations are not classified as dead merely because the current UI does not import them. | -| Legacy hotfixes | All six registered groups retain their existing exit conditions. The v0.2.4 tag still carries `backend/cmd/server/VERSION=0.2.3`, so the product version-alignment correction remains necessary. Billing, race, formatting, sticky-log, Grok-test and dependency fixes are not retired without an equivalent baseline and passing regressions. | +| Legacy hotfixes | All six registered groups retain their existing exit conditions. The v0.2.12 tag still carries `backend/cmd/server/VERSION=0.2.11`, so the product version-alignment correction remains necessary. Billing, race, formatting, sticky-log, Grok-test and dependency fixes are not retired without an equivalent baseline and passing regressions. | | Operations and historical evidence | Release/backup entry points live under `deploy/zero-one`; recovery material stays outside Git. Historical design evidence, completed plans and provenance records do not define current runtime behavior. | The active maintenance commands and release prerequisites are owned by @@ -166,18 +166,16 @@ edge image. Image rollback does not reverse a database migration; see 合回 `main`。每次同步同时更新本节的 tag 与完整提交 SHA。 主题改动保持集中,使新增上游页面继承设计系统,避免逐页分叉。 -当前 Upstream Baseline 是 `v0.2.4`,解引用源码提交为 -`5de5e2bed035d43591a2e10e51f420ef6a84eb98`,annotated tag object 为 -`d681d0798064ee0ffff376d19687d12f09fe600f`。本次通过真实双父合并引入 MiniMax -平台目录、长流 keepalive、跨实例缓存失效、持久化 cooldown、Grok 媒体资格、 -OpenAI Image 2.5 及网关、代理和管理界面修复;266 个上游变化路径的整合决定见 -[v0.2.4 升级记录](upgrades/v0.2.4.md)与对应 change map。 - -继续保留 Zero One 的分组与账号长上下文计费双重开关、定价快照隔离、 -历史上游声明证据、兑换领取证明、请求首 Token 和生图错误语义。当前 Provider -Account 成本只使用请求时冻结的上游声明有效倍率,缺失证据保持待核算;本地账号倍率 -继续用于调度和额度核算,不重算历史账。新增上游 Console 能力通过 v8/v10/v17 -恢复资源接入,并保持所有已发布历史 URL 不变。 +当前 Upstream Baseline 是 `v0.2.12`,解引用源码提交为 +`5106065716e494204fc0e8db16f68f6e9d576be0`。本次通过真实双父合并引入 +TypeSafe / Jev System One、充值赠金与折扣阶梯、密钥分组排序、账号优先级快捷调整, +以及验证码并发限制、哈希重置令牌、订单查询限流、上游错误脱敏和 Grok CLI 身份修复。 +151 个上游变化路径的决定见 [v0.2.12 升级记录](upgrades/v0.2.12.md)及对应 change map。 + +继续保留分组与账号长上下文计费双重开关、请求时的上游成本证据、兑换领取证明、 +邀请归属、站点首 Token、主动 V1 监控和已发布资源字节。普通源码组件吸收兼容变更, +Production Console 仍使用独立批准的恢复资源;新增源码控件不意味着已进入生产 UI。 +本次新增两条增量迁移,旧订单 `bonus_amount` 为零,原有业务列不改写。 Go 版本保持 `1.27.0`,`approved_backports` 为空。六组 legacy hotfix 继续按 各自退出条件审查:与上游重叠的业务修复已整合,未达到等价条件的精确路径保留。 diff --git a/docs/upgrades/v0.2.12-change-map.json b/docs/upgrades/v0.2.12-change-map.json new file mode 100644 index 000000000..17b6af3cf --- /dev/null +++ b/docs/upgrades/v0.2.12-change-map.json @@ -0,0 +1,776 @@ +{ + "schema_version": 1, + "previous_upstream": "96f4c115c9749078f90cbf210a01d39baf3f53b6", + "upstream": "5106065716e494204fc0e8db16f68f6e9d576be0", + "product": "6a540fca7db34220add8d4c6a2f71f1d8ea61617", + "merge": "a29346e2dd0406fa2c82bc3102ba804884f88f98", + "path_count": 151, + "summary": { + "adapt": 28, + "adopt": 75, + "defer-production-ui": 48 + }, + "rejected_policy_changes": [ + "revive retired product features", + "derive historical account cost from mutable local rates or later probes", + "replace approved production UI snapshots without approval", + "prune Docker volumes, bind mounts, Redis, PostgreSQL, environment files, or release locks" + ], + "changes": [ + { + "path": "README.md", + "decision": "adapt", + "reason": "recorded merge combines upstream behavior with product contracts" + }, + { + "path": "README_CN.md", + "decision": "adapt", + "reason": "recorded merge combines upstream behavior with product contracts" + }, + { + "path": "backend/cmd/server/VERSION", + "decision": "adopt", + "reason": "recorded merge matches upstream blob" + }, + { + "path": "backend/ent/migrate/schema.go", + "decision": "adapt", + "reason": "recorded merge combines upstream behavior with product contracts" + }, + { + "path": "backend/ent/mutation.go", + "decision": "adapt", + "reason": "recorded merge combines upstream behavior with product contracts" + }, + { + "path": "backend/ent/paymentorder.go", + "decision": "adopt", + "reason": "recorded merge matches upstream blob" + }, + { + "path": "backend/ent/paymentorder/paymentorder.go", + "decision": "adopt", + "reason": "recorded merge matches upstream blob" + }, + { + "path": "backend/ent/paymentorder/where.go", + "decision": "adopt", + "reason": "recorded merge matches upstream blob" + }, + { + "path": "backend/ent/paymentorder_create.go", + "decision": "adopt", + "reason": "recorded merge matches upstream blob" + }, + { + "path": "backend/ent/paymentorder_update.go", + "decision": "adopt", + "reason": "recorded merge matches upstream blob" + }, + { + "path": "backend/ent/runtime/runtime.go", + "decision": "adapt", + "reason": "recorded merge combines upstream behavior with product contracts" + }, + { + "path": "backend/ent/schema/payment_order.go", + "decision": "adopt", + "reason": "recorded merge matches upstream blob" + }, + { + "path": "backend/ent/schema/user_platform_quota.go", + "decision": "adopt", + "reason": "recorded merge matches upstream blob" + }, + { + "path": "backend/internal/domain/constants.go", + "decision": "adapt", + "reason": "recorded merge combines upstream behavior with product contracts" + }, + { + "path": "backend/internal/handler/admin/account_handler.go", + "decision": "adapt", + "reason": "recorded merge combines upstream behavior with product contracts" + }, + { + "path": "backend/internal/handler/admin/account_handler_available_models_test.go", + "decision": "adopt", + "reason": "recorded merge matches upstream blob" + }, + { + "path": "backend/internal/handler/admin/channel_handler.go", + "decision": "adopt", + "reason": "recorded merge matches upstream blob" + }, + { + "path": "backend/internal/handler/admin/channel_handler_test.go", + "decision": "adopt", + "reason": "recorded merge matches upstream blob" + }, + { + "path": "backend/internal/handler/admin/group_handler.go", + "decision": "adapt", + "reason": "recorded merge combines upstream behavior with product contracts" + }, + { + "path": "backend/internal/handler/admin/group_handler_platform_test.go", + "decision": "adopt", + "reason": "recorded merge matches upstream blob" + }, + { + "path": "backend/internal/handler/admin/payment_handler.go", + "decision": "adopt", + "reason": "recorded merge matches upstream blob" + }, + { + "path": "backend/internal/handler/admin/setting_handler.go", + "decision": "adapt", + "reason": "recorded merge combines upstream behavior with product contracts" + }, + { + "path": "backend/internal/handler/admin/setting_handler_recharge_bonus.go", + "decision": "adopt", + "reason": "recorded merge matches upstream blob" + }, + { + "path": "backend/internal/handler/admin/setting_handler_update.go", + "decision": "adapt", + "reason": "recorded merge combines upstream behavior with product contracts" + }, + { + "path": "backend/internal/handler/auth_oauth_pending_flow_test.go", + "decision": "adapt", + "reason": "recorded merge combines upstream behavior with product contracts" + }, + { + "path": "backend/internal/handler/dto/recharge_bonus_tiers.go", + "decision": "adopt", + "reason": "recorded merge matches upstream blob" + }, + { + "path": "backend/internal/handler/dto/settings.go", + "decision": "adapt", + "reason": "recorded merge combines upstream behavior with product contracts" + }, + { + "path": "backend/internal/handler/endpoint.go", + "decision": "adopt", + "reason": "recorded merge matches upstream blob" + }, + { + "path": "backend/internal/handler/endpoint_test.go", + "decision": "adopt", + "reason": "recorded merge matches upstream blob" + }, + { + "path": "backend/internal/handler/gateway_handler.go", + "decision": "adapt", + "reason": "recorded merge combines upstream behavior with product contracts" + }, + { + "path": "backend/internal/handler/gateway_handler_chat_completions.go", + "decision": "adapt", + "reason": "recorded merge combines upstream behavior with product contracts" + }, + { + "path": "backend/internal/handler/gateway_handler_responses.go", + "decision": "adapt", + "reason": "recorded merge combines upstream behavior with product contracts" + }, + { + "path": "backend/internal/handler/gateway_models_retrieve_test.go", + "decision": "adopt", + "reason": "recorded merge matches upstream blob" + }, + { + "path": "backend/internal/handler/gateway_models_test.go", + "decision": "adopt", + "reason": "recorded merge matches upstream blob" + }, + { + "path": "backend/internal/handler/gateway_systemone.go", + "decision": "adopt", + "reason": "recorded merge matches upstream blob" + }, + { + "path": "backend/internal/handler/gateway_systemone_test.go", + "decision": "adopt", + "reason": "recorded merge matches upstream blob" + }, + { + "path": "backend/internal/handler/payment_handler.go", + "decision": "adopt", + "reason": "recorded merge matches upstream blob" + }, + { + "path": "backend/internal/handler/user_handler_test.go", + "decision": "adopt", + "reason": "recorded merge matches upstream blob" + }, + { + "path": "backend/internal/model/error_passthrough_rule.go", + "decision": "adopt", + "reason": "recorded merge matches upstream blob" + }, + { + "path": "backend/internal/model/error_passthrough_rule_test.go", + "decision": "adopt", + "reason": "recorded merge matches upstream blob" + }, + { + "path": "backend/internal/pkg/typesafe/client.go", + "decision": "adopt", + "reason": "recorded merge matches upstream blob" + }, + { + "path": "backend/internal/pkg/typesafe/systemone.go", + "decision": "adopt", + "reason": "recorded merge matches upstream blob" + }, + { + "path": "backend/internal/pkg/typesafe/systemone_test.go", + "decision": "adopt", + "reason": "recorded merge matches upstream blob" + }, + { + "path": "backend/internal/pkg/xai/billing.go", + "decision": "adopt", + "reason": "recorded merge matches upstream blob" + }, + { + "path": "backend/internal/pkg/xai/billing_test.go", + "decision": "adopt", + "reason": "recorded merge matches upstream blob" + }, + { + "path": "backend/internal/pkg/xai/cli_identity.go", + "decision": "adopt", + "reason": "recorded merge matches upstream blob" + }, + { + "path": "backend/internal/pkg/xai/cli_identity_test.go", + "decision": "adopt", + "reason": "recorded merge matches upstream blob" + }, + { + "path": "backend/internal/repository/api_key_repo.go", + "decision": "adapt", + "reason": "recorded merge combines upstream behavior with product contracts" + }, + { + "path": "backend/internal/repository/api_key_repo_sort_test.go", + "decision": "adopt", + "reason": "recorded merge matches upstream blob" + }, + { + "path": "backend/internal/repository/email_cache.go", + "decision": "adopt", + "reason": "recorded merge matches upstream blob" + }, + { + "path": "backend/internal/repository/email_cache_atomic_test.go", + "decision": "adopt", + "reason": "recorded merge matches upstream blob" + }, + { + "path": "backend/internal/repository/http_upstream.go", + "decision": "adopt", + "reason": "recorded merge matches upstream blob" + }, + { + "path": "backend/internal/repository/http_upstream_test.go", + "decision": "adopt", + "reason": "recorded merge matches upstream blob" + }, + { + "path": "backend/internal/securityaudit/prompt_snapshot.go", + "decision": "adopt", + "reason": "recorded merge matches upstream blob" + }, + { + "path": "backend/internal/securityaudit/prompt_snapshot_test.go", + "decision": "adopt", + "reason": "recorded merge matches upstream blob" + }, + { + "path": "backend/internal/server/api_contract_test.go", + "decision": "adapt", + "reason": "recorded merge combines upstream behavior with product contracts" + }, + { + "path": "backend/internal/server/router.go", + "decision": "adapt", + "reason": "recorded merge combines upstream behavior with product contracts" + }, + { + "path": "backend/internal/server/routes/gateway.go", + "decision": "adapt", + "reason": "recorded merge combines upstream behavior with product contracts" + }, + { + "path": "backend/internal/server/routes/gateway_model_allowlist_test.go", + "decision": "adapt", + "reason": "recorded merge combines upstream behavior with product contracts" + }, + { + "path": "backend/internal/server/routes/payment.go", + "decision": "adopt", + "reason": "recorded merge matches upstream blob" + }, + { + "path": "backend/internal/server/routes/payment_public_rate_limit_test.go", + "decision": "adopt", + "reason": "recorded merge matches upstream blob" + }, + { + "path": "backend/internal/server/routes/prompt_audit_route_coverage_test.go", + "decision": "adopt", + "reason": "recorded merge matches upstream blob" + }, + { + "path": "backend/internal/service/account.go", + "decision": "adopt", + "reason": "recorded merge matches upstream blob" + }, + { + "path": "backend/internal/service/account_service.go", + "decision": "adopt", + "reason": "recorded merge matches upstream blob" + }, + { + "path": "backend/internal/service/account_test_service.go", + "decision": "adopt", + "reason": "recorded merge matches upstream blob" + }, + { + "path": "backend/internal/service/account_test_service_typesafe.go", + "decision": "adopt", + "reason": "recorded merge matches upstream blob" + }, + { + "path": "backend/internal/service/account_test_service_typesafe_test.go", + "decision": "adopt", + "reason": "recorded merge matches upstream blob" + }, + { + "path": "backend/internal/service/admin_account.go", + "decision": "adopt", + "reason": "recorded merge matches upstream blob" + }, + { + "path": "backend/internal/service/admin_group.go", + "decision": "adapt", + "reason": "recorded merge combines upstream behavior with product contracts" + }, + { + "path": "backend/internal/service/antigravity_gateway_gemini.go", + "decision": "adopt", + "reason": "recorded merge matches upstream blob" + }, + { + "path": "backend/internal/service/antigravity_upstream_error_sanitize.go", + "decision": "adopt", + "reason": "recorded merge matches upstream blob" + }, + { + "path": "backend/internal/service/antigravity_upstream_error_sanitize_test.go", + "decision": "adopt", + "reason": "recorded merge matches upstream blob" + }, + { + "path": "backend/internal/service/auth_service_email_bind_test.go", + "decision": "adopt", + "reason": "recorded merge matches upstream blob" + }, + { + "path": "backend/internal/service/auth_service_register_test.go", + "decision": "adapt", + "reason": "recorded merge combines upstream behavior with product contracts" + }, + { + "path": "backend/internal/service/billing_service.go", + "decision": "adapt", + "reason": "recorded merge combines upstream behavior with product contracts" + }, + { + "path": "backend/internal/service/billing_service_test.go", + "decision": "adopt", + "reason": "recorded merge matches upstream blob" + }, + { + "path": "backend/internal/service/channel_service.go", + "decision": "adopt", + "reason": "recorded merge matches upstream blob" + }, + { + "path": "backend/internal/service/channel_service_test.go", + "decision": "adopt", + "reason": "recorded merge matches upstream blob" + }, + { + "path": "backend/internal/service/composite_platform.go", + "decision": "adopt", + "reason": "recorded merge matches upstream blob" + }, + { + "path": "backend/internal/service/composite_platform_test.go", + "decision": "adopt", + "reason": "recorded merge matches upstream blob" + }, + { + "path": "backend/internal/service/content_moderation.go", + "decision": "adapt", + "reason": "recorded merge combines upstream behavior with product contracts" + }, + { + "path": "backend/internal/service/content_moderation_input.go", + "decision": "adopt", + "reason": "recorded merge matches upstream blob" + }, + { + "path": "backend/internal/service/content_moderation_systemone_test.go", + "decision": "adopt", + "reason": "recorded merge matches upstream blob" + }, + { + "path": "backend/internal/service/domain_constants.go", + "decision": "adapt", + "reason": "recorded merge combines upstream behavior with product contracts" + }, + { + "path": "backend/internal/service/email_service.go", + "decision": "adapt", + "reason": "recorded merge combines upstream behavior with product contracts" + }, + { + "path": "backend/internal/service/email_service_reset_token_test.go", + "decision": "adopt", + "reason": "recorded merge matches upstream blob" + }, + { + "path": "backend/internal/service/gateway_systemone.go", + "decision": "adopt", + "reason": "recorded merge matches upstream blob" + }, + { + "path": "backend/internal/service/gateway_systemone_test.go", + "decision": "adopt", + "reason": "recorded merge matches upstream blob" + }, + { + "path": "backend/internal/service/grok_upstream_headers.go", + "decision": "adopt", + "reason": "recorded merge matches upstream blob" + }, + { + "path": "backend/internal/service/grok_upstream_headers_test.go", + "decision": "adopt", + "reason": "recorded merge matches upstream blob" + }, + { + "path": "backend/internal/service/openai_gateway_grok.go", + "decision": "adopt", + "reason": "recorded merge matches upstream blob" + }, + { + "path": "backend/internal/service/payment_config_service.go", + "decision": "adopt", + "reason": "recorded merge matches upstream blob" + }, + { + "path": "backend/internal/service/payment_fulfillment.go", + "decision": "adapt", + "reason": "recorded merge combines upstream behavior with product contracts" + }, + { + "path": "backend/internal/service/payment_order.go", + "decision": "adopt", + "reason": "recorded merge matches upstream blob" + }, + { + "path": "backend/internal/service/payment_order_provider_snapshot_test.go", + "decision": "adopt", + "reason": "recorded merge matches upstream blob" + }, + { + "path": "backend/internal/service/payment_recharge_bonus.go", + "decision": "adopt", + "reason": "recorded merge matches upstream blob" + }, + { + "path": "backend/internal/service/payment_recharge_bonus_test.go", + "decision": "adopt", + "reason": "recorded merge matches upstream blob" + }, + { + "path": "backend/internal/service/payment_service.go", + "decision": "adopt", + "reason": "recorded merge matches upstream blob" + }, + { + "path": "backend/internal/service/scheduler_snapshot_service.go", + "decision": "adopt", + "reason": "recorded merge matches upstream blob" + }, + { + "path": "backend/internal/service/user_service.go", + "decision": "adapt", + "reason": "recorded merge combines upstream behavior with product contracts" + }, + { + "path": "backend/migrations/241_add_payment_order_bonus_amount.sql", + "decision": "adopt", + "reason": "recorded merge matches upstream blob" + }, + { + "path": "backend/migrations/241_add_typesafe_platform.sql", + "decision": "adopt", + "reason": "recorded merge matches upstream blob" + }, + { + "path": "backend/migrations/typesafe_platform_migration_test.go", + "decision": "adopt", + "reason": "recorded merge matches upstream blob" + }, + { + "path": "frontend/package.json", + "decision": "defer-production-ui", + "reason": "source-compatible; production snapshot remains separately approved" + }, + { + "path": "frontend/pnpm-lock.yaml", + "decision": "defer-production-ui", + "reason": "source-compatible; production snapshot remains separately approved" + }, + { + "path": "frontend/src/api/admin/settings.ts", + "decision": "defer-production-ui", + "reason": "source-compatible; production snapshot remains separately approved" + }, + { + "path": "frontend/src/api/admin/users.ts", + "decision": "defer-production-ui", + "reason": "source-compatible; production snapshot remains separately approved" + }, + { + "path": "frontend/src/components/account/AccountPriorityCell.vue", + "decision": "defer-production-ui", + "reason": "source-compatible; production snapshot remains separately approved" + }, + { + "path": "frontend/src/components/account/CreateAccountModal.vue", + "decision": "defer-production-ui", + "reason": "source-compatible; production snapshot remains separately approved" + }, + { + "path": "frontend/src/components/account/EditAccountModal.vue", + "decision": "defer-production-ui", + "reason": "source-compatible; production snapshot remains separately approved" + }, + { + "path": "frontend/src/components/account/__tests__/AccountPriorityCell.spec.ts", + "decision": "defer-production-ui", + "reason": "source-compatible; production snapshot remains separately approved" + }, + { + "path": "frontend/src/components/admin/payment/AdminOrderDetail.vue", + "decision": "defer-production-ui", + "reason": "source-compatible; production snapshot remains separately approved" + }, + { + "path": "frontend/src/components/admin/payment/AdminOrderTable.vue", + "decision": "defer-production-ui", + "reason": "source-compatible; production snapshot remains separately approved" + }, + { + "path": "frontend/src/components/admin/settings/RechargeBonusTierEditor.vue", + "decision": "defer-production-ui", + "reason": "source-compatible; production snapshot remains separately approved" + }, + { + "path": "frontend/src/components/admin/user/__tests__/UserPlatformQuotaModal.spec.ts", + "decision": "defer-production-ui", + "reason": "source-compatible; production snapshot remains separately approved" + }, + { + "path": "frontend/src/components/keys/UseKeyModal.vue", + "decision": "defer-production-ui", + "reason": "source-compatible; production snapshot remains separately approved" + }, + { + "path": "frontend/src/components/payment/AmountInput.vue", + "decision": "defer-production-ui", + "reason": "source-compatible; production snapshot remains separately approved" + }, + { + "path": "frontend/src/components/payment/OrderTable.vue", + "decision": "defer-production-ui", + "reason": "source-compatible; production snapshot remains separately approved" + }, + { + "path": "frontend/src/components/payment/PaymentStatusPanel.vue", + "decision": "defer-production-ui", + "reason": "source-compatible; production snapshot remains separately approved" + }, + { + "path": "frontend/src/components/payment/__tests__/AmountInput.spec.ts", + "decision": "defer-production-ui", + "reason": "source-compatible; production snapshot remains separately approved" + }, + { + "path": "frontend/src/components/user/UserPlatformQuotaCell.vue", + "decision": "defer-production-ui", + "reason": "source-compatible; production snapshot remains separately approved" + }, + { + "path": "frontend/src/components/user/__tests__/UserPlatformQuotaCell.spec.ts", + "decision": "defer-production-ui", + "reason": "source-compatible; production snapshot remains separately approved" + }, + { + "path": "frontend/src/components/user/dashboard/UserDashboardStats.vue", + "decision": "defer-production-ui", + "reason": "source-compatible; production snapshot remains separately approved" + }, + { + "path": "frontend/src/components/user/dashboard/__tests__/UserDashboardStats.spec.ts", + "decision": "defer-production-ui", + "reason": "source-compatible; production snapshot remains separately approved" + }, + { + "path": "frontend/src/composables/useModelWhitelist.ts", + "decision": "defer-production-ui", + "reason": "source-compatible; production snapshot remains separately approved" + }, + { + "path": "frontend/src/constants/__tests__/platforms.spec.ts", + "decision": "defer-production-ui", + "reason": "source-compatible; production snapshot remains separately approved" + }, + { + "path": "frontend/src/constants/platforms.ts", + "decision": "defer-production-ui", + "reason": "source-compatible; production snapshot remains separately approved" + }, + { + "path": "frontend/src/i18n/locales/en/admin/accounts.ts", + "decision": "defer-production-ui", + "reason": "source-compatible; production snapshot remains separately approved" + }, + { + "path": "frontend/src/i18n/locales/en/admin/overview.ts", + "decision": "defer-production-ui", + "reason": "source-compatible; production snapshot remains separately approved" + }, + { + "path": "frontend/src/i18n/locales/en/admin/settings.ts", + "decision": "defer-production-ui", + "reason": "source-compatible; production snapshot remains separately approved" + }, + { + "path": "frontend/src/i18n/locales/en/dashboard.ts", + "decision": "defer-production-ui", + "reason": "source-compatible; production snapshot remains separately approved" + }, + { + "path": "frontend/src/i18n/locales/en/misc.ts", + "decision": "defer-production-ui", + "reason": "source-compatible; production snapshot remains separately approved" + }, + { + "path": "frontend/src/i18n/locales/zh/admin/accounts.ts", + "decision": "defer-production-ui", + "reason": "source-compatible; production snapshot remains separately approved" + }, + { + "path": "frontend/src/i18n/locales/zh/admin/overview.ts", + "decision": "defer-production-ui", + "reason": "source-compatible; production snapshot remains separately approved" + }, + { + "path": "frontend/src/i18n/locales/zh/admin/settings.ts", + "decision": "defer-production-ui", + "reason": "source-compatible; production snapshot remains separately approved" + }, + { + "path": "frontend/src/i18n/locales/zh/dashboard.ts", + "decision": "defer-production-ui", + "reason": "source-compatible; production snapshot remains separately approved" + }, + { + "path": "frontend/src/i18n/locales/zh/misc.ts", + "decision": "defer-production-ui", + "reason": "source-compatible; production snapshot remains separately approved" + }, + { + "path": "frontend/src/types/index.ts", + "decision": "defer-production-ui", + "reason": "source-compatible; production snapshot remains separately approved" + }, + { + "path": "frontend/src/types/payment.ts", + "decision": "defer-production-ui", + "reason": "source-compatible; production snapshot remains separately approved" + }, + { + "path": "frontend/src/utils/__tests__/rechargeBonus.spec.ts", + "decision": "defer-production-ui", + "reason": "source-compatible; production snapshot remains separately approved" + }, + { + "path": "frontend/src/utils/keyGroupProviders.ts", + "decision": "defer-production-ui", + "reason": "source-compatible; production snapshot remains separately approved" + }, + { + "path": "frontend/src/utils/platformColors.ts", + "decision": "defer-production-ui", + "reason": "source-compatible; production snapshot remains separately approved" + }, + { + "path": "frontend/src/utils/rechargeBonus.ts", + "decision": "defer-production-ui", + "reason": "source-compatible; production snapshot remains separately approved" + }, + { + "path": "frontend/src/views/admin/AccountsView.vue", + "decision": "defer-production-ui", + "reason": "source-compatible; production snapshot remains separately approved" + }, + { + "path": "frontend/src/views/admin/ChannelsView.vue", + "decision": "defer-production-ui", + "reason": "source-compatible; production snapshot remains separately approved" + }, + { + "path": "frontend/src/views/admin/SettingsView.vue", + "decision": "defer-production-ui", + "reason": "source-compatible; production snapshot remains separately approved" + }, + { + "path": "frontend/src/views/admin/__tests__/SettingsView.spec.ts", + "decision": "defer-production-ui", + "reason": "source-compatible; production snapshot remains separately approved" + }, + { + "path": "frontend/src/views/user/KeysView.vue", + "decision": "defer-production-ui", + "reason": "source-compatible; production snapshot remains separately approved" + }, + { + "path": "frontend/src/views/user/PaymentResultView.vue", + "decision": "defer-production-ui", + "reason": "source-compatible; production snapshot remains separately approved" + }, + { + "path": "frontend/src/views/user/PaymentView.vue", + "decision": "defer-production-ui", + "reason": "source-compatible; production snapshot remains separately approved" + }, + { + "path": "frontend/src/views/user/__tests__/KeysView.spec.ts", + "decision": "defer-production-ui", + "reason": "source-compatible; production snapshot remains separately approved" + } + ] +} diff --git a/docs/upgrades/v0.2.12.md b/docs/upgrades/v0.2.12.md new file mode 100644 index 000000000..c6b4c412a --- /dev/null +++ b/docs/upgrades/v0.2.12.md @@ -0,0 +1,26 @@ +# v0.2.12 integration and release contract + +The Upstream Baseline is `v0.2.12` at `5106065716e494204fc0e8db16f68f6e9d576be0`. The pre-upgrade product is `6a540fca7db34220add8d4c6a2f71f1d8ea61617`; the recorded two-parent merge is `a29346e2dd0406fa2c82bc3102ba804884f88f98`. The deterministic [change map](v0.2.12-change-map.json) covers all 151 upstream paths since v0.2.11. This target stays fixed through verification and publication. + +The backend adopts native TypeSafe / Jev System One routing, validation, prompt audit and moderation, recharge bonus/discount tiers, API Key group sorting, Grok CLI identity 1.0.46, Antigravity error sanitization, atomic email verification attempts and rate-limited public payment verification. Balance-order affiliate rebates exclude gifted credit; subscription rebates retain their prior base. Existing orders have zero bonus, so historical credit and rebates are unchanged. Promotions stay disabled unless configured. + +## Integration ownership + +The shared platform catalog remains authoritative for group/channel/quota choices, including default/auth-source quota normalization. The monitor subset excludes TypeSafe because System One is not a conversational probe. Account editor wrappers remain small; TypeSafe creation and editing join the existing editor panels. Retired admin payment components stay absent. Source payment views keep the existing layout owner. Platform colors retain the product palette. The source key-usage animation rejects callbacks after unmount or a newer animation; this fixes the lifecycle failure exposed by the full suite. Password reset uses one hash-based Redis compare/delete owner; cache failure retains the product service-unavailable contract. Regression fixtures store hashes and cover concurrent single use, old plaintext rejection, expired/replaced tokens and storage failures. + +The Approved UI Snapshot and every published namespace stay byte-identical. Compatible upstream frontend source is adopted and typechecked, but new bonus configuration, quick priority and TypeSafe controls remain deferred in the production snapshot. Version-badge screenshot changes require desktop/mobile review. No generic upstream frontend bundle replaces the approved Console. + +## Data and rollback + +Exactly two new full-filename migrations are expected: + +- `241_add_payment_order_bonus_amount.sql`: adds a non-null decimal column defaulting to zero without modifying existing order columns. +- `241_add_typesafe_platform.sql`: extends quota/composite platform checks to allow TypeSafe; no existing rows are deleted or rewritten. + +Retain the complete prior migration ledger, historical business columns and serial sequences. The populated off-host restored database must pass new migration-only twice and new/old/new application login and financial continuity. PostgreSQL/Redis container identities and mounts stay unchanged in production. Use the existing release controller, encrypted off-host recovery point, Backend-first paired images, dedicated model probes and at least 30-minute observation. Never restore an old dump over newly accepted writes. + +Already-issued plaintext password-reset links become invalid after this security update; users request a new link. Account credentials, API Keys, passwords and balances are preserved. Axios stays pinned at the already patched 1.20.0; no dependency or audit exception is relaxed. The upstream tag reports 0.2.11 in VERSION, so the existing version-alignment correction reports 0.2.12. + +## Acceptance + +Run complete `make test`, product-policy, source build, frozen-generation, deployment/routing, security and Chromium gates. Require protected PR checks, exact-main product/security evidence, same-SHA multi-architecture images and signed restore/cutover/observation records. Actual test and production results live in this release's restricted recovery report outside Git. diff --git a/frontend/src/api/__tests__/settings.authSourceDefaults.spec.ts b/frontend/src/api/__tests__/settings.authSourceDefaults.spec.ts index deeb75825..37829e834 100644 --- a/frontend/src/api/__tests__/settings.authSourceDefaults.spec.ts +++ b/frontend/src/api/__tests__/settings.authSourceDefaults.spec.ts @@ -9,13 +9,19 @@ import { type DefaultPlatformQuotasMap, } from "@/api/admin/settings"; -/** 全 null 的 5 平台 map,用于断言归一化默认值 */ +/** 全 null 的 11 平台 map,用于断言归一化默认值 */ const allNullQuotas: DefaultPlatformQuotasMap = { anthropic: { daily: null, weekly: null, monthly: null }, openai: { daily: null, weekly: null, monthly: null }, gemini: { daily: null, weekly: null, monthly: null }, antigravity: { daily: null, weekly: null, monthly: null }, grok: { daily: null, weekly: null, monthly: null }, + typesafe: { daily: null, weekly: null, monthly: null }, + kimi: { daily: null, weekly: null, monthly: null }, + zhipu: { daily: null, weekly: null, monthly: null }, + deepseek: { daily: null, weekly: null, monthly: null }, + minimax: { daily: null, weekly: null, monthly: null }, + opencode_go: { daily: null, weekly: null, monthly: null }, } describe("admin settings auth source defaults helpers", () => { @@ -240,9 +246,9 @@ describe("normalizePlatformQuotasMap", () => { expect(result.grok).toEqual({ daily: null, weekly: null, monthly: null }); }); - it("无参数时返回全 5 平台全 null", () => { + it("无参数时返回全 11 平台全 null", () => { const result = normalizePlatformQuotasMap(); - expect(Object.keys(result)).toHaveLength(5); + expect(Object.keys(result)).toHaveLength(11); for (const v of Object.values(result)) { expect(v).toEqual({ daily: null, weekly: null, monthly: null }); } @@ -290,7 +296,7 @@ describe("sanitizePlatformQuotasMap", () => { it("缺失平台填充为全 null", () => { const result = sanitizePlatformQuotasMap({}); - expect(Object.keys(result)).toHaveLength(5); + expect(Object.keys(result)).toHaveLength(11); for (const v of Object.values(result)) { expect(v).toEqual({ daily: null, weekly: null, monthly: null }); } diff --git a/frontend/src/api/admin/settings.ts b/frontend/src/api/admin/settings.ts index 44620de47..0352713b0 100644 --- a/frontend/src/api/admin/settings.ts +++ b/frontend/src/api/admin/settings.ts @@ -4,12 +4,14 @@ */ import { apiClient } from "../client"; +import { CONCRETE_PLATFORM_OPTIONS } from "@/constants/platforms"; import type { CustomEndpoint, CustomMenuItem, LoginAgreementDocument, NotifyEmailEntry, } from "@/types"; +import type { RechargeBonusTier } from "@/utils/rechargeBonus"; export interface DefaultSubscriptionSetting { group_id: number; @@ -17,7 +19,7 @@ export interface DefaultSubscriptionSetting { } // ── 平台限额类型 ────────────────────────────────────────────────── -export type PlatformType = "anthropic" | "openai" | "gemini" | "antigravity" | "grok" +export type PlatformType = (typeof CONCRETE_PLATFORM_OPTIONS)[number]["value"] export type QuotaWindowType = "daily" | "weekly" | "monthly" /** 单平台三档限额;null = 不限制,undefined = 未填(等价 null) */ @@ -30,7 +32,7 @@ export interface PlatformQuotaLimits { /** 全平台默认限额 map(key = PlatformType) */ export type DefaultPlatformQuotasMap = Partial> -const PLATFORMS: PlatformType[] = ["anthropic", "openai", "gemini", "antigravity", "grok"] +const PLATFORMS: PlatformType[] = CONCRETE_PLATFORM_OPTIONS.map(({ value }) => value) export type SchedulingThresholdPlatformType = | "openai" @@ -683,6 +685,9 @@ export interface SystemSettings { payment_balance_recharge_multiplier: number; payment_subscription_usd_to_cny_rate: number; payment_recharge_fee_rate: number; + payment_recharge_bonus_tiers?: RechargeBonusTier[]; + payment_recharge_bonus_mode?: string; + payment_recharge_bonus_notice?: string; payment_load_balance_strategy: string; payment_product_name_prefix: string; payment_product_name_suffix: string; @@ -1013,6 +1018,9 @@ export interface UpdateSettingsRequest { payment_balance_recharge_multiplier?: number; payment_subscription_usd_to_cny_rate?: number; payment_recharge_fee_rate?: number; + payment_recharge_bonus_tiers?: RechargeBonusTier[]; + payment_recharge_bonus_mode?: string; + payment_recharge_bonus_notice?: string; payment_load_balance_strategy?: string; payment_product_name_prefix?: string; payment_product_name_suffix?: string; diff --git a/frontend/src/api/admin/users.ts b/frontend/src/api/admin/users.ts index 0b173ebca..0679120e1 100644 --- a/frontend/src/api/admin/users.ts +++ b/frontend/src/api/admin/users.ts @@ -335,7 +335,7 @@ export async function bindUserAuthIdentity( // Keep aligned with backend/internal/service/domain_constants.go AllowedQuotaPlatforms. export const PLATFORM_QUOTA_PLATFORMS = [ 'anthropic', 'openai', 'gemini', 'antigravity', 'grok', - 'kimi', 'zhipu', 'deepseek', 'minimax', 'opencode_go', + 'kimi', 'zhipu', 'deepseek', 'minimax', 'opencode_go', 'typesafe', ] as const export type PlatformQuotaPlatform = typeof PLATFORM_QUOTA_PLATFORMS[number] export type PlatformQuotaWindow = 'daily' | 'weekly' | 'monthly' diff --git a/frontend/src/components/account/AccountPriorityCell.vue b/frontend/src/components/account/AccountPriorityCell.vue new file mode 100644 index 000000000..b24934149 --- /dev/null +++ b/frontend/src/components/account/AccountPriorityCell.vue @@ -0,0 +1,182 @@ + + + diff --git a/frontend/src/components/account/__tests__/AccountPriorityCell.spec.ts b/frontend/src/components/account/__tests__/AccountPriorityCell.spec.ts new file mode 100644 index 000000000..85f7fe37e --- /dev/null +++ b/frontend/src/components/account/__tests__/AccountPriorityCell.spec.ts @@ -0,0 +1,92 @@ +import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest' +import { flushPromises, mount } from '@vue/test-utils' +import AccountPriorityCell from '../AccountPriorityCell.vue' +import type { Account } from '@/types' +import { update } from '@/api/admin/accounts' + +vi.mock('@/api/admin/accounts', () => ({ update: vi.fn() })) +vi.mock('vue-i18n', () => ({ useI18n: () => ({ t: (key: string) => key }) })) + +const account = (overrides: Partial = {}) => ({ + id: 7, name: 'Claude 1', platform: 'anthropic', type: 'oauth', priority: 3, + ...overrides, +}) as Account + +const mountCell = (value = account()) => mount(AccountPriorityCell, { props: { account: value } }) + +beforeEach(() => { + vi.useFakeTimers() + vi.mocked(update).mockReset().mockImplementation(async (id, req) => account({ id, priority: req.priority })) +}) +afterEach(() => { + vi.useRealTimers() +}) + +describe('AccountPriorityCell', () => { + it('batches rapid +/- clicks into a single priority-only update', async () => { + const wrapper = mountCell() + await wrapper.get('[data-testid="account-priority-increment"]').trigger('click') + await wrapper.get('[data-testid="account-priority-increment"]').trigger('click') + await wrapper.get('[data-testid="account-priority-decrement"]').trigger('click') + expect(wrapper.get('[data-testid="account-priority-value"]').text()).toBe('4') + expect(update).not.toHaveBeenCalled() + + await vi.runAllTimersAsync() + await flushPromises() + + expect(update).toHaveBeenCalledTimes(1) + expect(update).toHaveBeenCalledWith(7, { priority: 4 }) + expect(wrapper.emitted('updated')?.[0]?.[0]).toMatchObject({ id: 7, priority: 4 }) + }) + + it('does not go below 1', async () => { + const wrapper = mountCell(account({ priority: 1 })) + const dec = wrapper.get('[data-testid="account-priority-decrement"]') + expect(dec.attributes('disabled')).toBeDefined() + // 到达下限时按钮仍应随悬停显隐,而不是常驻半透明 + expect(dec.classes()).toContain('opacity-0') + expect(dec.classes().some(c => c.startsWith('disabled:opacity'))).toBe(false) + await dec.trigger('click') + await vi.runAllTimersAsync() + expect(update).not.toHaveBeenCalled() + }) + + it('saves a typed value on Enter and ignores unchanged input', async () => { + const wrapper = mountCell() + await wrapper.get('[data-testid="account-priority-value"]').trigger('click') + const input = wrapper.get('[data-testid="account-priority-input"]') + await input.setValue('12') + await input.trigger('keydown', { key: 'Enter' }) + await flushPromises() + expect(update).toHaveBeenCalledWith(7, { priority: 12 }) + + vi.mocked(update).mockClear() + await wrapper.setProps({ account: account({ priority: 12 }) }) + await wrapper.get('[data-testid="account-priority-value"]').trigger('click') + await wrapper.get('[data-testid="account-priority-input"]').trigger('keydown', { key: 'Enter' }) + await flushPromises() + expect(update).not.toHaveBeenCalled() + }) + + it('Escape cancels typing without saving', async () => { + const wrapper = mountCell() + await wrapper.get('[data-testid="account-priority-value"]').trigger('click') + const input = wrapper.get('[data-testid="account-priority-input"]') + await input.setValue('50') + await input.trigger('keydown', { key: 'Escape' }) + await flushPromises() + expect(update).not.toHaveBeenCalled() + expect(wrapper.get('[data-testid="account-priority-value"]').text()).toBe('3') + }) + + it('reverts and emits an error when the update fails', async () => { + vi.mocked(update).mockRejectedValueOnce(new Error('boom')) + const wrapper = mountCell() + await wrapper.get('[data-testid="account-priority-increment"]').trigger('click') + await vi.runAllTimersAsync() + await flushPromises() + expect(wrapper.get('[data-testid="account-priority-value"]').text()).toBe('3') + expect(wrapper.emitted('error')).toHaveLength(1) + expect(wrapper.emitted('updated')).toBeUndefined() + }) +}) diff --git a/frontend/src/components/account/editor/CreateAccountEditorPanel.vue b/frontend/src/components/account/editor/CreateAccountEditorPanel.vue index c2f93b8e8..56882b845 100644 --- a/frontend/src/components/account/editor/CreateAccountEditorPanel.vue +++ b/frontend/src/components/account/editor/CreateAccountEditorPanel.vue @@ -215,6 +215,15 @@ MiniMax + @@ -3875,6 +3884,8 @@ const apiKeyBaseUrlPlaceholder = computed(() => { return 'https://generativelanguage.googleapis.com' case 'grok': return 'https://api.x.ai/v1' + case 'typesafe': + return 'https://api.typesafe.ai' default: return 'https://api.anthropic.com' } @@ -3888,6 +3899,8 @@ const apiKeyValuePlaceholder = computed(() => { return 'AIza...' case 'grok': return 'xai-...' + case 'typesafe': + return 'ts-...' case 'kimi': return 'sk-...' case 'zhipu': @@ -4608,11 +4621,19 @@ watch( ? 'https://generativelanguage.googleapis.com' : newPlatform === 'grok' ? 'https://api.x.ai/v1' - : 'https://api.anthropic.com' + : newPlatform === 'typesafe' + ? 'https://api.typesafe.ai' + : 'https://api.anthropic.com' } // Clear model-related settings allowedModels.value = [] modelMappings.value = [] + if (newPlatform === 'typesafe') { + form.type = 'apikey' + accountCategory.value = 'apikey' + modelRestrictionMode.value = 'whitelist' + allowedModels.value = ['jev-latest'] + } // Antigravity: 默认使用映射模式并填充默认映射 if (newPlatform === 'antigravity') { antigravityModelRestrictionMode.value = 'mapping' @@ -5504,6 +5525,8 @@ const handleSubmit = async () => { ? 'https://generativelanguage.googleapis.com' : form.platform === 'grok' ? 'https://api.x.ai/v1' + : form.platform === 'typesafe' + ? 'https://api.typesafe.ai' : 'https://api.anthropic.com' // Build credentials with optional model mapping diff --git a/frontend/src/components/account/editor/EditAccountEditorPanel.vue b/frontend/src/components/account/editor/EditAccountEditorPanel.vue index fca6c687b..e7f5f2352 100644 --- a/frontend/src/components/account/editor/EditAccountEditorPanel.vue +++ b/frontend/src/components/account/editor/EditAccountEditorPanel.vue @@ -3578,6 +3578,7 @@ const defaultBaseUrl = computed(() => { if (props.account?.platform === 'openai') return 'https://api.openai.com' if (props.account?.platform === 'gemini') return 'https://generativelanguage.googleapis.com' if (props.account?.platform === 'grok') return 'https://api.x.ai/v1' + if (props.account?.platform === 'typesafe') return 'https://api.typesafe.ai' // CN 供应商:按当前模式/协议回落到官方预设(清空输入框提交时使用), // 不能落到 anthropic 默认值(会被当 CC base 拼出错误端点)。 if ( @@ -4005,6 +4006,8 @@ const syncFormFromAccount = (newAccount: Account | null) => { ? 'https://generativelanguage.googleapis.com' : newAccount.platform === 'grok' ? 'https://api.x.ai/v1' + : newAccount.platform === 'typesafe' + ? 'https://api.typesafe.ai' : newAccount.platform === 'kimi' || newAccount.platform === 'zhipu' || newAccount.platform === 'deepseek' @@ -4082,6 +4085,8 @@ const syncFormFromAccount = (newAccount: Account | null) => { ? 'https://generativelanguage.googleapis.com' : newAccount.platform === 'grok' ? 'https://api.x.ai/v1' + : newAccount.platform === 'typesafe' + ? 'https://api.typesafe.ai' : 'https://api.anthropic.com' editBaseUrl.value = platformDefaultUrl diff --git a/frontend/src/components/admin/monitor/MonitorFiltersBar.vue b/frontend/src/components/admin/monitor/MonitorFiltersBar.vue index 3f4c9f51f..0f61ef523 100644 --- a/frontend/src/components/admin/monitor/MonitorFiltersBar.vue +++ b/frontend/src/components/admin/monitor/MonitorFiltersBar.vue @@ -66,7 +66,7 @@ import { useI18n } from 'vue-i18n' import type { Provider } from '@/api/admin/channelMonitor' import Select from '@/components/common/Select.vue' import Icon from '@/components/icons/Icon.vue' -import { CONCRETE_PLATFORM_OPTIONS } from '@/constants/platforms' +import { MONITOR_PLATFORM_OPTIONS } from '@/constants/platforms' defineProps<{ loading: boolean @@ -87,7 +87,7 @@ const { t } = useI18n() const providerFilterOptions = computed(() => [ { value: '', label: t('admin.channelMonitor.allProviders') }, - ...CONCRETE_PLATFORM_OPTIONS.map(({ value }) => ({ + ...MONITOR_PLATFORM_OPTIONS.map(({ value }) => ({ value, label: t(`monitorCommon.providers.${value}`), })), diff --git a/frontend/src/components/admin/monitor/MonitorFormDialog.vue b/frontend/src/components/admin/monitor/MonitorFormDialog.vue index a92273b00..405e16c0d 100644 --- a/frontend/src/components/admin/monitor/MonitorFormDialog.vue +++ b/frontend/src/components/admin/monitor/MonitorFormDialog.vue @@ -301,7 +301,7 @@ import { DEFAULT_MINIMAX_ENDPOINT, DEFAULT_INTERVAL_SECONDS, } from '@/constants/channelMonitor' -import { CONCRETE_PLATFORM_OPTIONS } from '@/constants/platforms' +import { MONITOR_PLATFORM_OPTIONS } from '@/constants/platforms' import { estimateDailyProbeRequests } from '@/features/channel-monitor/probeBudget' const props = defineProps<{ @@ -489,7 +489,7 @@ interface ProviderOption { label: string } -const providerOptions = computed(() => CONCRETE_PLATFORM_OPTIONS.map(({ value }) => ({ +const providerOptions = computed(() => MONITOR_PLATFORM_OPTIONS.map(({ value }) => ({ value, label: t(`monitorCommon.providers.${value}`), }))) diff --git a/frontend/src/components/admin/settings/RechargeBonusTierEditor.vue b/frontend/src/components/admin/settings/RechargeBonusTierEditor.vue new file mode 100644 index 000000000..6ef92a9a4 --- /dev/null +++ b/frontend/src/components/admin/settings/RechargeBonusTierEditor.vue @@ -0,0 +1,297 @@ + + + diff --git a/frontend/src/components/admin/user/UserPlatformQuotaModal.vue b/frontend/src/components/admin/user/UserPlatformQuotaModal.vue index e0df90496..e7983318a 100644 --- a/frontend/src/components/admin/user/UserPlatformQuotaModal.vue +++ b/frontend/src/components/admin/user/UserPlatformQuotaModal.vue @@ -121,6 +121,7 @@ import { useAppStore } from '@/stores/app' import { adminAPI } from '@/api/admin' import type { AdminUser, PlatformQuotaItem, PlatformQuotaPlatform, PlatformQuotaWindow } from '@/types' import BaseDialog from '@/components/common/BaseDialog.vue' +import { CONCRETE_PLATFORM_OPTIONS } from '@/constants/platforms' const props = defineProps<{ show: boolean; user: AdminUser | null }>() const emit = defineEmits(['close', 'success']) @@ -128,7 +129,7 @@ const emit = defineEmits(['close', 'success']) const { t } = useI18n() const appStore = useAppStore() -const PLATFORMS: PlatformQuotaPlatform[] = ['anthropic', 'openai', 'gemini', 'antigravity', 'grok'] +const PLATFORMS: PlatformQuotaPlatform[] = CONCRETE_PLATFORM_OPTIONS.map(({ value }) => value) interface QuotaRow { platform: PlatformQuotaPlatform diff --git a/frontend/src/components/admin/user/__tests__/UserPlatformQuotaModal.spec.ts b/frontend/src/components/admin/user/__tests__/UserPlatformQuotaModal.spec.ts index 50740b7ce..97eb69887 100644 --- a/frontend/src/components/admin/user/__tests__/UserPlatformQuotaModal.spec.ts +++ b/frontend/src/components/admin/user/__tests__/UserPlatformQuotaModal.spec.ts @@ -1,6 +1,7 @@ import { describe, it, expect, vi, beforeEach } from 'vitest' import { mount, flushPromises } from '@vue/test-utils' import { createPinia, setActivePinia } from 'pinia' +import type { PlatformQuotaUpdateItem } from '@/types' const apiMocks = vi.hoisted(() => ({ showError: vi.fn(), @@ -75,7 +76,7 @@ beforeEach(() => { }) describe('UserPlatformQuotaModal', () => { - it.each([0, 4, 14])('does not turn a negative limit in input %s into unlimited', async (index) => { + it.each([0, 4, 14, 17])('does not turn a negative limit in input %s into unlimited', async (index) => { const w = await mountAndOpen() await w.findAll('input[type=number]')[index].setValue('-1') await w.findAll('button').find(b => b.text() === 'admin.users.platformQuota.save')!.trigger('click') @@ -101,16 +102,48 @@ describe('UserPlatformQuotaModal', () => { expect(apiMocks.getPlatformQuotas).toHaveBeenCalledWith(99) }) - it('空数据渲染 5 个 platform 行', async () => { + it('renders all eleven supported platforms with empty limits', async () => { const w = await mountAndOpen() - const html = w.html() - expect(html).toContain('anthropic') - expect(html).toContain('openai') - expect(html).toContain('gemini') - expect(html).toContain('antigravity') - expect(html).toContain('grok') + const rows = w.findAll('tbody tr') + expect(rows.map(row => row.find('td').text())).toEqual([ + 'anthropic', 'openai', 'gemini', 'antigravity', 'grok', + 'kimi', 'zhipu', 'deepseek', 'minimax', 'opencode_go', 'typesafe', + ]) + for (const row of rows) { + const inputs = row.findAll('input[type=number]') + expect(inputs).toHaveLength(3) + expect(inputs.map(input => input.element.value)).toEqual(['', '', '']) + } + w.unmount() }) + it.each(['kimi', 'zhipu', 'deepseek', 'minimax', 'opencode_go', 'typesafe'] as const)( + 'saves edits to %s without erasing existing platform limits', async (platform) => { + const existing: PlatformQuotaUpdateItem[] = [ + { platform: 'openai', daily_limit_usd: 10, weekly_limit_usd: 20, monthly_limit_usd: 100 }, + ...(['kimi', 'zhipu', 'deepseek', 'minimax', 'opencode_go', 'typesafe'] as const).map(p => ({ + platform: p, daily_limit_usd: 0, weekly_limit_usd: null, monthly_limit_usd: 50, + })), + ] + apiMocks.getPlatformQuotas.mockResolvedValueOnce({ platform_quotas: existing }) + const w = await mountAndOpen() + const row = w.findAll('tbody tr').find(r => r.find('td').text() === platform)! + const inputs = row.findAll('input[type=number]') + expect(inputs.map(input => input.element.value)).toEqual(['0', '', '50']) + await inputs[1].setValue('12.5') + await w.findAll('button').find(b => b.text() === 'admin.users.platformQuota.save')!.trigger('click') + await flushPromises() + const expected = existing.map(item => item.platform === platform + ? { ...item, weekly_limit_usd: 12.5 } + : item) + expect(apiMocks.updatePlatformQuotas).toHaveBeenCalledTimes(1) + expect(apiMocks.updatePlatformQuotas).toHaveBeenCalledWith(99, expect.arrayContaining(expected)) + expect(apiMocks.updatePlatformQuotas.mock.calls[0][1]).toHaveLength(11) + expect(w.emitted('success')).toHaveLength(1) + w.unmount() + }, + ) + it('已有数据正确填充 limit input', async () => { apiMocks.getPlatformQuotas.mockResolvedValueOnce({ platform_quotas: [ @@ -120,13 +153,13 @@ describe('UserPlatformQuotaModal', () => { }) const w = await mountAndOpen() const inputs = w.findAll('input[type=number]') - // 5 platforms × 3 windows = 15 inputs - expect(inputs.length).toBe(15) + // 11 platforms × 3 windows = 33 inputs + expect(inputs.length).toBe(33) // 第一个 input 是 anthropic.daily = 10 expect((inputs[0].element as HTMLInputElement).value).toBe('10') }) - it('保存提交完整 5 platform payload', async () => { + it('保存提交完整 11 platform payload', async () => { apiMocks.getPlatformQuotas.mockResolvedValueOnce({ platform_quotas: [ { platform: 'openai', daily_limit_usd: null, weekly_limit_usd: 20, monthly_limit_usd: null, @@ -143,7 +176,7 @@ describe('UserPlatformQuotaModal', () => { expect(apiMocks.updatePlatformQuotas).toHaveBeenCalledTimes(1) const [uid, payload] = apiMocks.updatePlatformQuotas.mock.calls[0] expect(uid).toBe(99) - expect(payload).toHaveLength(5) // 5 platforms always submitted + expect(payload).toHaveLength(11) // 11 platforms always submitted const openai = payload.find((p: any) => p.platform === 'openai') expect(openai.weekly_limit_usd).toBe(20) }) @@ -208,7 +241,7 @@ describe('UserPlatformQuotaModal', () => { it('未配置限额的平台重置按钮禁用并提示不可用', async () => { const w = await mountAndOpen() const resetBtns = w.findAll('button').filter((b) => b.text() === '↻') - expect(resetBtns.length).toBe(15) // 5 平台 × 3 窗口 + expect(resetBtns.length).toBe(33) // 11 平台 × 3 窗口 for (const b of resetBtns) { expect((b.element as HTMLButtonElement).disabled).toBe(true) expect(b.attributes('title')).toBe('admin.users.platformQuota.reset.unavailable') diff --git a/frontend/src/components/keys/UseKeyModal.vue b/frontend/src/components/keys/UseKeyModal.vue index c38c4384d..823770da9 100644 --- a/frontend/src/components/keys/UseKeyModal.vue +++ b/frontend/src/components/keys/UseKeyModal.vue @@ -251,6 +251,8 @@ const defaultClientTab = computed(() => { return 'gemini' case 'antigravity': return 'claude' + case 'typesafe': + return 'systemone' default: return 'claude' } @@ -368,6 +370,8 @@ const clientTabs = computed((): TabConfig[] => { { id: 'codex', label: t('keys.useKeyModal.cliTabs.codexCli'), icon: TerminalIcon }, { id: 'opencode', label: t('keys.useKeyModal.cliTabs.opencode'), icon: TerminalIcon } ] + case 'typesafe': + return [{ id: 'systemone', label: t('keys.useKeyModal.cliTabs.systemOne'), icon: TerminalIcon }] default: return [ { id: 'claude', label: t('keys.useKeyModal.cliTabs.claudeCode'), icon: TerminalIcon }, @@ -423,6 +427,8 @@ const platformDescription = computed(() => { return t('keys.useKeyModal.grok.codexDescription') } return t('keys.useKeyModal.grok.description') + case 'typesafe': + return t('keys.useKeyModal.typesafe.description') default: return t('keys.useKeyModal.description') } @@ -460,6 +466,8 @@ const platformNote = computed(() => { return t('keys.useKeyModal.grok.noteWindows') } return t('keys.useKeyModal.grok.note') + case 'typesafe': + return t('keys.useKeyModal.typesafe.note') default: return t('keys.useKeyModal.note') } @@ -519,7 +527,9 @@ const currentFiles = computed((): FileConfig[] => { ] case 'grok': return [generateOpenCodeConfig('grok', apiBase, apiKey)] - default: + case 'typesafe': + return [generateSystemOneCurl(baseRoot, apiKey)] + default: return [generateOpenCodeConfig('openai', apiBase, apiKey)] } } @@ -553,6 +563,46 @@ const currentFiles = computed((): FileConfig[] => { } }) +function generateSystemOneCurl(baseUrl: string, apiKey: string): FileConfig { + const endpoint = `${baseUrl}/v1/systemone` + const payload = `{ + "model": "jev-latest", + "state": "Text to evaluate", + "questions": { + "safety": { + "type": "noul", + "instructions": "Evaluate whether the text is unsafe" + } + } +}` + if (activeTab.value === 'powershell') { + return { + path: 'PowerShell', + content: `$headers = @{ Authorization = "Bearer ${apiKey}" } +$body = @' +${payload} +'@ +Invoke-RestMethod -Method Post -Uri "${endpoint}" -Headers $headers -ContentType "application/json" -Body $body` + } + } + if (activeTab.value === 'cmd') { + return { + path: 'Command Prompt', + content: `curl -X POST "${endpoint}" ^ + -H "Authorization: Bearer ${apiKey}" ^ + -H "Content-Type: application/json" ^ + --data "{\"model\":\"jev-latest\",\"state\":\"Text to evaluate\",\"questions\":{\"safety\":{\"type\":\"noul\",\"instructions\":\"Evaluate whether the text is unsafe\"}}}"` + } + } + return { + path: 'Terminal', + content: `curl -X POST "${endpoint}" \\ + -H "Authorization: Bearer ${apiKey}" \\ + -H "Content-Type: application/json" \\ + --data '${payload}'` + } +} + function generateAnthropicFiles(baseUrl: string, apiKey: string): FileConfig[] { let path: string let content: string diff --git a/frontend/src/components/payment/AmountInput.vue b/frontend/src/components/payment/AmountInput.vue index 0e0c9d339..a16029e94 100644 --- a/frontend/src/components/payment/AmountInput.vue +++ b/frontend/src/components/payment/AmountInput.vue @@ -5,20 +5,45 @@ -
+
@@ -48,16 +73,31 @@ diff --git a/frontend/src/views/admin/AccountsView.vue b/frontend/src/views/admin/AccountsView.vue index 88d316d17..2111df3f3 100644 --- a/frontend/src/views/admin/AccountsView.vue +++ b/frontend/src/views/admin/AccountsView.vue @@ -377,8 +377,12 @@ @probe="handleProbeUpstreamBilling(row)" /> -