mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-10-07 12:57:57 +08:00
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:
@@ -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
|
||||
}
|
||||
@@ -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
@@ -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
@@ -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()
|
||||
}
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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"`
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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())
|
||||
}
|
||||
@@ -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,
|
||||
)
|
||||
}
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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 展示结果,不影响白名单模型调用和账号调度。',
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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;
|
||||
};
|
||||
Reference in New Issue
Block a user