diff --git a/backend/cmd/profit-preview/main.go b/backend/cmd/profit-preview/main.go new file mode 100644 index 0000000000..6cb53d331f --- /dev/null +++ b/backend/cmd/profit-preview/main.go @@ -0,0 +1,219 @@ +// profit-preview 读取生产只读导出的 JSON(分组利润配置、账号倍率与探测状态、 +// 用户覆盖倍率、主力模型清单),复用线上 U/D/阈值判定做五平台离线预演。 +// +// 用法: +// +// go run ./cmd/profit-preview -input dump.json [-assume-enabled] [-json] +package main + +import ( + "encoding/json" + "flag" + "fmt" + "os" + "sort" + "strings" + "time" + + "github.com/Wei-Shaw/sub2api/internal/service" +) + +type inputGroup struct { + ID int64 `json:"id"` + Name string `json:"name"` + Platform string `json:"platform"` + RateMultiplier float64 `json:"rate_multiplier"` + SubscriptionType string `json:"subscription_type"` + ProfitControlEnabled bool `json:"profit_control_enabled"` + ProfitMinMargin float64 `json:"profit_min_margin"` + ProfitSafetyBuffer float64 `json:"profit_safety_buffer"` + PeakRateEnabled bool `json:"peak_rate_enabled"` + PeakStart string `json:"peak_start"` + PeakEnd string `json:"peak_end"` + PeakRateMultiplier float64 `json:"peak_rate_multiplier"` +} + +type inputAccount struct { + ID int64 `json:"id"` + Name string `json:"name"` + Platform string `json:"platform"` + Type string `json:"type"` + RateMultiplier *float64 `json:"rate_multiplier"` + Extra map[string]any `json:"extra"` + ModelMapping map[string]string `json:"model_mapping"` +} + +type inputEntry struct { + Group inputGroup `json:"group"` + Accounts []inputAccount `json:"accounts"` + UserOverrides map[string]*float64 `json:"user_overrides"` + Models []string `json:"models"` +} + +type inputDoc struct { + Groups []inputEntry `json:"groups"` +} + +func main() { + inputPath := flag.String("input", "", "生产只读导出 JSON 路径") + assumeEnabled := flag.Bool("assume-enabled", false, "把当前关闭的支持平台分组按保存配置视为已启用") + jsonOut := flag.Bool("json", false, "以 JSON 输出完整报告(默认输出可读表格)") + flag.Parse() + if *inputPath == "" { + fmt.Fprintln(os.Stderr, "usage: profit-preview -input dump.json [-assume-enabled] [-json]") + os.Exit(2) + } + raw, err := os.ReadFile(*inputPath) + if err != nil { + fmt.Fprintf(os.Stderr, "read input: %v\n", err) + os.Exit(1) + } + inputs, err := parsePreviewInputs(raw, *assumeEnabled) + if err != nil { + fmt.Fprintf(os.Stderr, "parse input: %v\n", err) + os.Exit(1) + } + + evalAt := time.Now() + reports := service.PreviewProfitAdmission(inputs, evalAt) + if len(reports) == 0 { + fmt.Fprintln(os.Stderr, "input produced no preview reports") + os.Exit(1) + } + + if *jsonOut { + enc := json.NewEncoder(os.Stdout) + enc.SetIndent("", " ") + if err := enc.Encode(map[string]any{"evaluated_at": evalAt, "reports": reports}); err != nil { + fmt.Fprintf(os.Stderr, "write output: %v\n", err) + os.Exit(1) + } + return + } + + fmt.Printf("利润门预演 @ %s(U=账号倍率;探测状态仅告警)\n", evalAt.Format(time.RFC3339)) + for _, report := range reports { + fmt.Printf("\n== 分组 %d %s [%s] ==\n", report.GroupID, report.GroupName, report.Platform) + fmt.Printf(" 利润门生效=%v 假定启用=%v | 默认 D=%.4f 阈值=%.4f | 最低有效 D=%.4f 阈值=%.4f\n", + report.EffectiveGate, report.AssumedEnabled, + report.DefaultD, report.ThresholdDefault, report.MinEffectiveD, report.ThresholdMinD) + counts := map[string]int{} + for _, v := range report.Verdicts { + counts[v.Class]++ + rate := "-" + if v.AccountRate != nil { + rate = fmt.Sprintf("%.4f", *v.AccountRate) + } + flags := make([]string, 0, 2) + if v.RejectedUnderMinD { + flags = append(flags, "最低有效D下拒绝") + } + if len(v.Warnings) > 0 { + flags = append(flags, strings.Join(v.Warnings, ",")) + } + suffix := "" + if len(flags) > 0 { + suffix = " [" + strings.Join(flags, "; ") + "]" + } + fmt.Printf(" 账号 %-4d %-24s 平台=%-12s U=%-8s 来源=%-19s %s%s\n", + v.AccountID, v.Name, v.Platform, rate, v.RateSource, v.Class, suffix) + } + fmt.Printf(" 分类合计: 准入=%d 利润不足=%d 倍率非法=%d\n", + counts[service.ProfitPreviewClassAdmitted], + counts[service.ProfitPreviewClassRejectedThreshold], + counts[service.ProfitPreviewClassRejectedInvalidRate]) + models := make([]string, 0, len(report.RemainingByModel)) + for model := range report.RemainingByModel { + models = append(models, model) + } + sort.Strings(models) + for _, model := range models { + fmt.Printf(" 模型 %-20s 利润门准入账号: 默认D=%d 最低有效D=%d\n", + model, report.RemainingByModel[model], report.RemainingByModelMinD[model]) + } + for _, model := range modelsWithZeroRemaining(report) { + fmt.Printf(" 警告: 模型 %s 启用后利润门准入账号为 0\n", model) + } + } +} + +func parsePreviewInputs(raw []byte, assumeEnabled bool) ([]service.ProfitPreviewGroupInput, error) { + var doc inputDoc + if err := json.Unmarshal(raw, &doc); err != nil { + return nil, err + } + if len(doc.Groups) == 0 { + return nil, fmt.Errorf("input contains no groups; check the export query and target configuration") + } + + inputs := make([]service.ProfitPreviewGroupInput, 0, len(doc.Groups)) + for i, entry := range doc.Groups { + if entry.Group.ID <= 0 || strings.TrimSpace(entry.Group.Platform) == "" { + return nil, fmt.Errorf("invalid group at index %d: id and platform are required", i) + } + group := &service.Group{ + ID: entry.Group.ID, + Name: entry.Group.Name, + Platform: entry.Group.Platform, + Status: service.StatusActive, + Hydrated: true, + RateMultiplier: entry.Group.RateMultiplier, + SubscriptionType: entry.Group.SubscriptionType, + ProfitControlEnabled: entry.Group.ProfitControlEnabled, + ProfitMinMargin: entry.Group.ProfitMinMargin, + ProfitSafetyBuffer: entry.Group.ProfitSafetyBuffer, + PeakRateEnabled: entry.Group.PeakRateEnabled, + PeakStart: entry.Group.PeakStart, + PeakEnd: entry.Group.PeakEnd, + PeakRateMultiplier: entry.Group.PeakRateMultiplier, + } + accounts := make([]*service.Account, 0, len(entry.Accounts)) + for _, a := range entry.Accounts { + account := &service.Account{ + ID: a.ID, + Name: a.Name, + Platform: a.Platform, + Type: a.Type, + RateMultiplier: a.RateMultiplier, + Extra: a.Extra, + } + if len(a.ModelMapping) > 0 { + mapping := make(map[string]any, len(a.ModelMapping)) + for k, v := range a.ModelMapping { + mapping[k] = v + } + account.Credentials = map[string]any{"model_mapping": mapping} + } + accounts = append(accounts, account) + } + overrides := make(map[int64]float64, len(entry.UserOverrides)) + for userID, rate := range entry.UserOverrides { + if rate == nil { + continue + } + var id int64 + if _, err := fmt.Sscan(userID, &id); err == nil && id > 0 { + overrides[id] = *rate + } + } + inputs = append(inputs, service.ProfitPreviewGroupInput{ + Group: group, + Accounts: accounts, + UserOverrides: overrides, + Models: entry.Models, + AssumeEnabled: assumeEnabled, + }) + } + return inputs, nil +} + +func modelsWithZeroRemaining(report service.ProfitPreviewGroupReport) []string { + var out []string + for model, count := range report.RemainingByModel { + if count == 0 { + out = append(out, model) + } + } + sort.Strings(out) + return out +} diff --git a/backend/cmd/profit-preview/main_test.go b/backend/cmd/profit-preview/main_test.go new file mode 100644 index 0000000000..5607ecb09e --- /dev/null +++ b/backend/cmd/profit-preview/main_test.go @@ -0,0 +1,56 @@ +package main + +import ( + "testing" + "time" + + "github.com/Wei-Shaw/sub2api/internal/service" + "github.com/stretchr/testify/require" +) + +func TestParsePreviewInputsIgnoresNullUserOverride(t *testing.T) { + raw := []byte(`{ + "groups": [{ + "group": { + "id": 50, + "name": "preview", + "platform": "openai", + "rate_multiplier": 0.5, + "subscription_type": "standard", + "profit_control_enabled": false, + "profit_min_margin": 0.1, + "profit_safety_buffer": 0 + }, + "accounts": [{ + "id": 1, + "name": "cheap", + "platform": "openai", + "type": "apikey", + "rate_multiplier": 0.2 + }], + "user_overrides": {"40": null, "41": 0.4}, + "models": ["gpt-test"] + }] + }`) + + inputs, err := parsePreviewInputs(raw, true) + require.NoError(t, err) + require.Len(t, inputs, 1) + require.Equal(t, map[int64]float64{41: 0.4}, inputs[0].UserOverrides) + require.True(t, inputs[0].AssumeEnabled) + + report := service.PreviewProfitAdmission(inputs, time.Date(2026, 1, 15, 8, 30, 0, 0, time.UTC))[0] + require.InDelta(t, 0.4, report.MinEffectiveD, 1e-12, "null 覆盖不能被解码成 0 倍率") + require.InDelta(t, 0.36, report.ThresholdMinD, 1e-12) +} + +func TestParsePreviewInputsRejectsEmptyGroups(t *testing.T) { + for _, raw := range [][]byte{ + []byte(`{"groups":null}`), + []byte(`{"groups":[]}`), + } { + inputs, err := parsePreviewInputs(raw, false) + require.ErrorContains(t, err, "input contains no groups") + require.Nil(t, inputs) + } +} diff --git a/backend/ent/client.go b/backend/ent/client.go index 174d0fb2fd..7cdb7394fd 100644 --- a/backend/ent/client.go +++ b/backend/ent/client.go @@ -6828,25 +6828,25 @@ type ( APIKey, Account, AccountGroup, Announcement, AnnouncementRead, AuthIdentity, AuthIdentityChannel, BatchImageEvent, BatchImageItem, BatchImageJob, ChannelMonitor, ChannelMonitorDailyRollup, ChannelMonitorHistory, - ChannelMonitorRequestTemplate, CompositeModelRoute, - ErrorPassthroughRule, Group, IdempotencyRecord, IdentityAdoptionDecision, - PaymentAuditLog, PaymentOrder, PaymentProviderInstance, PendingAuthSession, - PromoCode, PromoCodeUsage, Proxy, RedeemCode, SecuritySecret, Setting, - SubscriptionPlan, TLSFingerprintProfile, UsageCleanupTask, UsageLog, User, - UserAllowedGroup, UserAttributeDefinition, UserAttributeValue, - UserPlatformQuota, UserSubscription []ent.Hook + ChannelMonitorRequestTemplate, CompositeModelRoute, ErrorPassthroughRule, + Group, IdempotencyRecord, IdentityAdoptionDecision, PaymentAuditLog, + PaymentOrder, PaymentProviderInstance, PendingAuthSession, PromoCode, + PromoCodeUsage, Proxy, RedeemCode, SecuritySecret, Setting, SubscriptionPlan, + TLSFingerprintProfile, UsageCleanupTask, UsageLog, User, UserAllowedGroup, + UserAttributeDefinition, UserAttributeValue, UserPlatformQuota, + UserSubscription []ent.Hook } inters struct { APIKey, Account, AccountGroup, Announcement, AnnouncementRead, AuthIdentity, AuthIdentityChannel, BatchImageEvent, BatchImageItem, BatchImageJob, ChannelMonitor, ChannelMonitorDailyRollup, ChannelMonitorHistory, - ChannelMonitorRequestTemplate, CompositeModelRoute, - ErrorPassthroughRule, Group, IdempotencyRecord, IdentityAdoptionDecision, - PaymentAuditLog, PaymentOrder, PaymentProviderInstance, PendingAuthSession, - PromoCode, PromoCodeUsage, Proxy, RedeemCode, SecuritySecret, Setting, - SubscriptionPlan, TLSFingerprintProfile, UsageCleanupTask, UsageLog, User, - UserAllowedGroup, UserAttributeDefinition, UserAttributeValue, - UserPlatformQuota, UserSubscription []ent.Interceptor + ChannelMonitorRequestTemplate, CompositeModelRoute, ErrorPassthroughRule, + Group, IdempotencyRecord, IdentityAdoptionDecision, PaymentAuditLog, + PaymentOrder, PaymentProviderInstance, PendingAuthSession, PromoCode, + PromoCodeUsage, Proxy, RedeemCode, SecuritySecret, Setting, SubscriptionPlan, + TLSFingerprintProfile, UsageCleanupTask, UsageLog, User, UserAllowedGroup, + UserAttributeDefinition, UserAttributeValue, UserPlatformQuota, + UserSubscription []ent.Interceptor } ) diff --git a/backend/ent/group.go b/backend/ent/group.go index d915740125..81a3e349c0 100644 --- a/backend/ent/group.go +++ b/backend/ent/group.go @@ -123,6 +123,12 @@ type Group struct { MaxReasoningEffort string `json:"max_reasoning_effort,omitempty"` // OpenAI reasoning effort 自定义精确映射;先映射再应用上限 ReasoningEffortMappings []domain.ReasoningEffortMapping `json:"reasoning_effort_mappings,omitempty"` + // 是否启用利润控制:调度时仅允许账号计费倍率满足毛利率要求的账号进入候选池 + ProfitControlEnabled bool `json:"profit_control_enabled,omitempty"` + // 最低毛利率,小数(0.30=30%);账号准入条件为 U <= D*(1-margin-buffer) + ProfitMinMargin float64 `json:"profit_min_margin,omitempty"` + // 安全缓冲,小数;与 margin 相加后从下游倍率中扣除,默认 0 + ProfitSafetyBuffer float64 `json:"profit_safety_buffer,omitempty"` // Edges holds the relations/edges for other nodes in the graph. // The values are being populated by the GroupQuery when eager-loading is set. Edges GroupEdges `json:"edges"` @@ -231,9 +237,9 @@ func (*Group) scanValues(columns []string) ([]any, error) { switch columns[i] { case group.FieldModelRouting, group.FieldSupportedModelScopes, group.FieldMessagesDispatchModelConfig, group.FieldModelsListConfig, group.FieldReasoningEffortMappings: values[i] = new([]byte) - case group.FieldPeakRateEnabled, group.FieldIsExclusive, group.FieldAllowImageGeneration, group.FieldAllowBatchImageGeneration, group.FieldImageRateIndependent, group.FieldVideoRateIndependent, group.FieldClaudeCodeOnly, group.FieldModelRoutingEnabled, group.FieldMcpXMLInject, group.FieldAllowMessagesDispatch, group.FieldAllowLive, group.FieldRequireOauthOnly, group.FieldRequirePrivacySet: + case group.FieldPeakRateEnabled, group.FieldIsExclusive, group.FieldAllowImageGeneration, group.FieldAllowBatchImageGeneration, group.FieldImageRateIndependent, group.FieldVideoRateIndependent, group.FieldClaudeCodeOnly, group.FieldModelRoutingEnabled, group.FieldMcpXMLInject, group.FieldAllowMessagesDispatch, group.FieldAllowLive, group.FieldRequireOauthOnly, group.FieldRequirePrivacySet, group.FieldProfitControlEnabled: values[i] = new(sql.NullBool) - case group.FieldRateMultiplier, group.FieldPeakRateMultiplier, group.FieldDailyLimitUsd, group.FieldWeeklyLimitUsd, group.FieldMonthlyLimitUsd, group.FieldImageRateMultiplier, group.FieldImagePrice1k, group.FieldImagePrice2k, group.FieldImagePrice4k, group.FieldBatchImageDiscountMultiplier, group.FieldBatchImageHoldMultiplier, group.FieldVideoRateMultiplier, group.FieldVideoPrice480p, group.FieldVideoPrice720p, group.FieldVideoPrice1080p, group.FieldWebSearchPricePerCall: + case group.FieldRateMultiplier, group.FieldPeakRateMultiplier, group.FieldDailyLimitUsd, group.FieldWeeklyLimitUsd, group.FieldMonthlyLimitUsd, group.FieldImageRateMultiplier, group.FieldImagePrice1k, group.FieldImagePrice2k, group.FieldImagePrice4k, group.FieldBatchImageDiscountMultiplier, group.FieldBatchImageHoldMultiplier, group.FieldVideoRateMultiplier, group.FieldVideoPrice480p, group.FieldVideoPrice720p, group.FieldVideoPrice1080p, group.FieldWebSearchPricePerCall, group.FieldProfitMinMargin, group.FieldProfitSafetyBuffer: values[i] = new(sql.NullFloat64) case group.FieldID, group.FieldDefaultValidityDays, group.FieldFallbackGroupID, group.FieldFallbackGroupIDOnInvalidRequest, group.FieldSortOrder, group.FieldRpmLimit: values[i] = new(sql.NullInt64) @@ -599,6 +605,24 @@ func (_m *Group) assignValues(columns []string, values []any) error { return fmt.Errorf("unmarshal field reasoning_effort_mappings: %w", err) } } + case group.FieldProfitControlEnabled: + if value, ok := values[i].(*sql.NullBool); !ok { + return fmt.Errorf("unexpected type %T for field profit_control_enabled", values[i]) + } else if value.Valid { + _m.ProfitControlEnabled = value.Bool + } + case group.FieldProfitMinMargin: + if value, ok := values[i].(*sql.NullFloat64); !ok { + return fmt.Errorf("unexpected type %T for field profit_min_margin", values[i]) + } else if value.Valid { + _m.ProfitMinMargin = value.Float64 + } + case group.FieldProfitSafetyBuffer: + if value, ok := values[i].(*sql.NullFloat64); !ok { + return fmt.Errorf("unexpected type %T for field profit_safety_buffer", values[i]) + } else if value.Valid { + _m.ProfitSafetyBuffer = value.Float64 + } default: _m.selectValues.Set(columns[i], values[i]) } @@ -860,6 +884,15 @@ func (_m *Group) String() string { builder.WriteString(", ") builder.WriteString("reasoning_effort_mappings=") builder.WriteString(fmt.Sprintf("%v", _m.ReasoningEffortMappings)) + builder.WriteString(", ") + builder.WriteString("profit_control_enabled=") + builder.WriteString(fmt.Sprintf("%v", _m.ProfitControlEnabled)) + builder.WriteString(", ") + builder.WriteString("profit_min_margin=") + builder.WriteString(fmt.Sprintf("%v", _m.ProfitMinMargin)) + builder.WriteString(", ") + builder.WriteString("profit_safety_buffer=") + builder.WriteString(fmt.Sprintf("%v", _m.ProfitSafetyBuffer)) builder.WriteByte(')') return builder.String() } diff --git a/backend/ent/group/group.go b/backend/ent/group/group.go index 6ec01aed53..35d6de1336 100644 --- a/backend/ent/group/group.go +++ b/backend/ent/group/group.go @@ -120,6 +120,12 @@ const ( FieldMaxReasoningEffort = "max_reasoning_effort" // FieldReasoningEffortMappings holds the string denoting the reasoning_effort_mappings field in the database. FieldReasoningEffortMappings = "reasoning_effort_mappings" + // FieldProfitControlEnabled holds the string denoting the profit_control_enabled field in the database. + FieldProfitControlEnabled = "profit_control_enabled" + // FieldProfitMinMargin holds the string denoting the profit_min_margin field in the database. + FieldProfitMinMargin = "profit_min_margin" + // FieldProfitSafetyBuffer holds the string denoting the profit_safety_buffer field in the database. + FieldProfitSafetyBuffer = "profit_safety_buffer" // EdgeAPIKeys holds the string denoting the api_keys edge name in mutations. EdgeAPIKeys = "api_keys" // EdgeRedeemCodes holds the string denoting the redeem_codes edge name in mutations. @@ -247,6 +253,9 @@ var Columns = []string{ FieldRpmLimit, FieldMaxReasoningEffort, FieldReasoningEffortMappings, + FieldProfitControlEnabled, + FieldProfitMinMargin, + FieldProfitSafetyBuffer, } var ( @@ -366,6 +375,12 @@ var ( MaxReasoningEffortValidator func(string) error // DefaultReasoningEffortMappings holds the default value on creation for the "reasoning_effort_mappings" field. DefaultReasoningEffortMappings []domain.ReasoningEffortMapping + // DefaultProfitControlEnabled holds the default value on creation for the "profit_control_enabled" field. + DefaultProfitControlEnabled bool + // DefaultProfitMinMargin holds the default value on creation for the "profit_min_margin" field. + DefaultProfitMinMargin float64 + // DefaultProfitSafetyBuffer holds the default value on creation for the "profit_safety_buffer" field. + DefaultProfitSafetyBuffer float64 ) // OrderOption defines the ordering options for the Group queries. @@ -611,6 +626,21 @@ func ByMaxReasoningEffort(opts ...sql.OrderTermOption) OrderOption { return sql.OrderByField(FieldMaxReasoningEffort, opts...).ToFunc() } +// ByProfitControlEnabled orders the results by the profit_control_enabled field. +func ByProfitControlEnabled(opts ...sql.OrderTermOption) OrderOption { + return sql.OrderByField(FieldProfitControlEnabled, opts...).ToFunc() +} + +// ByProfitMinMargin orders the results by the profit_min_margin field. +func ByProfitMinMargin(opts ...sql.OrderTermOption) OrderOption { + return sql.OrderByField(FieldProfitMinMargin, opts...).ToFunc() +} + +// ByProfitSafetyBuffer orders the results by the profit_safety_buffer field. +func ByProfitSafetyBuffer(opts ...sql.OrderTermOption) OrderOption { + return sql.OrderByField(FieldProfitSafetyBuffer, opts...).ToFunc() +} + // ByAPIKeysCount orders the results by api_keys count. func ByAPIKeysCount(opts ...sql.OrderTermOption) OrderOption { return func(s *sql.Selector) { diff --git a/backend/ent/group/where.go b/backend/ent/group/where.go index 4ad100573c..64e52ba80a 100644 --- a/backend/ent/group/where.go +++ b/backend/ent/group/where.go @@ -290,6 +290,21 @@ func MaxReasoningEffort(v string) predicate.Group { return predicate.Group(sql.FieldEQ(FieldMaxReasoningEffort, v)) } +// ProfitControlEnabled applies equality check predicate on the "profit_control_enabled" field. It's identical to ProfitControlEnabledEQ. +func ProfitControlEnabled(v bool) predicate.Group { + return predicate.Group(sql.FieldEQ(FieldProfitControlEnabled, v)) +} + +// ProfitMinMargin applies equality check predicate on the "profit_min_margin" field. It's identical to ProfitMinMarginEQ. +func ProfitMinMargin(v float64) predicate.Group { + return predicate.Group(sql.FieldEQ(FieldProfitMinMargin, v)) +} + +// ProfitSafetyBuffer applies equality check predicate on the "profit_safety_buffer" field. It's identical to ProfitSafetyBufferEQ. +func ProfitSafetyBuffer(v float64) predicate.Group { + return predicate.Group(sql.FieldEQ(FieldProfitSafetyBuffer, v)) +} + // CreatedAtEQ applies the EQ predicate on the "created_at" field. func CreatedAtEQ(v time.Time) predicate.Group { return predicate.Group(sql.FieldEQ(FieldCreatedAt, v)) @@ -2190,6 +2205,96 @@ func MaxReasoningEffortContainsFold(v string) predicate.Group { return predicate.Group(sql.FieldContainsFold(FieldMaxReasoningEffort, v)) } +// ProfitControlEnabledEQ applies the EQ predicate on the "profit_control_enabled" field. +func ProfitControlEnabledEQ(v bool) predicate.Group { + return predicate.Group(sql.FieldEQ(FieldProfitControlEnabled, v)) +} + +// ProfitControlEnabledNEQ applies the NEQ predicate on the "profit_control_enabled" field. +func ProfitControlEnabledNEQ(v bool) predicate.Group { + return predicate.Group(sql.FieldNEQ(FieldProfitControlEnabled, v)) +} + +// ProfitMinMarginEQ applies the EQ predicate on the "profit_min_margin" field. +func ProfitMinMarginEQ(v float64) predicate.Group { + return predicate.Group(sql.FieldEQ(FieldProfitMinMargin, v)) +} + +// ProfitMinMarginNEQ applies the NEQ predicate on the "profit_min_margin" field. +func ProfitMinMarginNEQ(v float64) predicate.Group { + return predicate.Group(sql.FieldNEQ(FieldProfitMinMargin, v)) +} + +// ProfitMinMarginIn applies the In predicate on the "profit_min_margin" field. +func ProfitMinMarginIn(vs ...float64) predicate.Group { + return predicate.Group(sql.FieldIn(FieldProfitMinMargin, vs...)) +} + +// ProfitMinMarginNotIn applies the NotIn predicate on the "profit_min_margin" field. +func ProfitMinMarginNotIn(vs ...float64) predicate.Group { + return predicate.Group(sql.FieldNotIn(FieldProfitMinMargin, vs...)) +} + +// ProfitMinMarginGT applies the GT predicate on the "profit_min_margin" field. +func ProfitMinMarginGT(v float64) predicate.Group { + return predicate.Group(sql.FieldGT(FieldProfitMinMargin, v)) +} + +// ProfitMinMarginGTE applies the GTE predicate on the "profit_min_margin" field. +func ProfitMinMarginGTE(v float64) predicate.Group { + return predicate.Group(sql.FieldGTE(FieldProfitMinMargin, v)) +} + +// ProfitMinMarginLT applies the LT predicate on the "profit_min_margin" field. +func ProfitMinMarginLT(v float64) predicate.Group { + return predicate.Group(sql.FieldLT(FieldProfitMinMargin, v)) +} + +// ProfitMinMarginLTE applies the LTE predicate on the "profit_min_margin" field. +func ProfitMinMarginLTE(v float64) predicate.Group { + return predicate.Group(sql.FieldLTE(FieldProfitMinMargin, v)) +} + +// ProfitSafetyBufferEQ applies the EQ predicate on the "profit_safety_buffer" field. +func ProfitSafetyBufferEQ(v float64) predicate.Group { + return predicate.Group(sql.FieldEQ(FieldProfitSafetyBuffer, v)) +} + +// ProfitSafetyBufferNEQ applies the NEQ predicate on the "profit_safety_buffer" field. +func ProfitSafetyBufferNEQ(v float64) predicate.Group { + return predicate.Group(sql.FieldNEQ(FieldProfitSafetyBuffer, v)) +} + +// ProfitSafetyBufferIn applies the In predicate on the "profit_safety_buffer" field. +func ProfitSafetyBufferIn(vs ...float64) predicate.Group { + return predicate.Group(sql.FieldIn(FieldProfitSafetyBuffer, vs...)) +} + +// ProfitSafetyBufferNotIn applies the NotIn predicate on the "profit_safety_buffer" field. +func ProfitSafetyBufferNotIn(vs ...float64) predicate.Group { + return predicate.Group(sql.FieldNotIn(FieldProfitSafetyBuffer, vs...)) +} + +// ProfitSafetyBufferGT applies the GT predicate on the "profit_safety_buffer" field. +func ProfitSafetyBufferGT(v float64) predicate.Group { + return predicate.Group(sql.FieldGT(FieldProfitSafetyBuffer, v)) +} + +// ProfitSafetyBufferGTE applies the GTE predicate on the "profit_safety_buffer" field. +func ProfitSafetyBufferGTE(v float64) predicate.Group { + return predicate.Group(sql.FieldGTE(FieldProfitSafetyBuffer, v)) +} + +// ProfitSafetyBufferLT applies the LT predicate on the "profit_safety_buffer" field. +func ProfitSafetyBufferLT(v float64) predicate.Group { + return predicate.Group(sql.FieldLT(FieldProfitSafetyBuffer, v)) +} + +// ProfitSafetyBufferLTE applies the LTE predicate on the "profit_safety_buffer" field. +func ProfitSafetyBufferLTE(v float64) predicate.Group { + return predicate.Group(sql.FieldLTE(FieldProfitSafetyBuffer, v)) +} + // HasAPIKeys applies the HasEdge predicate on the "api_keys" edge. func HasAPIKeys() predicate.Group { return predicate.Group(func(s *sql.Selector) { diff --git a/backend/ent/group_create.go b/backend/ent/group_create.go index d12efb838a..54e34a32f6 100644 --- a/backend/ent/group_create.go +++ b/backend/ent/group_create.go @@ -725,6 +725,48 @@ func (_c *GroupCreate) SetReasoningEffortMappings(v []domain.ReasoningEffortMapp return _c } +// SetProfitControlEnabled sets the "profit_control_enabled" field. +func (_c *GroupCreate) SetProfitControlEnabled(v bool) *GroupCreate { + _c.mutation.SetProfitControlEnabled(v) + return _c +} + +// SetNillableProfitControlEnabled sets the "profit_control_enabled" field if the given value is not nil. +func (_c *GroupCreate) SetNillableProfitControlEnabled(v *bool) *GroupCreate { + if v != nil { + _c.SetProfitControlEnabled(*v) + } + return _c +} + +// SetProfitMinMargin sets the "profit_min_margin" field. +func (_c *GroupCreate) SetProfitMinMargin(v float64) *GroupCreate { + _c.mutation.SetProfitMinMargin(v) + return _c +} + +// SetNillableProfitMinMargin sets the "profit_min_margin" field if the given value is not nil. +func (_c *GroupCreate) SetNillableProfitMinMargin(v *float64) *GroupCreate { + if v != nil { + _c.SetProfitMinMargin(*v) + } + return _c +} + +// SetProfitSafetyBuffer sets the "profit_safety_buffer" field. +func (_c *GroupCreate) SetProfitSafetyBuffer(v float64) *GroupCreate { + _c.mutation.SetProfitSafetyBuffer(v) + return _c +} + +// SetNillableProfitSafetyBuffer sets the "profit_safety_buffer" field if the given value is not nil. +func (_c *GroupCreate) SetNillableProfitSafetyBuffer(v *float64) *GroupCreate { + if v != nil { + _c.SetProfitSafetyBuffer(*v) + } + return _c +} + // AddAPIKeyIDs adds the "api_keys" edge to the APIKey entity by IDs. func (_c *GroupCreate) AddAPIKeyIDs(ids ...int64) *GroupCreate { _c.mutation.AddAPIKeyIDs(ids...) @@ -998,6 +1040,18 @@ func (_c *GroupCreate) defaults() error { v := group.DefaultReasoningEffortMappings _c.mutation.SetReasoningEffortMappings(v) } + if _, ok := _c.mutation.ProfitControlEnabled(); !ok { + v := group.DefaultProfitControlEnabled + _c.mutation.SetProfitControlEnabled(v) + } + if _, ok := _c.mutation.ProfitMinMargin(); !ok { + v := group.DefaultProfitMinMargin + _c.mutation.SetProfitMinMargin(v) + } + if _, ok := _c.mutation.ProfitSafetyBuffer(); !ok { + v := group.DefaultProfitSafetyBuffer + _c.mutation.SetProfitSafetyBuffer(v) + } return nil } @@ -1156,6 +1210,15 @@ func (_c *GroupCreate) check() error { if _, ok := _c.mutation.ReasoningEffortMappings(); !ok { return &ValidationError{Name: "reasoning_effort_mappings", err: errors.New(`ent: missing required field "Group.reasoning_effort_mappings"`)} } + if _, ok := _c.mutation.ProfitControlEnabled(); !ok { + return &ValidationError{Name: "profit_control_enabled", err: errors.New(`ent: missing required field "Group.profit_control_enabled"`)} + } + if _, ok := _c.mutation.ProfitMinMargin(); !ok { + return &ValidationError{Name: "profit_min_margin", err: errors.New(`ent: missing required field "Group.profit_min_margin"`)} + } + if _, ok := _c.mutation.ProfitSafetyBuffer(); !ok { + return &ValidationError{Name: "profit_safety_buffer", err: errors.New(`ent: missing required field "Group.profit_safety_buffer"`)} + } return nil } @@ -1391,6 +1454,18 @@ func (_c *GroupCreate) createSpec() (*Group, *sqlgraph.CreateSpec) { _spec.SetField(group.FieldReasoningEffortMappings, field.TypeJSON, value) _node.ReasoningEffortMappings = value } + if value, ok := _c.mutation.ProfitControlEnabled(); ok { + _spec.SetField(group.FieldProfitControlEnabled, field.TypeBool, value) + _node.ProfitControlEnabled = value + } + if value, ok := _c.mutation.ProfitMinMargin(); ok { + _spec.SetField(group.FieldProfitMinMargin, field.TypeFloat64, value) + _node.ProfitMinMargin = value + } + if value, ok := _c.mutation.ProfitSafetyBuffer(); ok { + _spec.SetField(group.FieldProfitSafetyBuffer, field.TypeFloat64, value) + _node.ProfitSafetyBuffer = value + } if nodes := _c.mutation.APIKeysIDs(); len(nodes) > 0 { edge := &sqlgraph.EdgeSpec{ Rel: sqlgraph.O2M, @@ -2363,6 +2438,54 @@ func (u *GroupUpsert) UpdateReasoningEffortMappings() *GroupUpsert { return u } +// SetProfitControlEnabled sets the "profit_control_enabled" field. +func (u *GroupUpsert) SetProfitControlEnabled(v bool) *GroupUpsert { + u.Set(group.FieldProfitControlEnabled, v) + return u +} + +// UpdateProfitControlEnabled sets the "profit_control_enabled" field to the value that was provided on create. +func (u *GroupUpsert) UpdateProfitControlEnabled() *GroupUpsert { + u.SetExcluded(group.FieldProfitControlEnabled) + return u +} + +// SetProfitMinMargin sets the "profit_min_margin" field. +func (u *GroupUpsert) SetProfitMinMargin(v float64) *GroupUpsert { + u.Set(group.FieldProfitMinMargin, v) + return u +} + +// UpdateProfitMinMargin sets the "profit_min_margin" field to the value that was provided on create. +func (u *GroupUpsert) UpdateProfitMinMargin() *GroupUpsert { + u.SetExcluded(group.FieldProfitMinMargin) + return u +} + +// AddProfitMinMargin adds v to the "profit_min_margin" field. +func (u *GroupUpsert) AddProfitMinMargin(v float64) *GroupUpsert { + u.Add(group.FieldProfitMinMargin, v) + return u +} + +// SetProfitSafetyBuffer sets the "profit_safety_buffer" field. +func (u *GroupUpsert) SetProfitSafetyBuffer(v float64) *GroupUpsert { + u.Set(group.FieldProfitSafetyBuffer, v) + return u +} + +// UpdateProfitSafetyBuffer sets the "profit_safety_buffer" field to the value that was provided on create. +func (u *GroupUpsert) UpdateProfitSafetyBuffer() *GroupUpsert { + u.SetExcluded(group.FieldProfitSafetyBuffer) + return u +} + +// AddProfitSafetyBuffer adds v to the "profit_safety_buffer" field. +func (u *GroupUpsert) AddProfitSafetyBuffer(v float64) *GroupUpsert { + u.Add(group.FieldProfitSafetyBuffer, v) + return u +} + // UpdateNewValues updates the mutable fields using the new values that were set on create. // Using this option is equivalent to using: // @@ -3363,6 +3486,62 @@ func (u *GroupUpsertOne) UpdateReasoningEffortMappings() *GroupUpsertOne { }) } +// SetProfitControlEnabled sets the "profit_control_enabled" field. +func (u *GroupUpsertOne) SetProfitControlEnabled(v bool) *GroupUpsertOne { + return u.Update(func(s *GroupUpsert) { + s.SetProfitControlEnabled(v) + }) +} + +// UpdateProfitControlEnabled sets the "profit_control_enabled" field to the value that was provided on create. +func (u *GroupUpsertOne) UpdateProfitControlEnabled() *GroupUpsertOne { + return u.Update(func(s *GroupUpsert) { + s.UpdateProfitControlEnabled() + }) +} + +// SetProfitMinMargin sets the "profit_min_margin" field. +func (u *GroupUpsertOne) SetProfitMinMargin(v float64) *GroupUpsertOne { + return u.Update(func(s *GroupUpsert) { + s.SetProfitMinMargin(v) + }) +} + +// AddProfitMinMargin adds v to the "profit_min_margin" field. +func (u *GroupUpsertOne) AddProfitMinMargin(v float64) *GroupUpsertOne { + return u.Update(func(s *GroupUpsert) { + s.AddProfitMinMargin(v) + }) +} + +// UpdateProfitMinMargin sets the "profit_min_margin" field to the value that was provided on create. +func (u *GroupUpsertOne) UpdateProfitMinMargin() *GroupUpsertOne { + return u.Update(func(s *GroupUpsert) { + s.UpdateProfitMinMargin() + }) +} + +// SetProfitSafetyBuffer sets the "profit_safety_buffer" field. +func (u *GroupUpsertOne) SetProfitSafetyBuffer(v float64) *GroupUpsertOne { + return u.Update(func(s *GroupUpsert) { + s.SetProfitSafetyBuffer(v) + }) +} + +// AddProfitSafetyBuffer adds v to the "profit_safety_buffer" field. +func (u *GroupUpsertOne) AddProfitSafetyBuffer(v float64) *GroupUpsertOne { + return u.Update(func(s *GroupUpsert) { + s.AddProfitSafetyBuffer(v) + }) +} + +// UpdateProfitSafetyBuffer sets the "profit_safety_buffer" field to the value that was provided on create. +func (u *GroupUpsertOne) UpdateProfitSafetyBuffer() *GroupUpsertOne { + return u.Update(func(s *GroupUpsert) { + s.UpdateProfitSafetyBuffer() + }) +} + // Exec executes the query. func (u *GroupUpsertOne) Exec(ctx context.Context) error { if len(u.create.conflict) == 0 { @@ -4529,6 +4708,62 @@ func (u *GroupUpsertBulk) UpdateReasoningEffortMappings() *GroupUpsertBulk { }) } +// SetProfitControlEnabled sets the "profit_control_enabled" field. +func (u *GroupUpsertBulk) SetProfitControlEnabled(v bool) *GroupUpsertBulk { + return u.Update(func(s *GroupUpsert) { + s.SetProfitControlEnabled(v) + }) +} + +// UpdateProfitControlEnabled sets the "profit_control_enabled" field to the value that was provided on create. +func (u *GroupUpsertBulk) UpdateProfitControlEnabled() *GroupUpsertBulk { + return u.Update(func(s *GroupUpsert) { + s.UpdateProfitControlEnabled() + }) +} + +// SetProfitMinMargin sets the "profit_min_margin" field. +func (u *GroupUpsertBulk) SetProfitMinMargin(v float64) *GroupUpsertBulk { + return u.Update(func(s *GroupUpsert) { + s.SetProfitMinMargin(v) + }) +} + +// AddProfitMinMargin adds v to the "profit_min_margin" field. +func (u *GroupUpsertBulk) AddProfitMinMargin(v float64) *GroupUpsertBulk { + return u.Update(func(s *GroupUpsert) { + s.AddProfitMinMargin(v) + }) +} + +// UpdateProfitMinMargin sets the "profit_min_margin" field to the value that was provided on create. +func (u *GroupUpsertBulk) UpdateProfitMinMargin() *GroupUpsertBulk { + return u.Update(func(s *GroupUpsert) { + s.UpdateProfitMinMargin() + }) +} + +// SetProfitSafetyBuffer sets the "profit_safety_buffer" field. +func (u *GroupUpsertBulk) SetProfitSafetyBuffer(v float64) *GroupUpsertBulk { + return u.Update(func(s *GroupUpsert) { + s.SetProfitSafetyBuffer(v) + }) +} + +// AddProfitSafetyBuffer adds v to the "profit_safety_buffer" field. +func (u *GroupUpsertBulk) AddProfitSafetyBuffer(v float64) *GroupUpsertBulk { + return u.Update(func(s *GroupUpsert) { + s.AddProfitSafetyBuffer(v) + }) +} + +// UpdateProfitSafetyBuffer sets the "profit_safety_buffer" field to the value that was provided on create. +func (u *GroupUpsertBulk) UpdateProfitSafetyBuffer() *GroupUpsertBulk { + return u.Update(func(s *GroupUpsert) { + s.UpdateProfitSafetyBuffer() + }) +} + // Exec executes the query. func (u *GroupUpsertBulk) Exec(ctx context.Context) error { if u.create.err != nil { diff --git a/backend/ent/group_update.go b/backend/ent/group_update.go index 595565df35..c2fdea3f51 100644 --- a/backend/ent/group_update.go +++ b/backend/ent/group_update.go @@ -953,6 +953,62 @@ func (_u *GroupUpdate) AppendReasoningEffortMappings(v []domain.ReasoningEffortM return _u } +// SetProfitControlEnabled sets the "profit_control_enabled" field. +func (_u *GroupUpdate) SetProfitControlEnabled(v bool) *GroupUpdate { + _u.mutation.SetProfitControlEnabled(v) + return _u +} + +// SetNillableProfitControlEnabled sets the "profit_control_enabled" field if the given value is not nil. +func (_u *GroupUpdate) SetNillableProfitControlEnabled(v *bool) *GroupUpdate { + if v != nil { + _u.SetProfitControlEnabled(*v) + } + return _u +} + +// SetProfitMinMargin sets the "profit_min_margin" field. +func (_u *GroupUpdate) SetProfitMinMargin(v float64) *GroupUpdate { + _u.mutation.ResetProfitMinMargin() + _u.mutation.SetProfitMinMargin(v) + return _u +} + +// SetNillableProfitMinMargin sets the "profit_min_margin" field if the given value is not nil. +func (_u *GroupUpdate) SetNillableProfitMinMargin(v *float64) *GroupUpdate { + if v != nil { + _u.SetProfitMinMargin(*v) + } + return _u +} + +// AddProfitMinMargin adds value to the "profit_min_margin" field. +func (_u *GroupUpdate) AddProfitMinMargin(v float64) *GroupUpdate { + _u.mutation.AddProfitMinMargin(v) + return _u +} + +// SetProfitSafetyBuffer sets the "profit_safety_buffer" field. +func (_u *GroupUpdate) SetProfitSafetyBuffer(v float64) *GroupUpdate { + _u.mutation.ResetProfitSafetyBuffer() + _u.mutation.SetProfitSafetyBuffer(v) + return _u +} + +// SetNillableProfitSafetyBuffer sets the "profit_safety_buffer" field if the given value is not nil. +func (_u *GroupUpdate) SetNillableProfitSafetyBuffer(v *float64) *GroupUpdate { + if v != nil { + _u.SetProfitSafetyBuffer(*v) + } + return _u +} + +// AddProfitSafetyBuffer adds value to the "profit_safety_buffer" field. +func (_u *GroupUpdate) AddProfitSafetyBuffer(v float64) *GroupUpdate { + _u.mutation.AddProfitSafetyBuffer(v) + return _u +} + // AddAPIKeyIDs adds the "api_keys" edge to the APIKey entity by IDs. func (_u *GroupUpdate) AddAPIKeyIDs(ids ...int64) *GroupUpdate { _u.mutation.AddAPIKeyIDs(ids...) @@ -1544,6 +1600,21 @@ func (_u *GroupUpdate) sqlSave(ctx context.Context) (_node int, err error) { sqljson.Append(u, group.FieldReasoningEffortMappings, value) }) } + if value, ok := _u.mutation.ProfitControlEnabled(); ok { + _spec.SetField(group.FieldProfitControlEnabled, field.TypeBool, value) + } + if value, ok := _u.mutation.ProfitMinMargin(); ok { + _spec.SetField(group.FieldProfitMinMargin, field.TypeFloat64, value) + } + if value, ok := _u.mutation.AddedProfitMinMargin(); ok { + _spec.AddField(group.FieldProfitMinMargin, field.TypeFloat64, value) + } + if value, ok := _u.mutation.ProfitSafetyBuffer(); ok { + _spec.SetField(group.FieldProfitSafetyBuffer, field.TypeFloat64, value) + } + if value, ok := _u.mutation.AddedProfitSafetyBuffer(); ok { + _spec.AddField(group.FieldProfitSafetyBuffer, field.TypeFloat64, value) + } if _u.mutation.APIKeysCleared() { edge := &sqlgraph.EdgeSpec{ Rel: sqlgraph.O2M, @@ -2775,6 +2846,62 @@ func (_u *GroupUpdateOne) AppendReasoningEffortMappings(v []domain.ReasoningEffo return _u } +// SetProfitControlEnabled sets the "profit_control_enabled" field. +func (_u *GroupUpdateOne) SetProfitControlEnabled(v bool) *GroupUpdateOne { + _u.mutation.SetProfitControlEnabled(v) + return _u +} + +// SetNillableProfitControlEnabled sets the "profit_control_enabled" field if the given value is not nil. +func (_u *GroupUpdateOne) SetNillableProfitControlEnabled(v *bool) *GroupUpdateOne { + if v != nil { + _u.SetProfitControlEnabled(*v) + } + return _u +} + +// SetProfitMinMargin sets the "profit_min_margin" field. +func (_u *GroupUpdateOne) SetProfitMinMargin(v float64) *GroupUpdateOne { + _u.mutation.ResetProfitMinMargin() + _u.mutation.SetProfitMinMargin(v) + return _u +} + +// SetNillableProfitMinMargin sets the "profit_min_margin" field if the given value is not nil. +func (_u *GroupUpdateOne) SetNillableProfitMinMargin(v *float64) *GroupUpdateOne { + if v != nil { + _u.SetProfitMinMargin(*v) + } + return _u +} + +// AddProfitMinMargin adds value to the "profit_min_margin" field. +func (_u *GroupUpdateOne) AddProfitMinMargin(v float64) *GroupUpdateOne { + _u.mutation.AddProfitMinMargin(v) + return _u +} + +// SetProfitSafetyBuffer sets the "profit_safety_buffer" field. +func (_u *GroupUpdateOne) SetProfitSafetyBuffer(v float64) *GroupUpdateOne { + _u.mutation.ResetProfitSafetyBuffer() + _u.mutation.SetProfitSafetyBuffer(v) + return _u +} + +// SetNillableProfitSafetyBuffer sets the "profit_safety_buffer" field if the given value is not nil. +func (_u *GroupUpdateOne) SetNillableProfitSafetyBuffer(v *float64) *GroupUpdateOne { + if v != nil { + _u.SetProfitSafetyBuffer(*v) + } + return _u +} + +// AddProfitSafetyBuffer adds value to the "profit_safety_buffer" field. +func (_u *GroupUpdateOne) AddProfitSafetyBuffer(v float64) *GroupUpdateOne { + _u.mutation.AddProfitSafetyBuffer(v) + return _u +} + // AddAPIKeyIDs adds the "api_keys" edge to the APIKey entity by IDs. func (_u *GroupUpdateOne) AddAPIKeyIDs(ids ...int64) *GroupUpdateOne { _u.mutation.AddAPIKeyIDs(ids...) @@ -3396,6 +3523,21 @@ func (_u *GroupUpdateOne) sqlSave(ctx context.Context) (_node *Group, err error) sqljson.Append(u, group.FieldReasoningEffortMappings, value) }) } + if value, ok := _u.mutation.ProfitControlEnabled(); ok { + _spec.SetField(group.FieldProfitControlEnabled, field.TypeBool, value) + } + if value, ok := _u.mutation.ProfitMinMargin(); ok { + _spec.SetField(group.FieldProfitMinMargin, field.TypeFloat64, value) + } + if value, ok := _u.mutation.AddedProfitMinMargin(); ok { + _spec.AddField(group.FieldProfitMinMargin, field.TypeFloat64, value) + } + if value, ok := _u.mutation.ProfitSafetyBuffer(); ok { + _spec.SetField(group.FieldProfitSafetyBuffer, field.TypeFloat64, value) + } + if value, ok := _u.mutation.AddedProfitSafetyBuffer(); ok { + _spec.AddField(group.FieldProfitSafetyBuffer, field.TypeFloat64, value) + } if _u.mutation.APIKeysCleared() { edge := &sqlgraph.EdgeSpec{ Rel: sqlgraph.O2M, diff --git a/backend/ent/migrate/schema.go b/backend/ent/migrate/schema.go index 5d78a74338..0475bca3bf 100644 --- a/backend/ent/migrate/schema.go +++ b/backend/ent/migrate/schema.go @@ -947,6 +947,9 @@ var ( {Name: "rpm_limit", Type: field.TypeInt, Default: 0}, {Name: "max_reasoning_effort", Type: field.TypeString, Size: 20, Default: ""}, {Name: "reasoning_effort_mappings", Type: field.TypeJSON, SchemaType: map[string]string{"postgres": "jsonb"}}, + {Name: "profit_control_enabled", Type: field.TypeBool, Default: false}, + {Name: "profit_min_margin", Type: field.TypeFloat64, Default: 0, SchemaType: map[string]string{"postgres": "decimal(10,4)"}}, + {Name: "profit_safety_buffer", Type: field.TypeFloat64, Default: 0, SchemaType: map[string]string{"postgres": "decimal(10,4)"}}, } // GroupsTable holds the schema information for the "groups" table. GroupsTable = &schema.Table{ diff --git a/backend/ent/mutation.go b/backend/ent/mutation.go index cf9258ffa9..61ff4f9e6d 100644 --- a/backend/ent/mutation.go +++ b/backend/ent/mutation.go @@ -21922,6 +21922,11 @@ type GroupMutation struct { max_reasoning_effort *string reasoning_effort_mappings *[]domain.ReasoningEffortMapping appendreasoning_effort_mappings []domain.ReasoningEffortMapping + profit_control_enabled *bool + profit_min_margin *float64 + addprofit_min_margin *float64 + profit_safety_buffer *float64 + addprofit_safety_buffer *float64 clearedFields map[string]struct{} api_keys map[int64]struct{} removedapi_keys map[int64]struct{} @@ -24586,6 +24591,154 @@ func (m *GroupMutation) ResetReasoningEffortMappings() { m.appendreasoning_effort_mappings = nil } +// SetProfitControlEnabled sets the "profit_control_enabled" field. +func (m *GroupMutation) SetProfitControlEnabled(b bool) { + m.profit_control_enabled = &b +} + +// ProfitControlEnabled returns the value of the "profit_control_enabled" field in the mutation. +func (m *GroupMutation) ProfitControlEnabled() (r bool, exists bool) { + v := m.profit_control_enabled + if v == nil { + return + } + return *v, true +} + +// OldProfitControlEnabled returns the old "profit_control_enabled" field's value of the Group entity. +// If the Group object wasn't provided to the builder, the object is fetched from the database. +// An error is returned if the mutation operation is not UpdateOne, or the database query fails. +func (m *GroupMutation) OldProfitControlEnabled(ctx context.Context) (v bool, err error) { + if !m.op.Is(OpUpdateOne) { + return v, errors.New("OldProfitControlEnabled is only allowed on UpdateOne operations") + } + if m.id == nil || m.oldValue == nil { + return v, errors.New("OldProfitControlEnabled requires an ID field in the mutation") + } + oldValue, err := m.oldValue(ctx) + if err != nil { + return v, fmt.Errorf("querying old value for OldProfitControlEnabled: %w", err) + } + return oldValue.ProfitControlEnabled, nil +} + +// ResetProfitControlEnabled resets all changes to the "profit_control_enabled" field. +func (m *GroupMutation) ResetProfitControlEnabled() { + m.profit_control_enabled = nil +} + +// SetProfitMinMargin sets the "profit_min_margin" field. +func (m *GroupMutation) SetProfitMinMargin(f float64) { + m.profit_min_margin = &f + m.addprofit_min_margin = nil +} + +// ProfitMinMargin returns the value of the "profit_min_margin" field in the mutation. +func (m *GroupMutation) ProfitMinMargin() (r float64, exists bool) { + v := m.profit_min_margin + if v == nil { + return + } + return *v, true +} + +// OldProfitMinMargin returns the old "profit_min_margin" field's value of the Group entity. +// If the Group object wasn't provided to the builder, the object is fetched from the database. +// An error is returned if the mutation operation is not UpdateOne, or the database query fails. +func (m *GroupMutation) OldProfitMinMargin(ctx context.Context) (v float64, err error) { + if !m.op.Is(OpUpdateOne) { + return v, errors.New("OldProfitMinMargin is only allowed on UpdateOne operations") + } + if m.id == nil || m.oldValue == nil { + return v, errors.New("OldProfitMinMargin requires an ID field in the mutation") + } + oldValue, err := m.oldValue(ctx) + if err != nil { + return v, fmt.Errorf("querying old value for OldProfitMinMargin: %w", err) + } + return oldValue.ProfitMinMargin, nil +} + +// AddProfitMinMargin adds f to the "profit_min_margin" field. +func (m *GroupMutation) AddProfitMinMargin(f float64) { + if m.addprofit_min_margin != nil { + *m.addprofit_min_margin += f + } else { + m.addprofit_min_margin = &f + } +} + +// AddedProfitMinMargin returns the value that was added to the "profit_min_margin" field in this mutation. +func (m *GroupMutation) AddedProfitMinMargin() (r float64, exists bool) { + v := m.addprofit_min_margin + if v == nil { + return + } + return *v, true +} + +// ResetProfitMinMargin resets all changes to the "profit_min_margin" field. +func (m *GroupMutation) ResetProfitMinMargin() { + m.profit_min_margin = nil + m.addprofit_min_margin = nil +} + +// SetProfitSafetyBuffer sets the "profit_safety_buffer" field. +func (m *GroupMutation) SetProfitSafetyBuffer(f float64) { + m.profit_safety_buffer = &f + m.addprofit_safety_buffer = nil +} + +// ProfitSafetyBuffer returns the value of the "profit_safety_buffer" field in the mutation. +func (m *GroupMutation) ProfitSafetyBuffer() (r float64, exists bool) { + v := m.profit_safety_buffer + if v == nil { + return + } + return *v, true +} + +// OldProfitSafetyBuffer returns the old "profit_safety_buffer" field's value of the Group entity. +// If the Group object wasn't provided to the builder, the object is fetched from the database. +// An error is returned if the mutation operation is not UpdateOne, or the database query fails. +func (m *GroupMutation) OldProfitSafetyBuffer(ctx context.Context) (v float64, err error) { + if !m.op.Is(OpUpdateOne) { + return v, errors.New("OldProfitSafetyBuffer is only allowed on UpdateOne operations") + } + if m.id == nil || m.oldValue == nil { + return v, errors.New("OldProfitSafetyBuffer requires an ID field in the mutation") + } + oldValue, err := m.oldValue(ctx) + if err != nil { + return v, fmt.Errorf("querying old value for OldProfitSafetyBuffer: %w", err) + } + return oldValue.ProfitSafetyBuffer, nil +} + +// AddProfitSafetyBuffer adds f to the "profit_safety_buffer" field. +func (m *GroupMutation) AddProfitSafetyBuffer(f float64) { + if m.addprofit_safety_buffer != nil { + *m.addprofit_safety_buffer += f + } else { + m.addprofit_safety_buffer = &f + } +} + +// AddedProfitSafetyBuffer returns the value that was added to the "profit_safety_buffer" field in this mutation. +func (m *GroupMutation) AddedProfitSafetyBuffer() (r float64, exists bool) { + v := m.addprofit_safety_buffer + if v == nil { + return + } + return *v, true +} + +// ResetProfitSafetyBuffer resets all changes to the "profit_safety_buffer" field. +func (m *GroupMutation) ResetProfitSafetyBuffer() { + m.profit_safety_buffer = nil + m.addprofit_safety_buffer = nil +} + // AddAPIKeyIDs adds the "api_keys" edge to the APIKey entity by ids. func (m *GroupMutation) AddAPIKeyIDs(ids ...int64) { if m.api_keys == nil { @@ -24944,7 +25097,7 @@ func (m *GroupMutation) Type() string { // order to get all numeric fields that were incremented/decremented, call // AddedFields(). func (m *GroupMutation) Fields() []string { - fields := make([]string, 0, 52) + fields := make([]string, 0, 55) if m.created_at != nil { fields = append(fields, group.FieldCreatedAt) } @@ -25101,6 +25254,15 @@ func (m *GroupMutation) Fields() []string { if m.reasoning_effort_mappings != nil { fields = append(fields, group.FieldReasoningEffortMappings) } + if m.profit_control_enabled != nil { + fields = append(fields, group.FieldProfitControlEnabled) + } + if m.profit_min_margin != nil { + fields = append(fields, group.FieldProfitMinMargin) + } + if m.profit_safety_buffer != nil { + fields = append(fields, group.FieldProfitSafetyBuffer) + } return fields } @@ -25213,6 +25375,12 @@ func (m *GroupMutation) Field(name string) (ent.Value, bool) { return m.MaxReasoningEffort() case group.FieldReasoningEffortMappings: return m.ReasoningEffortMappings() + case group.FieldProfitControlEnabled: + return m.ProfitControlEnabled() + case group.FieldProfitMinMargin: + return m.ProfitMinMargin() + case group.FieldProfitSafetyBuffer: + return m.ProfitSafetyBuffer() } return nil, false } @@ -25326,6 +25494,12 @@ func (m *GroupMutation) OldField(ctx context.Context, name string) (ent.Value, e return m.OldMaxReasoningEffort(ctx) case group.FieldReasoningEffortMappings: return m.OldReasoningEffortMappings(ctx) + case group.FieldProfitControlEnabled: + return m.OldProfitControlEnabled(ctx) + case group.FieldProfitMinMargin: + return m.OldProfitMinMargin(ctx) + case group.FieldProfitSafetyBuffer: + return m.OldProfitSafetyBuffer(ctx) } return nil, fmt.Errorf("unknown Group field %s", name) } @@ -25699,6 +25873,27 @@ func (m *GroupMutation) SetField(name string, value ent.Value) error { } m.SetReasoningEffortMappings(v) return nil + case group.FieldProfitControlEnabled: + v, ok := value.(bool) + if !ok { + return fmt.Errorf("unexpected type %T for field %s", value, name) + } + m.SetProfitControlEnabled(v) + return nil + case group.FieldProfitMinMargin: + v, ok := value.(float64) + if !ok { + return fmt.Errorf("unexpected type %T for field %s", value, name) + } + m.SetProfitMinMargin(v) + return nil + case group.FieldProfitSafetyBuffer: + v, ok := value.(float64) + if !ok { + return fmt.Errorf("unexpected type %T for field %s", value, name) + } + m.SetProfitSafetyBuffer(v) + return nil } return fmt.Errorf("unknown Group field %s", name) } @@ -25770,6 +25965,12 @@ func (m *GroupMutation) AddedFields() []string { if m.addrpm_limit != nil { fields = append(fields, group.FieldRpmLimit) } + if m.addprofit_min_margin != nil { + fields = append(fields, group.FieldProfitMinMargin) + } + if m.addprofit_safety_buffer != nil { + fields = append(fields, group.FieldProfitSafetyBuffer) + } return fields } @@ -25820,6 +26021,10 @@ func (m *GroupMutation) AddedField(name string) (ent.Value, bool) { return m.AddedSortOrder() case group.FieldRpmLimit: return m.AddedRpmLimit() + case group.FieldProfitMinMargin: + return m.AddedProfitMinMargin() + case group.FieldProfitSafetyBuffer: + return m.AddedProfitSafetyBuffer() } return nil, false } @@ -25976,6 +26181,20 @@ func (m *GroupMutation) AddField(name string, value ent.Value) error { } m.AddRpmLimit(v) return nil + case group.FieldProfitMinMargin: + v, ok := value.(float64) + if !ok { + return fmt.Errorf("unexpected type %T for field %s", value, name) + } + m.AddProfitMinMargin(v) + return nil + case group.FieldProfitSafetyBuffer: + v, ok := value.(float64) + if !ok { + return fmt.Errorf("unexpected type %T for field %s", value, name) + } + m.AddProfitSafetyBuffer(v) + return nil } return fmt.Errorf("unknown Group numeric field %s", name) } @@ -26258,6 +26477,15 @@ func (m *GroupMutation) ResetField(name string) error { case group.FieldReasoningEffortMappings: m.ResetReasoningEffortMappings() return nil + case group.FieldProfitControlEnabled: + m.ResetProfitControlEnabled() + return nil + case group.FieldProfitMinMargin: + m.ResetProfitMinMargin() + return nil + case group.FieldProfitSafetyBuffer: + m.ResetProfitSafetyBuffer() + return nil } return fmt.Errorf("unknown Group field %s", name) } diff --git a/backend/ent/runtime/runtime.go b/backend/ent/runtime/runtime.go index 58d1a5f35e..e85660dcc1 100644 --- a/backend/ent/runtime/runtime.go +++ b/backend/ent/runtime/runtime.go @@ -1181,6 +1181,18 @@ func init() { groupDescReasoningEffortMappings := groupFields[48].Descriptor() // group.DefaultReasoningEffortMappings holds the default value on creation for the reasoning_effort_mappings field. group.DefaultReasoningEffortMappings = groupDescReasoningEffortMappings.Default.([]domain.ReasoningEffortMapping) + // groupDescProfitControlEnabled is the schema descriptor for profit_control_enabled field. + groupDescProfitControlEnabled := groupFields[49].Descriptor() + // group.DefaultProfitControlEnabled holds the default value on creation for the profit_control_enabled field. + group.DefaultProfitControlEnabled = groupDescProfitControlEnabled.Default.(bool) + // groupDescProfitMinMargin is the schema descriptor for profit_min_margin field. + groupDescProfitMinMargin := groupFields[50].Descriptor() + // group.DefaultProfitMinMargin holds the default value on creation for the profit_min_margin field. + group.DefaultProfitMinMargin = groupDescProfitMinMargin.Default.(float64) + // groupDescProfitSafetyBuffer is the schema descriptor for profit_safety_buffer field. + groupDescProfitSafetyBuffer := groupFields[51].Descriptor() + // group.DefaultProfitSafetyBuffer holds the default value on creation for the profit_safety_buffer field. + group.DefaultProfitSafetyBuffer = groupDescProfitSafetyBuffer.Default.(float64) idempotencyrecordMixin := schema.IdempotencyRecord{}.Mixin() idempotencyrecordMixinFields0 := idempotencyrecordMixin[0].Fields() _ = idempotencyrecordMixinFields0 diff --git a/backend/ent/schema/group.go b/backend/ent/schema/group.go index 4b6d056ef6..49c0f63151 100644 --- a/backend/ent/schema/group.go +++ b/backend/ent/schema/group.go @@ -234,6 +234,20 @@ func (Group) Fields() []ent.Field { Default([]domain.ReasoningEffortMapping{}). SchemaType(map[string]string{dialect.Postgres: "jsonb"}). Comment("OpenAI reasoning effort 自定义精确映射;先映射再应用上限"), + + // 分组利润控制(migration 191):openai/anthropic/gemini/grok/antigravity + // 的 token 分组可启用,composite 分组不能直接启用。 + field.Bool("profit_control_enabled"). + Default(false). + Comment("是否启用利润控制:调度时仅允许账号计费倍率满足毛利率要求的账号进入候选池"), + field.Float("profit_min_margin"). + SchemaType(map[string]string{dialect.Postgres: "decimal(10,4)"}). + Default(0). + Comment("最低毛利率,小数(0.30=30%);账号准入条件为 U <= D*(1-margin-buffer)"), + field.Float("profit_safety_buffer"). + SchemaType(map[string]string{dialect.Postgres: "decimal(10,4)"}). + Default(0). + Comment("安全缓冲,小数;与 margin 相加后从下游倍率中扣除,默认 0"), } } diff --git a/backend/internal/handler/admin/group_handler.go b/backend/internal/handler/admin/group_handler.go index 3646f5ab3a..11662cbbef 100644 --- a/backend/internal/handler/admin/group_handler.go +++ b/backend/internal/handler/admin/group_handler.go @@ -118,6 +118,9 @@ type CreateGroupRequest struct { PeakStart string `json:"peak_start"` PeakEnd string `json:"peak_end"` PeakRateMultiplier *float64 `json:"peak_rate_multiplier"` + ProfitControlEnabled bool `json:"profit_control_enabled"` + ProfitMinMargin *float64 `json:"profit_min_margin"` + ProfitSafetyBuffer *float64 `json:"profit_safety_buffer"` ImagePrice1K *float64 `json:"image_price_1k"` ImagePrice2K *float64 `json:"image_price_2k"` ImagePrice4K *float64 `json:"image_price_4k"` @@ -177,6 +180,9 @@ type UpdateGroupRequest struct { PeakStart *string `json:"peak_start"` PeakEnd *string `json:"peak_end"` PeakRateMultiplier *float64 `json:"peak_rate_multiplier"` + ProfitControlEnabled *bool `json:"profit_control_enabled"` + ProfitMinMargin *float64 `json:"profit_min_margin"` + ProfitSafetyBuffer *float64 `json:"profit_safety_buffer"` ImagePrice1K *float64 `json:"image_price_1k"` ImagePrice2K *float64 `json:"image_price_2k"` ImagePrice4K *float64 `json:"image_price_4k"` @@ -475,6 +481,11 @@ func (h *GroupHandler) Create(c *gin.Context) { return } + if err := service.ValidateProfitControlConfig(req.Platform, req.ProfitControlEnabled, float64ValueOrDefault(req.ProfitMinMargin, 0), float64ValueOrDefault(req.ProfitSafetyBuffer, 0)); err != nil { + response.BadRequest(c, err.Error()) + return + } + group, err := h.adminService.CreateGroup(c.Request.Context(), &service.CreateGroupInput{ Name: req.Name, Description: req.Description, @@ -497,6 +508,9 @@ func (h *GroupHandler) Create(c *gin.Context) { PeakStart: req.PeakStart, PeakEnd: req.PeakEnd, PeakRateMultiplier: req.PeakRateMultiplier, + ProfitControlEnabled: req.ProfitControlEnabled, + ProfitMinMargin: req.ProfitMinMargin, + ProfitSafetyBuffer: req.ProfitSafetyBuffer, ImagePrice1K: req.ImagePrice1K, ImagePrice2K: req.ImagePrice2K, ImagePrice4K: req.ImagePrice4K, @@ -616,6 +630,9 @@ func (h *GroupHandler) Update(c *gin.Context) { PeakStart: req.PeakStart, PeakEnd: req.PeakEnd, PeakRateMultiplier: req.PeakRateMultiplier, + ProfitControlEnabled: req.ProfitControlEnabled, + ProfitMinMargin: req.ProfitMinMargin, + ProfitSafetyBuffer: req.ProfitSafetyBuffer, ImagePrice1K: req.ImagePrice1K, ImagePrice2K: req.ImagePrice2K, ImagePrice4K: req.ImagePrice4K, diff --git a/backend/internal/handler/dto/mappers.go b/backend/internal/handler/dto/mappers.go index 04c811cedf..980f92deb9 100644 --- a/backend/internal/handler/dto/mappers.go +++ b/backend/internal/handler/dto/mappers.go @@ -193,6 +193,9 @@ func groupFromServiceBase(g *service.Group) Group { PeakStart: g.PeakStart, PeakEnd: g.PeakEnd, PeakRateMultiplier: g.PeakRateMultiplier, + ProfitControlEnabled: g.ProfitControlEnabled, + ProfitMinMargin: g.ProfitMinMargin, + ProfitSafetyBuffer: g.ProfitSafetyBuffer, ImagePrice1K: g.ImagePrice1K, ImagePrice2K: g.ImagePrice2K, ImagePrice4K: g.ImagePrice4K, diff --git a/backend/internal/handler/dto/types.go b/backend/internal/handler/dto/types.go index d5ba4e9db6..b990b57894 100644 --- a/backend/internal/handler/dto/types.go +++ b/backend/internal/handler/dto/types.go @@ -111,16 +111,20 @@ type Group struct { VideoRateIndependent bool `json:"video_rate_independent"` VideoRateMultiplier float64 `json:"video_rate_multiplier"` // 高峰时段倍率配置 - PeakRateEnabled bool `json:"peak_rate_enabled"` - PeakStart string `json:"peak_start"` - PeakEnd string `json:"peak_end"` - PeakRateMultiplier float64 `json:"peak_rate_multiplier"` - ImagePrice1K *float64 `json:"image_price_1k"` - ImagePrice2K *float64 `json:"image_price_2k"` - ImagePrice4K *float64 `json:"image_price_4k"` - VideoPrice480P *float64 `json:"video_price_480p"` - VideoPrice720P *float64 `json:"video_price_720p"` - VideoPrice1080P *float64 `json:"video_price_1080p"` + PeakRateEnabled bool `json:"peak_rate_enabled"` + PeakStart string `json:"peak_start"` + PeakEnd string `json:"peak_end"` + PeakRateMultiplier float64 `json:"peak_rate_multiplier"` + // 分组利润控制(五个 token 平台分组可启用;margin/buffer 为小数存储) + ProfitControlEnabled bool `json:"profit_control_enabled"` + ProfitMinMargin float64 `json:"profit_min_margin"` + ProfitSafetyBuffer float64 `json:"profit_safety_buffer"` + ImagePrice1K *float64 `json:"image_price_1k"` + ImagePrice2K *float64 `json:"image_price_2k"` + ImagePrice4K *float64 `json:"image_price_4k"` + VideoPrice480P *float64 `json:"video_price_480p"` + VideoPrice720P *float64 `json:"video_price_720p"` + VideoPrice1080P *float64 `json:"video_price_1080p"` // Codex alpha/search 网页搜索单次价格(USD/次);null 表示使用默认价 0.01 WebSearchPricePerCall *float64 `json:"web_search_price_per_call"` diff --git a/backend/internal/handler/gateway_handler.go b/backend/internal/handler/gateway_handler.go index 5c8fcf7129..9b4808d768 100644 --- a/backend/internal/handler/gateway_handler.go +++ b/backend/internal/handler/gateway_handler.go @@ -196,6 +196,8 @@ func (h *GatewayHandler) Messages(c *gin.Context) { setOpsRequestContext(c, reqModel, reqStream) setOpsEndpointContext(c, "", int16(service.RequestTypeFromLegacy(reqStream, false))) + pricingCtx, pricingAt := service.WithGatewayTokenRequestPricing(c.Request.Context()) + c.Request = c.Request.WithContext(pricingCtx) // 验证 model 必填 if reqModel == "" { @@ -421,9 +423,20 @@ func (h *GatewayHandler) Messages(c *gin.Context) { } // Slot acquired: no longer waiting in queue. releaseWait() - if err := h.gatewayService.BindStickySession(c.Request.Context(), apiKey.GroupID, sessionKey, account.ID); err != nil { - reqLog.Warn("gateway.bind_sticky_session_failed", zap.Int64("account_id", account.ID), zap.Error(err)) + } + latest, vetoed, reason := h.gatewayService.GatewayProfitControlVetoLatest(c.Request.Context(), account) + if vetoed { + if accountReleaseFunc != nil { + accountReleaseFunc() } + reqLog.Debug("gateway.account_slot_profit_vetoed", zap.Int64("account_id", account.ID), zap.String("reason", reason)) + fs.FailedAccountIDs[account.ID] = struct{}{} + continue + } + account = latest + selection.Account = latest + if err := h.gatewayService.BindStickySessionAfterProfitAdmission(c.Request.Context(), apiKey.GroupID, sessionKey, account.ID, sessionBoundAccountID); err != nil { + reqLog.Warn("gateway.bind_sticky_session_after_profit_admission_failed", zap.Int64("account_id", account.ID), zap.Error(err)) } // 账号槽位/等待计数需要在超时或断开时安全回收 accountReleaseFunc = wrapReleaseOnDone(c.Request.Context(), accountReleaseFunc) @@ -544,6 +557,7 @@ func (h *GatewayHandler) Messages(c *gin.Context) { User: apiKey.User, Account: account, Subscription: subscription, + PricingAt: pricingAt, InboundEndpoint: inboundEndpoint, UpstreamEndpoint: upstreamEndpoint, UserAgent: userAgent, @@ -721,13 +735,20 @@ func (h *GatewayHandler) Messages(c *gin.Context) { } // Slot acquired: no longer waiting in queue. releaseWait() - reqLog.Info("sticky.bind_after_wait", - zap.String("session_key", sessionKey), - zap.Int64("account_id", account.ID), - ) - if err := h.gatewayService.BindStickySession(c.Request.Context(), currentAPIKey.GroupID, sessionKey, account.ID); err != nil { - reqLog.Warn("gateway.bind_sticky_session_failed", zap.Int64("account_id", account.ID), zap.Error(err)) + } + latest, vetoed, reason := h.gatewayService.GatewayProfitControlVetoLatest(c.Request.Context(), account) + if vetoed { + if accountReleaseFunc != nil { + accountReleaseFunc() } + reqLog.Debug("gateway.account_slot_profit_vetoed", zap.Int64("account_id", account.ID), zap.String("reason", reason)) + fs.FailedAccountIDs[account.ID] = struct{}{} + continue + } + account = latest + selection.Account = latest + if err := h.gatewayService.BindStickySessionAfterProfitAdmission(c.Request.Context(), currentAPIKey.GroupID, sessionKey, account.ID, sessionBoundAccountID); err != nil { + reqLog.Warn("gateway.bind_sticky_session_after_profit_admission_failed", zap.Int64("account_id", account.ID), zap.Error(err)) } // 账号槽位/等待计数需要在超时或断开时安全回收 accountReleaseFunc = wrapReleaseOnDone(c.Request.Context(), accountReleaseFunc) @@ -864,6 +885,7 @@ func (h *GatewayHandler) Messages(c *gin.Context) { User: currentAPIKey.User, Account: account, Subscription: currentSubscription, + PricingAt: pricingAt, InboundEndpoint: inboundEndpoint, UpstreamEndpoint: upstreamEndpoint, UserAgent: userAgent, diff --git a/backend/internal/handler/gateway_handler_chat_completions.go b/backend/internal/handler/gateway_handler_chat_completions.go index 3bbac76adb..cc504235a3 100644 --- a/backend/internal/handler/gateway_handler_chat_completions.go +++ b/backend/internal/handler/gateway_handler_chat_completions.go @@ -93,6 +93,8 @@ func (h *GatewayHandler) ChatCompletions(c *gin.Context) { setOpsRequestContext(c, reqModel, reqStream) setOpsEndpointContext(c, "", int16(service.RequestTypeFromLegacy(reqStream, false))) + pricingCtx, pricingAt := service.WithGatewayTokenRequestPricing(c.Request.Context()) + c.Request = c.Request.WithContext(pricingCtx) // 解析渠道级模型映射 channelMapping, _ := h.gatewayService.ResolveChannelMappingAndRestrict(c.Request.Context(), apiKey.GroupID, reqModel) @@ -157,6 +159,7 @@ func (h *GatewayHandler) ChatCompletions(c *gin.Context) { if groupPlatform == service.PlatformGemini && selectionSessionHash != "" { selectionSessionHash = "gemini:" + selectionSessionHash } + sessionBoundAccountID, _ := h.gatewayService.GetCachedSessionAccountID(c.Request.Context(), apiKey.GroupID, selectionSessionHash) // 3. Account selection + failover loop fs := NewFailoverState(h.maxAccountSwitches, false) @@ -223,6 +226,20 @@ func (h *GatewayHandler) ChatCompletions(c *gin.Context) { return } } + latest, vetoed, reason := h.gatewayService.GatewayProfitControlVetoLatest(c.Request.Context(), account) + if vetoed { + if accountReleaseFunc != nil { + accountReleaseFunc() + } + reqLog.Debug("gateway.cc.account_slot_profit_vetoed", zap.Int64("account_id", account.ID), zap.String("reason", reason)) + fs.FailedAccountIDs[account.ID] = struct{}{} + continue + } + account = latest + selection.Account = latest + if err := h.gatewayService.BindStickySessionAfterProfitAdmission(c.Request.Context(), apiKey.GroupID, selectionSessionHash, account.ID, sessionBoundAccountID); err != nil { + reqLog.Warn("gateway.cc.bind_sticky_session_after_profit_admission_failed", zap.Int64("account_id", account.ID), zap.Error(err)) + } accountReleaseFunc = wrapReleaseOnDone(c.Request.Context(), accountReleaseFunc) if groupPlatform == service.PlatformGemini && account.Platform != service.PlatformGemini { @@ -318,6 +335,7 @@ func (h *GatewayHandler) ChatCompletions(c *gin.Context) { User: apiKey.User, Account: account, Subscription: subscription, + PricingAt: pricingAt, InboundEndpoint: inboundEndpoint, UpstreamEndpoint: upstreamEndpoint, UserAgent: userAgent, diff --git a/backend/internal/handler/gateway_handler_responses.go b/backend/internal/handler/gateway_handler_responses.go index 5653438613..e50f20529d 100644 --- a/backend/internal/handler/gateway_handler_responses.go +++ b/backend/internal/handler/gateway_handler_responses.go @@ -90,9 +90,13 @@ func (h *GatewayHandler) Responses(c *gin.Context) { setOpsRequestContext(c, reqModel, reqStream) setOpsEndpointContext(c, "", int16(service.RequestTypeFromLegacy(reqStream, false))) requestCtx := c.Request.Context() + pricingAt := time.Time{} if service.IsImageGenerationIntentForPlatform("/v1/responses", reqModel, body, openAICompatibleRequestPlatform(c.Request.Context(), apiKey)) { requestCtx = service.WithOpenAIImageGenerationIntent(requestCtx) + } else { + requestCtx, pricingAt = service.WithGatewayTokenRequestPricing(requestCtx) } + c.Request = c.Request.WithContext(requestCtx) // 解析渠道级模型映射 channelMapping, _ := h.gatewayService.ResolveChannelMappingAndRestrict(requestCtx, apiKey.GroupID, reqModel) @@ -157,6 +161,7 @@ func (h *GatewayHandler) Responses(c *gin.Context) { APIKeyID: apiKey.ID, } sessionHash := h.gatewayService.GenerateSessionHash(parsedReq) + sessionBoundAccountID, _ := h.gatewayService.GetCachedSessionAccountID(requestCtx, apiKey.GroupID, sessionHash) // 3. Account selection + failover loop fs := NewFailoverState(h.maxAccountSwitches, false) @@ -220,6 +225,20 @@ func (h *GatewayHandler) Responses(c *gin.Context) { return } } + latest, vetoed, reason := h.gatewayService.GatewayProfitControlVetoLatest(requestCtx, account) + if vetoed { + if accountReleaseFunc != nil { + accountReleaseFunc() + } + reqLog.Debug("gateway.responses.account_slot_profit_vetoed", zap.Int64("account_id", account.ID), zap.String("reason", reason)) + fs.FailedAccountIDs[account.ID] = struct{}{} + continue + } + account = latest + selection.Account = latest + if err := h.gatewayService.BindStickySessionAfterProfitAdmission(requestCtx, apiKey.GroupID, sessionHash, account.ID, sessionBoundAccountID); err != nil { + reqLog.Warn("gateway.responses.bind_sticky_session_after_profit_admission_failed", zap.Int64("account_id", account.ID), zap.Error(err)) + } accountReleaseFunc = wrapReleaseOnDone(c.Request.Context(), accountReleaseFunc) // 5. Forward request @@ -299,6 +318,7 @@ func (h *GatewayHandler) Responses(c *gin.Context) { User: apiKey.User, Account: account, Subscription: subscription, + PricingAt: pricingAt, InboundEndpoint: inboundEndpoint, UpstreamEndpoint: upstreamEndpoint, UserAgent: userAgent, diff --git a/backend/internal/handler/gemini_v1beta_handler.go b/backend/internal/handler/gemini_v1beta_handler.go index 3c60a74043..8caf862673 100644 --- a/backend/internal/handler/gemini_v1beta_handler.go +++ b/backend/internal/handler/gemini_v1beta_handler.go @@ -204,6 +204,8 @@ func (h *GatewayHandler) GeminiV1BetaModels(c *gin.Context) { setOpsRequestContext(c, modelName, stream) setOpsEndpointContext(c, "", int16(service.RequestTypeFromLegacy(stream, false))) + pricingCtx, pricingAt := service.WithGatewayTokenRequestPricing(c.Request.Context()) + c.Request = c.Request.WithContext(pricingCtx) if decision := h.checkSecurityAudit(c, reqLog, apiKey, authSubject, service.ContentModerationProtocolGemini, modelName, body); decision != nil && !decision.AllowNextStage { googleSecurityAuditError(c, decision) @@ -353,6 +355,7 @@ func (h *GatewayHandler) GeminiV1BetaModels(c *gin.Context) { // 判断是否真的绑定了粘性会话:有 sessionKey 且已经绑定到某个账号 hasBoundSession := sessionKey != "" && sessionBoundAccountID > 0 + profitStickyAccountID := sessionBoundAccountID cleanedForUnknownBinding := false fs := NewFailoverState(h.maxAccountSwitchesGemini, hasBoundSession) @@ -466,9 +469,20 @@ func (h *GatewayHandler) GeminiV1BetaModels(c *gin.Context) { geminiConcurrency.DecrementAccountWaitCount(c.Request.Context(), account.ID) accountWaitCounted = false } - if err := h.gatewayService.BindStickySession(c.Request.Context(), apiKey.GroupID, sessionKey, account.ID); err != nil { - reqLog.Warn("gemini.bind_sticky_session_failed", zap.Int64("account_id", account.ID), zap.Error(err)) + } + latest, vetoed, reason := h.gatewayService.GatewayProfitControlVetoLatest(c.Request.Context(), account) + if vetoed { + if accountReleaseFunc != nil { + accountReleaseFunc() } + reqLog.Debug("gemini.account_slot_profit_vetoed", zap.Int64("account_id", account.ID), zap.String("reason", reason)) + fs.FailedAccountIDs[account.ID] = struct{}{} + continue + } + account = latest + selection.Account = latest + if err := h.gatewayService.BindStickySessionAfterProfitAdmission(c.Request.Context(), apiKey.GroupID, sessionKey, account.ID, profitStickyAccountID); err != nil { + reqLog.Warn("gemini.bind_sticky_session_after_profit_admission_failed", zap.Int64("account_id", account.ID), zap.Error(err)) } // 账号槽位/等待计数需要在超时或断开时安全回收 accountReleaseFunc = wrapReleaseOnDone(c.Request.Context(), accountReleaseFunc) @@ -553,6 +567,7 @@ func (h *GatewayHandler) GeminiV1BetaModels(c *gin.Context) { User: apiKey.User, Account: account, Subscription: subscription, + PricingAt: pricingAt, InboundEndpoint: inboundEndpoint, UpstreamEndpoint: upstreamEndpoint, UserAgent: userAgent, diff --git a/backend/internal/handler/grok_media.go b/backend/internal/handler/grok_media.go index fd0697884e..fb9df0228e 100644 --- a/backend/internal/handler/grok_media.go +++ b/backend/internal/handler/grok_media.go @@ -293,8 +293,13 @@ func (h *OpenAIGatewayHandler) handleGrokMedia(c *gin.Context, endpoint service. sessionHash = ensureOpenAIPoolModeSessionHash(sessionHash, account) setOpsSelectedAccount(c, account.ID, account.Platform) - accountReleaseFunc, accountAcquired := h.acquireResponsesAccountSlot(c, apiKey.GroupID, sessionHash, selection, false, &streamStarted, reqLog) - if !accountAcquired { + accountReleaseFunc, slotResult := h.acquireResponsesAccountSlot(c, apiKey.GroupID, sessionHash, selection, false, &streamStarted, reqLog) + if slotResult == openAISlotAcquireProfitVetoed { + // Grok 分组非 openai 平台不装利润门,此分支实际不可达;防御性排除重选。 + failedAccountIDs[account.ID] = struct{}{} + continue + } + if slotResult != openAISlotAcquireOK { return } diff --git a/backend/internal/handler/openai_alpha_search.go b/backend/internal/handler/openai_alpha_search.go index c76b7d9fa5..186e85e544 100644 --- a/backend/internal/handler/openai_alpha_search.go +++ b/backend/internal/handler/openai_alpha_search.go @@ -113,6 +113,11 @@ func (h *OpenAIGatewayHandler) AlphaSearch(c *gin.Context) { var oauth429FailoverState service.OpenAIOAuth429FailoverState routingStart := time.Now() + // 分组利润控制:alpha search 文本入口请求级装门并固定 pricingAt + //(记录路径经 service.OpenAIPricingAtFromContext 从请求 ctx 回读)。 + asPricingCtx, _ := h.gatewayService.WithOpenAIRequestPricingContext(c.Request.Context(), apiKey.GroupID, false) + c.Request = c.Request.WithContext(asPricingCtx) + for { selection, _, err := h.gatewayService.SelectAccountWithSchedulerForCapability( c.Request.Context(), @@ -151,8 +156,13 @@ func (h *OpenAIGatewayHandler) AlphaSearch(c *gin.Context) { account := selection.Account setOpsSelectedAccount(c, account.ID, account.Platform) - accountRelease, acquired := h.acquireResponsesAccountSlot(c, apiKey.GroupID, sessionHash, selection, false, &streamStarted, reqLog) - if !acquired { + accountRelease, slotResult := h.acquireResponsesAccountSlot(c, apiKey.GroupID, sessionHash, selection, false, &streamStarted, reqLog) + if slotResult == openAISlotAcquireProfitVetoed { + // 利润终检否决:排除该账号重新选号,全池耗尽由下一轮选号报错。 + failedAccountIDs[account.ID] = struct{}{} + continue + } + if slotResult != openAISlotAcquireOK { return } service.SetOpsLatencyMs(c, service.OpsRoutingLatencyMsKey, time.Since(routingStart).Milliseconds()) @@ -255,6 +265,7 @@ func (h *OpenAIGatewayHandler) recordAlphaSearchUsage( QuotaPlatform: quotaPlatform, SessionID: sessionID, ChannelUsageFields: channelMapping.ToUsageFields(requestedModel, result.UpstreamModel), + PricingAt: service.OpenAIPricingAtFromContext(c.Request.Context()), }); err != nil { logger.L().With( zap.String("component", "handler.openai_gateway.alpha_search"), diff --git a/backend/internal/handler/openai_chat_completions.go b/backend/internal/handler/openai_chat_completions.go index 39c12253b1..8cdc577e0c 100644 --- a/backend/internal/handler/openai_chat_completions.go +++ b/backend/internal/handler/openai_chat_completions.go @@ -149,6 +149,10 @@ func (h *OpenAIGatewayHandler) ChatCompletions(c *gin.Context) { var lastFailoverErr *service.UpstreamFailoverError var oauth429FailoverState service.OpenAIOAuth429FailoverState + // 分组利润控制:chat completions 文本入口请求级装门并固定 pricingAt。 + ccPricingCtx, pricingAt := h.gatewayService.WithOpenAIRequestPricingContext(c.Request.Context(), apiKey.GroupID, false) + c.Request = c.Request.WithContext(ccPricingCtx) + for { if failoverClientGone(c) { return @@ -207,8 +211,13 @@ func (h *OpenAIGatewayHandler) ChatCompletions(c *gin.Context) { _ = scheduleDecision setOpsSelectedAccount(c, account.ID, account.Platform) - accountReleaseFunc, acquired := h.acquireResponsesAccountSlot(c, apiKey.GroupID, sessionHash, selection, reqStream, &streamStarted, reqLog) - if !acquired { + accountReleaseFunc, slotResult := h.acquireResponsesAccountSlot(c, apiKey.GroupID, sessionHash, selection, reqStream, &streamStarted, reqLog) + if slotResult == openAISlotAcquireProfitVetoed { + // 利润终检否决:排除该账号重新选号,全池耗尽由下一轮选号报错。 + failedAccountIDs[account.ID] = struct{}{} + continue + } + if slotResult != openAISlotAcquireOK { return } @@ -358,6 +367,7 @@ func (h *OpenAIGatewayHandler) ChatCompletions(c *gin.Context) { QuotaPlatform: quotaPlatform, SessionID: sessionID, ChannelUsageFields: clientRequestedUsageFields(c, channelMapping, reqModel, result.UpstreamModel), + PricingAt: pricingAt, CyberBlocked: cyberBlocked, }); err != nil { logger.L().With( diff --git a/backend/internal/handler/openai_embeddings.go b/backend/internal/handler/openai_embeddings.go index 3ed44b12b8..4fe168b757 100644 --- a/backend/internal/handler/openai_embeddings.go +++ b/backend/internal/handler/openai_embeddings.go @@ -116,6 +116,10 @@ func (h *OpenAIGatewayHandler) Embeddings(c *gin.Context) { } routingStart := time.Now() + // 分组利润控制:embeddings 文本入口请求级装门并固定 pricingAt。 + embPricingCtx, pricingAt := h.gatewayService.WithOpenAIRequestPricingContext(c.Request.Context(), apiKey.GroupID, false) + c.Request = c.Request.WithContext(embPricingCtx) + for { selection, _, err := h.gatewayService.SelectAccountWithSchedulerForCapability( c.Request.Context(), @@ -165,8 +169,13 @@ func (h *OpenAIGatewayHandler) Embeddings(c *gin.Context) { account := selection.Account setOpsSelectedAccount(c, account.ID, account.Platform) - accountReleaseFunc, accountAcquired := h.acquireResponsesAccountSlot(c, apiKey.GroupID, "", selection, false, &streamStarted, reqLog) - if !accountAcquired { + accountReleaseFunc, slotResult := h.acquireResponsesAccountSlot(c, apiKey.GroupID, "", selection, false, &streamStarted, reqLog) + if slotResult == openAISlotAcquireProfitVetoed { + // 利润终检否决:排除该账号重新选号,全池耗尽由下一轮选号报错。 + failedAccountIDs[account.ID] = struct{}{} + continue + } + if slotResult != openAISlotAcquireOK { return } @@ -260,6 +269,7 @@ func (h *OpenAIGatewayHandler) Embeddings(c *gin.Context) { QuotaPlatform: quotaPlatform, SessionID: sessionID, ChannelUsageFields: clientRequestedUsageFields(c, channelMapping, reqModel, result.UpstreamModel), + PricingAt: pricingAt, }); err != nil { logger.L().With( zap.String("component", "handler.openai_gateway.embeddings"), diff --git a/backend/internal/handler/openai_gateway_handler.go b/backend/internal/handler/openai_gateway_handler.go index 8602a684cb..d7fc3c8824 100644 --- a/backend/internal/handler/openai_gateway_handler.go +++ b/backend/internal/handler/openai_gateway_handler.go @@ -439,6 +439,12 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) { // 该判断已排除 Codex 被动 image_gen namespace,避免 CC-only 账号被误过滤(#4476)。 requiredCapability := openAIResponsesRequiredCapability(imageIntent, requestPlatform) + // 分组利润控制:请求级装配定价上下文——pricingAt 固定本请求的 + // D 与计费高峰因子,选号、槽位终检与全部 failover 重入共用同一门与阈值; + // 显式生图意图跳门(图片/视频不在利润门范围)。 + pricingCtx, pricingAt := h.gatewayService.WithOpenAIRequestPricingContext(c.Request.Context(), apiKey.GroupID, imageIntent) + c.Request = c.Request.WithContext(pricingCtx) + for { // Streaming Forward intentionally detaches the upstream request so usage can // be drained after a disconnect. Re-check the client context before every @@ -516,8 +522,13 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) { reqLog.Debug("openai.account_selected", zap.Int64("account_id", account.ID), zap.String("account_name", account.Name)) setOpsSelectedAccount(c, account.ID, account.Platform) - accountReleaseFunc, acquired := h.acquireResponsesAccountSlot(c, apiKey.GroupID, sessionHash, selection, reqStream, &streamStarted, reqLog) - if !acquired { + accountReleaseFunc, slotResult := h.acquireResponsesAccountSlot(c, apiKey.GroupID, sessionHash, selection, reqStream, &streamStarted, reqLog) + if slotResult == openAISlotAcquireProfitVetoed { + // 利润终检否决:排除该账号重新选号,全池耗尽由下一轮选号报错。 + failedAccountIDs[account.ID] = struct{}{} + continue + } + if slotResult != openAISlotAcquireOK { return } @@ -696,6 +707,7 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) { QuotaPlatform: quotaPlatform, SessionID: sessionID, ChannelUsageFields: clientRequestedUsageFields(c, channelMapping, reqModel, result.UpstreamModel), + PricingAt: pricingAt, CyberBlocked: cyberBlocked, }); err != nil { logger.L().With( @@ -994,6 +1006,10 @@ func (h *OpenAIGatewayHandler) Messages(c *gin.Context) { var oauth429FailoverState service.OpenAIOAuth429FailoverState effectiveMappedModel := preferredMappedModel + // 分组利润控制:Messages 文本入口同样请求级装门并固定 pricingAt。 + msgPricingCtx, pricingAt := h.gatewayService.WithOpenAIRequestPricingContext(c.Request.Context(), apiKey.GroupID, false) + c.Request = c.Request.WithContext(msgPricingCtx) + for { if failoverClientGone(c) { return @@ -1058,8 +1074,13 @@ func (h *OpenAIGatewayHandler) Messages(c *gin.Context) { _ = scheduleDecision setOpsSelectedAccount(c, account.ID, account.Platform) - accountReleaseFunc, acquired := h.acquireResponsesAccountSlot(c, apiKey.GroupID, sessionHash, selection, reqStream, &streamStarted, reqLog) - if !acquired { + accountReleaseFunc, slotResult := h.acquireResponsesAccountSlot(c, apiKey.GroupID, sessionHash, selection, reqStream, &streamStarted, reqLog) + if slotResult == openAISlotAcquireProfitVetoed { + // 利润终检否决:排除该账号重新选号,全池耗尽由下一轮选号报错。 + failedAccountIDs[account.ID] = struct{}{} + continue + } + if slotResult != openAISlotAcquireOK { return } @@ -1208,6 +1229,7 @@ func (h *OpenAIGatewayHandler) Messages(c *gin.Context) { QuotaPlatform: quotaPlatform, SessionID: sessionID, ChannelUsageFields: clientRequestedUsageFields(c, channelMappingMsg, reqModel, result.UpstreamModel), + PricingAt: pricingAt, CyberBlocked: cyberBlocked, }); err != nil { logger.L().With( @@ -1348,6 +1370,19 @@ func (h *OpenAIGatewayHandler) acquireResponsesUserSlot( return wrapReleaseOnDone(ctx, userReleaseFunc), true } +// openAISlotAcquireResult 是账号槽位获取的三态结果。 +type openAISlotAcquireResult int + +const ( + openAISlotAcquireOK openAISlotAcquireResult = iota + // openAISlotAcquireFailed:错误响应已写出,调用方直接 return。 + openAISlotAcquireFailed + // openAISlotAcquireProfitVetoed:槽位获取成功后利润终检否决。槽位已释放、 + // 未写任何响应;调用方应把该账号加入本请求排除集并重新选号,全池耗尽由 + // 下一轮选号返回标准 no available accounts。 + openAISlotAcquireProfitVetoed +) + func (h *OpenAIGatewayHandler) acquireResponsesAccountSlot( c *gin.Context, groupID *int64, @@ -1356,22 +1391,32 @@ func (h *OpenAIGatewayHandler) acquireResponsesAccountSlot( reqStream bool, streamStarted *bool, reqLog *zap.Logger, -) (func(), bool) { +) (func(), openAISlotAcquireResult) { if selection == nil || selection.Account == nil { markOpsRoutingCapacityLimited(c) h.handleStreamingAwareError(c, http.StatusServiceUnavailable, "api_error", "No available accounts", *streamStarted) - return nil, false + return nil, openAISlotAcquireFailed } ctx := c.Request.Context() account := selection.Account if selection.Acquired { - return wrapReleaseOnDone(ctx, selection.ReleaseFunc), true + latest, vetoed, reason := h.gatewayService.ProfitControlVetoLatest(ctx, account) + if vetoed { + if selection.ReleaseFunc != nil { + selection.ReleaseFunc() + } + reqLog.Debug("openai.account_slot_profit_vetoed", zap.Int64("account_id", account.ID), zap.String("reason", reason)) + return nil, openAISlotAcquireProfitVetoed + } + account = latest + selection.Account = latest + return wrapReleaseOnDone(ctx, selection.ReleaseFunc), openAISlotAcquireOK } if selection.WaitPlan == nil { markOpsRoutingCapacityLimited(c) h.handleStreamingAwareError(c, http.StatusServiceUnavailable, "api_error", "No available accounts", *streamStarted) - return nil, false + return nil, openAISlotAcquireFailed } fastReleaseFunc, fastAcquired, err := h.concurrencyHelper.TryAcquireAccountSlot( @@ -1382,13 +1427,25 @@ func (h *OpenAIGatewayHandler) acquireResponsesAccountSlot( if err != nil { reqLog.Warn("openai.account_slot_quick_acquire_failed", zap.Int64("account_id", account.ID), zap.Error(err)) h.handleConcurrencyError(c, err, "account", *streamStarted) - return nil, false + return nil, openAISlotAcquireFailed } if fastAcquired { - if err := h.gatewayService.BindStickySession(ctx, groupID, sessionHash, account.ID); err != nil { - reqLog.Warn("openai.bind_sticky_session_failed", zap.Int64("account_id", account.ID), zap.Error(err)) + // 分组利润控制:快速抢槽成功后终检。选号与抢槽之间账号 + // 倍率可能刷新,越线则释放槽位交由调用方排除重选,不绑定粘连。 + latest, vetoed, reason := h.gatewayService.ProfitControlVetoLatest(ctx, account) + if vetoed { + if fastReleaseFunc != nil { + fastReleaseFunc() + } + reqLog.Debug("openai.account_slot_profit_vetoed", zap.Int64("account_id", account.ID), zap.String("reason", reason)) + return nil, openAISlotAcquireProfitVetoed } - return wrapReleaseOnDone(ctx, fastReleaseFunc), true + account = latest + selection.Account = latest + if err := h.gatewayService.BindStickySessionAfterProfitAdmission(ctx, groupID, sessionHash, account.ID); err != nil { + reqLog.Warn("openai.bind_sticky_session_after_profit_admission_failed", zap.Int64("account_id", account.ID), zap.Error(err)) + } + return wrapReleaseOnDone(ctx, fastReleaseFunc), openAISlotAcquireOK } canWait, waitErr := h.concurrencyHelper.IncrementAccountWaitCount(ctx, account.ID, selection.WaitPlan.MaxWaiting) @@ -1400,7 +1457,7 @@ func (h *OpenAIGatewayHandler) acquireResponsesAccountSlot( zap.Int("max_waiting", selection.WaitPlan.MaxWaiting), ) h.handleStreamingAwareError(c, http.StatusTooManyRequests, "rate_limit_error", "Too many pending requests, please retry later", *streamStarted) - return nil, false + return nil, openAISlotAcquireFailed } accountWaitCounted := waitErr == nil && canWait @@ -1423,15 +1480,27 @@ func (h *OpenAIGatewayHandler) acquireResponsesAccountSlot( if err != nil { reqLog.Warn("openai.account_slot_acquire_failed", zap.Int64("account_id", account.ID), zap.Error(err)) h.handleConcurrencyError(c, err, "account", *streamStarted) - return nil, false + return nil, openAISlotAcquireFailed } // Slot acquired: no longer waiting in queue. releaseWait() - if err := h.gatewayService.BindStickySession(ctx, groupID, sessionHash, account.ID); err != nil { - reqLog.Warn("openai.bind_sticky_session_failed", zap.Int64("account_id", account.ID), zap.Error(err)) + // 分组利润控制:WaitPlan 排队成功后终检。排队期间账号倍率 + // 可能上调,越线则释放槽位交由调用方排除重选,不绑定粘连。 + latest, vetoed, reason := h.gatewayService.ProfitControlVetoLatest(ctx, account) + if vetoed { + if accountReleaseFunc != nil { + accountReleaseFunc() + } + reqLog.Debug("openai.account_slot_profit_vetoed", zap.Int64("account_id", account.ID), zap.String("reason", reason)) + return nil, openAISlotAcquireProfitVetoed } - return wrapReleaseOnDone(ctx, accountReleaseFunc), true + account = latest + selection.Account = latest + if err := h.gatewayService.BindStickySessionAfterProfitAdmission(ctx, groupID, sessionHash, account.ID); err != nil { + reqLog.Warn("openai.bind_sticky_session_after_profit_admission_failed", zap.Int64("account_id", account.ID), zap.Error(err)) + } + return wrapReleaseOnDone(ctx, accountReleaseFunc), openAISlotAcquireOK } // ResponsesWebSocket handles OpenAI Responses API WebSocket ingress endpoint @@ -1717,6 +1786,12 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) { requiredCapability = service.OpenAIEndpointCapabilityResponses } + // 分组利润控制:WS 桥按连接装配定价上下文并装门(选号与抢槽 + // 共用该 ctx);连接内不重选号,每个 turn 的计费共用同一 pricingAt——长 + // 连接跨峰谷边界属已知残余风险,由安全缓冲承担。显式生图意图跳门。 + wsPricingCtx, wsPricingAt := h.gatewayService.WithOpenAIRequestPricingContext(ctx, apiKey.GroupID, imageIntent) + ctx = wsPricingCtx + for { if ctx.Err() != nil { return @@ -1782,11 +1857,24 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) { closeOpenAIClientWS(wsConn, coderws.StatusTryAgainLater, "account is busy, please retry later") return } + // 分组利润控制:WS 快速抢槽成功后终检,越线则释放 + // 槽位、排除该账号重新选号,全池耗尽由下一轮选号关闭连接。 + latest, vetoed, reason := h.gatewayService.ProfitControlVetoLatest(ctx, account) + if vetoed { + if fastReleaseFunc != nil { + fastReleaseFunc() + } + reqLog.Debug("openai.websocket_account_slot_profit_vetoed", zap.Int64("account_id", account.ID), zap.String("reason", reason)) + failedAccountIDs[account.ID] = struct{}{} + continue + } + account = latest + selection.Account = latest accountReleaseFunc = fastReleaseFunc } currentAccountRelease = wrapReleaseOnDone(ctx, accountReleaseFunc) - if err := h.gatewayService.BindStickySession(ctx, apiKey.GroupID, sessionHash, account.ID); err != nil { - reqLog.Warn("openai.websocket_bind_sticky_session_failed", zap.Int64("account_id", account.ID), zap.Error(err)) + if err := h.gatewayService.BindStickySessionAfterProfitAdmission(ctx, apiKey.GroupID, sessionHash, account.ID); err != nil { + reqLog.Warn("openai.websocket_bind_sticky_session_after_profit_admission_failed", zap.Int64("account_id", account.ID), zap.Error(err)) } token, _, err := h.gatewayService.GetRequestCredential(ctx, c, account) @@ -1984,6 +2072,7 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) { QuotaPlatform: quotaPlatform, SessionID: sessionID, ChannelUsageFields: turnUsageFields, + PricingAt: wsPricingAt, CyberBlocked: cyberBlocked, }); err != nil { reqLog.Error("openai.websocket_record_usage_failed", diff --git a/backend/internal/handler/openai_images.go b/backend/internal/handler/openai_images.go index 72e2f462ad..721c206981 100644 --- a/backend/internal/handler/openai_images.go +++ b/backend/internal/handler/openai_images.go @@ -221,8 +221,13 @@ func (h *OpenAIGatewayHandler) Images(c *gin.Context) { reqLog.Debug("openai.images.account_selected", zap.Int64("account_id", account.ID), zap.String("account_name", account.Name)) setOpsSelectedAccount(c, account.ID, account.Platform) - accountReleaseFunc, acquired := h.acquireResponsesAccountSlot(c, apiKey.GroupID, sessionHash, selection, parsed.Stream, &streamStarted, reqLog) - if !acquired { + accountReleaseFunc, slotResult := h.acquireResponsesAccountSlot(c, apiKey.GroupID, sessionHash, selection, parsed.Stream, &streamStarted, reqLog) + if slotResult == openAISlotAcquireProfitVetoed { + // Images 调度不装利润门,此分支实际不可达;防御性排除重选。 + failedAccountIDs[account.ID] = struct{}{} + continue + } + if slotResult != openAISlotAcquireOK { return } diff --git a/backend/internal/handler/openai_profit_slot_recheck_test.go b/backend/internal/handler/openai_profit_slot_recheck_test.go new file mode 100644 index 0000000000..224631cf0c --- /dev/null +++ b/backend/internal/handler/openai_profit_slot_recheck_test.go @@ -0,0 +1,147 @@ +//go:build unit + +package handler + +// 槽位终检与生图跳门回归(handler 半程): +// - 槽位获取成功后的利润终检:越线账号释放槽位并要求调用方排除重选, +// 不写响应、不绑定粘连; +// - openAIResponsesRequiredCapability 的生图意图映射钉死(scheduler 的 +// 跳门条件依赖 CapabilityResponses ⇔ 显式生图意图这一耦合)。 + +import ( + "context" + "net/http/httptest" + "sync/atomic" + "testing" + "time" + + "github.com/Wei-Shaw/sub2api/internal/pkg/ctxkey" + "github.com/Wei-Shaw/sub2api/internal/service" + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" + "go.uber.org/zap" +) + +type profitCountingConcurrencyCache struct { + fakeConcurrencyCache + accountReleases atomic.Int64 +} + +func (c *profitCountingConcurrencyCache) ReleaseAccountSlot(context.Context, int64, string) error { + c.accountReleases.Add(1) + return nil +} + +func profitSlotTestAccount(id int64, rate float64) *service.Account { + now := time.Now() + return &service.Account{ + ID: id, + Platform: service.PlatformOpenAI, + Type: service.AccountTypeAPIKey, + Status: service.StatusActive, + Schedulable: true, + Concurrency: 2, + RateMultiplier: &rate, + Extra: map[string]any{ + "upstream_billing_probe": map[string]any{ + "status": service.UpstreamBillingProbeStatusOK, + "received_at": now.Add(-time.Minute), + "fresh_until": now.Add(30 * time.Minute), + "data": map[string]any{ + "billing_scope": "token", + "resolved_rate_multiplier": rate, + "peak_rate_enabled": false, + }, + }, + }, + } +} + +func profitSlotTestContext(t *testing.T, gw *service.OpenAIGatewayService, groupID int64, suppress bool) context.Context { + t.Helper() + group := &service.Group{ + ID: groupID, + Platform: service.PlatformOpenAI, + Status: service.StatusActive, + Hydrated: true, + RateMultiplier: 1.0, + SubscriptionType: service.SubscriptionTypeStandard, + ProfitControlEnabled: true, + ProfitMinMargin: 0.5, + } + base := context.WithValue(context.Background(), ctxkey.Group, group) + ctx, pricingAt := gw.WithOpenAIRequestPricingContext(base, &groupID, suppress) + require.False(t, pricingAt.IsZero()) + return ctx +} + +func TestAcquireResponsesAccountSlotProfitRecheck(t *testing.T) { + gin.SetMode(gin.TestMode) + gw := &service.OpenAIGatewayService{} + groupID := int64(50) + + newHandler := func(cache *profitCountingConcurrencyCache) *OpenAIGatewayHandler { + return &OpenAIGatewayHandler{ + gatewayService: gw, + concurrencyHelper: NewConcurrencyHelper(service.NewConcurrencyService(cache), SSEPingFormatClaude, 0), + } + } + newSelection := func(account *service.Account) *service.AccountSelectionResult { + return &service.AccountSelectionResult{ + Account: account, + Acquired: false, + WaitPlan: &service.AccountWaitPlan{AccountID: account.ID, MaxConcurrency: 2, Timeout: time.Second, MaxWaiting: 2}, + } + } + + t.Run("veto releases slot and requests reschedule without writing response", func(t *testing.T) { + cache := &profitCountingConcurrencyCache{} + h := newHandler(cache) + w := httptest.NewRecorder() + c, _ := gin.CreateTestContext(w) + c.Request = httptest.NewRequest("POST", "/v1/responses", nil).WithContext(profitSlotTestContext(t, gw, groupID, false)) + streamStarted := false + + release, result := h.acquireResponsesAccountSlot(c, &groupID, "", newSelection(profitSlotTestAccount(1, 0.8)), false, &streamStarted, zap.NewNop()) + require.Equal(t, openAISlotAcquireProfitVetoed, result) + require.Nil(t, release) + require.Zero(t, w.Body.Len(), "利润终检否决不得写出任何响应") + require.Equal(t, int64(1), cache.accountReleases.Load(), "否决后必须立即释放已获取的槽位") + }) + + t.Run("qualifying account acquires normally", func(t *testing.T) { + cache := &profitCountingConcurrencyCache{} + h := newHandler(cache) + w := httptest.NewRecorder() + c, _ := gin.CreateTestContext(w) + c.Request = httptest.NewRequest("POST", "/v1/responses", nil).WithContext(profitSlotTestContext(t, gw, groupID, false)) + streamStarted := false + + release, result := h.acquireResponsesAccountSlot(c, &groupID, "", newSelection(profitSlotTestAccount(2, 0.3)), false, &streamStarted, zap.NewNop()) + require.Equal(t, openAISlotAcquireOK, result) + require.NotNil(t, release) + release() + }) + + t.Run("image intent suppression keeps official behavior", func(t *testing.T) { + cache := &profitCountingConcurrencyCache{} + h := newHandler(cache) + w := httptest.NewRecorder() + c, _ := gin.CreateTestContext(w) + c.Request = httptest.NewRequest("POST", "/v1/responses", nil).WithContext(profitSlotTestContext(t, gw, groupID, true)) + streamStarted := false + + release, result := h.acquireResponsesAccountSlot(c, &groupID, "", newSelection(profitSlotTestAccount(3, 0.8)), false, &streamStarted, zap.NewNop()) + require.Equal(t, openAISlotAcquireOK, result, "生图意图跳门:过贵账号照常获取(图片边界不装门)") + require.NotNil(t, release) + release() + }) +} + +// scheduler 跳门条件依赖"CapabilityResponses 仅在显式生图意图时被要求"这一 +// 映射;后续若扩展该 capability 的用途,本测试失败提示同步收窄跳门条件。 +func TestOpenAIResponsesRequiredCapabilityPinsImageIntentMapping(t *testing.T) { + require.Equal(t, service.OpenAIEndpointCapabilityResponses, openAIResponsesRequiredCapability(true, service.PlatformOpenAI)) + require.Equal(t, service.OpenAIEndpointCapabilityChatCompletions, openAIResponsesRequiredCapability(false, service.PlatformOpenAI)) + require.Equal(t, service.OpenAIEndpointCapabilityChatCompletions, openAIResponsesRequiredCapability(true, service.PlatformGrok)) +} diff --git a/backend/internal/repository/api_key_repo.go b/backend/internal/repository/api_key_repo.go index 95d6012bcb..c2773996f4 100644 --- a/backend/internal/repository/api_key_repo.go +++ b/backend/internal/repository/api_key_repo.go @@ -210,6 +210,12 @@ func (r *apiKeyRepository) GetByKeyForAuth(ctx context.Context, key string) (*se group.FieldPeakStart, group.FieldPeakEnd, group.FieldPeakRateMultiplier, + // 分组利润控制:认证快照是调度门 enable 判定的直接来源, + // 漏选会让门静默失效;新增快照分组字段时必须同步本投影, + // 集成测试对账兜底。 + group.FieldProfitControlEnabled, + group.FieldProfitMinMargin, + group.FieldProfitSafetyBuffer, ) }). Only(ctx) @@ -987,6 +993,9 @@ func groupEntityToService(g *dbent.Group) *service.Group { PeakStart: g.PeakStart, PeakEnd: g.PeakEnd, PeakRateMultiplier: g.PeakRateMultiplier, + ProfitControlEnabled: g.ProfitControlEnabled, + ProfitMinMargin: g.ProfitMinMargin, + ProfitSafetyBuffer: g.ProfitSafetyBuffer, CreatedAt: g.CreatedAt, UpdatedAt: g.UpdatedAt, } diff --git a/backend/internal/repository/api_key_repo_profit_projection_integration_test.go b/backend/internal/repository/api_key_repo_profit_projection_integration_test.go new file mode 100644 index 0000000000..283b103d9b --- /dev/null +++ b/backend/internal/repository/api_key_repo_profit_projection_integration_test.go @@ -0,0 +1,60 @@ +//go:build integration + +package repository + +// 投影漏列回归(repository 半程):真实 PostgreSQL 上认证专用 +// 查询 GetByKeyForAuth 的分组显式投影必须携带利润控制与计费字段。该查询是 +// 认证快照(进而是利润门 enable 判定)的唯一数据来源,漏选任何快照分组字段 +// 都会让对应功能在真实流量上静默失效。新增快照分组字段时必须同步扩展 +// GetByKeyForAuth 的 WithGroup Select 并在本测试补断言。 + +import ( + "context" + "fmt" + "testing" + "time" + + "github.com/Wei-Shaw/sub2api/internal/service" + "github.com/stretchr/testify/require" +) + +func TestGetByKeyForAuthCarriesProfitControlProjection(t *testing.T) { + ctx := context.Background() + suffix := time.Now().UnixNano() + group := mustCreateGroup(t, integrationEntClient, &service.Group{ + Name: fmt.Sprintf("profit-proj-group-%d", suffix), + Platform: service.PlatformOpenAI, + RateMultiplier: 0.06, + ProfitControlEnabled: true, + ProfitMinMargin: 0.2, + ProfitSafetyBuffer: 0.05, + }) + user := mustCreateUser(t, integrationEntClient, &service.User{ + Email: fmt.Sprintf("profit-proj-%d@example.com", suffix), Concurrency: 5, + }) + groupID := group.ID + keyValue := fmt.Sprintf("sk-profit-proj-%d", suffix) + apiKeyRepo := NewAPIKeyRepository(integrationEntClient, integrationDB) + key := &service.APIKey{UserID: user.ID, GroupID: &groupID, Key: keyValue, Name: "profit-proj", Status: service.StatusActive} + require.NoError(t, apiKeyRepo.Create(ctx, key)) + t.Cleanup(func() { + _, err := integrationDB.ExecContext(ctx, "DELETE FROM auth_cache_invalidation_outbox WHERE cache_key = encode(sha256(convert_to($1, 'UTF8')), 'hex')", keyValue) + require.NoError(t, err) + _, err = integrationDB.ExecContext(ctx, "DELETE FROM api_keys WHERE id = $1", key.ID) + require.NoError(t, err) + _, err = integrationDB.ExecContext(ctx, "DELETE FROM users WHERE id = $1", user.ID) + require.NoError(t, err) + _, err = integrationDB.ExecContext(ctx, "DELETE FROM groups WHERE id = $1", group.ID) + require.NoError(t, err) + }) + + got, err := apiKeyRepo.GetByKeyForAuth(ctx, keyValue) + require.NoError(t, err) + require.NotNil(t, got.Group, "认证查询必须带出分组") + + require.Equal(t, service.PlatformOpenAI, got.Group.Platform) + require.InDelta(t, 0.06, got.Group.RateMultiplier, 1e-9) + require.True(t, got.Group.ProfitControlEnabled, "profit_control_enabled 必须进入认证投影(投影漏列会让门静默失效)") + require.InDelta(t, 0.2, got.Group.ProfitMinMargin, 1e-9) + require.InDelta(t, 0.05, got.Group.ProfitSafetyBuffer, 1e-9) +} diff --git a/backend/internal/repository/auth_cache_invalidation_profit_integration_test.go b/backend/internal/repository/auth_cache_invalidation_profit_integration_test.go new file mode 100644 index 0000000000..bf40042a33 --- /dev/null +++ b/backend/internal/repository/auth_cache_invalidation_profit_integration_test.go @@ -0,0 +1,98 @@ +//go:build integration + +package repository + +// migration 193 回归:groups 触发器的 durable 失效监视清单必须覆盖利润控制 +// 配置及 D 依赖的分组计价字段(正常后台保存走 InvalidateAuthCacheByGroupID 即时失效,触发器兜底 +// 直改 SQL / 更新与失效之间崩溃等 out-of-band 场景)。 + +import ( + "context" + "crypto/sha256" + "encoding/hex" + "fmt" + "testing" + "time" + + "github.com/Wei-Shaw/sub2api/internal/service" + "github.com/stretchr/testify/require" +) + +func TestAuthCacheInvalidationTrigger_ProfitControlColumns(t *testing.T) { + ctx := context.Background() + suffix := time.Now().UnixNano() + group := mustCreateGroup(t, integrationEntClient, &service.Group{ + Name: fmt.Sprintf("profit-trigger-group-%d", suffix), Platform: service.PlatformOpenAI, RateMultiplier: 1, + }) + user := mustCreateUser(t, integrationEntClient, &service.User{ + Email: fmt.Sprintf("profit-trigger-%d@example.com", suffix), Concurrency: 5, + }) + groupID := group.ID + keyValue := fmt.Sprintf("sk-profit-trigger-%d", suffix) + apiKeyRepo := NewAPIKeyRepository(integrationEntClient, integrationDB) + key := &service.APIKey{UserID: user.ID, GroupID: &groupID, Key: keyValue, Name: "profit-trigger", Status: service.StatusActive} + require.NoError(t, apiKeyRepo.Create(ctx, key)) + + sum := sha256.Sum256([]byte(keyValue)) + cacheKey := hex.EncodeToString(sum[:]) + clear := func() { + _, err := integrationDB.ExecContext(ctx, "DELETE FROM auth_cache_invalidation_outbox WHERE cache_key = $1", cacheKey) + require.NoError(t, err) + } + count := func() int { + var value int + require.NoError(t, integrationDB.QueryRowContext(ctx, + "SELECT COUNT(*) FROM auth_cache_invalidation_outbox WHERE cache_key = $1", cacheKey).Scan(&value)) + return value + } + clear() + t.Cleanup(clear) + t.Cleanup(func() { + _, err := integrationDB.ExecContext(ctx, "DELETE FROM api_keys WHERE id = $1", key.ID) + require.NoError(t, err) + _, err = integrationDB.ExecContext(ctx, "DELETE FROM users WHERE id = $1", user.ID) + require.NoError(t, err) + _, err = integrationDB.ExecContext(ctx, "DELETE FROM groups WHERE id = $1", group.ID) + require.NoError(t, err) + }) + + _, err := integrationDB.ExecContext(ctx, "UPDATE groups SET name = name || '-cosmetic' WHERE id = $1", group.ID) + require.NoError(t, err) + require.Zero(t, count(), "cosmetic 更新不得入队(既有语义回归)") + + _, err = integrationDB.ExecContext(ctx, "UPDATE groups SET profit_control_enabled = NOT profit_control_enabled WHERE id = $1", group.ID) + require.NoError(t, err) + require.Equal(t, 1, count(), "profit_control_enabled 变更必须入队") + clear() + + _, err = integrationDB.ExecContext(ctx, "UPDATE groups SET profit_min_margin = 0.3 WHERE id = $1", group.ID) + require.NoError(t, err) + require.Equal(t, 1, count(), "profit_min_margin 变更必须入队") + clear() + + _, err = integrationDB.ExecContext(ctx, "UPDATE groups SET profit_safety_buffer = 0.02 WHERE id = $1", group.ID) + require.NoError(t, err) + require.Equal(t, 1, count(), "profit_safety_buffer 变更必须入队") + clear() + + _, err = integrationDB.ExecContext(ctx, "UPDATE groups SET profit_min_margin = profit_min_margin WHERE id = $1", group.ID) + require.NoError(t, err) + require.Zero(t, count(), "利润字段无实际变化的 UPDATE 不得入队") + + for name, update := range map[string]string{ + "platform": "platform = 'anthropic'", + "subscription_type": "subscription_type = 'subscription'", + "rate_multiplier": "rate_multiplier = 0.9", + "peak_rate_enabled": "peak_rate_enabled = true", + "peak_start": "peak_start = '08:00'", + "peak_end": "peak_end = '09:00'", + "peak_rate_multiplier": "peak_rate_multiplier = 1.2", + } { + t.Run(name, func(t *testing.T) { + clear() + _, err := integrationDB.ExecContext(ctx, "UPDATE groups SET "+update+" WHERE id = $1", group.ID) + require.NoError(t, err) + require.Equal(t, 1, count(), name+" 变更必须入队") + }) + } +} diff --git a/backend/internal/repository/fixtures_integration_test.go b/backend/internal/repository/fixtures_integration_test.go index 48c33364c3..ca927be265 100644 --- a/backend/internal/repository/fixtures_integration_test.go +++ b/backend/internal/repository/fixtures_integration_test.go @@ -89,7 +89,10 @@ func mustCreateGroup(t *testing.T, client *dbent.Client, g *service.Group) *serv SetStatus(g.Status). SetSubscriptionType(g.SubscriptionType). SetRateMultiplier(g.RateMultiplier). - SetIsExclusive(g.IsExclusive) + SetIsExclusive(g.IsExclusive). + SetProfitControlEnabled(g.ProfitControlEnabled). + SetProfitMinMargin(g.ProfitMinMargin). + SetProfitSafetyBuffer(g.ProfitSafetyBuffer) if g.Description != "" { create.SetDescription(g.Description) } diff --git a/backend/internal/repository/gateway_cache.go b/backend/internal/repository/gateway_cache.go index d19c3cd512..6c345e1a78 100644 --- a/backend/internal/repository/gateway_cache.go +++ b/backend/internal/repository/gateway_cache.go @@ -4,6 +4,7 @@ import ( "context" "crypto/sha256" "encoding/hex" + "errors" "fmt" "strconv" "time" @@ -31,7 +32,14 @@ func buildSessionKey(groupID int64, sessionHash string) string { func (c *gatewayCache) GetSessionAccountID(ctx context.Context, groupID int64, sessionHash string) (int64, error) { key := buildSessionKey(groupID, sessionHash) - return c.rdb.Get(ctx, key).Int64() + accountID, err := c.rdb.Get(ctx, key).Int64() + if err != nil { + if errors.Is(err, redis.Nil) { + return 0, service.ErrStickySessionNotFound + } + return 0, err + } + return accountID, nil } func (c *gatewayCache) SetSessionAccountID(ctx context.Context, groupID int64, sessionHash string, accountID int64, ttl time.Duration) error { diff --git a/backend/internal/repository/gateway_cache_integration_test.go b/backend/internal/repository/gateway_cache_integration_test.go index 0eebc33f63..1cd4217758 100644 --- a/backend/internal/repository/gateway_cache_integration_test.go +++ b/backend/internal/repository/gateway_cache_integration_test.go @@ -8,7 +8,6 @@ import ( "time" "github.com/Wei-Shaw/sub2api/internal/service" - "github.com/redis/go-redis/v9" "github.com/stretchr/testify/require" "github.com/stretchr/testify/suite" ) @@ -25,7 +24,7 @@ func (s *GatewayCacheSuite) SetupTest() { func (s *GatewayCacheSuite) TestGetSessionAccountID_Missing() { _, err := s.cache.GetSessionAccountID(s.ctx, 1, "nonexistent") - require.True(s.T(), errors.Is(err, redis.Nil), "expected redis.Nil for missing session") + require.True(s.T(), errors.Is(err, service.ErrStickySessionNotFound), "expected ErrStickySessionNotFound for missing session") } func (s *GatewayCacheSuite) TestSetAndGetSessionAccountID() { @@ -88,7 +87,7 @@ func (s *GatewayCacheSuite) TestDeleteSessionAccountID() { require.NoError(s.T(), s.cache.DeleteSessionAccountID(s.ctx, groupID, sessionID), "DeleteSessionAccountID") _, err := s.cache.GetSessionAccountID(s.ctx, groupID, sessionID) - require.True(s.T(), errors.Is(err, redis.Nil), "expected redis.Nil after delete") + require.True(s.T(), errors.Is(err, service.ErrStickySessionNotFound), "expected ErrStickySessionNotFound after delete") } func (s *GatewayCacheSuite) TestGetSessionAccountID_CorruptedValue() { @@ -101,7 +100,7 @@ func (s *GatewayCacheSuite) TestGetSessionAccountID_CorruptedValue() { _, err := s.cache.GetSessionAccountID(s.ctx, groupID, sessionID) require.Error(s.T(), err, "expected error for corrupted value") - require.False(s.T(), errors.Is(err, redis.Nil), "expected parsing error, not redis.Nil") + require.False(s.T(), errors.Is(err, service.ErrStickySessionNotFound), "expected parsing error, not a miss") } func TestGatewayCacheSuite(t *testing.T) { diff --git a/backend/internal/repository/group_repo.go b/backend/internal/repository/group_repo.go index 50c9aafa5f..31724f7c55 100644 --- a/backend/internal/repository/group_repo.go +++ b/backend/internal/repository/group_repo.go @@ -102,7 +102,10 @@ func createGroupRecord(ctx context.Context, client *dbent.Client, groupIn *servi SetPeakRateEnabled(groupIn.PeakRateEnabled). SetPeakStart(groupIn.PeakStart). SetPeakEnd(groupIn.PeakEnd). - SetPeakRateMultiplier(groupIn.PeakRateMultiplier) + SetPeakRateMultiplier(groupIn.PeakRateMultiplier). + SetProfitControlEnabled(groupIn.ProfitControlEnabled). + SetProfitMinMargin(groupIn.ProfitMinMargin). + SetProfitSafetyBuffer(groupIn.ProfitSafetyBuffer) if groupIn.DuplicateOperationID != "" { builder = builder.SetDuplicateOperationID(groupIn.DuplicateOperationID) } @@ -268,7 +271,10 @@ func (r *groupRepository) Update(ctx context.Context, groupIn *service.Group) er SetPeakRateEnabled(groupIn.PeakRateEnabled). SetPeakStart(groupIn.PeakStart). SetPeakEnd(groupIn.PeakEnd). - SetPeakRateMultiplier(groupIn.PeakRateMultiplier) + SetPeakRateMultiplier(groupIn.PeakRateMultiplier). + SetProfitControlEnabled(groupIn.ProfitControlEnabled). + SetProfitMinMargin(groupIn.ProfitMinMargin). + SetProfitSafetyBuffer(groupIn.ProfitSafetyBuffer) // 显式处理可空字段:nil 需要 clear,非 nil 需要 set。 if groupIn.DailyLimitUSD != nil { diff --git a/backend/internal/server/api_contract_test.go b/backend/internal/server/api_contract_test.go index f6b1a06992..5b09d42442 100644 --- a/backend/internal/server/api_contract_test.go +++ b/backend/internal/server/api_contract_test.go @@ -354,6 +354,9 @@ func TestAPIContracts(t *testing.T) { "peak_start": "", "peak_end": "", "peak_rate_multiplier": 1, + "profit_control_enabled": false, + "profit_min_margin": 0, + "profit_safety_buffer": 0, "is_exclusive": false, "status": "active", "subscription_type": "standard", diff --git a/backend/internal/service/admin_group.go b/backend/internal/service/admin_group.go index 68822bd3d7..2d4a3f7fcb 100644 --- a/backend/internal/service/admin_group.go +++ b/backend/internal/service/admin_group.go @@ -375,6 +375,20 @@ func (s *adminServiceImpl) CreateGroup(ctx context.Context, input *CreateGroupIn return nil, err } + profitMinMargin := 0.0 + if input.ProfitMinMargin != nil { + profitMinMargin = *input.ProfitMinMargin + } + profitSafetyBuffer := 0.0 + if input.ProfitSafetyBuffer != nil { + profitSafetyBuffer = *input.ProfitSafetyBuffer + } + // 利润控制与高峰倍率同一收口顺序:先按平台归一化(不支持的平台重置),再校验。 + profitControlEnabled, profitMinMargin, profitSafetyBuffer := NormalizeProfitControlConfig(platform, input.ProfitControlEnabled, profitMinMargin, profitSafetyBuffer) + if err := ValidateProfitControlConfig(platform, profitControlEnabled, profitMinMargin, profitSafetyBuffer); err != nil { + return nil, err + } + // 校验降级分组 if input.FallbackGroupID != nil { if err := s.validateFallbackGroup(ctx, 0, *input.FallbackGroupID); err != nil { @@ -456,6 +470,9 @@ func (s *adminServiceImpl) CreateGroup(ctx context.Context, input *CreateGroupIn PeakStart: peakStart, PeakEnd: peakEnd, PeakRateMultiplier: peakRateMultiplier, + ProfitControlEnabled: profitControlEnabled, + ProfitMinMargin: profitMinMargin, + ProfitSafetyBuffer: profitSafetyBuffer, ImagePrice1K: imagePrice1K, ImagePrice2K: imagePrice2K, ImagePrice4K: imagePrice4K, @@ -708,6 +725,21 @@ func (s *adminServiceImpl) UpdateGroup(ctx context.Context, id int64, input *Upd if err := ValidatePeakRateConfig(group.SubscriptionType, group.PeakRateEnabled, group.PeakStart, group.PeakEnd, group.PeakRateMultiplier); err != nil { return nil, err } + if input.ProfitControlEnabled != nil { + group.ProfitControlEnabled = *input.ProfitControlEnabled + } + if input.ProfitMinMargin != nil { + group.ProfitMinMargin = *input.ProfitMinMargin + } + if input.ProfitSafetyBuffer != nil { + group.ProfitSafetyBuffer = *input.ProfitSafetyBuffer + } + // 利润控制与高峰同一收口:按合并后的最终平台归一化(转到不支持平台时静默重置), + // 再对合并后的最终配置统一校验,防止部分字段更新拼出非法组合入库。 + group.ProfitControlEnabled, group.ProfitMinMargin, group.ProfitSafetyBuffer = NormalizeProfitControlConfig(group.Platform, group.ProfitControlEnabled, group.ProfitMinMargin, group.ProfitSafetyBuffer) + if err := ValidateProfitControlConfig(group.Platform, group.ProfitControlEnabled, group.ProfitMinMargin, group.ProfitSafetyBuffer); err != nil { + return nil, err + } if input.ImagePrice1K != nil { group.ImagePrice1K = normalizePrice(input.ImagePrice1K) } diff --git a/backend/internal/service/admin_group_duplicate.go b/backend/internal/service/admin_group_duplicate.go index fcfe39326d..841af5e5eb 100644 --- a/backend/internal/service/admin_group_duplicate.go +++ b/backend/internal/service/admin_group_duplicate.go @@ -88,6 +88,9 @@ func cloneGroupForDuplicate(source *Group, operationID string) *Group { PeakStart: source.PeakStart, PeakEnd: source.PeakEnd, PeakRateMultiplier: source.PeakRateMultiplier, + ProfitControlEnabled: source.ProfitControlEnabled, + ProfitMinMargin: source.ProfitMinMargin, + ProfitSafetyBuffer: source.ProfitSafetyBuffer, IsExclusive: source.IsExclusive, Status: duplicateGroupInactiveStatus, DuplicateOperationID: operationID, diff --git a/backend/internal/service/admin_service.go b/backend/internal/service/admin_service.go index 7e69ef258d..a620f86f14 100644 --- a/backend/internal/service/admin_service.go +++ b/backend/internal/service/admin_service.go @@ -264,6 +264,10 @@ type CreateGroupInput struct { MaxReasoningEffort string // ReasoningEffortMappings OpenAI/Codex 推理强度精确映射。 ReasoningEffortMappings []ReasoningEffortMapping + // 分组利润控制(五个 token 平台分组可启用;margin/buffer 为小数,nil 按 0 处理) + ProfitControlEnabled bool + ProfitMinMargin *float64 + ProfitSafetyBuffer *float64 // 从指定分组复制账号(创建分组后在同一事务内绑定) CopyAccountsFromGroupIDs []int64 } @@ -325,6 +329,10 @@ type UpdateGroupInput struct { MaxReasoningEffort *string // ReasoningEffortMappings nil 表示不修改,空数组表示清空,非空数组表示替换。 ReasoningEffortMappings *[]ReasoningEffortMapping + // 分组利润控制(nil 表示不修改;margin/buffer 为小数) + ProfitControlEnabled *bool + ProfitMinMargin *float64 + ProfitSafetyBuffer *float64 // 从指定分组复制账号(同步操作:先清空当前分组的账号绑定,再绑定源分组的账号) CopyAccountsFromGroupIDs []int64 } diff --git a/backend/internal/service/api_key_auth_cache.go b/backend/internal/service/api_key_auth_cache.go index 924fc20fcd..a6f76beb80 100644 --- a/backend/internal/service/api_key_auth_cache.go +++ b/backend/internal/service/api_key_auth_cache.go @@ -114,6 +114,13 @@ type APIKeyAuthGroupSnapshot struct { PeakStart string `json:"peak_start"` PeakEnd string `json:"peak_end"` PeakRateMultiplier float64 `json:"peak_rate_multiplier"` + + // 分组利润控制:调度准入门按 schedulerSnapshot 实时读取分组配置, + // 不依赖本快照;此处随快照缓存只为保证 apiKey.Group 字段完整, + // 任何消费方都不会读到误导性的零值。 + ProfitControlEnabled bool `json:"profit_control_enabled"` + ProfitMinMargin float64 `json:"profit_min_margin"` + ProfitSafetyBuffer float64 `json:"profit_safety_buffer"` } // APIKeyAuthCacheEntry 缓存条目,支持负缓存 diff --git a/backend/internal/service/api_key_auth_cache_impl.go b/backend/internal/service/api_key_auth_cache_impl.go index fd5fda22c7..e834ce62de 100644 --- a/backend/internal/service/api_key_auth_cache_impl.go +++ b/backend/internal/service/api_key_auth_cache_impl.go @@ -14,7 +14,7 @@ import ( "github.com/dgraph-io/ristretto" ) -const apiKeyAuthSnapshotVersion = 17 // v17: include the OpenAI group Live gate +const apiKeyAuthSnapshotVersion = 18 // v18: include group profit control fields (force refresh of pre-fix snapshots) type apiKeyAuthCacheConfig struct { l1Size int @@ -420,6 +420,9 @@ func (s *APIKeyService) snapshotFromAPIKey(ctx context.Context, apiKey *APIKey) PeakStart: apiKey.Group.PeakStart, PeakEnd: apiKey.Group.PeakEnd, PeakRateMultiplier: apiKey.Group.PeakRateMultiplier, + ProfitControlEnabled: apiKey.Group.ProfitControlEnabled, + ProfitMinMargin: apiKey.Group.ProfitMinMargin, + ProfitSafetyBuffer: apiKey.Group.ProfitSafetyBuffer, } } return snapshot @@ -507,6 +510,9 @@ func (s *APIKeyService) snapshotToAPIKey(key string, snapshot *APIKeyAuthSnapsho PeakStart: snapshot.Group.PeakStart, PeakEnd: snapshot.Group.PeakEnd, PeakRateMultiplier: snapshot.Group.PeakRateMultiplier, + ProfitControlEnabled: snapshot.Group.ProfitControlEnabled, + ProfitMinMargin: snapshot.Group.ProfitMinMargin, + ProfitSafetyBuffer: snapshot.Group.ProfitSafetyBuffer, } } s.compileAPIKeyIPRules(apiKey) diff --git a/backend/internal/service/api_key_auth_cache_profit_test.go b/backend/internal/service/api_key_auth_cache_profit_test.go new file mode 100644 index 0000000000..65bef3176f --- /dev/null +++ b/backend/internal/service/api_key_auth_cache_profit_test.go @@ -0,0 +1,93 @@ +package service + +// 投影漏列回归(service 半程):认证快照 build → L2 JSON 序列化 +// → 反序列化 → 还原 apiKey.Group → 请求 ctx → 利润门解析,全链路保真。 +// repository 半程(真实 GetByKeyForAuth 投影)见 +// internal/repository/api_key_repo_profit_projection_integration_test.go。 + +import ( + "context" + "encoding/json" + "testing" + + "github.com/Wei-Shaw/sub2api/internal/pkg/ctxkey" + "github.com/stretchr/testify/require" +) + +func profitAuthTestAPIKey() *APIKey { + groupID := int64(50) + return &APIKey{ + ID: 82, + UserID: 40, + GroupID: &groupID, + Name: "profit-auth-roundtrip", + Status: StatusActive, + User: &User{ + ID: 40, + Email: "profit@test.local", + Status: StatusActive, + Concurrency: 5, + }, + Group: &Group{ + ID: groupID, + Name: "VIP-roundtrip", + Platform: PlatformOpenAI, + Status: StatusActive, + Hydrated: true, + RateMultiplier: 0.06, + SubscriptionType: SubscriptionTypeStandard, + PeakRateEnabled: false, + ProfitControlEnabled: true, + ProfitMinMargin: 0.2, + ProfitSafetyBuffer: 0.05, + }, + } +} + +// 快照构建 → L2 JSON 往返 → 还原 → 装门:利润字段必须全程保真,阈值与 +// 计费同源(0.06 × (1−0.25) = 0.045)。 +func TestAPIKeyAuthSnapshotProfitControlRoundtrip(t *testing.T) { + svc := &APIKeyService{} + apiKey := profitAuthTestAPIKey() + + snapshot := svc.snapshotFromAPIKey(context.Background(), apiKey) + require.NotNil(t, snapshot) + require.Equal(t, apiKeyAuthSnapshotVersion, snapshot.Version) + require.Equal(t, 18, snapshot.Version, "v18 起认证快照携带利润控制字段") + + // 模拟 L2 缓存的完整 JSON 往返(与 apiKeyCache.SetAuthCache/GetAuthCache 同构)。 + payload, err := json.Marshal(&APIKeyAuthCacheEntry{Snapshot: snapshot}) + require.NoError(t, err) + var restored APIKeyAuthCacheEntry + require.NoError(t, json.Unmarshal(payload, &restored)) + + materialized, used, err := svc.applyAuthCacheEntry(apiKey.Key, &restored) + require.NoError(t, err) + require.True(t, used) + require.NotNil(t, materialized.Group) + require.True(t, materialized.Group.Hydrated) + require.True(t, materialized.Group.ProfitControlEnabled) + require.InDelta(t, 0.2, materialized.Group.ProfitMinMargin, 1e-12) + require.InDelta(t, 0.05, materialized.Group.ProfitSafetyBuffer, 1e-12) + require.InDelta(t, 0.06, materialized.Group.RateMultiplier, 1e-12) + + // 中间件语义:materialized.Group 进请求 ctx → 门必须按快照配置装上。 + ctx := context.WithValue(context.Background(), ctxkey.Group, materialized.Group) + gwSvc := &OpenAIGatewayService{} + gate := gwSvc.resolveOpenAIProfitControlGate(ctx, materialized.GroupID) + require.NotNil(t, gate, "还原后的认证分组必须能装门(投影漏列时本断言最先失败)") + require.InDelta(t, 0.06*(1-0.25), gate.threshold, 1e-12) +} + +// 旧版本快照(v16 及更早,无利润字段保真保证)必须被淘汰回源,不得复用。 +func TestAPIKeyAuthSnapshotOldVersionEvicted(t *testing.T) { + svc := &APIKeyService{} + snapshot := svc.snapshotFromAPIKey(context.Background(), profitAuthTestAPIKey()) + require.NotNil(t, snapshot) + snapshot.Version = 16 + + materialized, used, err := svc.applyAuthCacheEntry("sk-old", &APIKeyAuthCacheEntry{Snapshot: snapshot}) + require.NoError(t, err) + require.False(t, used, "版本不匹配的缓存条目必须淘汰并回源重建") + require.Nil(t, materialized) +} diff --git a/backend/internal/service/gateway_profit_control.go b/backend/internal/service/gateway_profit_control.go new file mode 100644 index 0000000000..20e0e4d91f --- /dev/null +++ b/backend/internal/service/gateway_profit_control.go @@ -0,0 +1,113 @@ +package service + +import ( + "context" + "log/slog" + "math" + + "github.com/Wei-Shaw/sub2api/internal/pkg/ctxkey" +) + +// withGatewayProfitControlGate installs the gate only for explicitly marked +// token requests. This keeps media, metadata, and models-list paths outside +// the profit-control surface by construction. +func (s *GatewayService) withGatewayProfitControlGate(ctx context.Context, groupID *int64) context.Context { + if _, ok := gatewayTokenRequestPricingAtFromContext(ctx); !ok || groupID == nil || *groupID <= 0 { + return ctx + } + if existing, ok := ctx.Value(openAIProfitControlGateCtxKey{}).(*openAIProfitControlGate); ok && existing != nil && existing.groupID == *groupID { + return ctx + } + + group, err := s.resolveProfitControlGroup(ctx, *groupID) + if err != nil { + slog.Warn("profit_control_group_load_failed", "group_id", *groupID, "error", err) + return s.clearForeignProfitControlGate(ctx, groupID) + } + if group == nil || !group.ProfitControlEnabled || !profitControlPlatformSupported(group.Platform) { + return s.clearForeignProfitControlGate(ctx, groupID) + } + + pricingAt, _ := gatewayTokenRequestPricingAtFromContext(ctx) + billingGroup := gatewayTokenRequestBillingGroupFromContext(ctx) + if billingGroup == nil { + if ctxGroup, ok := ctx.Value(ctxkey.Group).(*Group); ok && IsGroupContextValid(ctxGroup) { + billingGroup = ctxGroup + } else { + billingGroup = group + } + } + + downstream := billingGroup.RateMultiplier + if userID, _ := ctx.Value(ctxkey.UserID).(int64); userID > 0 { + downstream = s.ResolveUserGroupRateMultiplier(ctx, userID, billingGroup.ID, billingGroup.RateMultiplier) + } + downstream *= billingGroup.PeakMultiplierAt(pricingAt) + threshold := downstream * (1 - group.ProfitMinMargin - group.ProfitSafetyBuffer) + if math.IsNaN(threshold) || math.IsInf(threshold, 0) || threshold < 0 { + threshold = 0 + } + + gate := &openAIProfitControlGate{ + groupID: group.ID, + platform: group.Platform, + threshold: threshold, + pricingAt: pricingAt, + } + openAIProfitControlObserverInstance.recordInstall(gate.groupID, gate.platform, gate.threshold) + return context.WithValue(ctx, openAIProfitControlGateCtxKey{}, gate) +} + +func (s *GatewayService) clearForeignProfitControlGate(ctx context.Context, groupID *int64) context.Context { + existing, ok := ctx.Value(openAIProfitControlGateCtxKey{}).(*openAIProfitControlGate) + if !ok || existing == nil || groupID == nil || existing.groupID == *groupID { + return ctx + } + return context.WithValue(ctx, openAIProfitControlGateCtxKey{}, (*openAIProfitControlGate)(nil)) +} + +func (s *GatewayService) resolveProfitControlGroup(ctx context.Context, groupID int64) (*Group, error) { + if group, ok := ctx.Value(ctxkey.Group).(*Group); ok && IsGroupContextValid(group) && group.ID == groupID { + return group, nil + } + if s.schedulerSnapshot != nil { + return s.schedulerSnapshot.GetGroupByID(ctx, groupID) + } + return s.resolveGroupByID(ctx, groupID) +} + +// GatewayProfitControlVetoLatest performs the terminal post-slot check against +// the latest scheduler snapshot. Snapshot read failures are deliberately +// fail-open to preserve availability, but are observable. +func (s *GatewayService) GatewayProfitControlVetoLatest(ctx context.Context, selected *Account) (*Account, bool, string) { + return profitControlVetoLatest(ctx, selected, s.schedulerSnapshot) +} + +func profitControlVetoLatest(ctx context.Context, selected *Account, snapshot *SchedulerSnapshotService) (*Account, bool, string) { + gate, _ := ctx.Value(openAIProfitControlGateCtxKey{}).(*openAIProfitControlGate) + if gate == nil || selected == nil { + return selected, false, "" + } + latest := selected + if snapshot != nil { + refreshed, err := snapshot.GetAccount(ctx, selected.ID) + if err != nil || refreshed == nil { + slog.Warn("profit_control_account_refresh_failed", "group_id", gate.groupID, "platform", gate.platform, "account_id", selected.ID, "error", err) + openAIProfitControlObserverInstance.recordRefreshFailure(gate.groupID, gate.platform, gate.threshold) + } else { + latest = refreshed + } + } + vetoed, reason := openAIProfitControlVetoReason(ctx, latest) + return latest, vetoed, reason +} + +func (s *GatewayService) isGatewayAccountProfitEligible(ctx context.Context, account *Account) bool { + vetoed, _ := openAIProfitControlVetoReason(ctx, account) + return !vetoed +} + +func gatewayProfitControlGateActive(ctx context.Context) bool { + gate, _ := ctx.Value(openAIProfitControlGateCtxKey{}).(*openAIProfitControlGate) + return gate != nil +} diff --git a/backend/internal/service/gateway_profit_control_v2_test.go b/backend/internal/service/gateway_profit_control_v2_test.go new file mode 100644 index 0000000000..075eeb17b7 --- /dev/null +++ b/backend/internal/service/gateway_profit_control_v2_test.go @@ -0,0 +1,385 @@ +//go:build unit + +package service + +import ( + "context" + "errors" + "testing" + + "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/Wei-Shaw/sub2api/internal/pkg/ctxkey" + "github.com/stretchr/testify/require" +) + +func gatewayProfitTestGroup(id int64, platform string) *Group { + return &Group{ + ID: id, + Name: "profit-" + platform, + Platform: platform, + Status: StatusActive, + Hydrated: true, + RateMultiplier: 0.5, + SubscriptionType: SubscriptionTypeStandard, + ProfitControlEnabled: true, + ProfitMinMargin: 0, + ProfitSafetyBuffer: 0, + } +} + +func gatewayProfitTestContext(group *Group) context.Context { + ctx := context.WithValue(context.Background(), ctxkey.Group, group) + ctx, _ = WithGatewayTokenRequestPricing(ctx) + return ctx +} + +func gatewayProfitTestAccount(id int64, platform string, rate float64, groupID int64) Account { + return Account{ + ID: id, + Name: "account", + Platform: platform, + Type: AccountTypeAPIKey, + Status: StatusActive, + Schedulable: true, + Concurrency: 2, + Priority: 1, + RateMultiplier: &rate, + AccountGroups: []AccountGroup{{AccountID: id, GroupID: groupID}}, + GroupIDs: []int64{groupID}, + } +} + +func TestGatewayProfitControlInstallsForFivePlatformsOnlyOnTokenRequests(t *testing.T) { + for _, platform := range []string{ + PlatformOpenAI, + PlatformAnthropic, + PlatformGemini, + PlatformGrok, + PlatformAntigravity, + } { + t.Run(platform, func(t *testing.T) { + group := gatewayProfitTestGroup(101, platform) + groupID := group.ID + svc := &GatewayService{} + + tokenCtx := svc.withGatewayProfitControlGate(gatewayProfitTestContext(group), &groupID) + gate, _ := tokenCtx.Value(openAIProfitControlGateCtxKey{}).(*openAIProfitControlGate) + require.NotNil(t, gate) + require.Equal(t, platform, gate.platform) + require.InDelta(t, 0.5, gate.threshold, 1e-12) + + metadataCtx := context.WithValue(context.Background(), ctxkey.Group, group) + metadataCtx = svc.withGatewayProfitControlGate(metadataCtx, &groupID) + gate, _ = metadataCtx.Value(openAIProfitControlGateCtxKey{}).(*openAIProfitControlGate) + require.Nil(t, gate, "未显式标记为 token 请求的入口不得装门") + }) + } +} + +func TestGatewayProfitControlCompositeBillingUsesScheduledMemberConfig(t *testing.T) { + billingGroup := &Group{ + ID: 201, + Platform: PlatformComposite, + Status: StatusActive, + Hydrated: true, + RateMultiplier: 0.4, + SubscriptionType: SubscriptionTypeStandard, + } + memberGroup := gatewayProfitTestGroup(202, PlatformAnthropic) + memberGroup.RateMultiplier = 99 + memberGroup.ProfitMinMargin = 0.25 + + ctx := context.WithValue(context.Background(), ctxkey.Group, billingGroup) + ctx, pricingAt := WithGatewayTokenRequestPricing(ctx) + svc := &GatewayService{ + schedulerSnapshot: NewSchedulerSnapshotService( + nil, + nil, + nil, + profitControlGroupRepo{group: memberGroup}, + nil, + ), + } + ctx = svc.withGatewayProfitControlGate(ctx, &memberGroup.ID) + gate, _ := ctx.Value(openAIProfitControlGateCtxKey{}).(*openAIProfitControlGate) + require.NotNil(t, gate) + require.Equal(t, memberGroup.ID, gate.groupID) + require.Equal(t, PlatformAnthropic, gate.platform) + require.Equal(t, pricingAt, gate.pricingAt) + require.InDelta(t, 0.4*(1-0.25), gate.threshold, 1e-12, "D 必须取 composite 计费父分组,margin 取被调度成员分组") +} + +func TestGatewayProfitControlGroupLoadFailureClearsForeignGate(t *testing.T) { + billingGroup := &Group{ + ID: 211, + Platform: PlatformComposite, + Status: StatusActive, + Hydrated: true, + RateMultiplier: 0.4, + SubscriptionType: SubscriptionTypeStandard, + } + targetGroupID := int64(212) + ctx := gatewayProfitTestContext(billingGroup) + ctx = context.WithValue(ctx, openAIProfitControlGateCtxKey{}, &openAIProfitControlGate{ + groupID: 210, + platform: PlatformAnthropic, + threshold: 0.1, + }) + svc := &GatewayService{ + schedulerSnapshot: NewSchedulerSnapshotService( + nil, + nil, + nil, + profitControlFailingGroupRepo{}, + nil, + ), + } + + ctx = svc.withGatewayProfitControlGate(ctx, &targetGroupID) + gate, ok := ctx.Value(openAIProfitControlGateCtxKey{}).(*openAIProfitControlGate) + require.True(t, ok) + require.Nil(t, gate, "加载新分组失败时必须清除其他分组遗留的门") + + account := gatewayProfitTestAccount(213, PlatformAnthropic, 0.8, targetGroupID) + require.True(t, svc.isGatewayAccountProfitEligible(ctx, &account), "配置读取失败按既定语义 fail-open") +} + +type profitControlFailingGroupRepo struct { + GroupRepository +} + +func (profitControlFailingGroupRepo) GetByID(context.Context, int64) (*Group, error) { + return nil, errors.New("group cache unavailable") +} + +func TestGatewayProfitControlLegacyMixedAndRoutedSelection(t *testing.T) { + t.Run("legacy single-platform selection", func(t *testing.T) { + group := gatewayProfitTestGroup(111, PlatformGrok) + cheap := gatewayProfitTestAccount(1, PlatformGrok, 0.2, group.ID) + expensive := gatewayProfitTestAccount(2, PlatformGrok, 0.8, group.ID) + repo := &mockAccountRepoForPlatform{ + accounts: []Account{expensive, cheap}, + accountsByID: map[int64]*Account{cheap.ID: &cheap, expensive.ID: &expensive}, + } + svc := &GatewayService{ + accountRepo: repo, + cache: &mockGatewayCacheForPlatform{}, + cfg: testConfig(), + } + + selected, err := svc.SelectAccountForModelWithExclusions( + gatewayProfitTestContext(group), &group.ID, "", "", nil, + ) + require.NoError(t, err) + require.Equal(t, cheap.ID, selected.ID) + + _, err = svc.SelectAccountForModelWithExclusions( + gatewayProfitTestContext(group), &group.ID, "", "", map[int64]struct{}{cheap.ID: {}}, + ) + require.Error(t, err) + require.ErrorIs(t, err, ErrNoAvailableAccounts) + }) + + t.Run("mixed routing filters the routed account", func(t *testing.T) { + group := gatewayProfitTestGroup(112, PlatformAnthropic) + group.ModelRoutingEnabled = true + group.ModelRouting = map[string][]int64{"claude-test": {2, 1}} + cheap := gatewayProfitTestAccount(1, PlatformAntigravity, 0.2, group.ID) + cheap.Extra = map[string]any{"mixed_scheduling": true} + cheap.Credentials = map[string]any{"model_mapping": map[string]any{"claude-test": "claude-test"}} + expensive := gatewayProfitTestAccount(2, PlatformAnthropic, 0.8, group.ID) + repo := &mockAccountRepoForPlatform{ + accounts: []Account{expensive, cheap}, + accountsByID: map[int64]*Account{cheap.ID: &cheap, expensive.ID: &expensive}, + } + svc := &GatewayService{ + accountRepo: repo, + cache: &mockGatewayCacheForPlatform{}, + cfg: testConfig(), + } + + selected, err := svc.SelectAccountForModelWithExclusions( + gatewayProfitTestContext(group), &group.ID, "", "claude-test", nil, + ) + require.NoError(t, err) + require.Equal(t, cheap.ID, selected.ID) + }) +} + +func TestGatewayProfitControlLoadAwareSelectionAndFailover(t *testing.T) { + group := gatewayProfitTestGroup(121, PlatformGrok) + cheap := gatewayProfitTestAccount(1, PlatformGrok, 0.2, group.ID) + expensive := gatewayProfitTestAccount(2, PlatformGrok, 0.8, group.ID) + repo := &mockAccountRepoForPlatform{ + accounts: []Account{expensive, cheap}, + accountsByID: map[int64]*Account{cheap.ID: &cheap, expensive.ID: &expensive}, + } + cfg := &config.Config{RunMode: config.RunModeStandard} + cfg.Gateway.Scheduling.LoadBatchEnabled = true + svc := &GatewayService{ + accountRepo: repo, + cache: &mockGatewayCacheForPlatform{}, + cfg: cfg, + concurrencyService: NewConcurrencyService(stubConcurrencyCache{}), + } + + result, err := svc.SelectAccountWithLoadAwareness( + gatewayProfitTestContext(group), &group.ID, "", "", nil, "", 0, + ) + require.NoError(t, err) + require.NotNil(t, result) + require.Equal(t, cheap.ID, result.Account.ID) + if result.ReleaseFunc != nil { + result.ReleaseFunc() + } + + result, err = svc.SelectAccountWithLoadAwareness( + gatewayProfitTestContext(group), + &group.ID, + "", + "", + map[int64]struct{}{cheap.ID: {}}, + "", + 0, + ) + require.Nil(t, result) + require.Error(t, err) + require.ErrorIs(t, err, ErrNoAvailableAccounts) +} + +func TestGatewayProfitControlStickyVetoKeepsBindingUntilRateRecovers(t *testing.T) { + group := gatewayProfitTestGroup(131, PlatformAnthropic) + expensive := gatewayProfitTestAccount(1, PlatformAnthropic, 0.8, group.ID) + cheap := gatewayProfitTestAccount(2, PlatformAnthropic, 0.2, group.ID) + repo := &mockAccountRepoForPlatform{ + accounts: []Account{expensive, cheap}, + accountsByID: map[int64]*Account{expensive.ID: &expensive, cheap.ID: &cheap}, + } + cache := &mockGatewayCacheForPlatform{ + sessionBindings: map[string]int64{"sticky-profit": expensive.ID}, + } + svc := &GatewayService{ + accountRepo: repo, + cache: cache, + cfg: testConfig(), + } + ctx := gatewayProfitTestContext(group) + + selected, err := svc.SelectAccountForModelWithExclusions(ctx, &group.ID, "sticky-profit", "", nil) + require.NoError(t, err) + require.Equal(t, cheap.ID, selected.ID) + require.Equal(t, expensive.ID, cache.sessionBindings["sticky-profit"], "候选过滤不得覆盖旧粘性绑定") + + require.NoError(t, svc.BindStickySessionAfterProfitAdmission( + svc.withGatewayProfitControlGate(ctx, &group.ID), + &group.ID, + "sticky-profit", + cheap.ID, + expensive.ID, + )) + require.Equal(t, expensive.ID, cache.sessionBindings["sticky-profit"], "终检通过的 fallback 账号也不得覆盖旧绑定") + require.Zero(t, cache.deletedSessions["sticky-profit"]) + + recovered := expensive + recoveredRate := 0.2 + recovered.RateMultiplier = &recoveredRate + repo.accounts[0] = recovered + repo.accountsByID[recovered.ID] = &repo.accounts[0] + + selected, err = svc.SelectAccountForModelWithExclusions(ctx, &group.ID, "sticky-profit", "", nil) + require.NoError(t, err) + require.Equal(t, recovered.ID, selected.ID, "倍率恢复后应重新命中原粘性账号") + require.Zero(t, cache.deletedSessions["sticky-profit"]) +} + +type gatewayProfitSnapshotCache struct { + SchedulerCache + account *Account + err error +} + +func (c *gatewayProfitSnapshotCache) GetAccount(context.Context, int64) (*Account, error) { + return c.account, c.err +} + +type gatewayProfitAccountRepo struct { + AccountRepository + account *Account + err error +} + +func (r gatewayProfitAccountRepo) GetByID(context.Context, int64) (*Account, error) { + return r.account, r.err +} + +func TestGatewayProfitControlTerminalRefreshUsesReplacementObject(t *testing.T) { + selected := gatewayProfitTestAccount(141, PlatformGemini, 0.2, 1) + replacement := selected + expensiveRate := 0.8 + replacement.RateMultiplier = &expensiveRate + + snapshot := NewSchedulerSnapshotService( + &gatewayProfitSnapshotCache{account: &replacement}, + nil, + gatewayProfitAccountRepo{}, + nil, + nil, + ) + ctx := context.WithValue(context.Background(), openAIProfitControlGateCtxKey{}, &openAIProfitControlGate{ + groupID: 1, + platform: PlatformGemini, + threshold: 0.5, + }) + + latest, vetoed, reason := profitControlVetoLatest(ctx, &selected, snapshot) + require.Same(t, &replacement, latest) + require.True(t, vetoed) + require.Equal(t, openAIProfitFilterReasonThreshold, reason) + require.InDelta(t, 0.2, *selected.RateMultiplier, 1e-12, "测试必须替换缓存对象,不能原地修改旧指针") +} + +func TestGatewayProfitControlTerminalRefreshFallsBackFromCacheToDatabase(t *testing.T) { + selected := gatewayProfitTestAccount(145, PlatformAnthropic, 0.2, 1) + replacement := selected + expensiveRate := 0.8 + replacement.RateMultiplier = &expensiveRate + + snapshot := NewSchedulerSnapshotService( + &gatewayProfitSnapshotCache{err: errors.New("cache unavailable")}, + nil, + gatewayProfitAccountRepo{account: &replacement}, + nil, + nil, + ) + ctx := context.WithValue(context.Background(), openAIProfitControlGateCtxKey{}, &openAIProfitControlGate{ + groupID: 1, + platform: PlatformAnthropic, + threshold: 0.5, + }) + + latest, vetoed, reason := profitControlVetoLatest(ctx, &selected, snapshot) + require.Same(t, &replacement, latest) + require.True(t, vetoed, "缓存读取失败时必须继续从数据库重读,不能直接使用选号旧对象") + require.Equal(t, openAIProfitFilterReasonThreshold, reason) +} + +func TestGatewayProfitControlTerminalRefreshFailureFallsBackToSelectedObject(t *testing.T) { + selected := gatewayProfitTestAccount(151, PlatformAntigravity, 0.2, 1) + snapshot := NewSchedulerSnapshotService( + &gatewayProfitSnapshotCache{err: errors.New("cache unavailable")}, + nil, + gatewayProfitAccountRepo{err: errors.New("database unavailable")}, + nil, + nil, + ) + ctx := context.WithValue(context.Background(), openAIProfitControlGateCtxKey{}, &openAIProfitControlGate{ + groupID: 1, + platform: PlatformAntigravity, + threshold: 0.5, + }) + + latest, vetoed, reason := profitControlVetoLatest(ctx, &selected, snapshot) + require.Same(t, &selected, latest) + require.False(t, vetoed) + require.Empty(t, reason) +} diff --git a/backend/internal/service/gateway_record_usage_test.go b/backend/internal/service/gateway_record_usage_test.go index 81aff58278..fda35a8e2a 100644 --- a/backend/internal/service/gateway_record_usage_test.go +++ b/backend/internal/service/gateway_record_usage_test.go @@ -345,6 +345,53 @@ func TestGatewayServiceRecordUsage_PeakRateAffectsTokenModeImageOutputTokens(t * require.InDelta(t, expectedActual, userRepo.lastAmount, 1e-12) } +func TestGatewayServiceRecordUsage_UsesExplicitPricingAtForPeakRate(t *testing.T) { + for _, platform := range []string{PlatformAnthropic, PlatformGemini, PlatformGrok, PlatformAntigravity} { + t.Run(platform, func(t *testing.T) { + groupID := int64(903) + usageRepo := &openAIRecordUsageLogRepoStub{inserted: true} + userRepo := &openAIRecordUsageUserRepoStub{} + svc := newGatewayRecordUsageServiceForTest(usageRepo, userRepo, &openAIRecordUsageSubRepoStub{}) + svc.resolver = newOpenAITokenImageChannelPricingResolverForTest(t, groupID, "gemini-image") + + pricingAt := time.Date(2026, time.January, 1, 0, 30, 0, 0, time.UTC) + err := svc.RecordUsage(context.Background(), &RecordUsageInput{ + Result: &ForwardResult{ + RequestID: "gateway_explicit_pricing_at_" + platform, + Model: "gemini-image", + ImageCount: 1, + Usage: ClaudeUsage{ + InputTokens: 1000, + OutputTokens: 600, + ImageOutputTokens: 100, + }, + }, + APIKey: &APIKey{ + ID: 803, + GroupID: i64p(groupID), + Group: &Group{ + ID: groupID, + Platform: platform, + RateMultiplier: 1.0, + SubscriptionType: SubscriptionTypeSubscription, + PeakRateEnabled: true, + PeakStart: "00:00", + PeakEnd: "01:00", + PeakRateMultiplier: 3.0, + }, + }, + User: &User{ID: 603}, + Account: &Account{ID: 703, Platform: platform}, + PricingAt: pricingAt, + }) + + require.NoError(t, err) + require.NotNil(t, usageRepo.lastLog) + require.Equal(t, 3.0, usageRepo.lastLog.RateMultiplier) + }) + } +} + func TestGatewayServiceRecordUsage_UsageLogWriteErrorDoesNotSkipBilling(t *testing.T) { usageRepo := &openAIRecordUsageLogRepoStub{inserted: false, err: MarkUsageLogCreateNotPersisted(context.Canceled)} userRepo := &openAIRecordUsageUserRepoStub{} diff --git a/backend/internal/service/gateway_request_pricing.go b/backend/internal/service/gateway_request_pricing.go new file mode 100644 index 0000000000..d86d9b8bdd --- /dev/null +++ b/backend/internal/service/gateway_request_pricing.go @@ -0,0 +1,55 @@ +package service + +import ( + "context" + "time" + + "github.com/Wei-Shaw/sub2api/internal/pkg/ctxkey" + "github.com/Wei-Shaw/sub2api/internal/pkg/timezone" +) + +type gatewayTokenRequestPricingAtCtxKey struct{} +type gatewayTokenRequestBillingGroupCtxKey struct{} + +// WithGatewayTokenRequestPricing marks a shared-gateway request as token billed +// and freezes the downstream pricing instant for its whole lifetime. Media and +// metadata-only handlers deliberately do not call this helper. +func WithGatewayTokenRequestPricing(ctx context.Context) (context.Context, time.Time) { + if ctx == nil { + ctx = context.Background() + } + pricingAt := timezone.Now() + ctx = context.WithValue(ctx, gatewayTokenRequestPricingAtCtxKey{}, pricingAt) + // 调度过程中可能因 fallback/composite 路由覆盖 ctxkey.Group;计费 D 仍必须 + // 使用认证时刻的父分组,和最终 RecordUsage 的计费归属保持一致。 + if group, ok := ctx.Value(ctxkey.Group).(*Group); ok && IsGroupContextValid(group) { + ctx = context.WithValue(ctx, gatewayTokenRequestBillingGroupCtxKey{}, group) + } + return ctx, pricingAt +} + +func gatewayTokenRequestPricingAtFromContext(ctx context.Context) (time.Time, bool) { + if ctx == nil { + return time.Time{}, false + } + pricingAt, ok := ctx.Value(gatewayTokenRequestPricingAtCtxKey{}).(time.Time) + return pricingAt, ok && !pricingAt.IsZero() +} + +// GatewayTokenRequestPricingAtFromContext exposes the frozen instant to +// handlers before they detach asynchronous usage-recording work. +func GatewayTokenRequestPricingAtFromContext(ctx context.Context) time.Time { + pricingAt, _ := gatewayTokenRequestPricingAtFromContext(ctx) + return pricingAt +} + +func gatewayTokenRequestBillingGroupFromContext(ctx context.Context) *Group { + if ctx == nil { + return nil + } + group, _ := ctx.Value(gatewayTokenRequestBillingGroupCtxKey{}).(*Group) + if IsGroupContextValid(group) { + return group + } + return nil +} diff --git a/backend/internal/service/gateway_request_pricing_test.go b/backend/internal/service/gateway_request_pricing_test.go new file mode 100644 index 0000000000..3f0faf14eb --- /dev/null +++ b/backend/internal/service/gateway_request_pricing_test.go @@ -0,0 +1,20 @@ +//go:build unit + +package service + +import ( + "context" + "testing" + + "github.com/stretchr/testify/require" +) + +func TestWithGatewayTokenRequestPricingMarksOnlyExplicitTokenRequests(t *testing.T) { + ctx, pricingAt := WithGatewayTokenRequestPricing(context.Background()) + + got, ok := gatewayTokenRequestPricingAtFromContext(ctx) + require.True(t, ok) + require.Equal(t, pricingAt, got) + require.Equal(t, pricingAt, GatewayTokenRequestPricingAtFromContext(ctx)) + require.True(t, GatewayTokenRequestPricingAtFromContext(context.Background()).IsZero()) +} diff --git a/backend/internal/service/gateway_scheduling.go b/backend/internal/service/gateway_scheduling.go index 99cfd02c4e..b90bd5a7ee 100644 --- a/backend/internal/service/gateway_scheduling.go +++ b/backend/internal/service/gateway_scheduling.go @@ -64,6 +64,7 @@ func (s *GatewayService) SelectAccountForModelWithExclusions(ctx context.Context // 无分组时只使用原生 anthropic 平台 platform = PlatformAnthropic } + ctx = s.withGatewayProfitControlGate(ctx, groupID) // Claude Code 限制可能已将 groupID 解析为 fallback group, // 渠道限制预检查必须使用解析后的分组。 @@ -116,6 +117,7 @@ func (s *GatewayService) SelectAccountWithLoadAwareness(ctx context.Context, gro return nil, err } ctx = s.withGroupContext(ctx, group) + ctx = s.withGatewayProfitControlGate(ctx, groupID) // Claude Code 限制可能已将 groupID 解析为 fallback group, // 渠道限制预检查必须使用解析后的分组。 @@ -284,6 +286,9 @@ func (s *GatewayService) SelectAccountWithLoadAwareness(ctx context.Context, gro } continue } + if !s.isGatewayAccountProfitEligible(ctx, account) { + continue + } if !s.isAccountAllowedForPlatform(account, platform, useMixed) { filteredPlatform++ continue @@ -339,6 +344,7 @@ func (s *GatewayService) SelectAccountWithLoadAwareness(ctx context.Context, gro var stickyCacheMissReason string gatePass := s.isAccountSchedulableForSelection(stickyAccount) && + s.isGatewayAccountProfitEligible(ctx, stickyAccount) && s.isAccountAllowedForPlatform(stickyAccount, platform, useMixed) && (requestedModel == "" || s.isModelSupportedByAccountWithContext(ctx, stickyAccount, requestedModel)) && s.isAccountSchedulableForModelSelection(ctx, stickyAccount, requestedModel) && @@ -470,7 +476,7 @@ func (s *GatewayService) SelectAccountWithLoadAwareness(ctx context.Context, gro continue } if sessionHash != "" && s.cache != nil { - _ = s.cache.SetSessionAccountID(ctx, derefGroupID(groupID), sessionHash, item.account.ID, stickySessionTTL) + _ = s.bindGatewayStickySessionDuringSelection(ctx, groupID, sessionHash, item.account.ID) } if s.debugModelRoutingEnabled() { logger.LegacyPrintf("service.gateway", "[ModelRoutingDebug] routed select: group_id=%v model=%s session=%s account=%d", derefGroupID(groupID), requestedModel, shortSessionHash(sessionHash), item.account.ID) @@ -523,6 +529,7 @@ func (s *GatewayService) SelectAccountWithLoadAwareness(ctx context.Context, gro // accounts 列表构建,账号一定在分组内。而 scheduler snapshot 缓存 // 反序列化后 AccountGroups 字段为空,导致 isAccountInGroup 永远返回 false。 platformOK := s.isAccountAllowedForPlatform(account, platform, useMixed) + profitOK := s.isGatewayAccountProfitEligible(ctx, account) modelSupported := requestedModel == "" || s.isModelSupportedByAccountWithContext(ctx, account, requestedModel) modelSchedulable := s.isAccountSchedulableForModelSelection(ctx, account, requestedModel) quotaOK := s.isAccountSchedulableForQuota(account) @@ -536,6 +543,7 @@ func (s *GatewayService) SelectAccountWithLoadAwareness(ctx context.Context, gro "clear_sticky", clearSticky, "schedulable", schedulable, "platform_ok", platformOK, + "profit_ok", profitOK, "model_supported", modelSupported, "model_schedulable", modelSchedulable, "quota_ok", quotaOK, @@ -543,7 +551,7 @@ func (s *GatewayService) SelectAccountWithLoadAwareness(ctx context.Context, gro "rpm_ok", rpmOK, ) - if !clearSticky && platformOK && modelSupported && modelSchedulable && quotaOK && windowCostOK && rpmOK && schedulable { + if !clearSticky && platformOK && profitOK && modelSupported && modelSchedulable && quotaOK && windowCostOK && rpmOK && schedulable { result, err := s.tryAcquireAccountSlot(ctx, accountID, account.Concurrency) if err == nil && result.Acquired { // 会话数量限制检查 @@ -639,6 +647,9 @@ func (s *GatewayService) SelectAccountWithLoadAwareness(ctx context.Context, gro if !s.isAccountSchedulableForSelection(acc) { continue } + if !s.isGatewayAccountProfitEligible(ctx, acc) { + continue + } if !s.isAccountAllowedForPlatform(acc, platform, useMixed) { continue } @@ -720,7 +731,7 @@ func (s *GatewayService) SelectAccountWithLoadAwareness(ctx context.Context, gro result.ReleaseFunc() // 释放槽位,继续尝试下一个账号 } else { if sessionHash != "" && s.cache != nil { - _ = s.cache.SetSessionAccountID(ctx, derefGroupID(groupID), sessionHash, selected.account.ID, stickySessionTTL) + _ = s.bindGatewayStickySessionDuringSelection(ctx, groupID, sessionHash, selected.account.ID) } return s.newSelectionResult(ctx, selected.account, true, result.ReleaseFunc, nil) } @@ -768,7 +779,7 @@ func (s *GatewayService) tryAcquireByLegacyOrder(ctx context.Context, candidates continue } if sessionHash != "" && s.cache != nil { - _ = s.cache.SetSessionAccountID(ctx, derefGroupID(groupID), sessionHash, acc.ID, stickySessionTTL) + _ = s.bindGatewayStickySessionDuringSelection(ctx, groupID, sessionHash, acc.ID) } selection, err := s.newSelectionResult(ctx, acc, true, result.ReleaseFunc, nil) if err != nil { @@ -1787,7 +1798,7 @@ func (s *GatewayService) selectAccountForModelWithPlatform(ctx context.Context, if clearSticky { _ = s.cache.DeleteSessionAccountID(ctx, derefGroupID(groupID), sessionHash) } - if !clearSticky && s.isAccountInGroup(account, groupID) && account.Platform == platform && (requestedModel == "" || s.isModelSupportedByAccountWithContext(ctx, account, requestedModel)) && s.isAccountSchedulableForModelSelection(ctx, account, requestedModel) && s.isAccountSchedulableForQuota(account) && s.isAccountSchedulableForWindowCost(ctx, account, true) && s.isAccountSchedulableForRPM(ctx, account, true) && !s.isStickyAccountUpstreamRestricted(ctx, groupID, account, requestedModel) { + if !clearSticky && s.isGatewayAccountProfitEligible(ctx, account) && s.isAccountInGroup(account, groupID) && account.Platform == platform && (requestedModel == "" || s.isModelSupportedByAccountWithContext(ctx, account, requestedModel)) && s.isAccountSchedulableForModelSelection(ctx, account, requestedModel) && s.isAccountSchedulableForQuota(account) && s.isAccountSchedulableForWindowCost(ctx, account, true) && s.isAccountSchedulableForRPM(ctx, account, true) && !s.isStickyAccountUpstreamRestricted(ctx, groupID, account, requestedModel) { if s.debugModelRoutingEnabled() { logger.LegacyPrintf("service.gateway", "[ModelRoutingDebug] legacy routed sticky hit: group_id=%v model=%s session=%s account=%d", derefGroupID(groupID), requestedModel, shortSessionHash(sessionHash), accountID) } @@ -1835,6 +1846,9 @@ func (s *GatewayService) selectAccountForModelWithPlatform(ctx context.Context, if !s.isAccountSchedulableForSelection(acc) { continue } + if !s.isGatewayAccountProfitEligible(ctx, acc) { + continue + } // require_privacy_set: 跳过 privacy 未设置的账号并标记异常 if schedGroup != nil && schedGroup.RequirePrivacySet && !acc.IsPrivacySet() { _ = s.accountRepo.SetError(ctx, acc.ID, @@ -1882,7 +1896,7 @@ func (s *GatewayService) selectAccountForModelWithPlatform(ctx context.Context, if selected != nil { if sessionHash != "" && s.cache != nil { - if err := s.cache.SetSessionAccountID(ctx, derefGroupID(groupID), sessionHash, selected.ID, stickySessionTTL); err != nil { + if err := s.bindGatewayStickySessionDuringSelection(ctx, groupID, sessionHash, selected.ID); err != nil { logger.LegacyPrintf("service.gateway", "set session account failed: session=%s account_id=%d err=%v", sessionHash, selected.ID, err) } } @@ -1906,7 +1920,7 @@ func (s *GatewayService) selectAccountForModelWithPlatform(ctx context.Context, if clearSticky { _ = s.cache.DeleteSessionAccountID(ctx, derefGroupID(groupID), sessionHash) } - if !clearSticky && s.isAccountInGroup(account, groupID) && account.Platform == platform && (requestedModel == "" || s.isModelSupportedByAccountWithContext(ctx, account, requestedModel)) && s.isAccountSchedulableForModelSelection(ctx, account, requestedModel) && s.isAccountSchedulableForQuota(account) && s.isAccountSchedulableForWindowCost(ctx, account, true) && s.isAccountSchedulableForRPM(ctx, account, true) { + if !clearSticky && s.isGatewayAccountProfitEligible(ctx, account) && s.isAccountInGroup(account, groupID) && account.Platform == platform && (requestedModel == "" || s.isModelSupportedByAccountWithContext(ctx, account, requestedModel)) && s.isAccountSchedulableForModelSelection(ctx, account, requestedModel) && s.isAccountSchedulableForQuota(account) && s.isAccountSchedulableForWindowCost(ctx, account, true) && s.isAccountSchedulableForRPM(ctx, account, true) { return account, nil } } @@ -1946,6 +1960,9 @@ func (s *GatewayService) selectAccountForModelWithPlatform(ctx context.Context, if !s.isAccountSchedulableForSelection(acc) { continue } + if !s.isGatewayAccountProfitEligible(ctx, acc) { + continue + } // require_privacy_set: 跳过 privacy 未设置的账号并标记异常 if schedGroup != nil && schedGroup.RequirePrivacySet && !acc.IsPrivacySet() { _ = s.accountRepo.SetError(ctx, acc.ID, @@ -2004,7 +2021,7 @@ func (s *GatewayService) selectAccountForModelWithPlatform(ctx context.Context, // 4. 建立粘性绑定 if sessionHash != "" && s.cache != nil { - if err := s.cache.SetSessionAccountID(ctx, derefGroupID(groupID), sessionHash, selected.ID, stickySessionTTL); err != nil { + if err := s.bindGatewayStickySessionDuringSelection(ctx, groupID, sessionHash, selected.ID); err != nil { logger.LegacyPrintf("service.gateway", "set session account failed: session=%s account_id=%d err=%v", sessionHash, selected.ID, err) } } @@ -2045,7 +2062,7 @@ func (s *GatewayService) selectAccountWithMixedScheduling(ctx context.Context, g if clearSticky { _ = s.cache.DeleteSessionAccountID(ctx, derefGroupID(groupID), sessionHash) } - if !clearSticky && s.isAccountInGroup(account, groupID) && (requestedModel == "" || s.isModelSupportedByAccountWithContext(ctx, account, requestedModel)) && s.isAccountSchedulableForModelSelection(ctx, account, requestedModel) && s.isAccountSchedulableForQuota(account) && s.isAccountSchedulableForWindowCost(ctx, account, true) && s.isAccountSchedulableForRPM(ctx, account, true) { + if !clearSticky && s.isGatewayAccountProfitEligible(ctx, account) && s.isAccountInGroup(account, groupID) && (requestedModel == "" || s.isModelSupportedByAccountWithContext(ctx, account, requestedModel)) && s.isAccountSchedulableForModelSelection(ctx, account, requestedModel) && s.isAccountSchedulableForQuota(account) && s.isAccountSchedulableForWindowCost(ctx, account, true) && s.isAccountSchedulableForRPM(ctx, account, true) { if account.Platform == nativePlatform || (account.Platform == PlatformAntigravity && account.IsMixedSchedulingEnabled()) { if s.debugModelRoutingEnabled() { logger.LegacyPrintf("service.gateway", "[ModelRoutingDebug] legacy mixed routed sticky hit: group_id=%v model=%s session=%s account=%d", derefGroupID(groupID), requestedModel, shortSessionHash(sessionHash), accountID) @@ -2091,6 +2108,9 @@ func (s *GatewayService) selectAccountWithMixedScheduling(ctx context.Context, g if !s.isAccountSchedulableForSelection(acc) { continue } + if !s.isGatewayAccountProfitEligible(ctx, acc) { + continue + } // require_privacy_set: 跳过 privacy 未设置的账号并标记异常 if schedGroup != nil && schedGroup.RequirePrivacySet && !acc.IsPrivacySet() { _ = s.accountRepo.SetError(ctx, acc.ID, @@ -2142,7 +2162,7 @@ func (s *GatewayService) selectAccountWithMixedScheduling(ctx context.Context, g if selected != nil { if sessionHash != "" && s.cache != nil { - if err := s.cache.SetSessionAccountID(ctx, derefGroupID(groupID), sessionHash, selected.ID, stickySessionTTL); err != nil { + if err := s.bindGatewayStickySessionDuringSelection(ctx, groupID, sessionHash, selected.ID); err != nil { logger.LegacyPrintf("service.gateway", "set session account failed: session=%s account_id=%d err=%v", sessionHash, selected.ID, err) } } @@ -2166,7 +2186,7 @@ func (s *GatewayService) selectAccountWithMixedScheduling(ctx context.Context, g if clearSticky { _ = s.cache.DeleteSessionAccountID(ctx, derefGroupID(groupID), sessionHash) } - if !clearSticky && s.isAccountInGroup(account, groupID) && (requestedModel == "" || s.isModelSupportedByAccountWithContext(ctx, account, requestedModel)) && s.isAccountSchedulableForModelSelection(ctx, account, requestedModel) && s.isAccountSchedulableForQuota(account) && s.isAccountSchedulableForWindowCost(ctx, account, true) && s.isAccountSchedulableForRPM(ctx, account, true) && !s.isStickyAccountUpstreamRestricted(ctx, groupID, account, requestedModel) { + if !clearSticky && s.isGatewayAccountProfitEligible(ctx, account) && s.isAccountInGroup(account, groupID) && (requestedModel == "" || s.isModelSupportedByAccountWithContext(ctx, account, requestedModel)) && s.isAccountSchedulableForModelSelection(ctx, account, requestedModel) && s.isAccountSchedulableForQuota(account) && s.isAccountSchedulableForWindowCost(ctx, account, true) && s.isAccountSchedulableForRPM(ctx, account, true) && !s.isStickyAccountUpstreamRestricted(ctx, groupID, account, requestedModel) { if account.Platform == nativePlatform || (account.Platform == PlatformAntigravity && account.IsMixedSchedulingEnabled()) { return account, nil } @@ -2203,6 +2223,9 @@ func (s *GatewayService) selectAccountWithMixedScheduling(ctx context.Context, g if !s.isAccountSchedulableForSelection(acc) { continue } + if !s.isGatewayAccountProfitEligible(ctx, acc) { + continue + } // require_privacy_set: 跳过 privacy 未设置的账号并标记异常 if schedGroup != nil && schedGroup.RequirePrivacySet && !acc.IsPrivacySet() { _ = s.accountRepo.SetError(ctx, acc.ID, @@ -2265,7 +2288,7 @@ func (s *GatewayService) selectAccountWithMixedScheduling(ctx context.Context, g // 4. 建立粘性绑定 if sessionHash != "" && s.cache != nil { - if err := s.cache.SetSessionAccountID(ctx, derefGroupID(groupID), sessionHash, selected.ID, stickySessionTTL); err != nil { + if err := s.bindGatewayStickySessionDuringSelection(ctx, groupID, sessionHash, selected.ID); err != nil { logger.LegacyPrintf("service.gateway", "set session account failed: session=%s account_id=%d err=%v", sessionHash, selected.ID, err) } } @@ -2281,6 +2304,8 @@ type selectionFailureStats struct { PlatformFiltered int ModelUnsupported int ModelRateLimited int + ProfitThreshold int + ProfitInvalidRate int SamplePlatformIDs []int64 SampleMappingIDs []int64 SampleRateLimitIDs []string @@ -2304,7 +2329,7 @@ func (s *GatewayService) logDetailedSelectionFailure( stats := s.collectSelectionFailureStats(ctx, accounts, requestedModel, platform, excludedIDs, allowMixedScheduling) logger.LegacyPrintf( "service.gateway", - "[SelectAccountDetailed] group_id=%v model=%s platform=%s session=%s total=%d eligible=%d excluded=%d unschedulable=%d platform_filtered=%d model_unsupported=%d model_rate_limited=%d sample_platform_filtered=%v sample_model_unsupported=%v sample_model_rate_limited=%v", + "[SelectAccountDetailed] group_id=%v model=%s platform=%s session=%s total=%d eligible=%d excluded=%d unschedulable=%d platform_filtered=%d model_unsupported=%d model_rate_limited=%d profit_threshold=%d profit_invalid_account_rate=%d sample_platform_filtered=%v sample_model_unsupported=%v sample_model_rate_limited=%v", derefGroupID(groupID), requestedModel, platform, @@ -2316,6 +2341,8 @@ func (s *GatewayService) logDetailedSelectionFailure( stats.PlatformFiltered, stats.ModelUnsupported, stats.ModelRateLimited, + stats.ProfitThreshold, + stats.ProfitInvalidRate, stats.SamplePlatformIDs, stats.SampleMappingIDs, stats.SampleRateLimitIDs, @@ -2353,6 +2380,10 @@ func (s *GatewayService) collectSelectionFailureStats( stats.ModelRateLimited++ remaining := acc.GetRateLimitRemainingTimeWithContext(ctx, requestedModel).Truncate(time.Second) stats.SampleRateLimitIDs = appendSelectionFailureRateSample(stats.SampleRateLimitIDs, acc.ID, remaining) + case openAIProfitFilterReasonThreshold: + stats.ProfitThreshold++ + case openAIProfitFilterReasonInvalidAccountRate: + stats.ProfitInvalidRate++ default: stats.Eligible++ } @@ -2397,6 +2428,9 @@ func (s *GatewayService) diagnoseSelectionFailure( Detail: fmt.Sprintf("remaining=%s", remaining), } } + if vetoed, reason := openAIProfitControlVetoReason(ctx, acc); vetoed { + return selectionFailureDiagnosis{Category: reason} + } return selectionFailureDiagnosis{Category: "eligible"} } @@ -2434,7 +2468,7 @@ func appendSelectionFailureRateSample(samples []string, accountID int64, remaini func summarizeSelectionFailureStats(stats selectionFailureStats) string { return fmt.Sprintf( - "total=%d eligible=%d excluded=%d unschedulable=%d platform_filtered=%d model_unsupported=%d model_rate_limited=%d", + "total=%d eligible=%d excluded=%d unschedulable=%d platform_filtered=%d model_unsupported=%d model_rate_limited=%d profit_threshold=%d profit_invalid_account_rate=%d", stats.Total, stats.Eligible, stats.Excluded, @@ -2442,6 +2476,8 @@ func summarizeSelectionFailureStats(stats selectionFailureStats) string { stats.PlatformFiltered, stats.ModelUnsupported, stats.ModelRateLimited, + stats.ProfitThreshold, + stats.ProfitInvalidRate, ) } diff --git a/backend/internal/service/gateway_service.go b/backend/internal/service/gateway_service.go index 2f2f836ddd..7717f7d34a 100644 --- a/backend/internal/service/gateway_service.go +++ b/backend/internal/service/gateway_service.go @@ -444,14 +444,24 @@ var allowedHeaders = map[string]bool{ "x-client-request-id": true, } +// ErrStickySessionNotFound is returned by GatewayCache.GetSessionAccountID +// when no binding exists for the session. It abstracts away the underlying +// cache implementation (e.g. redis.Nil), mirroring ErrRefreshTokenNotFound. +var ErrStickySessionNotFound = errors.New("sticky session not found") + // GatewayCache 定义网关服务的缓存操作接口。 // 提供粘性会话(Sticky Session)的存储、查询、刷新和删除功能。 // // GatewayCache defines cache operations for gateway service. // Provides sticky session storage, retrieval, refresh and deletion capabilities. type GatewayCache interface { - // GetSessionAccountID 获取粘性会话绑定的账号 ID - // Get the account ID bound to a sticky session + // GetSessionAccountID 获取粘性会话绑定的账号 ID;无绑定时返回 + // ErrStickySessionNotFound,使 service 层无需依赖具体缓存实现即可 + // 区分"未绑定"与真实读取失败。 + // Get the account ID bound to a sticky session. Returns + // ErrStickySessionNotFound when no binding exists so service code can + // distinguish a miss from a real read failure without importing the + // cache driver. GetSessionAccountID(ctx context.Context, groupID int64, sessionHash string) (int64, error) // SetSessionAccountID 设置粘性会话与账号的绑定关系 // Set the binding between sticky session and account @@ -873,6 +883,34 @@ func (s *GatewayService) BindStickySession(ctx context.Context, groupID *int64, return s.cache.SetSessionAccountID(ctx, derefGroupID(groupID), sessionHash, accountID, stickySessionTTL) } +// bindGatewayStickySessionDuringSelection preserves the normal eager sticky +// behavior unless a profit gate is installed. Profit-controlled requests bind +// only after the terminal post-slot check, otherwise a rejected candidate could +// overwrite a healthy pre-existing sticky binding. +func (s *GatewayService) bindGatewayStickySessionDuringSelection(ctx context.Context, groupID *int64, sessionHash string, accountID int64) error { + if gatewayProfitControlGateActive(ctx) { + return nil + } + return s.BindStickySession(ctx, groupID, sessionHash, accountID) +} + +// BindStickySessionAfterProfitAdmission records a successful terminal +// selection for a profit-controlled request without replacing the binding +// observed at request start. A temporarily ineligible sticky account remains +// bound and automatically becomes eligible again if its account rate recovers. +func (s *GatewayService) BindStickySessionAfterProfitAdmission(ctx context.Context, groupID *int64, sessionHash string, accountID, existingAccountID int64) error { + if sessionHash == "" || accountID <= 0 || s.cache == nil { + return nil + } + if !gatewayProfitControlGateActive(ctx) { + return nil + } + if existingAccountID > 0 && existingAccountID != accountID { + return nil + } + return s.BindStickySession(ctx, groupID, sessionHash, accountID) +} + // GetCachedSessionAccountID retrieves the account ID bound to a sticky session. // Returns 0 if no binding exists or on error. func (s *GatewayService) GetCachedSessionAccountID(ctx context.Context, groupID *int64, sessionHash string) (int64, error) { diff --git a/backend/internal/service/gateway_usage_billing.go b/backend/internal/service/gateway_usage_billing.go index 51b8031fc4..80de338be4 100644 --- a/backend/internal/service/gateway_usage_billing.go +++ b/backend/internal/service/gateway_usage_billing.go @@ -42,6 +42,7 @@ type RecordUsageInput struct { User *User Account *Account Subscription *UserSubscription // 可选:订阅信息 + PricingAt time.Time // token 售价固定时刻;零值保持既有的记录时刻语义 InboundEndpoint string // 入站端点(客户端请求路径) UpstreamEndpoint string // 上游端点(标准化后的上游路径) UserAgent string // 请求的 User-Agent @@ -569,6 +570,7 @@ func (s *GatewayService) RecordUsage(ctx context.Context, input *RecordUsageInpu User: input.User, Account: input.Account, Subscription: input.Subscription, + PricingAt: input.PricingAt, InboundEndpoint: input.InboundEndpoint, UpstreamEndpoint: input.UpstreamEndpoint, UserAgent: input.UserAgent, @@ -589,6 +591,7 @@ type RecordUsageLongContextInput struct { User *User Account *Account Subscription *UserSubscription // 可选:订阅信息 + PricingAt time.Time // token 售价固定时刻;零值保持既有的记录时刻语义 InboundEndpoint string // 入站端点(客户端请求路径) UpstreamEndpoint string // 上游端点(标准化后的上游路径) UserAgent string // 请求的 User-Agent @@ -612,6 +615,7 @@ func (s *GatewayService) RecordUsageWithLongContext(ctx context.Context, input * User: input.User, Account: input.Account, Subscription: input.Subscription, + PricingAt: input.PricingAt, InboundEndpoint: input.InboundEndpoint, UpstreamEndpoint: input.UpstreamEndpoint, UserAgent: input.UserAgent, @@ -635,6 +639,7 @@ type recordUsageCoreInput struct { User *User Account *Account Subscription *UserSubscription + PricingAt time.Time InboundEndpoint string UpstreamEndpoint string UserAgent string @@ -685,7 +690,11 @@ func (s *GatewayService) recordUsageCore(ctx context.Context, input *recordUsage } // token 倍率叠加高峰因子(token 计费含图片 token,图片按次倍率不受影响)。高峰因子按请求时刻现算, // 不并入上面的 getUserGroupRateMultiplier,以免污染 user:group 倍率缓存。 - multiplier, imageMultiplier := computePeakAwareMultipliers(apiKey, multiplier, timezone.Now()) + pricingAt := input.PricingAt + if pricingAt.IsZero() { + pricingAt = timezone.Now() + } + multiplier, imageMultiplier := computePeakAwareMultipliers(apiKey, multiplier, pricingAt) // 确定计费模型 concreteBillingModel := forwardResultBillingModel(result.Model, result.UpstreamModel) diff --git a/backend/internal/service/group.go b/backend/internal/service/group.go index b71e6b9e84..549a37e40c 100644 --- a/backend/internal/service/group.go +++ b/backend/internal/service/group.go @@ -3,6 +3,7 @@ package service import ( "errors" "fmt" + "math" "strings" "time" @@ -99,6 +100,14 @@ type Group struct { // ReasoningEffortMappings rewrites explicit request values before applying the ceiling. ReasoningEffortMappings []ReasoningEffortMapping + // 分组利润控制(五个 token 计费平台可启用)。 + // 调度准入条件:账号倍率 U 满足 U <= D*(1-margin-buffer), + // D 为请求用户当刻有效下游倍率(用户覆盖 ?? 分组默认,再乘高峰因子)。 + // 只过滤候选账号,不改变既有排序/评分/粘性/熔断。 + ProfitControlEnabled bool + ProfitMinMargin float64 // 最低毛利率,小数存储(0.30=30%) + ProfitSafetyBuffer float64 // 安全缓冲,小数,与 margin 相加后从 D 中扣除 + CreatedAt time.Time UpdatedAt time.Time @@ -336,3 +345,61 @@ func computePeakAwareMultipliers(apiKey *APIKey, base float64, now time.Time) (t text = base * peak return } + +// validProfitControlRatio 判定 margin/buffer 是否为可落库的合法小数:[0,1) 且非 NaN/Inf。 +func validProfitControlRatio(v float64) bool { + return !math.IsNaN(v) && !math.IsInf(v, 0) && v >= 0 && v < 1 +} + +// ValidateProfitControlConfig 是分组利润控制配置的唯一校验来源,handler 与 service 层共用。 +// enabled=true 时仅允许五个可计费平台分组;margin/buffer 各自 ∈ [0,1),且 margin+buffer < 1 +// (相加 >=1 时阈值 <=0,所有可核价账号都会被排除,视为配置错误而不是静默全黑)。 +// enabled=false 时放行(不关心平台),由 Normalize 兜底清洗数值。 +func ValidateProfitControlConfig(platform string, enabled bool, minMargin, safetyBuffer float64) error { + if !enabled { + return nil + } + if !profitControlPlatformSupported(platform) { + return errors.New("利润控制仅支持 openai、anthropic、gemini、grok、antigravity 平台分组") + } + if !validProfitControlRatio(minMargin) { + return fmt.Errorf("profit_min_margin 应为 [0,1) 的小数,got %v", minMargin) + } + if !validProfitControlRatio(safetyBuffer) { + return fmt.Errorf("profit_safety_buffer 应为 [0,1) 的小数,got %v", safetyBuffer) + } + if minMargin+safetyBuffer >= 1 { + return errors.New("profit_min_margin 与 profit_safety_buffer 之和必须小于 1,否则将排除全部账号") + } + return nil +} + +// NormalizeProfitControlConfig 归一化最终落库的利润控制配置,CreateGroup 与 UpdateGroup 共用(唯一收口): +// - 非五个平台分组不携带利润控制,一律重置为默认(关、0、0); +// - 支持平台关闭开关时保留合法数值(便于再次启用),清洗 NaN/Inf/越界脏值。 +// +// 与 ValidateProfitControlConfig 的分工同高峰倍率:先归一化、后校验, +// 使"openai 转其他平台"这类更新能静默清空利润配置而不是被校验拒绝。 +func NormalizeProfitControlConfig(platform string, enabled bool, minMargin, safetyBuffer float64) (bool, float64, float64) { + if !profitControlPlatformSupported(platform) { + return false, 0, 0 + } + if !enabled { + if !validProfitControlRatio(minMargin) { + minMargin = 0 + } + if !validProfitControlRatio(safetyBuffer) { + safetyBuffer = 0 + } + } + return enabled, minMargin, safetyBuffer +} + +func profitControlPlatformSupported(platform string) bool { + switch platform { + case PlatformOpenAI, PlatformAnthropic, PlatformGemini, PlatformGrok, PlatformAntigravity: + return true + default: + return false + } +} diff --git a/backend/internal/service/openai_account_scheduler.go b/backend/internal/service/openai_account_scheduler.go index c82e845a86..bf5ef751e9 100644 --- a/backend/internal/service/openai_account_scheduler.go +++ b/backend/internal/service/openai_account_scheduler.go @@ -406,7 +406,7 @@ func (s *defaultOpenAIAccountScheduler) Select( decision.SelectedAccountID = selection.Account.ID decision.SelectedAccountType = selection.Account.Type if req.SessionHash != "" { - _ = s.service.BindStickySession(ctx, req.GroupID, req.SessionHash, selection.Account.ID) + _ = s.service.bindOpenAIStickySessionDuringSelection(ctx, req.GroupID, req.SessionHash, selection.Account.ID) } return selection, decision, nil } @@ -1168,7 +1168,7 @@ func (s *defaultOpenAIAccountScheduler) tryAcquireOpenAISelectionOrderWithBudget } } if req.SessionHash != "" && !req.PreserveStickyBinding { - _ = s.service.BindStickySession(ctx, req.GroupID, req.SessionHash, fresh.ID) + _ = s.service.bindOpenAIStickySessionDuringSelection(ctx, req.GroupID, req.SessionHash, fresh.ID) } return &AccountSelectionResult{ Account: fresh, @@ -1247,7 +1247,7 @@ func (s *defaultOpenAIAccountScheduler) tryFallbackToWeightedSticky( } if result != nil && result.Acquired { if req.SessionHash != "" && !req.PreserveStickyBinding { - _ = s.service.BindStickySession(ctx, req.GroupID, req.SessionHash, account.ID) + _ = s.service.bindOpenAIStickySessionDuringSelection(ctx, req.GroupID, req.SessionHash, account.ID) } return &AccountSelectionResult{ Account: account, @@ -1712,6 +1712,11 @@ func (s *defaultOpenAIAccountScheduler) isAccountRequestCompatibleReason(ctx con if !accountSupportsOpenAICapabilities(account, req.RequiredCapability, req.RequiredImageCapability) { return false, "capability_mismatch" } + // 分组利润控制:不合格账号在候选过滤与抢槽后终检阶段即被排除, + // 排序/评分/粘性/熔断只在合格账号之间工作;named reason 进入 filter stats。 + if vetoed, reason := openAIProfitControlVetoReason(ctx, account); vetoed { + return false, reason + } return true, "" } @@ -2093,6 +2098,17 @@ func (s *OpenAIGatewayService) selectAccountWithSchedulerOnce( useUpstreamTokenCost bool, ) (*AccountSelectionResult, OpenAIAccountScheduleDecision, error) { ctx = s.withOpenAIQuotaAutoPauseContext(ctx) + // 分组利润控制:唯一文本调度入口的防御性装门。handler 文本 + // 入口已在请求开始经 WithOpenAIRequestPricingContext 装门并固定 pricingAt, + // 此处对同分组门直接复用(failover 重入阈值稳定),仅为不经 handler 装配的 + // 内部调用兜底。图片/视频调度不在利润门范围:requiredImageCapability 非空的 + // Images 调度不装门;requiredCapability == OpenAIEndpointCapabilityResponses + // 当前仅显式生图意图的 /v1/responses 设置(HTTP openAIResponsesRequiredCapability + // 与 WS 桥同款判定),同样不装门——若未来把该 capability 用于非生图流量, + // 需要同步收窄本条件(有测试钉死该映射)。 + if requiredImageCapability == "" && requiredCapability != OpenAIEndpointCapabilityResponses { + ctx = s.withOpenAIProfitControlGate(ctx, groupID) + } platform = normalizeOpenAICompatiblePlatform(platform) decision := OpenAIAccountScheduleDecision{} scheduler := s.getOpenAIAccountScheduler(ctx) diff --git a/backend/internal/service/openai_account_scheduler_test.go b/backend/internal/service/openai_account_scheduler_test.go index a7c51f6fa7..03f30c5b31 100644 --- a/backend/internal/service/openai_account_scheduler_test.go +++ b/backend/internal/service/openai_account_scheduler_test.go @@ -155,7 +155,7 @@ func (c *schedulerTestGatewayCache) GetSessionAccountID(ctx context.Context, gro if id, ok := c.sessionBindings[sessionHash]; ok { return id, nil } - return 0, errors.New("not found") + return 0, ErrStickySessionNotFound } func (c *schedulerTestGatewayCache) SetSessionAccountID(ctx context.Context, groupID int64, sessionHash string, accountID int64, ttl time.Duration) error { diff --git a/backend/internal/service/openai_gateway_scheduling.go b/backend/internal/service/openai_gateway_scheduling.go index 7c0ca8f033..ca9ce44b05 100644 --- a/backend/internal/service/openai_gateway_scheduling.go +++ b/backend/internal/service/openai_gateway_scheduling.go @@ -300,6 +300,11 @@ func isOpenAICompatibleAccountEligibleForRequest(ctx context.Context, account *A if requireCompact && openAICompactSupportTier(account) == 0 { return false } + // 分组利润控制:legacy 引擎的粘性/候选循环与 DB recheck 共用 + // 本判定,任何 fallback 都不能把利润不合格账号重新放回候选。 + if vetoed, _ := openAIProfitControlVetoReason(ctx, account); vetoed { + return false + } return true } @@ -842,7 +847,11 @@ func (s *OpenAIGatewayService) isBetterAccount(candidate, current *Account) bool // SelectAccountWithLoadAwareness selects an account with load-awareness and wait plan. func (s *OpenAIGatewayService) SelectAccountWithLoadAwareness(ctx context.Context, groupID *int64, sessionHash string, requestedModel string, excludedIDs map[int64]struct{}) (*AccountSelectionResult, error) { - return s.selectAccountWithLoadAwareness(s.withOpenAIQuotaAutoPauseContext(ctx), groupID, PlatformOpenAI, sessionHash, requestedModel, excludedIDs, false, "", true) + ctx = s.withOpenAIQuotaAutoPauseContext(ctx) + // 分组利润控制:legacy 公共入口同样装门,保证不经 + // selectAccountWithScheduler 的调用方也无法绕过利润准入。 + ctx = s.withOpenAIProfitControlGate(ctx, groupID) + return s.selectAccountWithLoadAwareness(ctx, groupID, PlatformOpenAI, sessionHash, requestedModel, excludedIDs, false, "", true) } func (s *OpenAIGatewayService) selectAccountWithLoadAwareness(ctx context.Context, groupID *int64, platform string, sessionHash string, requestedModel string, excludedIDs map[int64]struct{}, requireCompact bool, requiredCapability OpenAIEndpointCapability, useUpstreamTokenCost bool) (*AccountSelectionResult, error) { diff --git a/backend/internal/service/openai_gateway_usage.go b/backend/internal/service/openai_gateway_usage.go index 1244698333..d1b27a1ff7 100644 --- a/backend/internal/service/openai_gateway_usage.go +++ b/backend/internal/service/openai_gateway_usage.go @@ -33,6 +33,10 @@ type OpenAIRecordUsageInput struct { RequestPayloadHash string APIKeyService APIKeyQuotaUpdater QuotaPlatform string // user×platform quota platform resolved by the handler before async billing. + // PricingAt 是请求级定价时刻(请求开始捕获,与利润门的 D 同源):高峰因子 + // 按该时刻计算,保证同一请求从准入到扣费不中途变价。零值回退记录时刻 + //(既有行为),供未装配的路径(图片/异步/cyber 等)沿用。 + PricingAt time.Time // CyberBlocked 为 true 时把该用量行标记为 cyber(request_type=cyber),计费逻辑不变。 CyberBlocked bool ChannelUsageFields @@ -113,6 +117,15 @@ func (s *OpenAIGatewayService) ResolveUserGroupRateMultiplier(ctx context.Contex return resolver.Resolve(ctx, userID, groupID, groupDefaultMultiplier) } +// openAIUsagePricingAt 返回本次用量记录使用的定价时刻:优先请求级 PricingAt +// (与利润门 D 同源同刻),未装配时回退记录时刻(既有行为)。 +func openAIUsagePricingAt(input *OpenAIRecordUsageInput) time.Time { + if input != nil && !input.PricingAt.IsZero() { + return input.PricingAt + } + return timezone.Now() +} + // RecordUsage records usage and deducts balance func (s *OpenAIGatewayService) RecordUsage(ctx context.Context, input *OpenAIRecordUsageInput) error { if input == nil { @@ -159,10 +172,12 @@ func (s *OpenAIGatewayService) RecordUsage(ctx context.Context, input *OpenAIRec if apiKey.GroupID != nil && apiKey.Group != nil { multiplier = s.ResolveUserGroupRateMultiplier(ctx, user.ID, *apiKey.GroupID, apiKey.Group.RateMultiplier) } - // token 倍率叠加高峰因子(token 计费含图片 token,图片按次倍率不受影响)。高峰因子按请求时刻现算, - // 不并入上面的 Resolve,以免污染 user:group 倍率缓存。 + // token 倍率叠加高峰因子(token 计费含图片 token,图片按次倍率不受影响)。 + // 高峰因子按请求级 PricingAt 现算(与利润门 D 同源同刻,跨峰谷请求不中途 + // 变价);未装配 PricingAt 的路径回退记录时刻,保持既有行为。不并入上面的 + // Resolve,以免污染 user:group 倍率缓存。 baseMultiplier := multiplier - multiplier, imageMultiplier := computePeakAwareMultipliers(apiKey, baseMultiplier, timezone.Now()) + multiplier, imageMultiplier := computePeakAwareMultipliers(apiKey, baseMultiplier, openAIUsagePricingAt(input)) videoMultiplier := resolveVideoRateMultiplier(apiKey, baseMultiplier) var cost *CostBreakdown diff --git a/backend/internal/service/openai_profit_control.go b/backend/internal/service/openai_profit_control.go new file mode 100644 index 0000000000..2e42295467 --- /dev/null +++ b/backend/internal/service/openai_profit_control.go @@ -0,0 +1,366 @@ +package service + +// 分组利润控制(配套 migration 191/192 的 groups.profit_* 字段)。 +// +// 定位:利润控制是"候选准入过滤",只决定账号能否进入调度候选池;既有的排序、 +// 评分、粘性、熔断、负载均衡在合格账号之间照常工作,本文件不改变它们的行为。 +// +// 准入条件: +// +// U(尝试时刻) <= D(pricingAt) × (1 − profit_min_margin − profit_safety_buffer) +// +// - D(用户售价倍率)固定在请求开始的 pricingAt:同一请求的全部 failover 与 +// 最终扣费共用同一 D(RecordUsage 的高峰因子同样取 pricingAt),一个请求 +// 不会中途变价。D 与计费完全同源:按请求真实计费分组(ctxkey.Group,即 +// apiKey 自身分组;composite 请求为父分组)做 ResolveUserGroupRateMultiplier +// (用户-分组覆盖 ?? 分组默认)× Group.PeakMultiplierAt(pricingAt),绝不在 +// 用户有覆盖时退回分组默认;开关与 margin/buffer 则始终取被调度 +// openai/grok 分组。 +// - U(上游成本倍率)取 accounts.rate_multiplier。倍率可以由运营者手工维护, +// 也可以由上游倍率探测同步写回;利润门不再耦合探测协议、新鲜度或账号类型。 +// 0 是合法的免费上游倍率;nil、负数、NaN、Inf 属于非法数据并保守拒绝。 +// +// 装门点(gate 随 ctx 传播,请求内复用,覆盖等待/重试/failover/抢槽后终检): +// - handler 各文本入口经 WithOpenAIRequestPricingContext 在请求开始统一装门并 +// 固定 pricingAt;显式生图意图(responses image_generation)以抑制标记跳门, +// 图片/视频保持既有调度不装门。 +// - selectAccountWithScheduler 顶部:唯一文本调度入口的防御性装门(ctx 已有 +// 同分组门则复用,保证 failover 阈值稳定);requiredImageCapability != "" 或 +// requiredCapability == OpenAIEndpointCapabilityResponses(该值当前仅显式 +// 生图意图设置)不装门。 +// - 公开 SelectAccountWithLoadAwareness / SelectAccountByPreviousResponseID: +// 防御性装门,保证不经唯一入口的调用方无法绕过。 +// +// 否决点(消费 gate,任何 fallback 都无法把已排除账号重新放回): +// - defaultOpenAIAccountScheduler.isAccountRequestCompatibleReason:候选池 +// 过滤 + 调度器内抢槽后终检共用,named reason 进入 openAISelectionFilterStats。 +// - isOpenAICompatibleAccountEligibleForRequest:legacy 引擎与 DB recheck 共用。 +// - resolveAccountByPreviousResponseIDForCapability:previous_response 粘连 +// 两阶段校验;与 quota auto-pause 同语义,跳过复用但不删除绑定。 +// - handler 槽位获取后终检(OpenAIProfitControlVeto):快速抢槽与 WaitPlan +// 排队成功后复核,越线则释放槽位、加入本请求排除集重新选号,全池耗尽才 +// 返回标准 no available accounts。 +// +// 失败语义:分组配置读取失败时放行并告警(fail-open)。这是"配置系统故障时 +// 可用性优先"的显式取舍——该异常窗口内利润保证不成立,靠 WARN 与采样观测 +// 暴露,绝不把瞬时 DB 抖动放大成全站不可调度。 +// +// 可观测性:按分组和平台累计装门/threshold 否决/invalid-rate 否决/终检刷新 +// 失败计数,≥5 分钟采样输出一条 Info(profit_control_activity),无逐请求日志。 + +import ( + "context" + "errors" + "fmt" + "log/slog" + "math" + "sync" + "sync/atomic" + "time" + + "github.com/Wei-Shaw/sub2api/internal/pkg/ctxkey" + "github.com/Wei-Shaw/sub2api/internal/pkg/timezone" +) + +const ( + // profitControlRateEpsilon 吸收 decimal(10,4) 落库与浮点乘法的边界误差: + // U 与阈值的相对差在该量级内视为相等(即 U == 阈值 判定为合格)。 + profitControlRateEpsilon = 1e-9 + + // 进入 openAISelectionFilterStats 的排除原因,全排除时出现在 + // "no available accounts" 的内部统计摘要里(与 quota_auto_pause 等同通道)。 + openAIProfitFilterReasonThreshold = "profit_threshold" + openAIProfitFilterReasonInvalidAccountRate = "profit_invalid_account_rate" + + // profitControlActivityLogInterval 是按分组采样输出累计计数的最小间隔。 + profitControlActivityLogInterval = 5 * time.Minute +) + +type openAIProfitControlGateCtxKey struct{} + +// openAIProfitControlSuppressCtxKey 标记本请求显式跳过利润门(生图意图等 +// 利润门范围外流量)。所有装门点看到该标记后一律不装门,防止 service 层防御性 +// 装门把边界外流量重新拉回利润过滤。 +type openAIProfitControlSuppressCtxKey struct{} + +// openAIPricingAtCtxKey 携带请求级定价时刻 pricingAt:门的 D 与 RecordUsage +// 的高峰因子共用,保证一个请求从准入到扣费不中途变价。 +type openAIPricingAtCtxKey struct{} + +// openAIProfitControlGate 是一个请求的利润准入门。除 pricingAt 外全部为预计算 +// 标量:候选过滤热路径上每账号只做一次快照解码与一次浮点比较。 +type openAIProfitControlGate struct { + // groupID 是门配置来源的被调度分组;请求内按分组复用(failover 阈值稳定), + // composite 等跨分组调度切换分组时重新解析。 + groupID int64 + // platform 是利润配置所在分组的平台,用于按平台观测门是否真实生效。 + platform string + // threshold = D(pricingAt) × (1 − margin − buffer),账号倍率必须 <= 它。 + threshold float64 + // pricingAt 是本请求的统一定价时刻(D 侧)。 + pricingAt time.Time +} + +// WithOpenAIRequestPricingContext 在请求开始处装配请求级定价上下文:固定 +// pricingAt(返回给调用方,供 RecordUsage 入参共用同一时刻),并按分组安装 +// 利润门;suppressProfitGate 为 true(显式生图意图)时只固定 pricingAt、写入 +// 抑制标记,不装门。handler 各文本入口应在选号循环前调用一次。 +func (s *OpenAIGatewayService) WithOpenAIRequestPricingContext(ctx context.Context, groupID *int64, suppressProfitGate bool) (context.Context, time.Time) { + pricingAt := timezone.Now() + ctx = context.WithValue(ctx, openAIPricingAtCtxKey{}, pricingAt) + if suppressProfitGate { + return context.WithValue(ctx, openAIProfitControlSuppressCtxKey{}, struct{}{}), pricingAt + } + return s.withOpenAIProfitControlGate(ctx, groupID), pricingAt +} + +// openAIPricingAtFromContext 返回请求级定价时刻;未装配(内部调用、非文本 +// 入口)时 ok=false,调用方回退 timezone.Now() 保持既有行为。 +func openAIPricingAtFromContext(ctx context.Context) (time.Time, bool) { + pricingAt, ok := ctx.Value(openAIPricingAtCtxKey{}).(time.Time) + return pricingAt, ok && !pricingAt.IsZero() +} + +// OpenAIPricingAtFromContext 是 handler 侧读取请求级定价时刻的公开入口(未装配 +// 时为零值,RecordUsage 对零值回退记录时刻)。供跨函数传递 pricingAt 不便的 +// 记录路径直接从请求 ctx 取值。 +func OpenAIPricingAtFromContext(ctx context.Context) time.Time { + pricingAt, _ := openAIPricingAtFromContext(ctx) + return pricingAt +} + +// withOpenAIProfitControlGate 解析分组利润控制配置;启用时把预计算好的准入门 +// 装进 ctx。抑制标记、未启用/非 openai 分组/无法取到分组配置时原样返回 ctx +// (门不存在,全部否决点自动放行,既有行为零变化)。ctx 已有同分组门时直接 +// 复用:同一请求的全部 failover 重入共享同一阈值。 +func (s *OpenAIGatewayService) withOpenAIProfitControlGate(ctx context.Context, groupID *int64) context.Context { + if _, suppressed := ctx.Value(openAIProfitControlSuppressCtxKey{}).(struct{}); suppressed { + return ctx + } + if groupID != nil { + if existing, ok := ctx.Value(openAIProfitControlGateCtxKey{}).(*openAIProfitControlGate); ok && existing != nil && existing.groupID == *groupID { + return ctx + } + } + gate := s.resolveOpenAIProfitControlGate(ctx, groupID) + if gate == nil { + // 被调度分组无门(未启用/非 openai/配置读取失败)而 ctx 带着其他分组的 + // 请求门时清除之:门配置取被调度分组,父分组阈值不得泄漏到成员分组 + //(composite/模型路由等跨分组调度)。typed-nil 覆盖值由否决点按无门放行。 + if existing, ok := ctx.Value(openAIProfitControlGateCtxKey{}).(*openAIProfitControlGate); ok && existing != nil && groupID != nil && existing.groupID != *groupID { + return context.WithValue(ctx, openAIProfitControlGateCtxKey{}, (*openAIProfitControlGate)(nil)) + } + return ctx + } + openAIProfitControlObserverInstance.recordInstall(gate.groupID, gate.platform, gate.threshold) + return context.WithValue(ctx, openAIProfitControlGateCtxKey{}, gate) +} + +func (s *OpenAIGatewayService) resolveOpenAIProfitControlGate(ctx context.Context, groupID *int64) *openAIProfitControlGate { + if s == nil || groupID == nil || *groupID <= 0 { + return nil + } + // 门配置取被调度分组。直连请求(ctx 认证分组即调度分组,生产绝大多数流量) + // 直接复用 auth cache 分组,热路径零额外查询;composite 父分组路由到成员 + // 分组等 ID 不一致场景才回源仓库读取。auth 快照的分组字段完备性由 + // GetByKeyForAuth 投影 + 集成测试保证(防投影漏列导致门静默失效)。 + var group *Group + if ctxGroup, ok := ctx.Value(ctxkey.Group).(*Group); ok && IsGroupContextValid(ctxGroup) && ctxGroup.ID == *groupID { + group = ctxGroup + } else if s.schedulerSnapshot != nil { + loaded, err := s.schedulerSnapshot.GetGroupByID(ctx, *groupID) + if err != nil { + // fail-open:配置系统故障时可用性优先,该窗口内利润保证不成立, + // 依赖 WARN 暴露;不把瞬时 DB 抖动放大成全站不可调度。 + slog.Warn("profit_control_group_load_failed", "group_id", *groupID, "error", err) + return nil + } + group = loaded + } + if group == nil || !group.ProfitControlEnabled || + (group.Platform != PlatformOpenAI && group.Platform != PlatformGrok) { + return nil + } + + pricingAt, ok := openAIPricingAtFromContext(ctx) + if !ok { + pricingAt = timezone.Now() + } + // D 与计费完全同源(RecordUsage 组合):计费永远按 apiKey 自身分组 + //(composite 请求即父分组)的"用户覆盖 ?? 分组默认 × 高峰因子"计算, + // 因此优先取认证中间件放入 ctx 的分组;ctx 中无有效分组(内部调用)时 + // 退回调度分组组合,直连 openai 分组场景两者等价。 + billingGroup := group + if ctxGroup, ok := ctx.Value(ctxkey.Group).(*Group); ok && IsGroupContextValid(ctxGroup) { + billingGroup = ctxGroup + } + downstream := billingGroup.RateMultiplier + if userID, _ := ctx.Value(ctxkey.UserID).(int64); userID > 0 { + downstream = s.ResolveUserGroupRateMultiplier(ctx, userID, billingGroup.ID, billingGroup.RateMultiplier) + } + downstream *= billingGroup.PeakMultiplierAt(pricingAt) + + deduction := group.ProfitMinMargin + group.ProfitSafetyBuffer + threshold := downstream * (1 - deduction) + // Validate/Normalize 已保证 margin+buffer < 1;此处仅对存量脏数据兜底, + // 阈值非有限或为负时按 0 处理(等价于只放行免费上游)。 + if math.IsNaN(threshold) || math.IsInf(threshold, 0) || threshold < 0 { + threshold = 0 + } + return &openAIProfitControlGate{ + groupID: *groupID, + platform: group.Platform, + threshold: threshold, + pricingAt: pricingAt, + } +} + +// openAIProfitControlVetoReason 报告利润门是否否决该账号。ctx 中没有门 +// (分组未启用利润控制或本请求跳门)或账号为 nil 时一律放行。 +func openAIProfitControlVetoReason(ctx context.Context, account *Account) (bool, string) { + gate, _ := ctx.Value(openAIProfitControlGateCtxKey{}).(*openAIProfitControlGate) + if gate == nil || account == nil { + return false, "" + } + if account.RateMultiplier == nil || + math.IsNaN(*account.RateMultiplier) || + math.IsInf(*account.RateMultiplier, 0) || + *account.RateMultiplier < 0 { + openAIProfitControlObserverInstance.recordVeto(gate.groupID, gate.platform, gate.threshold, openAIProfitFilterReasonInvalidAccountRate) + return true, openAIProfitFilterReasonInvalidAccountRate + } + upstream := *account.RateMultiplier + if upstream-gate.threshold > profitControlRateEpsilon*math.Max(1, math.Abs(gate.threshold)) { + openAIProfitControlObserverInstance.recordVeto(gate.groupID, gate.platform, gate.threshold, openAIProfitFilterReasonThreshold) + return true, openAIProfitFilterReasonThreshold + } + return false, "" +} + +// OpenAIProfitControlVeto 是 handler 层槽位获取后终检的公开入口:语义与调度 +// 内否决点完全一致。ctx 必须是经 WithOpenAIRequestPricingContext 装配过的 +// 请求上下文,否则视为无门放行。 +func OpenAIProfitControlVeto(ctx context.Context, account *Account) (bool, string) { + return openAIProfitControlVetoReason(ctx, account) +} + +// ProfitControlVetoLatest performs the handler-side terminal check after a +// concurrency slot is actually acquired. The latest cached account replaces +// the selection snapshot when available, so a probe/manual rate change during +// wait time cannot pass on a stale pointer. +func (s *OpenAIGatewayService) ProfitControlVetoLatest(ctx context.Context, selected *Account) (*Account, bool, string) { + if s == nil { + return selected, false, "" + } + return profitControlVetoLatest(ctx, selected, s.schedulerSnapshot) +} + +// bindOpenAIStickySessionDuringSelection preserves the official eager binding +// behavior for requests without a profit gate. Profit-controlled requests bind +// only after the terminal post-slot check, so an account rejected after a rate +// refresh cannot become the new sticky target. +func (s *OpenAIGatewayService) bindOpenAIStickySessionDuringSelection(ctx context.Context, groupID *int64, sessionHash string, accountID int64) error { + if gatewayProfitControlGateActive(ctx) { + return nil + } + return s.BindStickySession(ctx, groupID, sessionHash, accountID) +} + +// BindStickySessionAfterProfitAdmission records the terminally admitted +// account without overwriting a different binding that already exists. A +// temporarily ineligible account therefore remains sticky and becomes +// eligible again automatically after its rate recovers. +func (s *OpenAIGatewayService) BindStickySessionAfterProfitAdmission(ctx context.Context, groupID *int64, sessionHash string, accountID int64) error { + if sessionHash == "" || accountID <= 0 || !gatewayProfitControlGateActive(ctx) { + return nil + } + existingAccountID, err := s.getStickySessionAccountID(ctx, groupID, sessionHash) + if err != nil && !errors.Is(err, ErrStickySessionNotFound) { + slog.Warn("profit_control_sticky_binding_read_failed", "group_id", derefGroupID(groupID), "account_id", accountID, "error", err) + return nil + } + if existingAccountID > 0 && existingAccountID != accountID { + return nil + } + return s.BindStickySession(ctx, groupID, sessionHash, accountID) +} + +// ---- 可观测性:按分组累计计数 + 采样日志(无逐请求输出) ---- + +type openAIProfitControlGroupStats struct { + installs atomic.Int64 + vetoThreshold atomic.Int64 + vetoInvalidRate atomic.Int64 + refreshFailures atomic.Int64 + lastLogUnixMilli atomic.Int64 +} + +type openAIProfitControlObserver struct { + groups sync.Map // "platform:groupID" -> *openAIProfitControlGroupStats +} + +var openAIProfitControlObserverInstance = &openAIProfitControlObserver{} + +func profitControlObserverKey(groupID int64, platform string) string { + return platform + ":" + fmt.Sprintf("%d", groupID) +} + +func (o *openAIProfitControlObserver) stats(groupID int64, platform string) *openAIProfitControlGroupStats { + key := profitControlObserverKey(groupID, platform) + if v, ok := o.groups.Load(key); ok { + if s, ok := v.(*openAIProfitControlGroupStats); ok { + return s + } + } + v, _ := o.groups.LoadOrStore(key, &openAIProfitControlGroupStats{}) + if s, ok := v.(*openAIProfitControlGroupStats); ok { + return s + } + // 不可达:map 中只存 *openAIProfitControlGroupStats;兜底返回独立实例避免 panic。 + return &openAIProfitControlGroupStats{} +} + +func (o *openAIProfitControlObserver) recordInstall(groupID int64, platform string, threshold float64) { + s := o.stats(groupID, platform) + s.installs.Add(1) + o.maybeLog(groupID, platform, threshold, s) +} + +func (o *openAIProfitControlObserver) recordVeto(groupID int64, platform string, threshold float64, reason string) { + s := o.stats(groupID, platform) + switch reason { + case openAIProfitFilterReasonThreshold: + s.vetoThreshold.Add(1) + case openAIProfitFilterReasonInvalidAccountRate: + s.vetoInvalidRate.Add(1) + } + o.maybeLog(groupID, platform, threshold, s) +} + +func (o *openAIProfitControlObserver) recordRefreshFailure(groupID int64, platform string, threshold float64) { + s := o.stats(groupID, platform) + s.refreshFailures.Add(1) + o.maybeLog(groupID, platform, threshold, s) +} + +// maybeLog 以 CAS 保证同分组 ≥ profitControlActivityLogInterval 才输出一条 +// 累计计数 Info;计数为进程内累计值,用于确认门在真实流量上生效及否决构成。 +func (o *openAIProfitControlObserver) maybeLog(groupID int64, platform string, threshold float64, s *openAIProfitControlGroupStats) { + now := time.Now().UnixMilli() + last := s.lastLogUnixMilli.Load() + if last != 0 && now-last < profitControlActivityLogInterval.Milliseconds() { + return + } + if !s.lastLogUnixMilli.CompareAndSwap(last, now) { + return + } + slog.Info("profit_control_activity", + "group_id", groupID, + "platform", platform, + "threshold", threshold, + "installs_total", s.installs.Load(), + "veto_threshold_total", s.vetoThreshold.Load(), + "veto_invalid_account_rate_total", s.vetoInvalidRate.Load(), + "refresh_failure_total", s.refreshFailures.Load(), + ) +} diff --git a/backend/internal/service/openai_profit_control_paths_test.go b/backend/internal/service/openai_profit_control_paths_test.go new file mode 100644 index 0000000000..58bb46bbcd --- /dev/null +++ b/backend/internal/service/openai_profit_control_paths_test.go @@ -0,0 +1,316 @@ +package service + +// 利润控制请求路径矩阵测试:证明所有文本调度路径都经过利润准入过滤, +// 且任何 fallback 都不能把已排除账号重新放回候选。 +// 覆盖:高级调度器候选池(openai_profit_control_test.go)、legacy 引擎、 +// previous_response WSv2 粘连(跳过复用但保留绑定 + 倍率恢复重粘连)、 +// failover 排除不回收、抢槽后终检、倍率恢复重新准入、 +// 用户覆盖倍率 D、composite 计费分组与调度分组分离。 + +import ( + "context" + "errors" + "testing" + "time" + + "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/Wei-Shaw/sub2api/internal/pkg/ctxkey" + "github.com/stretchr/testify/require" +) + +func profitControlWSAccount(id int64, rate float64, now time.Time) Account { + account := upstreamCostTestAccount(id, UpstreamBillingProbeStatusOK, rate, now.Add(-time.Minute), 30*time.Minute) + profitControlTestAccountWithRate(account, rate) + account.Status = StatusActive + account.Schedulable = true + account.Concurrency = 2 + account.Extra["openai_apikey_responses_websockets_v2_enabled"] = true + return *account +} + +// previous_response_id 粘连:利润不合格 → 跳过复用但不删绑定;倍率恢复 → 重新粘连。 +func TestProfitControl_PreviousResponseStickyVetoKeepsBinding(t *testing.T) { + ctx := profitControlTestCtx(profitControlTestGroup(23, 0.5, 0)) + groupID := int64(23) + now := time.Now() + expensive := profitControlWSAccount(31, 0.8, now) + + cache := &stubGatewayCache{} + store := NewOpenAIWSStateStore(cache) + svc := &OpenAIGatewayService{ + accountRepo: stubOpenAIAccountRepo{accounts: []Account{expensive}}, + cache: cache, + cfg: newOpenAIWSV2TestConfig(), + concurrencyService: NewConcurrencyService(stubConcurrencyCache{}), + openaiWSStateStore: store, + } + require.NoError(t, store.BindResponseAccount(ctx, groupID, "resp_profit", expensive.ID, time.Hour)) + + selection, err := svc.SelectAccountByPreviousResponseID(ctx, &groupID, "resp_profit", "gpt-5.1", nil, false) + require.NoError(t, err) + require.Nil(t, selection, "上游倍率 0.8 超过阈值 0.5 的账号不应继续命中 previous_response_id 粘连") + + // 利润不合格与 quota auto-pause 同为暂时状态:绑定必须保留。 + boundAccountID, getErr := store.GetResponseAccount(ctx, groupID, "resp_profit") + require.NoError(t, getErr) + require.Equal(t, expensive.ID, boundAccountID) + + // 上游倍率回落(探测刷新)后同一绑定重新可用。 + recovered := profitControlWSAccount(31, 0.3, time.Now()) + svc.accountRepo = stubOpenAIAccountRepo{accounts: []Account{recovered}} + selection, err = svc.SelectAccountByPreviousResponseID(ctx, &groupID, "resp_profit", "gpt-5.1", nil, false) + require.NoError(t, err) + require.NotNil(t, selection) + require.Equal(t, recovered.ID, selection.Account.ID) + if selection.ReleaseFunc != nil { + selection.ReleaseFunc() + } +} + +// legacy 引擎(高级调度器关闭):候选过滤、全排除错误语义与既有语义一致。 +func TestProfitControl_LegacyEngineFiltersCandidates(t *testing.T) { + now := time.Now() + cheap := upstreamCostTestAccount(41, UpstreamBillingProbeStatusOK, 0.3, now.Add(-time.Minute), 30*time.Minute) + expensive := upstreamCostTestAccount(42, UpstreamBillingProbeStatusOK, 0.8, now.Add(-time.Minute), 30*time.Minute) + profitControlTestAccountWithRate(cheap, 0.3) + profitControlTestAccountWithRate(expensive, 0.8) + for _, account := range []*Account{cheap, expensive} { + account.Status = StatusActive + account.Schedulable = true + account.Concurrency = 2 + } + svc := &OpenAIGatewayService{ + accountRepo: stubOpenAIAccountRepo{accounts: []Account{*cheap, *expensive}}, + cfg: &config.Config{}, + rateLimitService: newOpenAIAdvancedSchedulerRateLimitService("false"), + concurrencyService: NewConcurrencyService(stubConcurrencyCache{}), + } + groupID := int64(7) + + t.Run("legacy path only admits profitable accounts", func(t *testing.T) { + ctx := profitControlTestCtx(profitControlTestGroup(groupID, 0.5, 0)) + for i := 0; i < 5; i++ { + selection, _, err := svc.SelectAccountWithScheduler(ctx, &groupID, "", "", "gpt-test", nil, OpenAIUpstreamTransportAny, false) + require.NoError(t, err) + require.NotNil(t, selection) + require.Equal(t, cheap.ID, selection.Account.ID) + if selection.ReleaseFunc != nil { + selection.ReleaseFunc() + } + } + }) + + t.Run("legacy path all excluded returns standard error", func(t *testing.T) { + ctx := profitControlTestCtx(profitControlTestGroup(groupID, 0.7, 0.1)) + selection, _, err := svc.SelectAccountWithScheduler(ctx, &groupID, "", "", "gpt-test", nil, OpenAIUpstreamTransportAny, false) + require.Nil(t, selection) + require.Error(t, err) + require.True(t, errors.Is(err, ErrNoAvailableAccounts)) + }) + + t.Run("legacy path keeps official behavior when gate disabled", func(t *testing.T) { + group := profitControlTestGroup(groupID, 0.7, 0.1) + group.ProfitControlEnabled = false + selection, _, err := svc.SelectAccountWithScheduler(profitControlTestCtx(group), &groupID, "", "", "gpt-test", nil, OpenAIUpstreamTransportAny, false) + require.NoError(t, err) + require.NotNil(t, selection) + if selection.ReleaseFunc != nil { + selection.ReleaseFunc() + } + }) +} + +// failover:可盈利账号因失败被排除后,剩余不合格账号不得被"放回"候选。 +func TestProfitControl_FailoverDoesNotReadmitExcluded(t *testing.T) { + resetOpenAIAdvancedSchedulerSettingCacheForTest() + defer resetOpenAIAdvancedSchedulerSettingCacheForTest() + + now := time.Now() + cheap := upstreamCostTestAccount(51, UpstreamBillingProbeStatusOK, 0.3, now.Add(-time.Minute), 30*time.Minute) + expensive := upstreamCostTestAccount(52, UpstreamBillingProbeStatusOK, 0.8, now.Add(-time.Minute), 30*time.Minute) + profitControlTestAccountWithRate(cheap, 0.3) + profitControlTestAccountWithRate(expensive, 0.8) + for _, account := range []*Account{cheap, expensive} { + account.Status = StatusActive + account.Schedulable = true + account.Concurrency = 2 + } + cache := &upstreamCostTrackingConcurrencyCache{loadMap: map[int64]*AccountLoadInfo{ + cheap.ID: {AccountID: cheap.ID}, + expensive.ID: {AccountID: expensive.ID}, + }} + svc := &OpenAIGatewayService{ + accountRepo: schedulerTestOpenAIAccountRepo{accounts: []Account{*cheap, *expensive}}, + cfg: &config.Config{}, + rateLimitService: newOpenAIAdvancedSchedulerRateLimitService("true"), + concurrencyService: NewConcurrencyService(cache), + } + groupID := int64(7) + ctx := profitControlTestCtx(profitControlTestGroup(groupID, 0.5, 0)) + + // 模拟 failover:上一轮失败的 cheap 已进入 excludedIDs,仅剩 expensive 不合格。 + excluded := map[int64]struct{}{cheap.ID: {}} + selection, _, err := svc.SelectAccountWithScheduler(ctx, &groupID, "", "", "gpt-test", excluded, OpenAIUpstreamTransportAny, false) + require.Nil(t, selection, "failover 后不得回收利润不合格账号") + require.Error(t, err) + require.True(t, errors.Is(err, ErrNoAvailableAccounts)) + require.Contains(t, err.Error(), openAIProfitFilterReasonThreshold+"=1") +} + +// 抢槽后终检:候选构建后才变得不合格的账号(状态竞态)在取得槽位前被拦截。 +func TestProfitControl_PostSlotRecheckVetoes(t *testing.T) { + now := time.Now() + expensive := upstreamCostTestAccount(61, UpstreamBillingProbeStatusOK, 0.8, now.Add(-time.Minute), 30*time.Minute) + profitControlTestAccountWithRate(expensive, 0.8) + expensive.Status = StatusActive + expensive.Schedulable = true + expensive.Concurrency = 2 + + cache := &upstreamCostTrackingConcurrencyCache{} + scheduler := &defaultOpenAIAccountScheduler{service: &OpenAIGatewayService{ + concurrencyService: NewConcurrencyService(cache), + }} + ctx := context.WithValue(context.Background(), openAIProfitControlGateCtxKey{}, &openAIProfitControlGate{ + threshold: 0.5, + pricingAt: now, + }) + selectionOrder := []openAIAccountCandidateScore{{ + account: expensive, + loadInfo: &AccountLoadInfo{AccountID: expensive.ID}, + }} + + selection, _, err := scheduler.tryAcquireOpenAISelectionOrder(ctx, OpenAIAccountScheduleRequest{Platform: PlatformOpenAI}, selectionOrder) + require.NoError(t, err) + require.Nil(t, selection, "抢槽终检必须拦截候选构建后才不合格的账号") + require.Equal(t, cache.totalAcquires(), cache.releaseCount(expensive.ID), "被拦截账号不得泄漏并发槽位") +} + +// 倍率恢复:探测刷新回落到阈值内后,此前被排除的账号重新参与调度。 +func TestProfitControl_RateRecoveryReadmitsAccount(t *testing.T) { + resetOpenAIAdvancedSchedulerSettingCacheForTest() + defer resetOpenAIAdvancedSchedulerSettingCacheForTest() + + now := time.Now() + expensive := upstreamCostTestAccount(71, UpstreamBillingProbeStatusOK, 0.8, now.Add(-time.Minute), 30*time.Minute) + profitControlTestAccountWithRate(expensive, 0.8) + expensive.Status = StatusActive + expensive.Schedulable = true + expensive.Concurrency = 2 + cache := &upstreamCostTrackingConcurrencyCache{loadMap: map[int64]*AccountLoadInfo{ + expensive.ID: {AccountID: expensive.ID}, + }} + svc := &OpenAIGatewayService{ + accountRepo: schedulerTestOpenAIAccountRepo{accounts: []Account{*expensive}}, + cfg: &config.Config{}, + rateLimitService: newOpenAIAdvancedSchedulerRateLimitService("true"), + concurrencyService: NewConcurrencyService(cache), + } + groupID := int64(7) + ctx := profitControlTestCtx(profitControlTestGroup(groupID, 0.5, 0)) + + selection, _, err := svc.SelectAccountWithScheduler(ctx, &groupID, "", "", "gpt-test", nil, OpenAIUpstreamTransportAny, false) + require.Nil(t, selection) + require.Error(t, err) + + // 同步/手工写回:账号倍率回落到 0.3(阈值 0.5 内)后自动恢复参与。 + recovered := upstreamCostTestAccount(71, UpstreamBillingProbeStatusOK, 0.3, time.Now().Add(-time.Minute), 30*time.Minute) + profitControlTestAccountWithRate(recovered, 0.3) + recovered.Status = StatusActive + recovered.Schedulable = true + recovered.Concurrency = 2 + svc.accountRepo = schedulerTestOpenAIAccountRepo{accounts: []Account{*recovered}} + + selection, _, err = svc.SelectAccountWithScheduler(ctx, &groupID, "", "", "gpt-test", nil, OpenAIUpstreamTransportAny, false) + require.NoError(t, err) + require.NotNil(t, selection) + require.Equal(t, recovered.ID, selection.Account.ID) + if selection.ReleaseFunc != nil { + selection.ReleaseFunc() + } +} + +type profitControlUserRateRepo struct { + UserGroupRateRepository + rate *float64 +} + +func (r profitControlUserRateRepo) GetByUserAndGroup(context.Context, int64, int64) (*float64, error) { + return r.rate, nil +} + +// D 必须取请求用户的真实倍率:有用户覆盖时用覆盖值,绝不退回分组默认。 +func TestProfitControl_GateUsesUserOverrideRate(t *testing.T) { + override := 0.5 + svc := &OpenAIGatewayService{ + userGroupRateResolver: newUserGroupRateResolver( + profitControlUserRateRepo{rate: &override}, nil, time.Minute, nil, "test.profit", + ), + } + groupID := int64(7) + group := profitControlTestGroup(groupID, 0, 0) + group.RateMultiplier = 2.0 + + ctx := context.WithValue(profitControlTestCtx(group), ctxkey.UserID, int64(42)) + gate := svc.resolveOpenAIProfitControlGate(ctx, &groupID) + require.NotNil(t, gate) + require.InDelta(t, 0.5, gate.threshold, 1e-12, "阈值必须基于用户覆盖倍率 0.5,而不是分组默认 2.0") + + // 无用户身份(内部调用)时按分组默认倍率计算。 + gate = svc.resolveOpenAIProfitControlGate(profitControlTestCtx(group), &groupID) + require.NotNil(t, gate) + require.InDelta(t, 2.0, gate.threshold, 1e-12) +} + +type profitControlGroupRepo struct { + GroupRepository + group *Group +} + +func (r profitControlGroupRepo) GetByID(context.Context, int64) (*Group, error) { + return r.group, nil +} + +// composite 路由:门配置取被调度成员分组,D 取请求真实计费分组(ctx 认证分组)。 +func TestProfitControl_CompositeUsesBillingGroupRate(t *testing.T) { + memberGroupID := int64(7) + memberGroup := profitControlTestGroup(memberGroupID, 0.5, 0) + memberGroup.RateMultiplier = 99 // 若 D 误取成员分组倍率,阈值会是 49.5 + + billingGroup := &Group{ + ID: 1001, + Platform: PlatformComposite, + Status: StatusActive, + Hydrated: true, + RateMultiplier: 1.0, + } + svc := &OpenAIGatewayService{ + schedulerSnapshot: &SchedulerSnapshotService{groupRepo: profitControlGroupRepo{group: memberGroup}}, + } + + ctx := profitControlTestCtx(billingGroup) + gate := svc.resolveOpenAIProfitControlGate(ctx, &memberGroupID) + require.NotNil(t, gate) + require.InDelta(t, 0.5, gate.threshold, 1e-12, "D 必须来自计费分组(composite 父分组)倍率 1.0") +} + +// legacy 引擎与 DB recheck 共用的资格判定直接覆盖利润门。 +func TestProfitControl_EligibilityFunctionVetoes(t *testing.T) { + now := time.Now() + gateCtx := context.WithValue(context.Background(), openAIProfitControlGateCtxKey{}, &openAIProfitControlGate{ + threshold: 0.5, + pricingAt: now, + }) + cheap := upstreamCostTestAccount(81, UpstreamBillingProbeStatusOK, 0.3, now.Add(-time.Minute), 30*time.Minute) + expensive := upstreamCostTestAccount(82, UpstreamBillingProbeStatusOK, 0.8, now.Add(-time.Minute), 30*time.Minute) + profitControlTestAccountWithRate(cheap, 0.3) + profitControlTestAccountWithRate(expensive, 0.8) + for _, account := range []*Account{cheap, expensive} { + account.Status = StatusActive + account.Schedulable = true + } + + require.True(t, isOpenAICompatibleAccountEligibleForRequest(gateCtx, cheap, PlatformOpenAI, "", false, "")) + require.False(t, isOpenAICompatibleAccountEligibleForRequest(gateCtx, expensive, PlatformOpenAI, "", false, "")) + // 无门时保持既有行为。 + require.True(t, isOpenAICompatibleAccountEligibleForRequest(context.Background(), expensive, PlatformOpenAI, "", false, "")) +} diff --git a/backend/internal/service/openai_profit_control_pricing_test.go b/backend/internal/service/openai_profit_control_pricing_test.go new file mode 100644 index 0000000000..f155a1f5cb --- /dev/null +++ b/backend/internal/service/openai_profit_control_pricing_test.go @@ -0,0 +1,201 @@ +package service + +// 请求级定价与利润门回归:请求级 pricingAt 定价上下文、门复用(failover 阈值稳定)、 +// 生图意图跳门、U 使用账号倍率且与探测新鲜度解耦、 +// 用量记录定价时刻取值。 + +import ( + "context" + "testing" + "time" + + "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/Wei-Shaw/sub2api/internal/pkg/timezone" + "github.com/stretchr/testify/require" +) + +// WithOpenAIRequestPricingContext:装门 + 固定 pricingAt;抑制标记跳门且防御性 +// 装门无法把门加回来。 +func TestProfitControl_RequestPricingContext(t *testing.T) { + svc := &OpenAIGatewayService{} + groupID := int64(61) + now := time.Now() + expensive := upstreamCostTestAccount(1, UpstreamBillingProbeStatusOK, 0.8, now.Add(-time.Minute), 30*time.Minute) + profitControlTestAccountWithRate(expensive, 0.8) + + t.Run("installs gate and pricing instant", func(t *testing.T) { + base := profitControlTestCtx(profitControlTestGroup(groupID, 0.5, 0)) + ctx, pricingAt := svc.WithOpenAIRequestPricingContext(base, &groupID, false) + require.False(t, pricingAt.IsZero()) + require.Equal(t, pricingAt, OpenAIPricingAtFromContext(ctx)) + vetoed, reason := OpenAIProfitControlVeto(ctx, expensive) + require.True(t, vetoed) + require.Equal(t, openAIProfitFilterReasonThreshold, reason) + }) + + t.Run("image intent suppresses gate everywhere", func(t *testing.T) { + base := profitControlTestCtx(profitControlTestGroup(groupID, 0.5, 0)) + ctx, pricingAt := svc.WithOpenAIRequestPricingContext(base, &groupID, true) + require.False(t, pricingAt.IsZero(), "跳门时 pricingAt 仍需固定供计费共用") + vetoed, _ := OpenAIProfitControlVeto(ctx, expensive) + require.False(t, vetoed) + // service 层防御性装门也必须被抑制标记挡住。 + reCtx := svc.withOpenAIProfitControlGate(ctx, &groupID) + vetoed, _ = OpenAIProfitControlVeto(reCtx, expensive) + require.False(t, vetoed) + }) +} + +// failover 重入复用同一门:请求中途分组配置变化不得改变本请求阈值。 +func TestProfitControl_GateReuseKeepsThresholdAcrossFailover(t *testing.T) { + svc := &OpenAIGatewayService{} + groupID := int64(62) + group := profitControlTestGroup(groupID, 0.5, 0) + ctx := svc.withOpenAIProfitControlGate(profitControlTestCtx(group), &groupID) + gate, ok := ctx.Value(openAIProfitControlGateCtxKey{}).(*openAIProfitControlGate) + require.True(t, ok) + require.InDelta(t, 0.5, gate.threshold, 1e-12) + + // 模拟请求进行中管理员改配置(ctx 分组为同一指针,与 auth 快照语义一致)。 + group.ProfitMinMargin = 0.9 + reCtx := svc.withOpenAIProfitControlGate(ctx, &groupID) + reGate, ok := reCtx.Value(openAIProfitControlGateCtxKey{}).(*openAIProfitControlGate) + require.True(t, ok) + require.Same(t, gate, reGate, "failover 重入必须复用同一门,阈值不得中途变化") + + // 换分组(composite/模型路由成员调度)重新解析;成员分组无门时必须清除 + // 父分组门,阈值不得跨组泄漏。 + otherID := int64(63) + otherCtx := svc.withOpenAIProfitControlGate(reCtx, &otherID) + otherGate, _ := otherCtx.Value(openAIProfitControlGateCtxKey{}).(*openAIProfitControlGate) + require.Nil(t, otherGate, "成员分组未启用利润控制时父分组门必须清除") + now := time.Now() + expensive := upstreamCostTestAccount(8, UpstreamBillingProbeStatusOK, 0.9, now.Add(-time.Minute), 30*time.Minute) + vetoed, _ := openAIProfitControlVetoReason(otherCtx, expensive) + require.False(t, vetoed) +} + +// D 固定在 pricingAt:高峰因子按请求开始时刻计算,与"当前时刻"无关。 +func TestProfitControl_PricingAtFixesDownstreamPeakFactor(t *testing.T) { + svc := &OpenAIGatewayService{} + groupID := int64(64) + group := profitControlTestGroup(groupID, 0, 0) + group.SubscriptionType = SubscriptionTypeSubscription + group.PeakRateEnabled = true + group.PeakRateMultiplier = 3.0 + + pricingAt := time.Date(2026, time.January, 15, 8, 30, 0, 0, timezone.Location()) + outsideWindow := time.Date(2026, time.January, 15, 10, 30, 0, 0, timezone.Location()) + group.PeakStart = "08:00" + group.PeakEnd = "09:00" + require.Equal(t, 1.0, group.PeakMultiplierAt(outsideWindow), "构造前提:对照时刻不在窗口内") + require.Equal(t, 3.0, group.PeakMultiplierAt(pricingAt), "构造前提:pricingAt 在窗口内") + + ctx := context.WithValue(profitControlTestCtx(group), openAIPricingAtCtxKey{}, pricingAt) + gate := svc.resolveOpenAIProfitControlGate(ctx, &groupID) + require.NotNil(t, gate) + require.InDelta(t, 3.0, gate.threshold, 1e-9, "阈值必须用 pricingAt 时刻的高峰因子(1.0×3.0×(1-0))") + require.Equal(t, pricingAt, gate.pricingAt) +} + +// U 只取账号倍率:探测快照内容和新鲜度不再直接参与利润判断。 +func TestProfitControl_UsesAccountRateInsteadOfProbeSnapshot(t *testing.T) { + gate := &openAIProfitControlGate{threshold: 0.5, pricingAt: time.Now().Add(-12 * time.Hour)} + ctx := context.WithValue(context.Background(), openAIProfitControlGateCtxKey{}, gate) + account := upstreamCostTestAccount(9, UpstreamBillingProbeStatusOK, 0.1, time.Now().Add(-3*time.Hour), 30*time.Minute) + profitControlTestAccountWithRate(account, 0.8) + vetoed, reason := openAIProfitControlVetoReason(ctx, account) + require.True(t, vetoed) + require.Equal(t, openAIProfitFilterReasonThreshold, reason) +} + +// 显式生图意图(requiredCapability=Responses)在唯一调度入口跳门(图片边界不装门)。 +func TestProfitControl_ResponsesImageIntentSkipsGateAtScheduler(t *testing.T) { + now := time.Now() + expensive := upstreamCostTestAccount(51, UpstreamBillingProbeStatusOK, 0.8, now.Add(-time.Minute), 30*time.Minute) + expensive.Status = StatusActive + expensive.Schedulable = true + expensive.Concurrency = 2 + svc := &OpenAIGatewayService{ + accountRepo: stubOpenAIAccountRepo{accounts: []Account{*expensive}}, + cfg: &config.Config{}, + rateLimitService: newOpenAIAdvancedSchedulerRateLimitService("true"), + concurrencyService: NewConcurrencyService(stubConcurrencyCache{}), + } + groupID := int64(77) + ctx := profitControlTestCtx(profitControlTestGroup(groupID, 0.5, 0)) + + _, _, err := svc.SelectAccountWithSchedulerForCapability(ctx, &groupID, "", "", "gpt-test", nil, OpenAIUpstreamTransportAny, OpenAIEndpointCapabilityChatCompletions, false, false, true) + require.ErrorIs(t, err, ErrNoAvailableAccounts, "文本能力必须过利润门") + + selection, _, err := svc.SelectAccountWithSchedulerForCapability(ctx, &groupID, "", "", "gpt-test", nil, OpenAIUpstreamTransportAny, OpenAIEndpointCapabilityResponses, false, false, true) + require.NoError(t, err, "生图意图不装门,保持既有调度") + require.NotNil(t, selection) + require.Equal(t, expensive.ID, selection.Account.ID) + if selection.ReleaseFunc != nil { + selection.ReleaseFunc() + } +} + +// 账号倍率缺失一律视为非法保守拒绝;手工或同步维护了倍率的任意账号类型都按 +// 同一阈值判断(OAuth 与 API Key 无差别)。 +func TestProfitControl_AccountRateSemantics(t *testing.T) { + now := time.Now() + missing := upstreamCostTestOAuthAccount(2) + manualOAuth := profitControlTestAccountWithRate(upstreamCostTestOAuthAccount(3), 0.3) + expensive := profitControlTestAccountWithRate(upstreamCostTestAccount(4, UpstreamBillingProbeStatusOK, 0.1, now.Add(-3*time.Hour), 30*time.Minute), 0.8) + + group := profitControlTestGroup(77, 0.5, 0) + group.RateMultiplier = 1 + base := context.WithValue(profitControlTestCtx(group), openAIPricingAtCtxKey{}, now) + gate := (&OpenAIGatewayService{}).resolveOpenAIProfitControlGate(base, &group.ID) + require.NotNil(t, gate) + gateCtx := context.WithValue(base, openAIProfitControlGateCtxKey{}, gate) + + vetoed, reason := openAIProfitControlVetoReason(gateCtx, missing) + require.True(t, vetoed, "缺失账号倍率必须保守拒绝") + require.Equal(t, openAIProfitFilterReasonInvalidAccountRate, reason) + + vetoed, _ = openAIProfitControlVetoReason(gateCtx, manualOAuth) + require.False(t, vetoed, "手工维护的 OAuth 倍率应正常准入") + + vetoed, reason = openAIProfitControlVetoReason(gateCtx, expensive) + require.True(t, vetoed) + require.Equal(t, openAIProfitFilterReasonThreshold, reason) +} + +// 用量记录定价时刻:优先请求级 PricingAt,未装配回退记录时刻。 +func TestOpenAIUsagePricingAt(t *testing.T) { + fixed := time.Now().Add(-2 * time.Hour) + require.Equal(t, fixed, openAIUsagePricingAt(&OpenAIRecordUsageInput{PricingAt: fixed})) + fallback := openAIUsagePricingAt(&OpenAIRecordUsageInput{}) + require.WithinDuration(t, timezone.Now(), fallback, 5*time.Second) + require.WithinDuration(t, timezone.Now(), openAIUsagePricingAt(nil), 5*time.Second) +} + +func TestOpenAIProfitControlStickyBindingOccursOnlyAfterTerminalAdmission(t *testing.T) { + groupID := int64(81) + expensiveID := int64(901) + cheapID := int64(902) + const sessionHash = "profit-sticky" + const cacheKey = "openai:" + sessionHash + cache := &schedulerTestGatewayCache{ + sessionBindings: map[string]int64{cacheKey: expensiveID}, + } + svc := &OpenAIGatewayService{cache: cache} + ctx := context.WithValue(context.Background(), openAIProfitControlGateCtxKey{}, &openAIProfitControlGate{ + groupID: groupID, + platform: PlatformOpenAI, + threshold: 0.5, + }) + + require.NoError(t, svc.bindOpenAIStickySessionDuringSelection(ctx, &groupID, sessionHash, cheapID)) + require.Equal(t, expensiveID, cache.sessionBindings[cacheKey], "选号阶段不得覆盖原粘性绑定") + + require.NoError(t, svc.BindStickySessionAfterProfitAdmission(ctx, &groupID, sessionHash, cheapID)) + require.Equal(t, expensiveID, cache.sessionBindings[cacheKey], "终检通过的 fallback 账号不得覆盖原粘性绑定") + + cache.sessionBindings[cacheKey] = 0 + require.NoError(t, svc.BindStickySessionAfterProfitAdmission(ctx, &groupID, sessionHash, cheapID)) + require.Equal(t, cheapID, cache.sessionBindings[cacheKey], "无既有绑定时应在终检通过后建立粘性") +} diff --git a/backend/internal/service/openai_profit_control_test.go b/backend/internal/service/openai_profit_control_test.go new file mode 100644 index 0000000000..c0fb51e64d --- /dev/null +++ b/backend/internal/service/openai_profit_control_test.go @@ -0,0 +1,311 @@ +package service + +import ( + "context" + "errors" + "math" + + "testing" + "time" + + "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/Wei-Shaw/sub2api/internal/pkg/ctxkey" + "github.com/Wei-Shaw/sub2api/internal/pkg/timezone" + "github.com/stretchr/testify/require" +) + +func profitControlTestGroup(id int64, margin, buffer float64) *Group { + return &Group{ + ID: id, + Platform: PlatformOpenAI, + Status: StatusActive, + Hydrated: true, + RateMultiplier: 1.0, + SubscriptionType: SubscriptionTypeStandard, + ProfitControlEnabled: true, + ProfitMinMargin: margin, + ProfitSafetyBuffer: buffer, + } +} + +func profitControlTestCtx(group *Group) context.Context { + return context.WithValue(context.Background(), ctxkey.Group, group) +} + +func profitControlTestAccountWithRate(account *Account, rate float64) *Account { + account.RateMultiplier = &rate + return account +} + +func TestResolveOpenAIProfitControlGate(t *testing.T) { + svc := &OpenAIGatewayService{} + groupID := int64(7) + + t.Run("nil group id yields no gate", func(t *testing.T) { + require.Nil(t, svc.resolveOpenAIProfitControlGate(context.Background(), nil)) + }) + + t.Run("no ctx group and no snapshot yields no gate", func(t *testing.T) { + require.Nil(t, svc.resolveOpenAIProfitControlGate(context.Background(), &groupID)) + }) + + t.Run("disabled group yields no gate", func(t *testing.T) { + group := profitControlTestGroup(groupID, 0.3, 0) + group.ProfitControlEnabled = false + require.Nil(t, svc.resolveOpenAIProfitControlGate(profitControlTestCtx(group), &groupID)) + }) + + t.Run("non openai or grok platform yields no gate even if enabled", func(t *testing.T) { + group := profitControlTestGroup(groupID, 0.3, 0) + group.Platform = PlatformAnthropic + require.Nil(t, svc.resolveOpenAIProfitControlGate(profitControlTestCtx(group), &groupID)) + }) + + t.Run("grok group routed through openai handler installs gate", func(t *testing.T) { + group := profitControlTestGroup(groupID, 0.3, 0.05) + group.Platform = PlatformGrok + group.RateMultiplier = 0.5 + gate := svc.resolveOpenAIProfitControlGate(profitControlTestCtx(group), &groupID) + require.NotNil(t, gate) + require.Equal(t, PlatformGrok, gate.platform) + require.InDelta(t, 0.5*(1-0.35), gate.threshold, 1e-12) + }) + + t.Run("ctx group id mismatch without snapshot yields no gate", func(t *testing.T) { + group := profitControlTestGroup(groupID+1, 0.3, 0) + require.Nil(t, svc.resolveOpenAIProfitControlGate(profitControlTestCtx(group), &groupID)) + }) + + t.Run("threshold composes margin and buffer from downstream rate", func(t *testing.T) { + group := profitControlTestGroup(groupID, 0.3, 0.05) + group.RateMultiplier = 2.0 + gate := svc.resolveOpenAIProfitControlGate(profitControlTestCtx(group), &groupID) + require.NotNil(t, gate) + require.InDelta(t, 2.0*(1-0.35), gate.threshold, 1e-12) + require.Equal(t, PlatformOpenAI, gate.platform) + require.False(t, gate.pricingAt.IsZero()) + require.Equal(t, groupID, gate.groupID) + }) + + t.Run("threshold applies peak factor exactly like billing", func(t *testing.T) { + group := profitControlTestGroup(groupID, 0.5, 0) + group.SubscriptionType = SubscriptionTypeSubscription + group.PeakRateEnabled = true + group.PeakStart = "00:00" + group.PeakEnd = "23:59" + group.PeakRateMultiplier = 3.0 + gate := svc.resolveOpenAIProfitControlGate(profitControlTestCtx(group), &groupID) + require.NotNil(t, gate) + expected := group.RateMultiplier * group.PeakMultiplierAt(timezone.Now()) * 0.5 + require.InDelta(t, expected, gate.threshold, 1e-9) + require.Equal(t, PlatformOpenAI, gate.platform) + }) +} + +func TestOpenAIProfitControlVetoReason(t *testing.T) { + now := time.Now() + gateCtx := func(threshold float64) context.Context { + return context.WithValue(context.Background(), openAIProfitControlGateCtxKey{}, &openAIProfitControlGate{ + threshold: threshold, + pricingAt: now, + }) + } + + t.Run("no gate admits everything", func(t *testing.T) { + vetoed, reason := openAIProfitControlVetoReason(context.Background(), upstreamCostTestOAuthAccount(1)) + require.False(t, vetoed) + require.Empty(t, reason) + }) + + t.Run("fresh rate below threshold admits", func(t *testing.T) { + account := profitControlTestAccountWithRate(upstreamCostTestAccount(1, UpstreamBillingProbeStatusOK, 99, now.Add(-time.Minute), 30*time.Minute), 0.5) + vetoed, _ := openAIProfitControlVetoReason(gateCtx(0.7), account) + require.False(t, vetoed) + }) + + t.Run("rate exactly at threshold admits via epsilon", func(t *testing.T) { + account := profitControlTestAccountWithRate(upstreamCostTestAccount(1, UpstreamBillingProbeStatusOK, 99, now.Add(-time.Minute), 30*time.Minute), 0.7) + vetoed, _ := openAIProfitControlVetoReason(gateCtx(0.7), account) + require.False(t, vetoed) + }) + + t.Run("rate within float noise above threshold admits", func(t *testing.T) { + account := profitControlTestAccountWithRate(upstreamCostTestAccount(1, UpstreamBillingProbeStatusOK, 99, now.Add(-time.Minute), 30*time.Minute), 0.7+1e-12) + vetoed, _ := openAIProfitControlVetoReason(gateCtx(0.7), account) + require.False(t, vetoed) + }) + + t.Run("rate above threshold is vetoed", func(t *testing.T) { + account := profitControlTestAccountWithRate(upstreamCostTestAccount(1, UpstreamBillingProbeStatusOK, 0.1, now.Add(-time.Minute), 30*time.Minute), 0.8) + vetoed, reason := openAIProfitControlVetoReason(gateCtx(0.7), account) + require.True(t, vetoed) + require.Equal(t, openAIProfitFilterReasonThreshold, reason) + }) + + t.Run("zero threshold only admits free upstream", func(t *testing.T) { + free := profitControlTestAccountWithRate(upstreamCostTestAccount(1, UpstreamBillingProbeStatusOK, 99, now.Add(-time.Minute), 30*time.Minute), 0) + vetoed, _ := openAIProfitControlVetoReason(gateCtx(0), free) + require.False(t, vetoed) + paid := profitControlTestAccountWithRate(upstreamCostTestAccount(2, UpstreamBillingProbeStatusOK, 0, now.Add(-time.Minute), 30*time.Minute), 0.01) + vetoed, reason := openAIProfitControlVetoReason(gateCtx(0), paid) + require.True(t, vetoed) + require.Equal(t, openAIProfitFilterReasonThreshold, reason) + }) + + t.Run("missing account rate is invalid", func(t *testing.T) { + vetoed, reason := openAIProfitControlVetoReason(gateCtx(0.7), upstreamCostTestOAuthAccount(1)) + require.True(t, vetoed) + require.Equal(t, openAIProfitFilterReasonInvalidAccountRate, reason) + }) + + t.Run("oauth account with manual rate is priceable", func(t *testing.T) { + account := profitControlTestAccountWithRate(upstreamCostTestOAuthAccount(1), 0.2) + vetoed, _ := openAIProfitControlVetoReason(gateCtx(0.7), account) + require.False(t, vetoed) + }) + + t.Run("stale probe does not affect manual account rate", func(t *testing.T) { + account := profitControlTestAccountWithRate(upstreamCostTestAccount(1, UpstreamBillingProbeStatusOK, 99, now.Add(-3*time.Hour), 30*time.Minute), 0.1) + vetoed, _ := openAIProfitControlVetoReason(gateCtx(0.7), account) + require.False(t, vetoed) + }) + + t.Run("negative and non-finite rates are invalid", func(t *testing.T) { + for _, rate := range []float64{-1, math.NaN(), math.Inf(1)} { + account := profitControlTestAccountWithRate(upstreamCostTestOAuthAccount(1), rate) + vetoed, reason := openAIProfitControlVetoReason(gateCtx(0.7), account) + require.True(t, vetoed) + require.Equal(t, openAIProfitFilterReasonInvalidAccountRate, reason) + } + }) +} + +func TestProfitControlSchedulerFiltersCandidates(t *testing.T) { + resetOpenAIAdvancedSchedulerSettingCacheForTest() + defer resetOpenAIAdvancedSchedulerSettingCacheForTest() + + now := time.Now() + cheap := upstreamCostTestAccount(1, UpstreamBillingProbeStatusOK, 0.3, now.Add(-time.Minute), 30*time.Minute) + expensive := upstreamCostTestAccount(2, UpstreamBillingProbeStatusOK, 0.8, now.Add(-time.Minute), 30*time.Minute) + oauth := upstreamCostTestOAuthAccount(3) + profitControlTestAccountWithRate(cheap, 0.3) + profitControlTestAccountWithRate(expensive, 0.8) + for _, account := range []*Account{cheap, expensive, oauth} { + account.Status = StatusActive + account.Schedulable = true + account.Concurrency = 5 + } + cache := &upstreamCostTrackingConcurrencyCache{loadMap: map[int64]*AccountLoadInfo{ + cheap.ID: {AccountID: cheap.ID}, + expensive.ID: {AccountID: expensive.ID}, + oauth.ID: {AccountID: oauth.ID}, + }} + cfg := &config.Config{} + svc := &OpenAIGatewayService{ + accountRepo: schedulerTestOpenAIAccountRepo{accounts: []Account{*cheap, *expensive, *oauth}}, + cfg: cfg, + rateLimitService: newOpenAIAdvancedSchedulerRateLimitService("true"), + concurrencyService: NewConcurrencyService(cache), + } + groupID := int64(7) + + t.Run("unprofitable and invalid-rate accounts never win", func(t *testing.T) { + // margin 0.5 → 阈值 0.5:expensive(0.8) 超阈值、oauth 倍率缺失,仅 cheap 可选。 + ctx := profitControlTestCtx(profitControlTestGroup(groupID, 0.5, 0)) + for i := 0; i < 5; i++ { + selection, _, err := svc.SelectAccountWithScheduler(ctx, &groupID, "", "", "gpt-test", nil, OpenAIUpstreamTransportAny, false) + require.NoError(t, err) + require.NotNil(t, selection) + require.Equal(t, cheap.ID, selection.Account.ID) + if selection.ReleaseFunc != nil { + selection.ReleaseFunc() + } + } + }) + + t.Run("all excluded surfaces standard no-available error with profit reasons", func(t *testing.T) { + // margin+buffer 0.8 → 阈值 0.2:cheap/expensive 超阈值,oauth 倍率非法。 + ctx := profitControlTestCtx(profitControlTestGroup(groupID, 0.7, 0.1)) + selection, _, err := svc.SelectAccountWithScheduler(ctx, &groupID, "", "", "gpt-test", nil, OpenAIUpstreamTransportAny, false) + require.Nil(t, selection) + require.Error(t, err) + require.True(t, errors.Is(err, ErrNoAvailableAccounts)) + require.Contains(t, err.Error(), openAIProfitFilterReasonThreshold+"=2") + require.Contains(t, err.Error(), openAIProfitFilterReasonInvalidAccountRate+"=1") + }) + + t.Run("manually rated oauth account is admitted", func(t *testing.T) { + // 阈值 0.2 排除两个 API Key;OAuth 手工倍率 0.1 可参与调度。 + profitControlTestAccountWithRate(oauth, 0.1) + svc.accountRepo = schedulerTestOpenAIAccountRepo{accounts: []Account{*cheap, *expensive, *oauth}} + ctx := profitControlTestCtx(profitControlTestGroup(groupID, 0.7, 0.1)) + selection, _, err := svc.SelectAccountWithScheduler(ctx, &groupID, "", "", "gpt-test", nil, OpenAIUpstreamTransportAny, false) + require.NoError(t, err) + require.NotNil(t, selection) + require.Equal(t, oauth.ID, selection.Account.ID) + if selection.ReleaseFunc != nil { + selection.ReleaseFunc() + } + }) + + t.Run("gate disabled keeps official behavior", func(t *testing.T) { + group := profitControlTestGroup(groupID, 0.7, 0.1) + group.ProfitControlEnabled = false + selection, _, err := svc.SelectAccountWithScheduler(profitControlTestCtx(group), &groupID, "", "", "gpt-test", nil, OpenAIUpstreamTransportAny, false) + require.NoError(t, err) + require.NotNil(t, selection) + if selection.ReleaseFunc != nil { + selection.ReleaseFunc() + } + }) +} + +func TestValidateProfitControlConfig(t *testing.T) { + require.NoError(t, ValidateProfitControlConfig(PlatformAnthropic, false, 0, 0)) + for _, platform := range []string{PlatformOpenAI, PlatformAnthropic, PlatformGemini, PlatformGrok, PlatformAntigravity} { + require.NoError(t, ValidateProfitControlConfig(platform, true, 0.3, 0.05)) + require.NoError(t, ValidateProfitControlConfig(platform, true, 0, 0)) + } + + require.Error(t, ValidateProfitControlConfig(PlatformComposite, true, 0.3, 0)) + require.Error(t, ValidateProfitControlConfig(PlatformOpenAI, true, -0.1, 0)) + require.Error(t, ValidateProfitControlConfig(PlatformOpenAI, true, 1.0, 0)) + require.Error(t, ValidateProfitControlConfig(PlatformOpenAI, true, 0, 1.0)) + require.Error(t, ValidateProfitControlConfig(PlatformOpenAI, true, 0.6, 0.4)) +} + +func TestNormalizeProfitControlConfig(t *testing.T) { + t.Run("unsupported platform resets everything", func(t *testing.T) { + enabled, margin, buffer := NormalizeProfitControlConfig(PlatformComposite, true, 0.3, 0.1) + require.False(t, enabled) + require.Zero(t, margin) + require.Zero(t, buffer) + }) + + t.Run("all five platforms retain configuration", func(t *testing.T) { + for _, platform := range []string{PlatformOpenAI, PlatformAnthropic, PlatformGemini, PlatformGrok, PlatformAntigravity} { + enabled, margin, buffer := NormalizeProfitControlConfig(platform, true, 0.3, 0.1) + require.True(t, enabled) + require.InDelta(t, 0.3, margin, 1e-12) + require.InDelta(t, 0.1, buffer, 1e-12) + } + }) + + t.Run("openai disabled keeps legal values and cleans dirty ones", func(t *testing.T) { + enabled, margin, buffer := NormalizeProfitControlConfig(PlatformOpenAI, false, 0.3, 0.05) + require.False(t, enabled) + require.InDelta(t, 0.3, margin, 1e-12) + require.InDelta(t, 0.05, buffer, 1e-12) + + _, margin, buffer = NormalizeProfitControlConfig(PlatformOpenAI, false, -1, 1.5) + require.Zero(t, margin) + require.Zero(t, buffer) + }) + + t.Run("openai enabled passes through for validation", func(t *testing.T) { + enabled, margin, buffer := NormalizeProfitControlConfig(PlatformOpenAI, true, 0.3, 0.05) + require.True(t, enabled) + require.InDelta(t, 0.3, margin, 1e-12) + require.InDelta(t, 0.05, buffer, 1e-12) + }) +} diff --git a/backend/internal/service/openai_ws_forwarder_support.go b/backend/internal/service/openai_ws_forwarder_support.go index 91b9ef5cb1..681d5d82c2 100644 --- a/backend/internal/service/openai_ws_forwarder_support.go +++ b/backend/internal/service/openai_ws_forwarder_support.go @@ -384,6 +384,9 @@ func (s *OpenAIGatewayService) SelectAccountByPreviousResponseID( excludedIDs map[int64]struct{}, requireCompact bool, ) (*AccountSelectionResult, error) { + // 分组利润控制:公共入口装门,保证不经 selectAccountWithScheduler + // 的调用方也无法绕过利润准入(scheduler 内部路径已在唯一调度入口装门)。 + ctx = s.withOpenAIProfitControlGate(ctx, groupID) return s.selectAccountByPreviousResponseIDForCapability(ctx, groupID, previousResponseID, requestedModel, excludedIDs, "", requireCompact) } @@ -509,6 +512,12 @@ func (s *OpenAIGatewayService) resolveAccountByPreviousResponseIDForCapability( if paused, _ := shouldAutoPauseOpenAIAccountByQuota(ctx, account); paused { return 0, nil, "", nil } + // 分组利润控制:与 quota auto-pause 同语义——利润不合格是暂时 + // 状态(上游倍率/高峰随时间变化),只跳过本次复用、落回普通调度,不删除 + // 绑定(倍率恢复后可继续按 previous_response_id 粘连)。 + if vetoed, _ := openAIProfitControlVetoReason(ctx, account); vetoed { + return 0, nil, "", nil + } if s.schedulerSnapshot != nil && s.accountRepo != nil { latest, latestErr := s.accountRepo.GetByID(ctx, account.ID) if latestErr != nil || latest == nil { @@ -532,6 +541,10 @@ func (s *OpenAIGatewayService) resolveAccountByPreviousResponseIDForCapability( if paused, _ := shouldAutoPauseOpenAIAccountByQuota(ctx, latest); paused { return 0, nil, "", nil } + // 利润门对最新账号状态复检一次,语义同上:跳过复用、不删绑定。 + if vetoed, _ := openAIProfitControlVetoReason(ctx, latest); vetoed { + return 0, nil, "", nil + } if s.isOpenAIAccountRequestRuntimeBlocked(latest, requestedModel) { _ = store.DeleteResponseAccount(ctx, derefGroupID(groupID), responseID) return 0, nil, "", nil diff --git a/backend/internal/service/profit_preview.go b/backend/internal/service/profit_preview.go new file mode 100644 index 0000000000..5cb8a78a89 --- /dev/null +++ b/backend/internal/service/profit_preview.go @@ -0,0 +1,223 @@ +package service + +// 分组利润控制离线预演:使用生产只读导出的分组配置、账号倍率、探测状态、 +// 用户覆盖倍率与主力模型清单,复用线上 U/D/epsilon 语义推演准入。 +// 探测状态仅用于解释账号倍率来源是否健康,不影响准入结论。 + +import ( + "math" + "sort" + "time" +) + +const ( + ProfitPreviewRateSourceManual = "manual" + ProfitPreviewRateSourceUpstreamProbe = "upstream_probe_sync" + + ProfitPreviewWarningProbeMissing = "probe_snapshot_missing" + ProfitPreviewWarningProbeStale = "probe_snapshot_stale" + ProfitPreviewWarningProbeFailed = "probe_failed" + ProfitPreviewWarningProbeUnsupported = "probe_unsupported" + ProfitPreviewWarningManualRateOne = "manual_rate_1_suspected_unmaintained" +) + +// ProfitPreviewGroupInput 是一个分组的预演输入。 +type ProfitPreviewGroupInput struct { + Group *Group + // Accounts 为该分组绑定且状态可调度的账号,允许 Anthropic/Gemini 分组中 + // 参与 mixed scheduling 的 Antigravity 账号。 + Accounts []*Account + // UserOverrides 为该分组所有用户覆盖倍率(user_id → 倍率)。 + UserOverrides map[int64]float64 + // Models 为该分组近期主力模型清单。 + Models []string + // AssumeEnabled 在分组当前关闭时仍按其保存配置启用利润门,供部署前执行 + // “先关组、预演、逐组重开”。 + AssumeEnabled bool +} + +// 预演账号分类与线上否决原因一一对应。 +const ( + ProfitPreviewClassAdmitted = "admitted" + ProfitPreviewClassRejectedThreshold = "rejected_threshold" + ProfitPreviewClassRejectedInvalidRate = "rejected_invalid_account_rate" +) + +// ProfitPreviewAccountVerdict 是单账号的预演结论。 +type ProfitPreviewAccountVerdict struct { + AccountID int64 `json:"account_id"` + Name string `json:"name"` + Platform string `json:"platform"` + Class string `json:"class"` + AccountRate *float64 `json:"account_rate,omitempty"` + RateSource string `json:"rate_source"` + Warnings []string `json:"warnings,omitempty"` + // RejectedUnderMinD 表示账号在最低有效 D(用户覆盖最坏口径)下会被排除。 + RejectedUnderMinD bool `json:"rejected_under_min_d,omitempty"` + // SupportedModels 为该账号支持的主力模型子集。 + SupportedModels []string `json:"supported_models,omitempty"` +} + +// ProfitPreviewGroupReport 是单分组的预演报告。 +type ProfitPreviewGroupReport struct { + GroupID int64 `json:"group_id"` + GroupName string `json:"group_name"` + Platform string `json:"platform"` + EffectiveGate bool `json:"effective_gate"` + AssumedEnabled bool `json:"assumed_enabled,omitempty"` + // DefaultD / MinEffectiveD:分组默认 D 与最低有效 D(含 evalAt 高峰因子; + // MinEffectiveD = min(分组默认, 全部非空用户覆盖) × 高峰因子)。 + DefaultD float64 `json:"default_d"` + MinEffectiveD float64 `json:"min_effective_d"` + // ThresholdDefault / ThresholdMinD:两档 D 对应的准入阈值。 + ThresholdDefault float64 `json:"threshold_default"` + ThresholdMinD float64 `json:"threshold_min_d"` + // RemainingByModel 仅表示利润门准入账号数,不模拟健康、冷却、限流或槽位。 + RemainingByModel map[string]int `json:"profit_admitted_by_model"` + RemainingByModelMinD map[string]int `json:"profit_admitted_by_model_min_d"` + Verdicts []ProfitPreviewAccountVerdict `json:"verdicts"` +} + +// PreviewProfitAdmission 对五大平台分组推演利润门准入结果。未启用的分组仅在 +// AssumeEnabled=true 时按保存配置执行门;Composite 分组始终不直接安装利润门。 +func PreviewProfitAdmission(inputs []ProfitPreviewGroupInput, evalAt time.Time) []ProfitPreviewGroupReport { + reports := make([]ProfitPreviewGroupReport, 0, len(inputs)) + for _, in := range inputs { + if in.Group == nil { + continue + } + group := in.Group + effectiveGate := (group.ProfitControlEnabled || in.AssumeEnabled) && profitControlPlatformSupported(group.Platform) + report := ProfitPreviewGroupReport{ + GroupID: group.ID, + GroupName: group.Name, + Platform: group.Platform, + EffectiveGate: effectiveGate, + AssumedEnabled: in.AssumeEnabled && !group.ProfitControlEnabled && effectiveGate, + RemainingByModel: make(map[string]int, len(in.Models)), + RemainingByModelMinD: make(map[string]int, len(in.Models)), + } + for _, model := range in.Models { + report.RemainingByModel[model] = 0 + report.RemainingByModelMinD[model] = 0 + } + + peak := group.PeakMultiplierAt(evalAt) + defaultD := group.RateMultiplier * peak + minRate := group.RateMultiplier + for _, override := range in.UserOverrides { + if math.IsNaN(override) || math.IsInf(override, 0) || override < 0 { + continue + } + if override < minRate { + minRate = override + } + } + minD := minRate * peak + deduction := group.ProfitMinMargin + group.ProfitSafetyBuffer + thresholdDefault := clampProfitPreviewThreshold(defaultD * (1 - deduction)) + thresholdMinD := clampProfitPreviewThreshold(minD * (1 - deduction)) + report.DefaultD = defaultD + report.MinEffectiveD = minD + report.ThresholdDefault = thresholdDefault + report.ThresholdMinD = thresholdMinD + + for _, account := range in.Accounts { + if account == nil { + continue + } + verdict := previewAccountProfitAdmission(account, effectiveGate, thresholdDefault, thresholdMinD, evalAt) + admittedDefault := verdict.Class == ProfitPreviewClassAdmitted + admittedMinD := admittedDefault && !verdict.RejectedUnderMinD + for _, model := range in.Models { + if !account.IsModelSupported(model) { + continue + } + verdict.SupportedModels = append(verdict.SupportedModels, model) + if admittedDefault { + report.RemainingByModel[model]++ + } + if admittedMinD { + report.RemainingByModelMinD[model]++ + } + } + sort.Strings(verdict.SupportedModels) + report.Verdicts = append(report.Verdicts, verdict) + } + sort.Slice(report.Verdicts, func(i, j int) bool { return report.Verdicts[i].AccountID < report.Verdicts[j].AccountID }) + reports = append(reports, report) + } + return reports +} + +func previewAccountProfitAdmission( + account *Account, + effectiveGate bool, + thresholdDefault float64, + thresholdMinD float64, + evalAt time.Time, +) ProfitPreviewAccountVerdict { + verdict := ProfitPreviewAccountVerdict{ + AccountID: account.ID, + Name: account.Name, + Platform: account.Platform, + RateSource: ProfitPreviewRateSourceManual, + } + if enabled, _ := account.Extra[UpstreamBillingRateSyncEnabledExtraKey].(bool); enabled { + verdict.RateSource = ProfitPreviewRateSourceUpstreamProbe + verdict.Warnings = append(verdict.Warnings, profitPreviewProbeWarnings(account, evalAt)...) + } else if account.RateMultiplier != nil && *account.RateMultiplier == 1 { + verdict.Warnings = append(verdict.Warnings, ProfitPreviewWarningManualRateOne) + } + + validRate := account.RateMultiplier != nil && + !math.IsNaN(*account.RateMultiplier) && + !math.IsInf(*account.RateMultiplier, 0) && + *account.RateMultiplier >= 0 + if validRate { + rate := *account.RateMultiplier + verdict.AccountRate = &rate + } + switch { + case !effectiveGate: + verdict.Class = ProfitPreviewClassAdmitted + case !validRate: + verdict.Class = ProfitPreviewClassRejectedInvalidRate + case profitPreviewOverThreshold(*account.RateMultiplier, thresholdDefault): + verdict.Class = ProfitPreviewClassRejectedThreshold + default: + verdict.Class = ProfitPreviewClassAdmitted + verdict.RejectedUnderMinD = profitPreviewOverThreshold(*account.RateMultiplier, thresholdMinD) + } + return verdict +} + +func profitPreviewProbeWarnings(account *Account, evalAt time.Time) []string { + snapshot := decodeUpstreamBillingProbeSnapshot(account.Extra) + if snapshot == nil { + return []string{ProfitPreviewWarningProbeMissing} + } + switch snapshot.Status { + case UpstreamBillingProbeStatusFailed: + return []string{ProfitPreviewWarningProbeFailed} + case UpstreamBillingProbeStatusUnsupported: + return []string{ProfitPreviewWarningProbeUnsupported} + case UpstreamBillingProbeStatusOK: + if snapshot.FreshUntil == nil || !evalAt.Before(*snapshot.FreshUntil) { + return []string{ProfitPreviewWarningProbeStale} + } + } + return nil +} + +func clampProfitPreviewThreshold(threshold float64) float64 { + if math.IsNaN(threshold) || math.IsInf(threshold, 0) || threshold < 0 { + return 0 + } + return threshold +} + +// profitPreviewOverThreshold 与线上否决点共用同一 epsilon 边界语义。 +func profitPreviewOverThreshold(upstream, threshold float64) bool { + return upstream-threshold > profitControlRateEpsilon*math.Max(1, math.Abs(threshold)) +} diff --git a/backend/internal/service/profit_preview_test.go b/backend/internal/service/profit_preview_test.go new file mode 100644 index 0000000000..1030ef41aa --- /dev/null +++ b/backend/internal/service/profit_preview_test.go @@ -0,0 +1,146 @@ +package service + +import ( + "testing" + "time" + + "github.com/stretchr/testify/require" +) + +func TestPreviewProfitAdmissionUsesAccountRatesAndPreinitializesModels(t *testing.T) { + now := time.Now() + group := profitControlTestGroup(50, 0.2, 0) + group.Name = "VIP-preview" + + cheap := profitControlTestAccountWithRate( + upstreamCostTestAccount(1, UpstreamBillingProbeStatusOK, 1.0, now.Add(-3*time.Hour), 30*time.Minute), + 0.5, + ) + cheap.Name = "cheap" + cheap.Extra[UpstreamBillingRateSyncEnabledExtraKey] = true + cheap.Credentials = map[string]any{"model_mapping": map[string]any{"gpt-sol": "gpt-sol"}} + + boundary := profitControlTestAccountWithRate( + upstreamCostTestAccount(2, UpstreamBillingProbeStatusOK, 0.8, now.Add(-time.Minute), 30*time.Minute), + 0.8, + ) + boundary.Name = "boundary" + boundary.Credentials = map[string]any{"model_mapping": map[string]any{ + "gpt-sol": "gpt-sol", + "gpt-luna": "gpt-luna", + }} + + expensive := profitControlTestAccountWithRate( + upstreamCostTestAccount(3, UpstreamBillingProbeStatusOK, 0.2, now.Add(-time.Minute), 30*time.Minute), + 1.0, + ) + expensive.Name = "expensive" + + invalid := upstreamCostTestOAuthAccount(4) + invalid.Name = "invalid" + + reports := PreviewProfitAdmission([]ProfitPreviewGroupInput{{ + Group: group, + Accounts: []*Account{cheap, boundary, expensive, invalid}, + UserOverrides: map[int64]float64{40: 0.5}, + Models: []string{"gpt-sol", "gpt-luna", "gpt-no-account"}, + }}, now) + require.Len(t, reports, 1) + report := reports[0] + + require.True(t, report.EffectiveGate) + require.InDelta(t, 1.0, report.DefaultD, 1e-12) + require.InDelta(t, 0.8, report.ThresholdDefault, 1e-12) + require.InDelta(t, 0.5, report.MinEffectiveD, 1e-12) + require.InDelta(t, 0.4, report.ThresholdMinD, 1e-12) + + byID := map[int64]ProfitPreviewAccountVerdict{} + for _, verdict := range report.Verdicts { + byID[verdict.AccountID] = verdict + } + require.Equal(t, ProfitPreviewClassAdmitted, byID[cheap.ID].Class) + require.Equal(t, ProfitPreviewRateSourceUpstreamProbe, byID[cheap.ID].RateSource) + require.Contains(t, byID[cheap.ID].Warnings, ProfitPreviewWarningProbeStale) + require.True(t, byID[cheap.ID].RejectedUnderMinD) + + require.Equal(t, ProfitPreviewClassAdmitted, byID[boundary.ID].Class, "U == 阈值按 epsilon 语义准入") + require.Equal(t, ProfitPreviewClassRejectedThreshold, byID[expensive.ID].Class) + require.Contains(t, byID[expensive.ID].Warnings, ProfitPreviewWarningManualRateOne) + require.Equal(t, ProfitPreviewClassRejectedInvalidRate, byID[invalid.ID].Class) + + require.Equal(t, 2, report.RemainingByModel["gpt-sol"]) + require.Equal(t, 1, report.RemainingByModel["gpt-luna"]) + require.Equal(t, 0, report.RemainingByModel["gpt-no-account"]) + _, present := report.RemainingByModel["gpt-no-account"] + require.True(t, present, "全部账号均不支持时也必须显式保留 0,供 CLI 发出警告") +} + +func TestPreviewProfitAdmissionAssumeEnabled(t *testing.T) { + now := time.Now() + group := profitControlTestGroup(51, 0, 0) + group.Platform = PlatformOpenAI + group.RateMultiplier = 0.5 + group.ProfitControlEnabled = false + + cheapAccount := profitControlTestAccountWithRate( + upstreamCostTestAccount(1, UpstreamBillingProbeStatusOK, 0.9, now.Add(-time.Minute), 30*time.Minute), + 0.2, + ) + expensiveAccount := profitControlTestAccountWithRate( + upstreamCostTestAccount(2, UpstreamBillingProbeStatusOK, 0.2, now.Add(-time.Minute), 30*time.Minute), + 0.9, + ) + + withoutAssume := PreviewProfitAdmission([]ProfitPreviewGroupInput{{ + Group: group, + Accounts: []*Account{cheapAccount, expensiveAccount}, + Models: []string{"gpt-test"}, + }}, now)[0] + require.False(t, withoutAssume.EffectiveGate) + require.Equal(t, ProfitPreviewClassAdmitted, withoutAssume.Verdicts[0].Class) + require.Equal(t, ProfitPreviewClassAdmitted, withoutAssume.Verdicts[1].Class) + + withAssume := PreviewProfitAdmission([]ProfitPreviewGroupInput{{ + Group: group, + Accounts: []*Account{cheapAccount, expensiveAccount}, + Models: []string{"gpt-test"}, + AssumeEnabled: true, + }}, now)[0] + require.True(t, withAssume.EffectiveGate) + require.True(t, withAssume.AssumedEnabled) + + byID := map[int64]ProfitPreviewAccountVerdict{} + for _, verdict := range withAssume.Verdicts { + byID[verdict.AccountID] = verdict + } + require.Equal(t, ProfitPreviewClassAdmitted, byID[cheapAccount.ID].Class, + "账号倍率 0.2 <= 阈值 0.5,探测快照的高倍率不参与准入判断") + require.Equal(t, ProfitPreviewClassRejectedThreshold, byID[expensiveAccount.ID].Class, + "账号倍率 0.9 > 阈值 0.5,探测快照的低倍率不能替代账号倍率") +} + +func TestPreviewProfitAdmissionSupportsFivePlatforms(t *testing.T) { + for i, platform := range []string{ + PlatformOpenAI, + PlatformAnthropic, + PlatformGemini, + PlatformGrok, + PlatformAntigravity, + } { + group := profitControlTestGroup(int64(100+i), 0, 0) + group.Platform = platform + rate := 0.2 + account := &Account{ + ID: int64(200 + i), + Platform: platform, + RateMultiplier: &rate, + } + report := PreviewProfitAdmission([]ProfitPreviewGroupInput{{ + Group: group, + Accounts: []*Account{account}, + Models: []string{"model"}, + }}, time.Now())[0] + require.True(t, report.EffectiveGate, platform) + require.Equal(t, ProfitPreviewClassAdmitted, report.Verdicts[0].Class, platform) + } +} diff --git a/backend/migrations/192_group_profit_control.sql b/backend/migrations/192_group_profit_control.sql new file mode 100644 index 0000000000..072b3c5db1 --- /dev/null +++ b/backend/migrations/192_group_profit_control.sql @@ -0,0 +1,9 @@ +-- Per-group profit control for scheduling admission. +-- Admission rule at request time: an account qualifies iff its cost multiplier +-- U (accounts.rate_multiplier) satisfies U <= D * (1 - margin - buffer), where +-- D is the requester's effective downstream multiplier at the request's +-- pricing instant. +ALTER TABLE groups + ADD COLUMN IF NOT EXISTS profit_control_enabled BOOLEAN NOT NULL DEFAULT FALSE, + ADD COLUMN IF NOT EXISTS profit_min_margin DECIMAL(10,4) NOT NULL DEFAULT 0, + ADD COLUMN IF NOT EXISTS profit_safety_buffer DECIMAL(10,4) NOT NULL DEFAULT 0; diff --git a/backend/migrations/193_group_profit_control_auth_cache_invalidation.sql b/backend/migrations/193_group_profit_control_auth_cache_invalidation.sql new file mode 100644 index 0000000000..f32f6e6f8b --- /dev/null +++ b/backend/migrations/193_group_profit_control_auth_cache_invalidation.sql @@ -0,0 +1,47 @@ +-- Profit-control fields are part of the API-key auth snapshot and gate the +-- scheduling admission filter; the profit threshold D additionally depends on +-- group pricing and peak-window fields. Extend the durable invalidation +-- trigger so out-of-band group edits (direct SQL, crash between update and +-- app-level invalidation) cannot leave cached snapshots using stale +-- profit-control inputs. Normal admin saves already invalidate via +-- InvalidateAuthCacheByGroupID; this trigger is the durable backstop. Based on +-- the latest function body from 186_group_auth_cache_image_generation.sql. + +CREATE OR REPLACE FUNCTION enqueue_group_auth_cache_invalidation() +RETURNS TRIGGER +LANGUAGE plpgsql +AS $$ +DECLARE + target_group_id BIGINT; +BEGIN + target_group_id := OLD.id; + IF TG_OP = 'UPDATE' + AND OLD.status IS NOT DISTINCT FROM NEW.status + AND OLD.is_exclusive IS NOT DISTINCT FROM NEW.is_exclusive + AND OLD.allow_image_generation IS NOT DISTINCT FROM NEW.allow_image_generation + AND OLD.platform IS NOT DISTINCT FROM NEW.platform + AND OLD.subscription_type IS NOT DISTINCT FROM NEW.subscription_type + AND OLD.rate_multiplier IS NOT DISTINCT FROM NEW.rate_multiplier + AND OLD.peak_rate_enabled IS NOT DISTINCT FROM NEW.peak_rate_enabled + AND OLD.peak_start IS NOT DISTINCT FROM NEW.peak_start + AND OLD.peak_end IS NOT DISTINCT FROM NEW.peak_end + AND OLD.peak_rate_multiplier IS NOT DISTINCT FROM NEW.peak_rate_multiplier + AND OLD.profit_control_enabled IS NOT DISTINCT FROM NEW.profit_control_enabled + AND OLD.profit_min_margin IS NOT DISTINCT FROM NEW.profit_min_margin + AND OLD.profit_safety_buffer IS NOT DISTINCT FROM NEW.profit_safety_buffer + AND OLD.deleted_at IS NOT DISTINCT FROM NEW.deleted_at THEN + RETURN NEW; + END IF; + + INSERT INTO auth_cache_invalidation_outbox (cache_key) + SELECT encode(sha256(convert_to(k.key, 'UTF8')), 'hex') + FROM api_keys AS k + WHERE k.group_id = target_group_id + AND k.deleted_at IS NULL + AND k.key <> ''; + IF TG_OP = 'DELETE' THEN + RETURN OLD; + END IF; + RETURN NEW; +END; +$$; diff --git a/frontend/src/i18n/locales/en/admin/overview.ts b/frontend/src/i18n/locales/en/admin/overview.ts index 67887a9f45..462fa3cb27 100644 --- a/frontend/src/i18n/locales/en/admin/overview.ts +++ b/frontend/src/i18n/locales/en/admin/overview.ts @@ -1006,6 +1006,18 @@ export default { peakMultiplier: 'Peak multiplier', multiplierHint: 'Applies to token billing multiplier; image tokens in token billing are also affected. 0 means peak token requests are billed at 0x.' }, + profitControl: { + enable: 'Enable profit control', + enabledHint: 'Scheduling only admits accounts whose account multiplier ≤ the request\'s effective downstream multiplier × (1 − min margin − safety buffer). Account multipliers may be maintained manually or synchronized from probes; existing ordering, stickiness and breakers keep working among qualified accounts. Image/video scheduling is not covered yet.', + disabledHint: 'When disabled, scheduling does no profit filtering: accounts whose account multiplier exceeds the downstream multiplier can still be selected, which may produce loss-making requests.', + minMargin: 'Min gross margin (%)', + minMarginHint: 'Percent input, e.g. 30 means 30%; stored as a decimal on the backend', + safetyBuffer: 'Safety buffer (%)', + safetyBufferHint: 'Added to min margin and deducted from the downstream multiplier; defaults to 0', + marginRangeError: 'Min gross margin must be between 0 and 100 (exclusive)', + bufferRangeError: 'Safety buffer must be between 0 and 100 (exclusive)', + sumTooHigh: 'Min gross margin plus safety buffer must be less than 100%, otherwise every account would be excluded' + }, modelsList: { title: 'Custom /v1/models Model List', hint: 'Only changes the /v1/models response. Whitelist model calls and account routing are unchanged.', diff --git a/frontend/src/i18n/locales/zh/admin/overview.ts b/frontend/src/i18n/locales/zh/admin/overview.ts index 112577b894..3fa15e3a2e 100644 --- a/frontend/src/i18n/locales/zh/admin/overview.ts +++ b/frontend/src/i18n/locales/zh/admin/overview.ts @@ -1003,6 +1003,18 @@ export default { peakMultiplier: '高峰倍率', multiplierHint: '作用于 token 计费倍率;token 计费的图片 token 同样适用,0 表示高峰 token 请求按 0 倍计费' }, + profitControl: { + enable: '启用利润控制', + enabledHint: '调度时仅允许"账号倍率 ≤ 请求实际下游倍率 ×(1 − 最低毛利率 − 安全缓冲)"的账号进入候选池;账号倍率可手工维护或由探测同步,既有排序、粘性与熔断在合格账号间照常工作。图片/视频调度暂不参与。', + disabledHint: '关闭后调度不做利润过滤,账号倍率高于下游倍率的账号也会被选中,可能产生亏损请求。', + minMargin: '最低毛利率(%)', + minMarginHint: '百分比输入,如 30 表示 30%;后端按小数存储', + safetyBuffer: '安全缓冲(%)', + safetyBufferHint: '与最低毛利率相加后从下游倍率中扣除,默认 0', + marginRangeError: '最低毛利率应在 0 到 100 之间(不含 100)', + bufferRangeError: '安全缓冲应在 0 到 100 之间(不含 100)', + sumTooHigh: '最低毛利率与安全缓冲之和必须小于 100%,否则将排除全部账号' + }, modelsList: { title: '自定义 /v1/models 模型列表', hint: '仅影响 /v1/models 展示结果,不影响白名单模型调用和账号调度。', diff --git a/frontend/src/types/index.ts b/frontend/src/types/index.ts index 272e8e8e97..25b41a6f59 100644 --- a/frontend/src/types/index.ts +++ b/frontend/src/types/index.ts @@ -549,6 +549,10 @@ export interface Group { peak_start: string peak_end: string peak_rate_multiplier: number + // 分组利润控制(openai/anthropic/gemini/grok/antigravity 分组可启用;margin/buffer 为小数存储) + profit_control_enabled: boolean + profit_min_margin: number + profit_safety_buffer: number // Claude Code 客户端限制 claude_code_only: boolean fallback_group_id: number | null @@ -741,6 +745,10 @@ export interface CreateGroupRequest { peak_start?: string peak_end?: string peak_rate_multiplier?: number + // 分组利润控制(五个 token 平台;margin/buffer 为小数) + profit_control_enabled?: boolean + profit_min_margin?: number + profit_safety_buffer?: number claude_code_only?: boolean fallback_group_id?: number | null fallback_group_id_on_invalid_request?: number | null @@ -792,6 +800,10 @@ export interface UpdateGroupRequest { peak_start?: string peak_end?: string peak_rate_multiplier?: number + // 分组利润控制(五个 token 平台;margin/buffer 为小数) + profit_control_enabled?: boolean + profit_min_margin?: number + profit_safety_buffer?: number claude_code_only?: boolean fallback_group_id?: number | null fallback_group_id_on_invalid_request?: number | null diff --git a/frontend/src/views/admin/GroupsView.vue b/frontend/src/views/admin/GroupsView.vue index 54375581b9..9ff7d4a2ab 100644 --- a/frontend/src/views/admin/GroupsView.vue +++ b/frontend/src/views/admin/GroupsView.vue @@ -1152,6 +1152,56 @@ + +
+ +

