feat(scheduler): per-group profit control for token account admission

Group pricing (rate multiplier, peak windows, per-user overrides) and
account cost (accounts.rate_multiplier) already live side by side, but
nothing stops the scheduler from handing a request to an account whose
cost multiplier exceeds what the group's pricing can profitably serve.
Add an opt-in per-group profit gate that filters scheduling candidates
by a margin rule, while ordering, scoring, stickiness and breakers keep
working unchanged among qualified accounts.

Admission rule: an account qualifies iff U <= D * (1 - min_margin -
safety_buffer) within a small relative epsilon, where U is
accounts.rate_multiplier (0 is legal; missing/negative/NaN/Inf are
conservatively rejected as invalid) and D is the requester's effective
downstream multiplier (user-group override ?? group default, times the
group peak factor) frozen at the request's pricing instant.

- groups gain profit_control_enabled / profit_min_margin /
  profit_safety_buffer (migration 191); the durable auth-cache
  invalidation trigger additionally watches the profit and pricing
  columns (migration 192) so out-of-band group edits cannot leave
  stale auth snapshots; GetByKeyForAuth explicitly projects the new
  columns and the API-key auth snapshot version is bumped to force a
  refresh of pre-existing snapshots
- request-level pricing instant: token entry points install pricingAt
  into ctx; the profit threshold D and the RecordUsage peak factor
  read the same instant, so one request never changes price mid-flight
  across waits/retries/failover (media and unwired paths keep the
  existing record-time semantics)
- the gate covers token requests on openai, anthropic, gemini, grok
  and antigravity groups: OpenAI-family handlers via
  WithOpenAIRequestPricingContext (responses incl. WS bridge, chat
  completions, messages, embeddings, alpha search), the shared gateway
  via WithGatewayTokenRequestPricing (messages, chat completions,
  responses, gemini model actions); composite groups cannot enable it
  directly; image/video/models/usage/count_tokens stay ungated and an
  explicit image-generation intent suppresses the gate end to end
- post-slot recheck: after a slot is acquired the account is re-read
  via SchedulerSnapshotService.GetAccount (scheduler cache, then DB;
  only when both fail the check fails open with WARN + metric); a
  vetoed account releases its slot and joins the request's exclusion
  set for reselection; sticky bindings are written only after the
  final check passes, and an over-threshold sticky account is skipped,
  not deleted, so it comes back once its rate recovers
- sticky-session cache contract: GatewayCache.GetSessionAccountID now
  returns ErrStickySessionNotFound on a miss (mapped from redis.Nil in
  the repository implementation, mirroring ErrRefreshTokenNotFound) so
  the profit sticky path can distinguish "no binding yet" from a real
  read failure without importing the cache driver in service code
- cross-group re-entry (composite parent -> member group) resolves the
  gate against the member group and clears a stale parent gate instead
  of letting a foreign threshold veto accounts
- per-platform/group activity counters (installs, threshold vetoes,
  invalid-rate vetoes, refresh failures) for observability
- admin UI: profit-control section on the five platforms' group forms
  with percent input, validation and platform-switch reset; group
  create/update/duplicate normalize and validate the config at a
  single choke point
- cmd/profit-preview: offline what-if tool that replays the production
  admission semantics over an exported config/account/override/model
  dump, reports per-model admitted-account counts under the default
  and the worst-case (lowest user override) D, and surfaces probe-sync
  staleness as warnings without affecting admission