+ {{ + createForm.profit_control_enabled + ? t("admin.groups.profitControl.enabledHint") + : t("admin.groups.profitControl.disabledHint") + }} +

+
+
+ + +
+
+ + +
+
+
+
@@ -2707,6 +2757,56 @@
+ +
+ +

+ {{ + editForm.profit_control_enabled + ? t("admin.groups.profitControl.enabledHint") + : t("admin.groups.profitControl.disabledHint") + }} +

+
+
+ + +
+
+ + +
+
+
+
@@ -4095,6 +4195,13 @@ import { } from "./groupsModelsList"; import { createModelsListCandidatesTracker } from "./groupsModelsListCandidates"; import { normalizeSupportedModelScopesForPlatform } from "./groupsSupportedModelScopes"; +import { + isProfitControlPlatform, + profitPercentToDecimal, + profitDecimalToPercent, + validateProfitControlFormState, + type ProfitControlFormState, +} from "./groupsProfitControl"; import { normalizeReasoningEffortForPlatform, reasoningEffortMappingsToAPI, @@ -4618,6 +4725,10 @@ const createForm = reactive({ peak_start: "", peak_end: "", peak_rate_multiplier: 1.0, + // 分组利润控制(五个 token 平台);界面按百分比输入,提交时转小数 + profit_control_enabled: false, + profit_min_margin_percent: 0, + profit_safety_buffer_percent: 0, // Claude Code 客户端限制(仅 anthropic 平台使用) claude_code_only: false, fallback_group_id: null as number | null, @@ -4968,6 +5079,10 @@ const editForm = reactive({ peak_start: "", peak_end: "", peak_rate_multiplier: 1.0, + // 分组利润控制(五个 token 平台);界面按百分比输入,提交时转小数 + profit_control_enabled: false, + profit_min_margin_percent: 0, + profit_safety_buffer_percent: 0, // Claude Code 客户端限制(仅 anthropic 平台使用) claude_code_only: false, fallback_group_id: null as number | null, @@ -5412,6 +5527,9 @@ const closeCreateModal = () => { createForm.peak_start = ""; createForm.peak_end = ""; createForm.peak_rate_multiplier = 1.0; + createForm.profit_control_enabled = false; + createForm.profit_min_margin_percent = 0; + createForm.profit_safety_buffer_percent = 0; createForm.claude_code_only = false; createForm.fallback_group_id = null; createForm.fallback_group_id_on_invalid_request = null; @@ -5459,6 +5577,19 @@ const normalizeRateMultiplier = ( return Number.isFinite(parsed) && parsed >= 0 ? parsed : 1; }; +// 利润控制表单辅助(换算与校验逻辑见 groupsProfitControl.ts,便于单测)。 +const percentToDecimal = profitPercentToDecimal; +const decimalToPercent = profitDecimalToPercent; + +const validateProfitControlForm = (form: ProfitControlFormState): boolean => { + const errorKey = validateProfitControlFormState(form); + if (errorKey) { + appStore.showError(t(`admin.groups.profitControl.${errorKey}`)); + return false; + } + return true; +}; + const handleCreateGroup = async () => { if (!createForm.name.trim()) { appStore.showError(t("admin.groups.nameRequired")); @@ -5471,6 +5602,9 @@ const handleCreateGroup = async () => { ) { return; } + if (!validateProfitControlForm(createForm)) { + return; + } submitting.value = true; try { // 构建请求数据,包含模型路由配置 @@ -5506,7 +5640,17 @@ const handleCreateGroup = async () => { reasoning_effort_mappings: reasoningEffortMappingsToAPI( createForm.reasoning_effort_mappings, ), + // 利润控制:界面百分比转小数提交;仅五个 token 平台可启用 + profit_control_enabled: + isProfitControlPlatform(createForm.platform) && + createForm.profit_control_enabled, + profit_min_margin: percentToDecimal(createForm.profit_min_margin_percent), + profit_safety_buffer: percentToDecimal( + createForm.profit_safety_buffer_percent, + ), }; + delete (requestData as Record).profit_min_margin_percent; + delete (requestData as Record).profit_safety_buffer_percent; // v-model.number 清空输入框时产生 "",转为 null 让后端设为无限制 const emptyToNull = (v: any) => (v === "" ? null : v); requestData.daily_limit_usd = emptyToNull(requestData.daily_limit_usd); @@ -5594,6 +5738,13 @@ const handleEdit = async (group: AdminGroup) => { editForm.peak_start = group.peak_start ?? ""; editForm.peak_end = group.peak_end ?? ""; editForm.peak_rate_multiplier = group.peak_rate_multiplier ?? 1.0; + editForm.profit_control_enabled = group.profit_control_enabled ?? false; + editForm.profit_min_margin_percent = decimalToPercent( + group.profit_min_margin ?? 0, + ); + editForm.profit_safety_buffer_percent = decimalToPercent( + group.profit_safety_buffer ?? 0, + ); editForm.claude_code_only = group.claude_code_only || false; editForm.fallback_group_id = group.fallback_group_id; editForm.fallback_group_id_on_invalid_request = @@ -5654,6 +5805,9 @@ const closeEditModal = () => { editForm.peak_start = ""; editForm.peak_end = ""; editForm.peak_rate_multiplier = 1.0; + editForm.profit_control_enabled = false; + editForm.profit_min_margin_percent = 0; + editForm.profit_safety_buffer_percent = 0; editForm.video_rate_independent = false; editForm.video_rate_multiplier = 1; editForm.video_price_480p = null; @@ -5678,6 +5832,9 @@ const handleUpdateGroup = async () => { ) { return; } + if (!validateProfitControlForm(editForm)) { + return; + } submitting.value = true; try { @@ -5720,7 +5877,17 @@ const handleUpdateGroup = async () => { reasoning_effort_mappings: reasoningEffortMappingsToAPI( editForm.reasoning_effort_mappings, ), + // 利润控制:界面百分比转小数提交;仅五个 token 平台可启用 + profit_control_enabled: + isProfitControlPlatform(editForm.platform) && + editForm.profit_control_enabled, + profit_min_margin: percentToDecimal(editForm.profit_min_margin_percent), + profit_safety_buffer: percentToDecimal( + editForm.profit_safety_buffer_percent, + ), }; + delete (payload as Record).profit_min_margin_percent; + delete (payload as Record).profit_safety_buffer_percent; // v-model.number 清空输入框时产生 "",转为 null 让后端设为无限制 const emptyToNull = (v: any) => (v === "" ? null : v); payload.daily_limit_usd = emptyToNull(payload.daily_limit_usd); @@ -6068,6 +6235,11 @@ watch( resetMessagesDispatchFormState(createForm); createForm.allow_live = false; } + if (!isProfitControlPlatform(newVal)) { + createForm.profit_control_enabled = false; + createForm.profit_min_margin_percent = 0; + createForm.profit_safety_buffer_percent = 0; + } createForm.max_reasoning_effort = normalizeReasoningEffortForPlatform( newVal, createForm.max_reasoning_effort, @@ -6111,6 +6283,11 @@ watch( resetMessagesDispatchFormState(editForm); editForm.allow_live = false; } + if (!isProfitControlPlatform(newVal)) { + editForm.profit_control_enabled = false; + editForm.profit_min_margin_percent = 0; + editForm.profit_safety_buffer_percent = 0; + } editForm.max_reasoning_effort = normalizeReasoningEffortForPlatform( newVal, editForm.max_reasoning_effort, diff --git a/frontend/src/views/admin/__tests__/groupsProfitControl.spec.ts b/frontend/src/views/admin/__tests__/groupsProfitControl.spec.ts new file mode 100644 index 0000000000..21e458281e --- /dev/null +++ b/frontend/src/views/admin/__tests__/groupsProfitControl.spec.ts @@ -0,0 +1,153 @@ +import { describe, expect, it } from "vitest"; + +import { + profitDecimalToPercent, + profitPercentToDecimal, + validateProfitControlFormState, + type ProfitControlFormState, +} from "../groupsProfitControl"; + +const formState = ( + overrides: Partial = {}, +): ProfitControlFormState => ({ + platform: "openai", + profit_control_enabled: true, + profit_min_margin_percent: 30, + profit_safety_buffer_percent: 0, + ...overrides, +}); + +describe("profitPercentToDecimal", () => { + it("converts percent input to backend decimal", () => { + expect(profitPercentToDecimal(30)).toBe(0.3); + expect(profitPercentToDecimal(5)).toBe(0.05); + expect(profitPercentToDecimal(33.33)).toBe(0.3333); + expect(profitPercentToDecimal(99.99)).toBe(0.9999); + }); + + it("rounds to four decimal places matching decimal(10,4) storage", () => { + expect(profitPercentToDecimal(33.333)).toBe(0.3333); + expect(profitPercentToDecimal(0.005)).toBe(0.0001); + }); + + it("treats empty, invalid and non-positive input as zero", () => { + expect(profitPercentToDecimal("")).toBe(0); + expect(profitPercentToDecimal(null)).toBe(0); + expect(profitPercentToDecimal(undefined)).toBe(0); + expect(profitPercentToDecimal("abc")).toBe(0); + expect(profitPercentToDecimal(-5)).toBe(0); + expect(profitPercentToDecimal(0)).toBe(0); + }); +}); + +describe("profitDecimalToPercent", () => { + it("converts backend decimal to percent without float tail noise", () => { + expect(profitDecimalToPercent(0.3)).toBe(30); + expect(profitDecimalToPercent(0.05)).toBe(5); + expect(profitDecimalToPercent(0.3333)).toBe(33.33); + expect(profitDecimalToPercent(0.9999)).toBe(99.99); + }); + + it("treats missing and non-positive values as zero", () => { + expect(profitDecimalToPercent(null)).toBe(0); + expect(profitDecimalToPercent(undefined)).toBe(0); + expect(profitDecimalToPercent(0)).toBe(0); + expect(profitDecimalToPercent(-0.3)).toBe(0); + }); + + it("round-trips representative storage values", () => { + for (const decimal of [0.05, 0.1, 0.3, 0.3333, 0.5, 0.75, 0.9999]) { + expect(profitPercentToDecimal(profitDecimalToPercent(decimal))).toBe( + decimal, + ); + } + }); +}); + +describe("validateProfitControlFormState", () => { + it("passes valid enabled configurations", () => { + expect(validateProfitControlFormState(formState())).toBeNull(); + expect( + validateProfitControlFormState( + formState({ + profit_min_margin_percent: 0, + profit_safety_buffer_percent: 0, + }), + ), + ).toBeNull(); + expect( + validateProfitControlFormState( + formState({ + profit_min_margin_percent: 60, + profit_safety_buffer_percent: 39.99, + }), + ), + ).toBeNull(); + }); + + it("validates all five supported platforms and skips unsupported platforms", () => { + expect( + validateProfitControlFormState( + formState({ profit_control_enabled: false, profit_min_margin_percent: 200 }), + ), + ).toBeNull(); + expect( + validateProfitControlFormState( + formState({ platform: "anthropic", profit_min_margin_percent: 200 }), + ), + ).toBe("marginRangeError"); + for (const platform of ["openai", "anthropic", "gemini", "grok", "antigravity"]) { + expect(validateProfitControlFormState(formState({ platform }))).toBeNull(); + } + expect( + validateProfitControlFormState( + formState({ platform: "composite", profit_min_margin_percent: 200 }), + ), + ).toBeNull(); + }); + + it("treats empty inputs as zero", () => { + expect( + validateProfitControlFormState( + formState({ + profit_min_margin_percent: "", + profit_safety_buffer_percent: null, + }), + ), + ).toBeNull(); + }); + + it("rejects out-of-range margin and buffer", () => { + expect( + validateProfitControlFormState( + formState({ profit_min_margin_percent: 100 }), + ), + ).toBe("marginRangeError"); + expect( + validateProfitControlFormState( + formState({ profit_min_margin_percent: -1 }), + ), + ).toBe("marginRangeError"); + expect( + validateProfitControlFormState( + formState({ profit_safety_buffer_percent: 100 }), + ), + ).toBe("bufferRangeError"); + expect( + validateProfitControlFormState( + formState({ profit_safety_buffer_percent: -0.1 }), + ), + ).toBe("bufferRangeError"); + }); + + it("rejects margin plus buffer reaching 100 percent", () => { + expect( + validateProfitControlFormState( + formState({ + profit_min_margin_percent: 60, + profit_safety_buffer_percent: 40, + }), + ), + ).toBe("sumTooHigh"); + }); +}); diff --git a/frontend/src/views/admin/groupsProfitControl.ts b/frontend/src/views/admin/groupsProfitControl.ts new file mode 100644 index 0000000000..6bb55761ed --- /dev/null +++ b/frontend/src/views/admin/groupsProfitControl.ts @@ -0,0 +1,56 @@ +// 分组利润控制表单辅助:百分比 <-> 小数换算与提交前校验。 +// 后端按小数存 decimal(10,4)(0.30 = 30%),界面按百分比输入展示; +// 固定 4 位小数精度,避免 0.3 * 100 = 30.000000000000004 之类的浮点尾数回显。 + +export const profitPercentToDecimal = ( + value: number | string | null | undefined, +): number => { + const parsed = Number(value); + if (!Number.isFinite(parsed) || parsed <= 0) { + return 0; + } + return Math.round(parsed * 100) / 10000; +}; + +export const profitDecimalToPercent = ( + value: number | null | undefined, +): number => { + const parsed = Number(value); + if (!Number.isFinite(parsed) || parsed <= 0) { + return 0; + } + return Math.round(parsed * 1e6) / 1e4; +}; + +export type ProfitControlFormState = { + platform: string; + profit_control_enabled: boolean; + profit_min_margin_percent: number | string | null; + profit_safety_buffer_percent: number | string | null; +}; + +export const isProfitControlPlatform = (platform: string): boolean => + ["openai", "anthropic", "gemini", "grok", "antigravity"].includes(platform); + +// 提交前校验:margin/buffer 各自 ∈ [0,100),且相加 < 100(否则阈值 <= 0, +// 所有可核价账号都会被排除)。返回 null 表示通过,否则返回错误信息的 i18n key +//(相对 admin.groups.profitControl 前缀)。仅支持平台且开关开启时才校验。 +export const validateProfitControlFormState = ( + form: ProfitControlFormState, +): string | null => { + if (!isProfitControlPlatform(form.platform) || !form.profit_control_enabled) { + return null; + } + const margin = Number(form.profit_min_margin_percent || 0); + const buffer = Number(form.profit_safety_buffer_percent || 0); + if (!Number.isFinite(margin) || margin < 0 || margin >= 100) { + return "marginRangeError"; + } + if (!Number.isFinite(buffer) || buffer < 0 || buffer >= 100) { + return "bufferRangeError"; + } + if (margin + buffer >= 100) { + return "sumTooHigh"; + } + return null; +};