Tests: service unit coverage for gate resolution/veto/threshold
epsilon/pricing instant/suppress marker/scheduler filtering and
post-slot recheck (incl. -race on the profit surface), unit-tagged
handler slot-recheck and capability-mapping regressions, sqlmock and
real-PostgreSQL integration regressions for the GetByKeyForAuth
projection and the migration-192 trigger watch list, API contract
update, and frontend specs for the five-platform form helpers.
This commit is contained in:
Brisbanehuang
2026-08-01 22:39:31 +08:00
committed by shaw
parent 0b6b4ea956
commit 20ad5ec506
68 changed files with 4752 additions and 100 deletions
+219
View File
@@ -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
}
+56
View File
@@ -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)
}
}
+14 -14
View File
@@ -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
}
)
+35 -2
View File
@@ -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()
}
+30
View File
@@ -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) {
+105
View File
@@ -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) {
+235
View File
@@ -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 {
+142
View File
@@ -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,
+3
View File
@@ -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{
+229 -1
View File
@@ -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)
}
+12
View File
@@ -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
+14
View File
@@ -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"),
}
}
@@ -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,
+3
View File
@@ -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,
+14 -10
View File
@@ -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"`
+30 -8
View File
@@ -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,
@@ -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,
@@ -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,
@@ -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,
+7 -2
View File
@@ -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
}
@@ -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"),
@@ -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(
+12 -2
View File
@@ -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"),
@@ -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",
+7 -2
View File
@@ -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
}
@@ -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))
}
@@ -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,
}
@@ -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)
}
@@ -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+" 变更必须入队")
})
}
}
@@ -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)
}
+9 -1
View File
@@ -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 {
@@ -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) {
+8 -2
View File
@@ -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 {
@@ -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",
+32
View File
@@ -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)
}
@@ -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,
@@ -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
}
@@ -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 缓存条目,支持负缓存
@@ -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)
@@ -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)
}
@@ -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
}
@@ -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)
}
@@ -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{}
@@ -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
}
@@ -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())
}
+50 -14
View File
@@ -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,
)
}
+40 -2
View File
@@ -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) {
@@ -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)
+67
View File
@@ -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
}
}
@@ -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)
@@ -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 {
@@ -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) {
@@ -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
@@ -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(),
)
}
@@ -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, ""))
}
@@ -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], "无既有绑定时应在终检通过后建立粘性")
}
@@ -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)
})
}
@@ -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
+223
View File
@@ -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))
}
@@ -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)
}
}
@@ -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;
@@ -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;
$$;
@@ -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.',
@@ -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 展示结果,不影响白名单模型调用和账号调度。',
+12
View File
@@ -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
+177
View File
@@ -1152,6 +1152,56 @@
</div>
</div>
<!-- 分组利润控制(五个平台 token 请求) -->
<div v-if="isProfitControlPlatform(createForm.platform)" class="border-t pt-4">
<label class="flex items-center gap-2 text-sm text-gray-700 dark:text-gray-300">
<input
v-model="createForm.profit_control_enabled"
type="checkbox"
class="rounded border-gray-300 text-blue-600 focus:ring-blue-500"
/>
<span>{{ t("admin.groups.profitControl.enable") }}</span>
</label>
<p class="mb-3 mt-1.5 text-xs text-gray-500 dark:text-gray-400">
{{
createForm.profit_control_enabled
? t("admin.groups.profitControl.enabledHint")
: t("admin.groups.profitControl.disabledHint")
}}
</p>
<div
v-if="createForm.profit_control_enabled"
class="mb-3 grid grid-cols-1 gap-3 sm:grid-cols-2"
>
<div>
<label class="input-label">{{ t("admin.groups.profitControl.minMargin") }}</label>
<input
v-model.number="createForm.profit_min_margin_percent"
type="number"
step="0.1"
min="0"
max="99.99"
class="input"
placeholder="0"
:title="t('admin.groups.profitControl.minMarginHint')"
/>
</div>
<div>
<label class="input-label">{{ t("admin.groups.profitControl.safetyBuffer") }}</label>
<input
v-model.number="createForm.profit_safety_buffer_percent"
type="number"
step="0.1"
min="0"
max="99.99"
class="input"
placeholder="0"
:title="t('admin.groups.profitControl.safetyBufferHint')"
/>
</div>
</div>
</div>
<!-- 支持的模型系列(仅 antigravity 平台) -->
<div v-if="createForm.platform === 'antigravity'" class="border-t pt-4">
<div class="mb-1.5 flex items-center gap-1">
@@ -2707,6 +2757,56 @@
</div>
</div>
<!-- 分组利润控制(五个平台 token 请求) -->
<div v-if="isProfitControlPlatform(editForm.platform)" class="border-t pt-4">
<label class="flex items-center gap-2 text-sm text-gray-700 dark:text-gray-300">
<input
v-model="editForm.profit_control_enabled"
type="checkbox"
class="rounded border-gray-300 text-blue-600 focus:ring-blue-500"
/>
<span>{{ t("admin.groups.profitControl.enable") }}</span>
</label>
<p class="mb-3 mt-1.5 text-xs text-gray-500 dark:text-gray-400">
{{
editForm.profit_control_enabled
? t("admin.groups.profitControl.enabledHint")
: t("admin.groups.profitControl.disabledHint")
}}
</p>
<div
v-if="editForm.profit_control_enabled"
class="mb-3 grid grid-cols-1 gap-3 sm:grid-cols-2"
>
<div>
<label class="input-label">{{ t("admin.groups.profitControl.minMargin") }}</label>
<input
v-model.number="editForm.profit_min_margin_percent"
type="number"
step="0.1"
min="0"
max="99.99"
class="input"
placeholder="0"
:title="t('admin.groups.profitControl.minMarginHint')"
/>
</div>
<div>
<label class="input-label">{{ t("admin.groups.profitControl.safetyBuffer") }}</label>
<input
v-model.number="editForm.profit_safety_buffer_percent"
type="number"
step="0.1"
min="0"
max="99.99"
class="input"
placeholder="0"
:title="t('admin.groups.profitControl.safetyBufferHint')"
/>
</div>
</div>
</div>
<!-- 支持的模型系列(仅 antigravity 平台) -->
<div v-if="editForm.platform === 'antigravity'" class="border-t pt-4">
<div class="mb-1.5 flex items-center gap-1">
@@ -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<string, unknown>).profit_min_margin_percent;
delete (requestData as Record<string, unknown>).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<string, unknown>).profit_min_margin_percent;
delete (payload as Record<string, unknown>).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,
@@ -0,0 +1,153 @@
import { describe, expect, it } from "vitest";
import {
profitDecimalToPercent,
profitPercentToDecimal,
validateProfitControlFormState,
type ProfitControlFormState,
} from "../groupsProfitControl";
const formState = (
overrides: Partial<ProfitControlFormState> = {},
): 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");
});
});
@@ -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;
};