Merge branch 'release/custom-0.1.102' into release/custom-0.1.103

# Conflicts:
#	backend/go.sum
This commit is contained in:
erio
2026-03-18 22:43:48 +08:00
164 changed files with 15982 additions and 825 deletions
+2
View File
@@ -17,6 +17,7 @@ jobs:
go-version-file: backend/go.mod
check-latest: false
cache: true
cache-dependency-path: backend/go.sum
- name: Verify Go version
run: |
go version | grep -q 'go1.26.1'
@@ -36,6 +37,7 @@ jobs:
go-version-file: backend/go.mod
check-latest: false
cache: true
cache-dependency-path: backend/go.sum
- name: Verify Go version
run: |
go version | grep -q 'go1.26.1'
+10 -1
View File
@@ -78,6 +78,7 @@ Desktop.ini
# ===================
tmp/
temp/
logs/
*.tmp
*.temp
*.log
@@ -127,9 +128,17 @@ deploy/docker-compose.override.yml
.gocache/
vite.config.js
docs/*
!docs/ACCOUNT_SCHEDULING_FLOW.md
.serena/
# ===================
# 压测工具
# ===================
tools/loadtest/
# Antigravity Manager
Antigravity-Manager/
antigravity_projectid_fix.patch
.codex/
frontend/coverage/
aicodex
output/
+1303
View File
File diff suppressed because it is too large Load Diff
+1303
View File
File diff suppressed because it is too large Load Diff
+1 -1
View File
@@ -1 +1 @@
0.1.88
0.1.102.5
+24 -8
View File
@@ -144,7 +144,7 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) {
rpmCache := repository.NewRPMCache(redisClient)
groupCapacityService := service.NewGroupCapacityService(accountRepository, groupRepository, concurrencyService, sessionLimitCache, rpmCache)
groupHandler := admin.NewGroupHandler(adminService, dashboardService, groupCapacityService)
accountHandler := admin.NewAccountHandler(adminService, oAuthService, openAIOAuthService, geminiOAuthService, antigravityOAuthService, rateLimitService, accountUsageService, accountTestService, concurrencyService, crsSyncService, sessionLimitCache, rpmCache, compositeTokenCacheInvalidator)
accountHandler := admin.NewAccountHandler(adminService, oAuthService, openAIOAuthService, geminiOAuthService, antigravityOAuthService, rateLimitService, accountUsageService, accountTestService, concurrencyService, crsSyncService, sessionLimitCache, rpmCache, compositeTokenCacheInvalidator, gatewayCache)
adminAnnouncementHandler := admin.NewAnnouncementHandler(announcementService)
dataManagementService := service.NewDataManagementService()
dataManagementHandler := admin.NewDataManagementHandler(dataManagementService)
@@ -177,11 +177,15 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) {
opsSystemLogSink := service.ProvideOpsSystemLogSink(opsRepository)
opsService := service.NewOpsService(opsRepository, settingRepository, configConfig, accountRepository, userRepository, concurrencyService, gatewayService, openAIGatewayService, geminiMessagesCompatService, antigravityGatewayService, opsSystemLogSink)
soraS3Storage := service.NewSoraS3Storage(settingService)
settingService.SetOnS3UpdateCallback(soraS3Storage.RefreshClient)
soraGDriveStorage := service.NewSoraGDriveStorage(settingService)
soraStorageRouter := service.NewSoraStorageRouter(settingService, soraS3Storage, soraGDriveStorage)
settingService.SetOnS3UpdateCallback(soraStorageRouter.RefreshAll)
soraGenerationRepository := repository.NewSoraGenerationRepository(db)
soraQuotaService := service.NewSoraQuotaService(userRepository, groupRepository, settingService)
soraGenerationService := service.NewSoraGenerationService(soraGenerationRepository, soraS3Storage, soraQuotaService)
settingHandler := admin.NewSettingHandler(settingService, emailService, turnstileService, opsService, soraS3Storage)
soraGenerationService := service.NewSoraGenerationService(soraGenerationRepository, soraStorageRouter, soraQuotaService)
settingHandler := admin.NewSettingHandler(settingService, emailService, turnstileService, opsService, soraS3Storage, soraGDriveStorage, soraGenerationService)
soraGDriveOAuthService := service.NewSoraGDriveOAuthService(settingService)
gdriveOAuthHandler := admin.NewGDriveOAuthHandler(settingService, soraGDriveOAuthService, soraGDriveStorage)
opsHandler := admin.NewOpsHandler(opsService)
updateCache := repository.NewUpdateCache(redisClient)
gitHubReleaseClient := repository.ProvideGitHubReleaseClient(configConfig)
@@ -207,7 +211,7 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) {
scheduledTestResultRepository := repository.NewScheduledTestResultRepository(db)
scheduledTestService := service.ProvideScheduledTestService(scheduledTestPlanRepository, scheduledTestResultRepository)
scheduledTestHandler := admin.NewScheduledTestHandler(scheduledTestService)
adminHandlers := handler.ProvideAdminHandlers(dashboardHandler, adminUserHandler, groupHandler, accountHandler, adminAnnouncementHandler, dataManagementHandler, backupHandler, oAuthHandler, openAIOAuthHandler, geminiOAuthHandler, antigravityOAuthHandler, proxyHandler, adminRedeemHandler, promoHandler, settingHandler, opsHandler, systemHandler, adminSubscriptionHandler, adminUsageHandler, userAttributeHandler, errorPassthroughHandler, adminAPIKeyHandler, scheduledTestHandler)
adminHandlers := handler.ProvideAdminHandlers(dashboardHandler, adminUserHandler, groupHandler, accountHandler, adminAnnouncementHandler, dataManagementHandler, backupHandler, oAuthHandler, openAIOAuthHandler, geminiOAuthHandler, antigravityOAuthHandler, gdriveOAuthHandler, proxyHandler, adminRedeemHandler, promoHandler, settingHandler, opsHandler, systemHandler, adminSubscriptionHandler, adminUsageHandler, userAttributeHandler, errorPassthroughHandler, adminAPIKeyHandler, scheduledTestHandler)
usageRecordWorkerPool := service.NewUsageRecordWorkerPool(configConfig)
userMsgQueueCache := repository.NewUserMsgQueueCache(redisClient)
userMessageQueueService := service.ProvideUserMessageQueueService(userMsgQueueCache, rpmCache, configConfig)
@@ -216,13 +220,18 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) {
soraSDKClient := service.ProvideSoraSDKClient(configConfig, httpUpstream, openAITokenProvider, accountRepository, soraAccountRepository)
soraMediaStorage := service.ProvideSoraMediaStorage(configConfig)
soraGatewayService := service.NewSoraGatewayService(soraSDKClient, rateLimitService, httpUpstream, configConfig)
soraClientHandler := handler.NewSoraClientHandler(soraGenerationService, soraQuotaService, soraS3Storage, soraGatewayService, gatewayService, soraMediaStorage, apiKeyService)
soraClientHandler := handler.NewSoraClientHandler(soraGenerationService, soraQuotaService, soraStorageRouter, soraGatewayService, gatewayService, soraMediaStorage, apiKeyService)
soraTaskRepository := repository.NewSoraTaskRepository(db)
soraTaskService := service.NewSoraTaskService(soraTaskRepository, accountRepository, soraSDKClient, httpUpstream)
soraTaskWorker := service.NewSoraTaskWorker(soraTaskService, accountRepository, soraStorageRouter, soraMediaStorage, 60*time.Second)
soraTaskWorker.Start()
soraVideosHandler := handler.NewSoraVideosHandler(soraTaskService, gatewayService, soraStorageRouter, soraMediaStorage, soraGatewayService)
soraGatewayHandler := handler.NewSoraGatewayHandler(gatewayService, soraGatewayService, concurrencyService, billingCacheService, usageRecordWorkerPool, configConfig)
handlerSettingHandler := handler.ProvideSettingHandler(settingService, buildInfo)
totpHandler := handler.NewTotpHandler(totpService)
idempotencyCoordinator := service.ProvideIdempotencyCoordinator(idempotencyRepository, configConfig)
idempotencyCleanupService := service.ProvideIdempotencyCleanupService(idempotencyRepository, configConfig)
handlers := handler.ProvideHandlers(authHandler, userHandler, apiKeyHandler, usageHandler, redeemHandler, subscriptionHandler, announcementHandler, adminHandlers, gatewayHandler, openAIGatewayHandler, soraGatewayHandler, soraClientHandler, handlerSettingHandler, totpHandler, idempotencyCoordinator, idempotencyCleanupService)
handlers := handler.ProvideHandlers(authHandler, userHandler, apiKeyHandler, usageHandler, redeemHandler, subscriptionHandler, announcementHandler, adminHandlers, gatewayHandler, openAIGatewayHandler, soraGatewayHandler, soraClientHandler, soraVideosHandler, handlerSettingHandler, totpHandler, idempotencyCoordinator, idempotencyCleanupService)
jwtAuthMiddleware := middleware.NewJWTAuthMiddleware(authService, userService)
adminAuthMiddleware := middleware.NewAdminAuthMiddleware(authService, userService, settingService)
apiKeyAuthMiddleware := middleware.NewAPIKeyAuthMiddleware(apiKeyService, subscriptionService, configConfig)
@@ -238,7 +247,7 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) {
accountExpiryService := service.ProvideAccountExpiryService(accountRepository)
subscriptionExpiryService := service.ProvideSubscriptionExpiryService(userSubscriptionRepository)
scheduledTestRunnerService := service.ProvideScheduledTestRunnerService(scheduledTestPlanRepository, scheduledTestService, accountTestService, rateLimitService, configConfig)
v := provideCleanup(client, redisClient, opsMetricsCollector, opsAggregationService, opsAlertEvaluatorService, opsCleanupService, opsScheduledReportService, opsSystemLogSink, soraMediaCleanupService, schedulerSnapshotService, tokenRefreshService, accountExpiryService, subscriptionExpiryService, usageCleanupService, idempotencyCleanupService, pricingService, emailQueueService, billingCacheService, usageRecordWorkerPool, subscriptionService, oAuthService, openAIOAuthService, geminiOAuthService, antigravityOAuthService, openAIGatewayService, scheduledTestRunnerService, backupService)
v := provideCleanup(client, redisClient, opsMetricsCollector, opsAggregationService, opsAlertEvaluatorService, opsCleanupService, opsScheduledReportService, opsSystemLogSink, soraMediaCleanupService, schedulerSnapshotService, tokenRefreshService, accountExpiryService, subscriptionExpiryService, usageCleanupService, idempotencyCleanupService, pricingService, emailQueueService, billingCacheService, usageRecordWorkerPool, subscriptionService, oAuthService, openAIOAuthService, geminiOAuthService, antigravityOAuthService, openAIGatewayService, scheduledTestRunnerService, backupService, soraTaskWorker)
application := &Application{
Server: httpServer,
Cleanup: v,
@@ -292,6 +301,7 @@ func provideCleanup(
openAIGateway *service.OpenAIGatewayService,
scheduledTestRunner *service.ScheduledTestRunnerService,
backupSvc *service.BackupService,
soraTaskWorker *service.SoraTaskWorker,
) func() {
return func() {
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
@@ -433,6 +443,12 @@ func provideCleanup(
}
return nil
}},
{"SoraTaskWorker", func() error {
if soraTaskWorker != nil {
soraTaskWorker.Stop()
}
return nil
}},
}
infraSteps := []cleanupStep{
+1
View File
@@ -76,6 +76,7 @@ func TestProvideCleanup_WithMinimalDependencies_NoPanic(t *testing.T) {
nil, // openAIGateway
nil, // scheduledTestRunner
nil, // backupSvc
nil, // soraTaskWorker
)
require.NotPanics(t, func() {
+20 -9
View File
@@ -62,26 +62,28 @@ type Group struct {
SoraVideoPricePerRequestHd *float64 `json:"sora_video_price_per_request_hd,omitempty"`
// SoraStorageQuotaBytes holds the value of the "sora_storage_quota_bytes" field.
SoraStorageQuotaBytes int64 `json:"sora_storage_quota_bytes,omitempty"`
// 是否仅允许 Claude Code 客户端
// allow Claude Code client only
ClaudeCodeOnly bool `json:"claude_code_only,omitempty"`
// 非 Claude Code 请求降级使用的分组 ID
// fallback group for non-Claude-Code requests
FallbackGroupID *int64 `json:"fallback_group_id,omitempty"`
// 无效请求兜底使用的分组 ID
// fallback group for invalid request
FallbackGroupIDOnInvalidRequest *int64 `json:"fallback_group_id_on_invalid_request,omitempty"`
// 模型路由配置:模型模式 -> 优先账号ID列表
// model routing config: pattern -> account ids
ModelRouting map[string][]int64 `json:"model_routing,omitempty"`
// 是否启用模型路由配置
// whether model routing is enabled
ModelRoutingEnabled bool `json:"model_routing_enabled,omitempty"`
// 是否注入 MCP XML 调用协议提示词(仅 antigravity 平台)
// whether MCP XML prompt injection is enabled
McpXMLInject bool `json:"mcp_xml_inject,omitempty"`
// 支持的模型系列:claude, gemini_text, gemini_image
// supported model scopes: claude, gemini_text, gemini_image
SupportedModelScopes []string `json:"supported_model_scopes,omitempty"`
// 分组显示排序,数值越小越靠前
// group display order, lower comes first
SortOrder int `json:"sort_order,omitempty"`
// 是否允许 /v1/messages 调度到此 OpenAI 分组
AllowMessagesDispatch bool `json:"allow_messages_dispatch,omitempty"`
// 默认映射模型 ID,当账号级映射找不到时使用此值
DefaultMappedModel string `json:"default_mapped_model,omitempty"`
// simulate claude usage as claude-max style (1h cache write)
SimulateClaudeMaxEnabled bool `json:"simulate_claude_max_enabled,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"`
@@ -190,7 +192,7 @@ func (*Group) scanValues(columns []string) ([]any, error) {
switch columns[i] {
case group.FieldModelRouting, group.FieldSupportedModelScopes:
values[i] = new([]byte)
case group.FieldIsExclusive, group.FieldClaudeCodeOnly, group.FieldModelRoutingEnabled, group.FieldMcpXMLInject, group.FieldAllowMessagesDispatch:
case group.FieldIsExclusive, group.FieldClaudeCodeOnly, group.FieldModelRoutingEnabled, group.FieldMcpXMLInject, group.FieldAllowMessagesDispatch, group.FieldSimulateClaudeMaxEnabled:
values[i] = new(sql.NullBool)
case group.FieldRateMultiplier, group.FieldDailyLimitUsd, group.FieldWeeklyLimitUsd, group.FieldMonthlyLimitUsd, group.FieldImagePrice1k, group.FieldImagePrice2k, group.FieldImagePrice4k, group.FieldSoraImagePrice360, group.FieldSoraImagePrice540, group.FieldSoraVideoPricePerRequest, group.FieldSoraVideoPricePerRequestHd:
values[i] = new(sql.NullFloat64)
@@ -431,6 +433,12 @@ func (_m *Group) assignValues(columns []string, values []any) error {
} else if value.Valid {
_m.DefaultMappedModel = value.String
}
case group.FieldSimulateClaudeMaxEnabled:
if value, ok := values[i].(*sql.NullBool); !ok {
return fmt.Errorf("unexpected type %T for field simulate_claude_max_enabled", values[i])
} else if value.Valid {
_m.SimulateClaudeMaxEnabled = value.Bool
}
default:
_m.selectValues.Set(columns[i], values[i])
}
@@ -630,6 +638,9 @@ func (_m *Group) String() string {
builder.WriteString(", ")
builder.WriteString("default_mapped_model=")
builder.WriteString(_m.DefaultMappedModel)
builder.WriteString(", ")
builder.WriteString("simulate_claude_max_enabled=")
builder.WriteString(fmt.Sprintf("%v", _m.SimulateClaudeMaxEnabled))
builder.WriteByte(')')
return builder.String()
}
+10
View File
@@ -79,6 +79,8 @@ const (
FieldAllowMessagesDispatch = "allow_messages_dispatch"
// FieldDefaultMappedModel holds the string denoting the default_mapped_model field in the database.
FieldDefaultMappedModel = "default_mapped_model"
// FieldSimulateClaudeMaxEnabled holds the string denoting the simulate_claude_max_enabled field in the database.
FieldSimulateClaudeMaxEnabled = "simulate_claude_max_enabled"
// 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.
@@ -186,6 +188,7 @@ var Columns = []string{
FieldSortOrder,
FieldAllowMessagesDispatch,
FieldDefaultMappedModel,
FieldSimulateClaudeMaxEnabled,
}
var (
@@ -259,6 +262,8 @@ var (
DefaultDefaultMappedModel string
// DefaultMappedModelValidator is a validator for the "default_mapped_model" field. It is called by the builders before save.
DefaultMappedModelValidator func(string) error
// DefaultSimulateClaudeMaxEnabled holds the default value on creation for the "simulate_claude_max_enabled" field.
DefaultSimulateClaudeMaxEnabled bool
)
// OrderOption defines the ordering options for the Group queries.
@@ -419,6 +424,11 @@ func ByDefaultMappedModel(opts ...sql.OrderTermOption) OrderOption {
return sql.OrderByField(FieldDefaultMappedModel, opts...).ToFunc()
}
// BySimulateClaudeMaxEnabled orders the results by the simulate_claude_max_enabled field.
func BySimulateClaudeMaxEnabled(opts ...sql.OrderTermOption) OrderOption {
return sql.OrderByField(FieldSimulateClaudeMaxEnabled, opts...).ToFunc()
}
// ByAPIKeysCount orders the results by api_keys count.
func ByAPIKeysCount(opts ...sql.OrderTermOption) OrderOption {
return func(s *sql.Selector) {
+15
View File
@@ -205,6 +205,11 @@ func DefaultMappedModel(v string) predicate.Group {
return predicate.Group(sql.FieldEQ(FieldDefaultMappedModel, v))
}
// SimulateClaudeMaxEnabled applies equality check predicate on the "simulate_claude_max_enabled" field. It's identical to SimulateClaudeMaxEnabledEQ.
func SimulateClaudeMaxEnabled(v bool) predicate.Group {
return predicate.Group(sql.FieldEQ(FieldSimulateClaudeMaxEnabled, 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))
@@ -1555,6 +1560,16 @@ func DefaultMappedModelContainsFold(v string) predicate.Group {
return predicate.Group(sql.FieldContainsFold(FieldDefaultMappedModel, v))
}
// SimulateClaudeMaxEnabledEQ applies the EQ predicate on the "simulate_claude_max_enabled" field.
func SimulateClaudeMaxEnabledEQ(v bool) predicate.Group {
return predicate.Group(sql.FieldEQ(FieldSimulateClaudeMaxEnabled, v))
}
// SimulateClaudeMaxEnabledNEQ applies the NEQ predicate on the "simulate_claude_max_enabled" field.
func SimulateClaudeMaxEnabledNEQ(v bool) predicate.Group {
return predicate.Group(sql.FieldNEQ(FieldSimulateClaudeMaxEnabled, v))
}
// HasAPIKeys applies the HasEdge predicate on the "api_keys" edge.
func HasAPIKeys() predicate.Group {
return predicate.Group(func(s *sql.Selector) {
+65
View File
@@ -452,6 +452,20 @@ func (_c *GroupCreate) SetNillableDefaultMappedModel(v *string) *GroupCreate {
return _c
}
// SetSimulateClaudeMaxEnabled sets the "simulate_claude_max_enabled" field.
func (_c *GroupCreate) SetSimulateClaudeMaxEnabled(v bool) *GroupCreate {
_c.mutation.SetSimulateClaudeMaxEnabled(v)
return _c
}
// SetNillableSimulateClaudeMaxEnabled sets the "simulate_claude_max_enabled" field if the given value is not nil.
func (_c *GroupCreate) SetNillableSimulateClaudeMaxEnabled(v *bool) *GroupCreate {
if v != nil {
_c.SetSimulateClaudeMaxEnabled(*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...)
@@ -649,6 +663,10 @@ func (_c *GroupCreate) defaults() error {
v := group.DefaultDefaultMappedModel
_c.mutation.SetDefaultMappedModel(v)
}
if _, ok := _c.mutation.SimulateClaudeMaxEnabled(); !ok {
v := group.DefaultSimulateClaudeMaxEnabled
_c.mutation.SetSimulateClaudeMaxEnabled(v)
}
return nil
}
@@ -730,6 +748,9 @@ func (_c *GroupCreate) check() error {
return &ValidationError{Name: "default_mapped_model", err: fmt.Errorf(`ent: validator failed for field "Group.default_mapped_model": %w`, err)}
}
}
if _, ok := _c.mutation.SimulateClaudeMaxEnabled(); !ok {
return &ValidationError{Name: "simulate_claude_max_enabled", err: errors.New(`ent: missing required field "Group.simulate_claude_max_enabled"`)}
}
return nil
}
@@ -885,6 +906,10 @@ func (_c *GroupCreate) createSpec() (*Group, *sqlgraph.CreateSpec) {
_spec.SetField(group.FieldDefaultMappedModel, field.TypeString, value)
_node.DefaultMappedModel = value
}
if value, ok := _c.mutation.SimulateClaudeMaxEnabled(); ok {
_spec.SetField(group.FieldSimulateClaudeMaxEnabled, field.TypeBool, value)
_node.SimulateClaudeMaxEnabled = value
}
if nodes := _c.mutation.APIKeysIDs(); len(nodes) > 0 {
edge := &sqlgraph.EdgeSpec{
Rel: sqlgraph.O2M,
@@ -1599,6 +1624,18 @@ func (u *GroupUpsert) UpdateDefaultMappedModel() *GroupUpsert {
return u
}
// SetSimulateClaudeMaxEnabled sets the "simulate_claude_max_enabled" field.
func (u *GroupUpsert) SetSimulateClaudeMaxEnabled(v bool) *GroupUpsert {
u.Set(group.FieldSimulateClaudeMaxEnabled, v)
return u
}
// UpdateSimulateClaudeMaxEnabled sets the "simulate_claude_max_enabled" field to the value that was provided on create.
func (u *GroupUpsert) UpdateSimulateClaudeMaxEnabled() *GroupUpsert {
u.SetExcluded(group.FieldSimulateClaudeMaxEnabled)
return u
}
// UpdateNewValues updates the mutable fields using the new values that were set on create.
// Using this option is equivalent to using:
//
@@ -2295,6 +2332,20 @@ func (u *GroupUpsertOne) UpdateDefaultMappedModel() *GroupUpsertOne {
})
}
// SetSimulateClaudeMaxEnabled sets the "simulate_claude_max_enabled" field.
func (u *GroupUpsertOne) SetSimulateClaudeMaxEnabled(v bool) *GroupUpsertOne {
return u.Update(func(s *GroupUpsert) {
s.SetSimulateClaudeMaxEnabled(v)
})
}
// UpdateSimulateClaudeMaxEnabled sets the "simulate_claude_max_enabled" field to the value that was provided on create.
func (u *GroupUpsertOne) UpdateSimulateClaudeMaxEnabled() *GroupUpsertOne {
return u.Update(func(s *GroupUpsert) {
s.UpdateSimulateClaudeMaxEnabled()
})
}
// Exec executes the query.
func (u *GroupUpsertOne) Exec(ctx context.Context) error {
if len(u.create.conflict) == 0 {
@@ -3157,6 +3208,20 @@ func (u *GroupUpsertBulk) UpdateDefaultMappedModel() *GroupUpsertBulk {
})
}
// SetSimulateClaudeMaxEnabled sets the "simulate_claude_max_enabled" field.
func (u *GroupUpsertBulk) SetSimulateClaudeMaxEnabled(v bool) *GroupUpsertBulk {
return u.Update(func(s *GroupUpsert) {
s.SetSimulateClaudeMaxEnabled(v)
})
}
// UpdateSimulateClaudeMaxEnabled sets the "simulate_claude_max_enabled" field to the value that was provided on create.
func (u *GroupUpsertBulk) UpdateSimulateClaudeMaxEnabled() *GroupUpsertBulk {
return u.Update(func(s *GroupUpsert) {
s.UpdateSimulateClaudeMaxEnabled()
})
}
// Exec executes the query.
func (u *GroupUpsertBulk) Exec(ctx context.Context) error {
if u.create.err != nil {
+34
View File
@@ -653,6 +653,20 @@ func (_u *GroupUpdate) SetNillableDefaultMappedModel(v *string) *GroupUpdate {
return _u
}
// SetSimulateClaudeMaxEnabled sets the "simulate_claude_max_enabled" field.
func (_u *GroupUpdate) SetSimulateClaudeMaxEnabled(v bool) *GroupUpdate {
_u.mutation.SetSimulateClaudeMaxEnabled(v)
return _u
}
// SetNillableSimulateClaudeMaxEnabled sets the "simulate_claude_max_enabled" field if the given value is not nil.
func (_u *GroupUpdate) SetNillableSimulateClaudeMaxEnabled(v *bool) *GroupUpdate {
if v != nil {
_u.SetSimulateClaudeMaxEnabled(*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...)
@@ -1149,6 +1163,9 @@ func (_u *GroupUpdate) sqlSave(ctx context.Context) (_node int, err error) {
if value, ok := _u.mutation.DefaultMappedModel(); ok {
_spec.SetField(group.FieldDefaultMappedModel, field.TypeString, value)
}
if value, ok := _u.mutation.SimulateClaudeMaxEnabled(); ok {
_spec.SetField(group.FieldSimulateClaudeMaxEnabled, field.TypeBool, value)
}
if _u.mutation.APIKeysCleared() {
edge := &sqlgraph.EdgeSpec{
Rel: sqlgraph.O2M,
@@ -2081,6 +2098,20 @@ func (_u *GroupUpdateOne) SetNillableDefaultMappedModel(v *string) *GroupUpdateO
return _u
}
// SetSimulateClaudeMaxEnabled sets the "simulate_claude_max_enabled" field.
func (_u *GroupUpdateOne) SetSimulateClaudeMaxEnabled(v bool) *GroupUpdateOne {
_u.mutation.SetSimulateClaudeMaxEnabled(v)
return _u
}
// SetNillableSimulateClaudeMaxEnabled sets the "simulate_claude_max_enabled" field if the given value is not nil.
func (_u *GroupUpdateOne) SetNillableSimulateClaudeMaxEnabled(v *bool) *GroupUpdateOne {
if v != nil {
_u.SetSimulateClaudeMaxEnabled(*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...)
@@ -2607,6 +2638,9 @@ func (_u *GroupUpdateOne) sqlSave(ctx context.Context) (_node *Group, err error)
if value, ok := _u.mutation.DefaultMappedModel(); ok {
_spec.SetField(group.FieldDefaultMappedModel, field.TypeString, value)
}
if value, ok := _u.mutation.SimulateClaudeMaxEnabled(); ok {
_spec.SetField(group.FieldSimulateClaudeMaxEnabled, field.TypeBool, value)
}
if _u.mutation.APIKeysCleared() {
edge := &sqlgraph.EdgeSpec{
Rel: sqlgraph.O2M,
+1
View File
@@ -410,6 +410,7 @@ var (
{Name: "sort_order", Type: field.TypeInt, Default: 0},
{Name: "allow_messages_dispatch", Type: field.TypeBool, Default: false},
{Name: "default_mapped_model", Type: field.TypeString, Size: 100, Default: ""},
{Name: "simulate_claude_max_enabled", Type: field.TypeBool, Default: false},
}
// GroupsTable holds the schema information for the "groups" table.
GroupsTable = &schema.Table{
+55 -1
View File
@@ -8252,6 +8252,7 @@ type GroupMutation struct {
addsort_order *int
allow_messages_dispatch *bool
default_mapped_model *string
simulate_claude_max_enabled *bool
clearedFields map[string]struct{}
api_keys map[int64]struct{}
removedapi_keys map[int64]struct{}
@@ -10068,6 +10069,42 @@ func (m *GroupMutation) ResetDefaultMappedModel() {
m.default_mapped_model = nil
}
// SetSimulateClaudeMaxEnabled sets the "simulate_claude_max_enabled" field.
func (m *GroupMutation) SetSimulateClaudeMaxEnabled(b bool) {
m.simulate_claude_max_enabled = &b
}
// SimulateClaudeMaxEnabled returns the value of the "simulate_claude_max_enabled" field in the mutation.
func (m *GroupMutation) SimulateClaudeMaxEnabled() (r bool, exists bool) {
v := m.simulate_claude_max_enabled
if v == nil {
return
}
return *v, true
}
// OldSimulateClaudeMaxEnabled returns the old "simulate_claude_max_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) OldSimulateClaudeMaxEnabled(ctx context.Context) (v bool, err error) {
if !m.op.Is(OpUpdateOne) {
return v, errors.New("OldSimulateClaudeMaxEnabled is only allowed on UpdateOne operations")
}
if m.id == nil || m.oldValue == nil {
return v, errors.New("OldSimulateClaudeMaxEnabled requires an ID field in the mutation")
}
oldValue, err := m.oldValue(ctx)
if err != nil {
return v, fmt.Errorf("querying old value for OldSimulateClaudeMaxEnabled: %w", err)
}
return oldValue.SimulateClaudeMaxEnabled, nil
}
// ResetSimulateClaudeMaxEnabled resets all changes to the "simulate_claude_max_enabled" field.
func (m *GroupMutation) ResetSimulateClaudeMaxEnabled() {
m.simulate_claude_max_enabled = nil
}
// AddAPIKeyIDs adds the "api_keys" edge to the APIKey entity by ids.
func (m *GroupMutation) AddAPIKeyIDs(ids ...int64) {
if m.api_keys == nil {
@@ -10426,7 +10463,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, 32)
fields := make([]string, 0, 33)
if m.created_at != nil {
fields = append(fields, group.FieldCreatedAt)
}
@@ -10523,6 +10560,9 @@ func (m *GroupMutation) Fields() []string {
if m.default_mapped_model != nil {
fields = append(fields, group.FieldDefaultMappedModel)
}
if m.simulate_claude_max_enabled != nil {
fields = append(fields, group.FieldSimulateClaudeMaxEnabled)
}
return fields
}
@@ -10595,6 +10635,8 @@ func (m *GroupMutation) Field(name string) (ent.Value, bool) {
return m.AllowMessagesDispatch()
case group.FieldDefaultMappedModel:
return m.DefaultMappedModel()
case group.FieldSimulateClaudeMaxEnabled:
return m.SimulateClaudeMaxEnabled()
}
return nil, false
}
@@ -10668,6 +10710,8 @@ func (m *GroupMutation) OldField(ctx context.Context, name string) (ent.Value, e
return m.OldAllowMessagesDispatch(ctx)
case group.FieldDefaultMappedModel:
return m.OldDefaultMappedModel(ctx)
case group.FieldSimulateClaudeMaxEnabled:
return m.OldSimulateClaudeMaxEnabled(ctx)
}
return nil, fmt.Errorf("unknown Group field %s", name)
}
@@ -10901,6 +10945,13 @@ func (m *GroupMutation) SetField(name string, value ent.Value) error {
}
m.SetDefaultMappedModel(v)
return nil
case group.FieldSimulateClaudeMaxEnabled:
v, ok := value.(bool)
if !ok {
return fmt.Errorf("unexpected type %T for field %s", value, name)
}
m.SetSimulateClaudeMaxEnabled(v)
return nil
}
return fmt.Errorf("unknown Group field %s", name)
}
@@ -11334,6 +11385,9 @@ func (m *GroupMutation) ResetField(name string) error {
case group.FieldDefaultMappedModel:
m.ResetDefaultMappedModel()
return nil
case group.FieldSimulateClaudeMaxEnabled:
m.ResetSimulateClaudeMaxEnabled()
return nil
}
return fmt.Errorf("unknown Group field %s", name)
}
+4
View File
@@ -463,6 +463,10 @@ func init() {
group.DefaultDefaultMappedModel = groupDescDefaultMappedModel.Default.(string)
// group.DefaultMappedModelValidator is a validator for the "default_mapped_model" field. It is called by the builders before save.
group.DefaultMappedModelValidator = groupDescDefaultMappedModel.Validators[0].(func(string) error)
// groupDescSimulateClaudeMaxEnabled is the schema descriptor for simulate_claude_max_enabled field.
groupDescSimulateClaudeMaxEnabled := groupFields[29].Descriptor()
// group.DefaultSimulateClaudeMaxEnabled holds the default value on creation for the simulate_claude_max_enabled field.
group.DefaultSimulateClaudeMaxEnabled = groupDescSimulateClaudeMaxEnabled.Default.(bool)
idempotencyrecordMixin := schema.IdempotencyRecord{}.Mixin()
idempotencyrecordMixinFields0 := idempotencyrecordMixin[0].Fields()
_ = idempotencyrecordMixinFields0
+11 -23
View File
@@ -33,8 +33,6 @@ func (Group) Mixin() []ent.Mixin {
func (Group) Fields() []ent.Field {
return []ent.Field{
// 唯一约束通过部分索引实现(WHERE deleted_at IS NULL),支持软删除后重用
// 见迁移文件 016_soft_delete_partial_unique_indexes.sql
field.String("name").
MaxLen(100).
NotEmpty(),
@@ -51,7 +49,6 @@ func (Group) Fields() []ent.Field {
MaxLen(20).
Default(domain.StatusActive),
// Subscription-related fields (added by migration 003)
field.String("platform").
MaxLen(50).
Default(domain.PlatformAnthropic),
@@ -73,7 +70,6 @@ func (Group) Fields() []ent.Field {
field.Int("default_validity_days").
Default(30),
// 图片生成计费配置(antigravity 和 gemini 平台使用)
field.Float("image_price_1k").
Optional().
Nillable().
@@ -87,7 +83,6 @@ func (Group) Fields() []ent.Field {
Nillable().
SchemaType(map[string]string{dialect.Postgres: "decimal(20,8)"}),
// Sora 按次计费配置(阶段 1)
field.Float("sora_image_price_360").
Optional().
Nillable().
@@ -109,45 +104,38 @@ func (Group) Fields() []ent.Field {
field.Int64("sora_storage_quota_bytes").
Default(0),
// Claude Code 客户端限制 (added by migration 029)
field.Bool("claude_code_only").
Default(false).
Comment("是否仅允许 Claude Code 客户端"),
Comment("allow Claude Code client only"),
field.Int64("fallback_group_id").
Optional().
Nillable().
Comment("非 Claude Code 请求降级使用的分组 ID"),
Comment("fallback group for non-Claude-Code requests"),
field.Int64("fallback_group_id_on_invalid_request").
Optional().
Nillable().
Comment("无效请求兜底使用的分组 ID"),
Comment("fallback group for invalid request"),
// 模型路由配置 (added by migration 040)
field.JSON("model_routing", map[string][]int64{}).
Optional().
SchemaType(map[string]string{dialect.Postgres: "jsonb"}).
Comment("模型路由配置:模型模式 -> 优先账号ID列表"),
// 模型路由开关 (added by migration 041)
Comment("model routing config: pattern -> account ids"),
field.Bool("model_routing_enabled").
Default(false).
Comment("是否启用模型路由配置"),
Comment("whether model routing is enabled"),
// MCP XML 协议注入开关 (added by migration 042)
field.Bool("mcp_xml_inject").
Default(true).
Comment("是否注入 MCP XML 调用协议提示词(仅 antigravity 平台)"),
Comment("whether MCP XML prompt injection is enabled"),
// 支持的模型系列 (added by migration 046)
field.JSON("supported_model_scopes", []string{}).
Default([]string{"claude", "gemini_text", "gemini_image"}).
SchemaType(map[string]string{dialect.Postgres: "jsonb"}).
Comment("支持的模型系列:claude, gemini_text, gemini_image"),
Comment("supported model scopes: claude, gemini_text, gemini_image"),
// 分组排序 (added by migration 052)
field.Int("sort_order").
Default(0).
Comment("分组显示排序,数值越小越靠前"),
Comment("group display order, lower comes first"),
// OpenAI Messages 调度配置 (added by migration 069)
field.Bool("allow_messages_dispatch").
@@ -157,6 +145,9 @@ func (Group) Fields() []ent.Field {
MaxLen(100).
Default("").
Comment("默认映射模型 ID,当账号级映射找不到时使用此值"),
field.Bool("simulate_claude_max_enabled").
Default(false).
Comment("simulate claude usage as claude-max style (1h cache write)"),
}
}
@@ -172,14 +163,11 @@ func (Group) Edges() []ent.Edge {
edge.From("allowed_users", User.Type).
Ref("allowed_groups").
Through("user_allowed_groups", UserAllowedGroup.Type),
// 注意:fallback_group_id 直接作为字段使用,不定义 edge
// 这样允许多个分组指向同一个降级分组(M2O 关系)
}
}
func (Group) Indexes() []ent.Index {
return []ent.Index{
// name 字段已在 Fields() 中声明 Unique(),无需重复索引
index.Fields("status"),
index.Fields("platform"),
index.Fields("subscription_type"),
+14 -1
View File
@@ -22,6 +22,8 @@ require (
github.com/imroc/req/v3 v3.57.0
github.com/lib/pq v1.10.9
github.com/patrickmn/go-cache v2.1.0+incompatible
github.com/pkoukk/tiktoken-go v0.1.8
github.com/pkoukk/tiktoken-go-loader v0.0.2
github.com/pquerna/otp v1.5.0
github.com/redis/go-redis/v9 v9.17.2
github.com/refraction-networking/utls v1.8.2
@@ -37,8 +39,10 @@ require (
go.uber.org/zap v1.24.0
golang.org/x/crypto v0.48.0
golang.org/x/net v0.49.0
golang.org/x/oauth2 v0.30.0
golang.org/x/sync v0.19.0
golang.org/x/term v0.40.0
google.golang.org/api v0.153.0
gopkg.in/natefinch/lumberjack.v2 v2.2.1
gopkg.in/yaml.v3 v3.0.1
modernc.org/sqlite v1.44.3
@@ -46,6 +50,7 @@ require (
require (
ariga.io/atlas v0.32.1-0.20250325101103-175b25e1c1b9 // indirect
cloud.google.com/go/compute/metadata v0.7.0 // indirect
dario.cat/mergo v1.0.2 // indirect
github.com/Azure/go-ansiterm v0.0.0-20210617225240-d185dfc1b5a1 // indirect
github.com/Microsoft/go-winio v0.6.2 // indirect
@@ -87,6 +92,7 @@ require (
github.com/davecgh/go-spew v1.1.2-0.20180830191138-d8f796af33cc // indirect
github.com/dgryski/go-rendezvous v0.0.0-20200823014737-9f7001d12a5f // indirect
github.com/distribution/reference v0.6.0 // indirect
github.com/dlclark/regexp2 v1.10.0 // indirect
github.com/docker/docker v28.5.1+incompatible // indirect
github.com/docker/go-connections v0.6.0 // indirect
github.com/docker/go-units v0.5.0 // indirect
@@ -105,8 +111,13 @@ require (
github.com/go-playground/universal-translator v0.18.1 // indirect
github.com/go-playground/validator/v10 v10.14.0 // indirect
github.com/goccy/go-json v0.10.2 // indirect
github.com/golang/groupcache v0.0.0-20210331224755-41bb18bfe9da // indirect
github.com/golang/protobuf v1.5.4 // indirect
github.com/google/go-cmp v0.7.0 // indirect
github.com/google/go-querystring v1.1.0 // indirect
github.com/google/s2a-go v0.1.7 // indirect
github.com/googleapis/enterprise-certificate-proxy v0.3.2 // indirect
github.com/googleapis/gax-go/v2 v2.12.0 // indirect
github.com/grpc-ecosystem/grpc-gateway/v2 v2.27.3 // indirect
github.com/hashicorp/hcl v1.0.0 // indirect
github.com/hashicorp/hcl/v2 v2.18.1 // indirect
@@ -162,11 +173,11 @@ require (
github.com/yusufpapurcu/wmi v1.2.4 // indirect
github.com/zclconf/go-cty v1.14.4 // indirect
github.com/zclconf/go-cty-yaml v1.1.0 // indirect
go.opencensus.io v0.24.0 // indirect
go.opentelemetry.io/auto/sdk v1.1.0 // indirect
go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp v0.49.0 // indirect
go.opentelemetry.io/otel v1.37.0 // indirect
go.opentelemetry.io/otel/metric v1.37.0 // indirect
go.opentelemetry.io/otel/sdk v1.37.0 // indirect
go.opentelemetry.io/otel/trace v1.37.0 // indirect
go.uber.org/atomic v1.10.0 // indirect
go.uber.org/automaxprocs v1.6.0 // indirect
@@ -176,6 +187,8 @@ require (
golang.org/x/mod v0.32.0 // indirect
golang.org/x/sys v0.41.0 // indirect
golang.org/x/text v0.34.0 // indirect
google.golang.org/genproto/googleapis/api v0.0.0-20250929231259-57b25ae835d4 // indirect
google.golang.org/genproto/googleapis/rpc v0.0.0-20250929231259-57b25ae835d4 // indirect
google.golang.org/grpc v1.75.1 // indirect
google.golang.org/protobuf v1.36.10 // indirect
gopkg.in/ini.v1 v1.67.0 // indirect
+104 -1
View File
@@ -1,5 +1,8 @@
ariga.io/atlas v0.32.1-0.20250325101103-175b25e1c1b9 h1:E0wvcUXTkgyN4wy4LGtNzMNGMytJN8afmIWXJVMi4cc=
ariga.io/atlas v0.32.1-0.20250325101103-175b25e1c1b9/go.mod h1:Oe1xWPuu5q9LzyrWfbZmEZxFYeu4BHTyzfjeW2aZp/w=
cloud.google.com/go v0.26.0/go.mod h1:aQUYkXzVsufM+DwF1aE+0xfcU+56JwCaLick0ClmMTw=
cloud.google.com/go/compute/metadata v0.7.0 h1:PBWF+iiAerVNe8UCHxdOt6eHLVc3ydFeOCw78U8ytSU=
cloud.google.com/go/compute/metadata v0.7.0/go.mod h1:j5MvL9PprKL39t166CoB1uVHfQMs4tFQZZcKwksXUjo=
dario.cat/mergo v1.0.2 h1:85+piFYR1tMbRrLcDwR18y4UKJ3aH1Tbzi24VRW1TK8=
dario.cat/mergo v1.0.2/go.mod h1:E/hbnu0NxMFBjpMIE34DRGLWqDy0g5FuKDhCb31ngxA=
entgo.io/ent v0.14.5 h1:Rj2WOYJtCkWyFo6a+5wB3EfBRP0rnx1fMk6gGA0UUe4=
@@ -8,6 +11,7 @@ github.com/AdaLogics/go-fuzz-headers v0.0.0-20240806141605-e8a1dd7889d6 h1:He8af
github.com/AdaLogics/go-fuzz-headers v0.0.0-20240806141605-e8a1dd7889d6/go.mod h1:8o94RPi1/7XTJvwPpRSzSUedZrtlirdB3r9Z20bi2f8=
github.com/Azure/go-ansiterm v0.0.0-20210617225240-d185dfc1b5a1 h1:UQHMgLO+TxOElx5B5HZ4hJQsoJ/PvUvKRhJHDQXO8P8=
github.com/Azure/go-ansiterm v0.0.0-20210617225240-d185dfc1b5a1/go.mod h1:xomTg63KZ2rFqZQzSB4Vz2SUXa1BpHTVz9L5PTmPC4E=
github.com/BurntSushi/toml v0.3.1/go.mod h1:xHWCNGjB5oqiDr8zfno3MHue2Ht5sIBksp03qcyfWMU=
github.com/DATA-DOG/go-sqlmock v1.5.2 h1:OcvFkGmslmlZibjAjaHm3L//6LiuBgolP7OputlJIzU=
github.com/DATA-DOG/go-sqlmock v1.5.2/go.mod h1:88MAG/4G7SMwSE3CeA0ZKzrT5CiOU3OJ+JlNzwDqpNU=
github.com/DouDOU-start/go-sora2api v1.1.0 h1:PxWiukK77StiHxEngOFwT1rKUn9oTAJJTl07wQUXwiU=
@@ -89,11 +93,14 @@ github.com/bytedance/sonic v1.9.1 h1:6iJ6NqdoxCDr6mbY8h18oSO+cShGSMRGCEo7F2h0x8s
github.com/bytedance/sonic v1.9.1/go.mod h1:i736AoUSYt75HyZLoJW9ERYxcy6eaN6h4BZXU064P/U=
github.com/cenkalti/backoff/v4 v4.3.0 h1:MyRJ/UdXutAwSAT+s3wNd7MfTIcy71VQueUuFK343L8=
github.com/cenkalti/backoff/v4 v4.3.0/go.mod h1:Y3VNntkOUPxTVeUxJ/G5vcM//AlwfmyYozVcomhLiZE=
github.com/census-instrumentation/opencensus-proto v0.2.1/go.mod h1:f6KPmirojxKA12rnyqOA5BBL4O983OfeGPqjHWSTneU=
github.com/cespare/xxhash/v2 v2.3.0 h1:UL815xU9SqsFlibzuggzjXhog7bL6oX9BbNZnL2UFvs=
github.com/cespare/xxhash/v2 v2.3.0/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs=
github.com/chenzhuoyu/base64x v0.0.0-20211019084208-fb5309c8db06/go.mod h1:DH46F32mSOjUmXrMHnKwZdA8wcEefY7UVqBKYGjpdQY=
github.com/chenzhuoyu/base64x v0.0.0-20221115062448-fe3a3abad311 h1:qSGYFH7+jGhDF8vLC+iwCD4WpbV1EBDSzWkJODFLams=
github.com/chenzhuoyu/base64x v0.0.0-20221115062448-fe3a3abad311/go.mod h1:b583jCggY9gE99b6G5LEC39OIiVsWj+R97kbl5odCEk=
github.com/client9/misspell v0.3.4/go.mod h1:qj6jICC3Q7zFZvVWo7KLAzC3yx5G7kyvSDkc90ppPyw=
github.com/cncf/udpa/go v0.0.0-20191209042840-269d4d468f6f/go.mod h1:M8M6+tZqaGXZJjfX53e64911xZQV5JYwmTeXPW+k8Sc=
github.com/coder/websocket v1.8.14 h1:9L0p0iKiNOibykf283eHkKUHHrpG7f65OE3BhhO7v9g=
github.com/coder/websocket v1.8.14/go.mod h1:NX3SzP+inril6yawo5CQXx8+fk145lPDC6pumgx0mVg=
github.com/containerd/errdefs v1.0.0 h1:tg5yIfIlQIrxYtu9ajqY42W3lpS19XqdxRQeEwYG8PI=
@@ -120,6 +127,8 @@ github.com/dgryski/go-rendezvous v0.0.0-20200823014737-9f7001d12a5f h1:lO4WD4F/r
github.com/dgryski/go-rendezvous v0.0.0-20200823014737-9f7001d12a5f/go.mod h1:cuUVRXasLTGF7a8hSLbxyZXjz+1KgoB3wDUb6vlszIc=
github.com/distribution/reference v0.6.0 h1:0IXCQ5g4/QMHHkarYzh5l+u8T3t73zM5QvfrDyIgxBk=
github.com/distribution/reference v0.6.0/go.mod h1:BbU0aIcezP1/5jX/8MP0YiH4SdvB5Y4f/wlDRiLyi3E=
github.com/dlclark/regexp2 v1.10.0 h1:+/GIL799phkJqYW+3YbOd8LCcbHzT0Pbo8zl70MHsq0=
github.com/dlclark/regexp2 v1.10.0/go.mod h1:DHkYz0B9wPfa6wondMfaivmHpzrQ3v9q8cnmRbL6yW8=
github.com/docker/docker v28.5.1+incompatible h1:Bm8DchhSD2J6PsFzxC35TZo4TLGR2PdW/E69rU45NhM=
github.com/docker/docker v28.5.1+incompatible/go.mod h1:eEKB0N0r5NX/I1kEveEz05bcu8tLC/8azJZsviup8Sk=
github.com/docker/go-connections v0.6.0 h1:LlMG9azAe1TqfR7sO+NJttz1gy6KO7VJBh+pMmjSD94=
@@ -130,6 +139,10 @@ github.com/dustin/go-humanize v1.0.1 h1:GzkhY7T5VNhEkwH0PVJgjz+fX1rhBrR7pRT3mDkp
github.com/dustin/go-humanize v1.0.1/go.mod h1:Mu1zIs6XwVuF/gI1OepvI0qD18qycQx+mFykh5fBlto=
github.com/ebitengine/purego v0.8.4 h1:CF7LEKg5FFOsASUj0+QwaXf8Ht6TlFxg09+S9wz0omw=
github.com/ebitengine/purego v0.8.4/go.mod h1:iIjxzd6CiRiOG0UyXP+V1+jWqUXVjPKLAI0mRfJZTmQ=
github.com/envoyproxy/go-control-plane v0.9.0/go.mod h1:YTl/9mNaCwkRvm6d1a2C3ymFceY/DCBVvsKhRF0iEA4=
github.com/envoyproxy/go-control-plane v0.9.1-0.20191026205805-5f8ba28d4473/go.mod h1:YTl/9mNaCwkRvm6d1a2C3ymFceY/DCBVvsKhRF0iEA4=
github.com/envoyproxy/go-control-plane v0.9.4/go.mod h1:6rpuAdCZL397s3pYoYcLgu1mIlRU8Am5FuJP05cCM98=
github.com/envoyproxy/protoc-gen-validate v0.1.0/go.mod h1:iSmxcyjqTsJpI2R4NaDN7+kN2VEUnK/pcBlmesArF7c=
github.com/fatih/color v1.18.0 h1:S8gINlzdQ840/4pfAwic/ZE0djQEH3wM94VfqLTZcOM=
github.com/fatih/color v1.18.0/go.mod h1:4FelSpRwEGDpQ12mAdzqdOukCy4u8WUtOY6lkT/6HfU=
github.com/felixge/httpsnoop v1.0.4 h1:NFTV2Zj1bL4mc9sqWACXbQFVBBg2W3GPvqp8/ESS2Wg=
@@ -167,7 +180,29 @@ github.com/goccy/go-json v0.10.2 h1:CrxCmQqYDkv1z7lO7Wbh2HN93uovUHgrECaO5ZrCXAU=
github.com/goccy/go-json v0.10.2/go.mod h1:6MelG93GURQebXPDq3khkgXZkazVtN9CRI+MGFi0w8I=
github.com/golang-jwt/jwt/v5 v5.2.2 h1:Rl4B7itRWVtYIHFrSNd7vhTiz9UpLdi6gZhZ3wEeDy8=
github.com/golang-jwt/jwt/v5 v5.2.2/go.mod h1:pqrtFR0X4osieyHYxtmOUWsAWrfe1Q5UVIyoH402zdk=
github.com/golang/glog v0.0.0-20160126235308-23def4e6c14b/go.mod h1:SBH7ygxi8pfUlaOkMMuAQtPIUF8ecWP5IEl/CR7VP2Q=
github.com/golang/groupcache v0.0.0-20200121045136-8c9f03a8e57e/go.mod h1:cIg4eruTrX1D+g88fzRXU5OdNfaM+9IcxsU14FzY7Hc=
github.com/golang/groupcache v0.0.0-20210331224755-41bb18bfe9da h1:oI5xCqsCo564l8iNU+DwB5epxmsaqB+rhGL0m5jtYqE=
github.com/golang/groupcache v0.0.0-20210331224755-41bb18bfe9da/go.mod h1:cIg4eruTrX1D+g88fzRXU5OdNfaM+9IcxsU14FzY7Hc=
github.com/golang/mock v1.1.1/go.mod h1:oTYuIxOrZwtPieC+H1uAHpcLFnEyAGVDL/k47Jfbm0A=
github.com/golang/protobuf v1.2.0/go.mod h1:6lQm79b+lXiMfvg/cZm0SGofjICqVBUtrP5yJMmIC1U=
github.com/golang/protobuf v1.3.2/go.mod h1:6lQm79b+lXiMfvg/cZm0SGofjICqVBUtrP5yJMmIC1U=
github.com/golang/protobuf v1.4.0-rc.1/go.mod h1:ceaxUfeHdC40wWswd/P6IGgMaK3YpKi5j83Wpe3EHw8=
github.com/golang/protobuf v1.4.0-rc.1.0.20200221234624-67d41d38c208/go.mod h1:xKAWHe0F5eneWXFV3EuXVDTCmh+JuBKY0li0aMyXATA=
github.com/golang/protobuf v1.4.0-rc.2/go.mod h1:LlEzMj4AhA7rCAGe4KMBDvJI+AwstrUpVNzEA03Pprs=
github.com/golang/protobuf v1.4.0-rc.4.0.20200313231945-b860323f09d0/go.mod h1:WU3c8KckQ9AFe+yFwt9sWVRKCVIyN9cPHBJSNnbL67w=
github.com/golang/protobuf v1.4.0/go.mod h1:jodUvKwWbYaEsadDk5Fwe5c77LiNKVO9IDvqG2KuDX0=
github.com/golang/protobuf v1.4.1/go.mod h1:U8fpvMrcmy5pZrNK1lt4xCsGvpyWQ/VVv6QDs8UjoX8=
github.com/golang/protobuf v1.4.3/go.mod h1:oDoupMAO8OvCJWAcko0GGGIgR6R6ocIYbsSw735rRwI=
github.com/golang/protobuf v1.5.4 h1:i7eJL8qZTpSEXOPTxNKhASYpMn+8e5Q6AdndVa1dWek=
github.com/golang/protobuf v1.5.4/go.mod h1:lnTiLA8Wa4RWRcIUkrtSVa5nRhsEGBg48fD6rSs7xps=
github.com/google/go-cmp v0.2.0/go.mod h1:oXzfMopK8JAjlY9xF4vHSVASa0yLyX7SntLO5aqRK0M=
github.com/google/go-cmp v0.3.0/go.mod h1:8QqcDgzrUqlUb/G2PQTWiueGozuR1884gddMywk6iLU=
github.com/google/go-cmp v0.3.1/go.mod h1:8QqcDgzrUqlUb/G2PQTWiueGozuR1884gddMywk6iLU=
github.com/google/go-cmp v0.4.0/go.mod h1:v8dTdLbMG2kIc/vJvl+f65V22dbkXbowE6jgT/gNBxE=
github.com/google/go-cmp v0.5.0/go.mod h1:v8dTdLbMG2kIc/vJvl+f65V22dbkXbowE6jgT/gNBxE=
github.com/google/go-cmp v0.5.2/go.mod h1:v8dTdLbMG2kIc/vJvl+f65V22dbkXbowE6jgT/gNBxE=
github.com/google/go-cmp v0.5.3/go.mod h1:v8dTdLbMG2kIc/vJvl+f65V22dbkXbowE6jgT/gNBxE=
github.com/google/go-cmp v0.5.6/go.mod h1:v8dTdLbMG2kIc/vJvl+f65V22dbkXbowE6jgT/gNBxE=
github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8=
github.com/google/go-cmp v0.7.0/go.mod h1:pXiqmnSA92OHEEa9HXL2W4E7lf9JzCmGVUdgjX3N/iU=
@@ -176,10 +211,17 @@ github.com/google/go-querystring v1.1.0/go.mod h1:Kcdr2DB4koayq7X8pmAG4sNG59So17
github.com/google/gofuzz v1.0.0/go.mod h1:dBl0BpW6vV/+mYPU4Po3pmUjxk6FQPldtuIdl/M65Eg=
github.com/google/pprof v0.0.0-20250317173921-a4b03ec1a45e h1:ijClszYn+mADRFY17kjQEVQ1XRhq2/JR1M3sGqeJoxs=
github.com/google/pprof v0.0.0-20250317173921-a4b03ec1a45e/go.mod h1:boTsfXsheKC2y+lKOCMpSfarhxDeIzfZG1jqGcPl3cA=
github.com/google/s2a-go v0.1.7 h1:60BLSyTrOV4/haCDW4zb1guZItoSq8foHCXrAnjBo/o=
github.com/google/s2a-go v0.1.7/go.mod h1:50CgR4k1jNlWBu4UfS4AcfhVe1r6pdZPygJ3R8F0Qdw=
github.com/google/uuid v1.1.2/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo=
github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0=
github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo=
github.com/google/wire v0.7.0 h1:JxUKI6+CVBgCO2WToKy/nQk0sS+amI9z9EjVmdaocj4=
github.com/google/wire v0.7.0/go.mod h1:n6YbUQD9cPKTnHXEBN2DXlOp/mVADhVErcMFb0v3J18=
github.com/googleapis/enterprise-certificate-proxy v0.3.2 h1:Vie5ybvEvT75RniqhfFxPRy3Bf7vr3h0cechB90XaQs=
github.com/googleapis/enterprise-certificate-proxy v0.3.2/go.mod h1:VLSiSSBs/ksPL8kq3OBOQ6WRI2QnaFynd1DCjZ62+V0=
github.com/googleapis/gax-go/v2 v2.12.0 h1:A+gCJKdRfqXkr+BIRGtZLibNXf0m1f9E4HG56etFpas=
github.com/googleapis/gax-go/v2 v2.12.0/go.mod h1:y+aIqrI5eb1YGMVJfuV3185Ts/D7qKpsEkdD5+I6QGU=
github.com/gorilla/websocket v1.5.3 h1:saDtZ6Pbx/0u+bgYQ3q96pZgCzfhKXGPqt7kZ72aNNg=
github.com/gorilla/websocket v1.5.3/go.mod h1:YR8l580nyteQvAITg2hZ9XVh4b55+EU/adAjf1fMHhE=
github.com/grpc-ecosystem/grpc-gateway/v2 v2.27.3 h1:NmZ1PKzSTQbuGHw9DGPFomqkkLWMC+vZCkfs+FHv1Vg=
@@ -273,6 +315,10 @@ github.com/pelletier/go-toml/v2 v2.2.2 h1:aYUidT7k73Pcl9nb2gScu7NSrKCSHIDE89b3+6
github.com/pelletier/go-toml/v2 v2.2.2/go.mod h1:1t835xjRzz80PqgE6HHgN2JOsmgYu/h4qDAS4n929Rs=
github.com/pkg/errors v0.9.1 h1:FEBLx1zS214owpjy7qsBeixbURkuhQAwrK5UwLGTwt4=
github.com/pkg/errors v0.9.1/go.mod h1:bwawxfHBFNV+L2hUp1rHADufV3IMtnDRdf1r5NINEl0=
github.com/pkoukk/tiktoken-go v0.1.8 h1:85ENo+3FpWgAACBaEUVp+lctuTcYUO7BtmfhlN/QTRo=
github.com/pkoukk/tiktoken-go v0.1.8/go.mod h1:9NiV+i9mJKGj1rYOT+njbv+ZwA/zJxYdewGl6qVatpg=
github.com/pkoukk/tiktoken-go-loader v0.0.2 h1:LUKws63GV3pVHwH1srkBplBv+7URgmOmhSkRxsIvsK4=
github.com/pkoukk/tiktoken-go-loader v0.0.2/go.mod h1:4mIkYyZooFlnenDlormIo6cd5wrlUKNr97wp9nGgEKo=
github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4=
github.com/pmezard/go-difflib v1.0.1-0.20181226105442-5d4384ee4fb2 h1:Jamvg5psRIccs7FGNTlIRMkT8wgtp5eCXdBlqhYGL6U=
github.com/pmezard/go-difflib v1.0.1-0.20181226105442-5d4384ee4fb2/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4=
@@ -282,6 +328,7 @@ github.com/pquerna/otp v1.5.0 h1:NMMR+WrmaqXU4EzdGJEE1aUUI0AMRzsp96fFFWNPwxs=
github.com/pquerna/otp v1.5.0/go.mod h1:dkJfzwRKNiegxyNb54X/3fLwhCynbMspSyWKnvi1AEg=
github.com/prashantv/gostub v1.1.0 h1:BTyx3RfQjRHnUWaGF9oQos79AlQ5k8WNktv7VGvVH4g=
github.com/prashantv/gostub v1.1.0/go.mod h1:A5zLQHz7ieHGG7is6LLXLz7I8+3LZzsrV0P1IAHhP5U=
github.com/prometheus/client_model v0.0.0-20190812154241-14fe0d1b01d4/go.mod h1:xMI15A0UPsDsEKsMN9yxemIoYk6Tm2C1GtYGdfGttqA=
github.com/quic-go/qpack v0.6.0 h1:g7W+BMYynC1LbYLSqRt8PBg5Tgwxn214ZZR34VIOjz8=
github.com/quic-go/qpack v0.6.0/go.mod h1:lUpLKChi8njB4ty2bFLX2x4gzDqXwUpaO1DP9qMDZII=
github.com/quic-go/quic-go v0.57.1 h1:25KAAR9QR8KZrCZRThWMKVAwGoiHIrNbT72ULHTuI10=
@@ -370,6 +417,8 @@ github.com/zclconf/go-cty-yaml v1.1.0 h1:nP+jp0qPHv2IhUVqmQSzjvqAWcObN0KBkUl2rWB
github.com/zclconf/go-cty-yaml v1.1.0/go.mod h1:9YLUH4g7lOhVWqUbctnVlZ5KLpg7JAprQNgxSZ1Gyxs=
github.com/zeromicro/go-zero v1.9.4 h1:aRLFoISqAYijABtkbliQC5SsI5TbizJpQvoHc9xup8k=
github.com/zeromicro/go-zero v1.9.4/go.mod h1:a17JOTch25SWxBcUgJZYps60hygK3pIYdw7nGwlcS38=
go.opencensus.io v0.24.0 h1:y73uSU6J157QMP2kn2r30vwW1A2W2WFwSCGnAVxeaD0=
go.opencensus.io v0.24.0/go.mod h1:vNK8G9p7aAivkbmorf4v+7Hgx+Zs0yY+0fOtgBfjQKo=
go.opentelemetry.io/auto/sdk v1.1.0 h1:cH53jehLUN6UFLY71z+NDOiNJqDdPRaXzTel0sJySYA=
go.opentelemetry.io/auto/sdk v1.1.0/go.mod h1:3wSPjt5PWp2RhlCcmmOial7AvC4DQqZb7a7wCow3W8A=
go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp v0.49.0 h1:jq9TW8u3so/bN+JPT166wjOI6/vQPF6Xe7nMNIltagk=
@@ -384,6 +433,8 @@ go.opentelemetry.io/otel/metric v1.37.0 h1:mvwbQS5m0tbmqML4NqK+e3aDiO02vsf/Wgbsd
go.opentelemetry.io/otel/metric v1.37.0/go.mod h1:04wGrZurHYKOc+RKeye86GwKiTb9FKm1WHtO+4EVr2E=
go.opentelemetry.io/otel/sdk v1.37.0 h1:ItB0QUqnjesGRvNcmAcU0LyvkVyGJ2xftD29bWdDvKI=
go.opentelemetry.io/otel/sdk v1.37.0/go.mod h1:VredYzxUvuo2q3WRcDnKDjbdvmO0sCzOvVAiY+yUkAg=
go.opentelemetry.io/otel/sdk/metric v1.37.0 h1:90lI228XrB9jCMuSdA0673aubgRobVZFhbjxHHspCPc=
go.opentelemetry.io/otel/sdk/metric v1.37.0/go.mod h1:cNen4ZWfiD37l5NhS+Keb5RXVWZWpRE+9WyVCpbo5ps=
go.opentelemetry.io/otel/trace v1.37.0 h1:HLdcFNbRQBE2imdSEgm/kwqmQj1Or1l/7bW6mxVK7z4=
go.opentelemetry.io/otel/trace v1.37.0/go.mod h1:TlgrlQ+PtQO5XFerSPUYG0JSgGyryXewPGyayAWSBS0=
go.opentelemetry.io/proto/otlp v1.3.1 h1:TrMUixzpM0yuc/znrFTP9MMRh8trP93mkCiDVeXrui0=
@@ -403,18 +454,40 @@ go.uber.org/zap v1.24.0/go.mod h1:2kMP+WWQ8aoFoedH3T2sq6iJ2yDWpHbP0f6MQbS9Gkg=
golang.org/x/arch v0.0.0-20210923205945-b76863e36670/go.mod h1:5om86z9Hs0C8fWVUuoMHwpExlXzs5Tkyp9hOrfG7pp8=
golang.org/x/arch v0.3.0 h1:02VY4/ZcO/gBOH6PUaoiptASxtXU10jazRCP865E97k=
golang.org/x/arch v0.3.0/go.mod h1:5om86z9Hs0C8fWVUuoMHwpExlXzs5Tkyp9hOrfG7pp8=
golang.org/x/crypto v0.0.0-20190308221718-c2843e01d9a2/go.mod h1:djNgcEr1/C05ACkg1iLfiJU5Ep61QUkGW8qpdssI0+w=
golang.org/x/crypto v0.0.0-20200622213623-75b288015ac9/go.mod h1:LzIPMQfyMNhhGPhUkYOs5KpL4U8rLKemX1yGLhDgUto=
golang.org/x/crypto v0.48.0 h1:/VRzVqiRSggnhY7gNRxPauEQ5Drw9haKdM0jqfcCFts=
golang.org/x/crypto v0.48.0/go.mod h1:r0kV5h3qnFPlQnBSrULhlsRfryS2pmewsg+XfMgkVos=
golang.org/x/exp v0.0.0-20190121172915-509febef88a4/go.mod h1:CJ0aWSM057203Lf6IL+f9T1iT9GByDxfZKAQTCR3kQA=
golang.org/x/exp v0.0.0-20251023183803-a4bb9ffd2546 h1:mgKeJMpvi0yx/sU5GsxQ7p6s2wtOnGAHZWCHUM4KGzY=
golang.org/x/exp v0.0.0-20251023183803-a4bb9ffd2546/go.mod h1:j/pmGrbnkbPtQfxEe5D0VQhZC6qKbfKifgD0oM7sR70=
golang.org/x/lint v0.0.0-20181026193005-c67002cb31c3/go.mod h1:UVdnD1Gm6xHRNCYTkRU2/jEulfH38KcIWyp/GAMgvoE=
golang.org/x/lint v0.0.0-20190227174305-5b3e6a55c961/go.mod h1:wehouNa3lNwaWXcvxsM5YxQ5yQlVC4a0KAMCusXpPoU=
golang.org/x/lint v0.0.0-20190313153728-d0100b6bd8b3/go.mod h1:6SW0HCj/g11FgYtHlgUYUwCkIfeOF89ocIRzGO/8vkc=
golang.org/x/mod v0.32.0 h1:9F4d3PHLljb6x//jOyokMv3eX+YDeepZSEo3mFJy93c=
golang.org/x/mod v0.32.0/go.mod h1:SgipZ/3h2Ci89DlEtEXWUk/HteuRin+HHhN+WbNhguU=
golang.org/x/net v0.0.0-20180724234803-3673e40ba225/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4=
golang.org/x/net v0.0.0-20180826012351-8a410e7b638d/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4=
golang.org/x/net v0.0.0-20190213061140-3a22650c66bd/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4=
golang.org/x/net v0.0.0-20190311183353-d8887717615a/go.mod h1:t9HGtf8HONx5eT2rtn7q6eTqICYqUVnKs3thJo3Qplg=
golang.org/x/net v0.0.0-20190404232315-eb5bcb51f2a3/go.mod h1:t9HGtf8HONx5eT2rtn7q6eTqICYqUVnKs3thJo3Qplg=
golang.org/x/net v0.0.0-20201110031124-69a78807bb2b/go.mod h1:sp8m0HH+o8qH0wwXwYZr8TS3Oi6o0r6Gce1SSxlDquU=
golang.org/x/net v0.0.0-20211104170005-ce137452f963/go.mod h1:9nx3DQGgdP8bBQD5qxJ1jj9UTztislL4KSBs9R2vV5Y=
golang.org/x/net v0.49.0 h1:eeHFmOGUTtaaPSGNmjBKpbng9MulQsJURQUAfUwY++o=
golang.org/x/net v0.49.0/go.mod h1:/ysNB2EvaqvesRkuLAyjI1ycPZlQHM3q01F02UY/MV8=
golang.org/x/oauth2 v0.0.0-20180821212333-d2e6202438be/go.mod h1:N/0e6XlmueqKjAGxoOufVs8QHGRruUQn6yWY3a++T0U=
golang.org/x/oauth2 v0.30.0 h1:dnDm7JmhM45NNpd8FDDeLhK6FwqbOf4MLCM9zb1BOHI=
golang.org/x/oauth2 v0.30.0/go.mod h1:B++QgG3ZKulg6sRPGD/mqlHQs5rB3Ml9erfeDY7xKlU=
golang.org/x/sync v0.0.0-20180314180146-1d60e4601c6f/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
golang.org/x/sync v0.0.0-20181108010431-42b317875d0f/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
golang.org/x/sync v0.0.0-20190423024810-112230192c58/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
golang.org/x/sync v0.19.0 h1:vV+1eWNmZ5geRlYjzm2adRgW2/mcpevXNg50YZtPCE4=
golang.org/x/sync v0.19.0/go.mod h1:9KTHXmSnoGruLpwFjVSX0lNNA75CykiMECbovNTZqGI=
golang.org/x/sys v0.0.0-20180830151530-49385e6e1522/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
golang.org/x/sys v0.0.0-20190412213103-97732733099d/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
golang.org/x/sys v0.0.0-20190916202348-b4ddaad3f8a3/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
golang.org/x/sys v0.0.0-20200930185726-fdedc70b468f/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
golang.org/x/sys v0.0.0-20201119102817-f84b799fce68/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
golang.org/x/sys v0.0.0-20201204225414-ed752295db88/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
golang.org/x/sys v0.0.0-20210423082822-04245dca01da/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
@@ -430,22 +503,50 @@ golang.org/x/sys v0.41.0/go.mod h1:OgkHotnGiDImocRcuBABYBEXf8A9a87e/uXjp9XT3ks=
golang.org/x/term v0.0.0-20201126162022-7de9c90e9dd1/go.mod h1:bj7SfCRtBDWHUb9snDiAeCFNEtKQo2Wmx5Cou7ajbmo=
golang.org/x/term v0.40.0 h1:36e4zGLqU4yhjlmxEaagx2KuYbJq3EwY8K943ZsHcvg=
golang.org/x/term v0.40.0/go.mod h1:w2P8uVp06p2iyKKuvXIm7N/y0UCRt3UfJTfZ7oOpglM=
golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ=
golang.org/x/text v0.3.3/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ=
golang.org/x/text v0.3.6/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ=
golang.org/x/text v0.34.0 h1:oL/Qq0Kdaqxa1KbNeMKwQq0reLCCaFtqu2eNuSeNHbk=
golang.org/x/text v0.34.0/go.mod h1:homfLqTYRFyVYemLBFl5GgL/DWEiH5wcsQ5gSh1yziA=
golang.org/x/time v0.12.0 h1:ScB/8o8olJvc+CQPWrK3fPZNfh7qgwCrY0zJmoEQLSE=
golang.org/x/time v0.12.0/go.mod h1:CDIdPxbZBQxdj6cxyCIdrNogrJKMJ7pr37NYpMcMDSg=
golang.org/x/tools v0.0.0-20180917221912-90fa682c2a6e/go.mod h1:n7NCudcB/nEzxVGmLbDWY5pfWTLqBcC2KZ6jyYvM4mQ=
golang.org/x/tools v0.0.0-20190114222345-bf090417da8b/go.mod h1:n7NCudcB/nEzxVGmLbDWY5pfWTLqBcC2KZ6jyYvM4mQ=
golang.org/x/tools v0.0.0-20190226205152-f727befe758c/go.mod h1:9Yl7xja0Znq3iFh3HoIrodX9oNMXvdceNzlUR8zjMvY=
golang.org/x/tools v0.0.0-20190311212946-11955173bddd/go.mod h1:LCzVGOaR6xXOjkQ3onu1FJEFr0SW1gC7cKk1uF8kGRs=
golang.org/x/tools v0.0.0-20190524140312-2c0ae7006135/go.mod h1:RgjU9mgBXZiqYHBnxXauZ1Gv1EHHAz9KjViQ78xBX0Q=
golang.org/x/tools v0.41.0 h1:a9b8iMweWG+S0OBnlU36rzLp20z1Rp10w+IY2czHTQc=
golang.org/x/tools v0.41.0/go.mod h1:XSY6eDqxVNiYgezAVqqCeihT4j1U2CCsqvH3WhQpnlg=
golang.org/x/xerrors v0.0.0-20191204190536-9bdfabe68543/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
google.golang.org/genproto v0.0.0-20231106174013-bbf56f31fb17 h1:wpZ8pe2x1Q3f2KyT5f8oP/fa9rHAKgFPr/HZdNuS+PQ=
gonum.org/v1/gonum v0.16.0 h1:5+ul4Swaf3ESvrOnidPp4GZbzf0mxVQpDCYUQE7OJfk=
gonum.org/v1/gonum v0.16.0/go.mod h1:fef3am4MQ93R2HHpKnLk4/Tbh/s0+wqD5nfa6Pnwy4E=
google.golang.org/api v0.153.0 h1:N1AwGhielyKFaUqH07/ZSIQR3uNPcV7NVw0vj+j4iR4=
google.golang.org/api v0.153.0/go.mod h1:3qNJX5eOmhiWYc67jRA/3GsDw97UFb5ivv7Y2PrriAY=
google.golang.org/appengine v1.1.0/go.mod h1:EbEs0AVv82hx2wNQdGPgUI5lhzA/G0D9YwlJXL52JkM=
google.golang.org/appengine v1.4.0/go.mod h1:xpcJRLb0r/rnEns0DIKYYv+WjYCduHsrkT7/EB5XEv4=
google.golang.org/genproto v0.0.0-20180817151627-c66870c02cf8/go.mod h1:JiN7NxoALGmiZfu7CAH4rXhgtRTLTxftemlI0sWmxmc=
google.golang.org/genproto v0.0.0-20190819201941-24fa4b261c55/go.mod h1:DMBHOl98Agz4BDEuKkezgsaosCRResVns1a3J2ZsMNc=
google.golang.org/genproto v0.0.0-20200526211855-cb27e3aa2013/go.mod h1:NbSheEEYHJ7i3ixzK3sjbqSGDJWnxyFXZblF3eUsNvo=
google.golang.org/genproto/googleapis/api v0.0.0-20250929231259-57b25ae835d4 h1:8XJ4pajGwOlasW+L13MnEGA8W4115jJySQtVfS2/IBU=
google.golang.org/genproto/googleapis/api v0.0.0-20250929231259-57b25ae835d4/go.mod h1:NnuHhy+bxcg30o7FnVAZbXsPHUDQ9qKWAQKCD7VxFtk=
google.golang.org/genproto/googleapis/rpc v0.0.0-20250929231259-57b25ae835d4 h1:i8QOKZfYg6AbGVZzUAY3LrNWCKF8O6zFisU9Wl9RER4=
google.golang.org/genproto/googleapis/rpc v0.0.0-20250929231259-57b25ae835d4/go.mod h1:HSkG/KdJWusxU1F6CNrwNDjBMgisKxGnc5dAZfT0mjQ=
google.golang.org/grpc v1.19.0/go.mod h1:mqu4LbDTu4XGKhr4mRzUsmM4RtVoemTSY81AxZiDr8c=
google.golang.org/grpc v1.23.0/go.mod h1:Y5yQAOtifL1yxbo5wqy6BxZv8vAUGQwXBOALyacEbxg=
google.golang.org/grpc v1.25.1/go.mod h1:c3i+UQWmh7LiEpx4sFZnkU36qjEYZ0imhYfXVyQciAY=
google.golang.org/grpc v1.27.0/go.mod h1:qbnxyOmOxrQa7FizSgH+ReBfzJrCY1pSN7KXBS8abTk=
google.golang.org/grpc v1.33.2/go.mod h1:JMHMWHQWaTccqQQlmk3MJZS+GWXOdAesneDmEnv2fbc=
google.golang.org/grpc v1.75.1 h1:/ODCNEuf9VghjgO3rqLcfg8fiOP0nSluljWFlDxELLI=
google.golang.org/grpc v1.75.1/go.mod h1:JtPAzKiq4v1xcAB2hydNlWI2RnF85XXcV0mhKXr2ecQ=
google.golang.org/protobuf v0.0.0-20200109180630-ec00e32a8dfd/go.mod h1:DFci5gLYBciE7Vtevhsrf46CRTquxDuWsQurQQe4oz8=
google.golang.org/protobuf v0.0.0-20200221191635-4d8936d0db64/go.mod h1:kwYJMbMJ01Woi6D6+Kah6886xMZcty6N08ah7+eCXa0=
google.golang.org/protobuf v0.0.0-20200228230310-ab0ca4ff8a60/go.mod h1:cfTl7dwQJ+fmap5saPgwCLgHXTUD7jkjRqWcaiX5VyM=
google.golang.org/protobuf v1.20.1-0.20200309200217-e05f789c0967/go.mod h1:A+miEFZTKqfCUM6K7xSMQL9OKL/b6hQv+e19PK+JZNE=
google.golang.org/protobuf v1.21.0/go.mod h1:47Nbq4nVaFHyn7ilMalzfO3qCViNmqZ2kzikPIcrTAo=
google.golang.org/protobuf v1.22.0/go.mod h1:EGpADcykh3NcUnDUJcl1+ZksZNG86OlYog2l/sGQquU=
google.golang.org/protobuf v1.23.0/go.mod h1:EGpADcykh3NcUnDUJcl1+ZksZNG86OlYog2l/sGQquU=
google.golang.org/protobuf v1.23.1-0.20200526195155-81db48ad09cc/go.mod h1:EGpADcykh3NcUnDUJcl1+ZksZNG86OlYog2l/sGQquU=
google.golang.org/protobuf v1.25.0/go.mod h1:9JNX74DMeImyA3h4bdi1ymwjUzf21/xIlbajtzgsN7c=
google.golang.org/protobuf v1.36.10 h1:AYd7cD/uASjIL6Q9LiTjz8JLcrh/88q5UObnmY3aOOE=
google.golang.org/protobuf v1.36.10/go.mod h1:HTf+CrKn2C3g5S8VImy6tdcUvCska2kB7j23XfzDpco=
gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0=
@@ -460,6 +561,8 @@ gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA=
gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM=
gotest.tools/v3 v3.5.2 h1:7koQfIKdy+I8UTetycgUqXWSDwpgv193Ka+qRsmBY8Q=
gotest.tools/v3 v3.5.2/go.mod h1:LtdLGcnqToBH83WByAAi/wiwSFCArdFIUV/xxN4pcjA=
honnef.co/go/tools v0.0.0-20190102054323-c2f93a96b099/go.mod h1:rf3lG4BRIbNafJWhAfAdb/ePZxsR/4RtNHQocxwk9r4=
honnef.co/go/tools v0.0.0-20190523083050-ea95bdfd59fc/go.mod h1:rf3lG4BRIbNafJWhAfAdb/ePZxsR/4RtNHQocxwk9r4=
modernc.org/cc/v4 v4.27.1 h1:9W30zRlYrefrDV2JE2O8VDtJ1yPGownxciz5rrbQZis=
modernc.org/cc/v4 v4.27.1/go.mod h1:uVtb5OGqUKpoLWhqwNQo/8LwvoiEBLvZXIQ/SmO6mL0=
modernc.org/ccgo/v4 v4.30.1 h1:4r4U1J6Fhj98NKfSjnPUN7Ze2c6MnAdL0hWw6+LrJpc=
@@ -0,0 +1,254 @@
package admin
import (
"context"
"log/slog"
"strconv"
"strings"
"time"
"github.com/Wei-Shaw/sub2api/internal/pkg/response"
"github.com/Wei-Shaw/sub2api/internal/service"
"github.com/gin-gonic/gin"
)
// 账号亲和调度 API 与契约说明见 docs/ACCOUNT_SCHEDULING_FLOW.md 。
// GetAffinityClients returns the list of affinity clients for an account with last active timestamps.
// GET /api/v1/admin/accounts/:id/affinity-clients
func (h *AccountHandler) GetAffinityClients(c *gin.Context) {
accountID, err := strconv.ParseInt(c.Param("id"), 10, 64)
if err != nil {
response.BadRequest(c, "Invalid account ID")
return
}
account, err := h.adminService.GetAccount(c.Request.Context(), accountID)
if err != nil {
response.ErrorFrom(c, err)
return
}
if !account.IsAffinityEnabled() {
response.Success(c, []service.AffinityClient{})
return
}
if h.gatewayCache == nil || len(account.GroupIDs) == 0 {
response.Success(c, []service.AffinityClient{})
return
}
clients, err := h.gatewayCache.GetAccountAffinityClientsWithScores(
c.Request.Context(), accountID, account.GroupIDs, service.ClientAffinityTTL,
)
if err != nil {
response.Success(c, []service.AffinityClient{})
return
}
response.Success(c, clients)
}
// clearAccountAffinity 清除指定账号在所有分组的亲和记录。
func (h *AccountHandler) clearAccountAffinity(ctx context.Context, accountID int64, groupIDs []int64) {
if h.gatewayCache == nil || len(groupIDs) == 0 {
return
}
if err := h.gatewayCache.ClearAccountAffinity(ctx, accountID, groupIDs); err != nil {
// 清理失败不影响主流程,记录日志即可
slog.Warn("clear account affinity failed",
"account_id", accountID,
"error", err,
)
}
}
// countUniqueUsersFromAffinityMembers 从 "{userID}/{clientID}" 格式的成员列表中计算唯一用户数。
func countUniqueUsersFromAffinityMembers(members []string) int64 {
users := make(map[string]struct{}, len(members))
for _, m := range members {
if idx := strings.Index(m, "/"); idx > 0 {
users[m[:idx]] = struct{}{}
}
}
return int64(len(users))
}
// AffinityDetailClient 亲和详情中的单个客户端信息
type AffinityDetailClient struct {
ClientID string `json:"client_id"`
LastActive time.Time `json:"last_active"`
IsPinned bool `json:"is_pinned"`
}
// AffinityDetailUser 亲和详情中按用户分组的信息
type AffinityDetailUser struct {
UserID int64 `json:"user_id"`
UserEmail string `json:"user_email"`
ClientCount int `json:"client_count"`
IsPinned bool `json:"is_pinned"`
Clients []AffinityDetailClient `json:"clients"`
}
// AffinityDetailsResponse 亲和详情响应
type AffinityDetailsResponse struct {
Users []AffinityDetailUser `json:"users"`
TotalUsers int `json:"total_users"`
TotalClients int `json:"total_clients"`
PinnedUsers []int64 `json:"pinned_users"`
}
type affinityStateSnapshot struct {
enabled bool
groupIDs []int64
}
func (h *AccountHandler) captureAffinityStates(ctx context.Context, accountIDs []int64) map[int64]affinityStateSnapshot {
states := make(map[int64]affinityStateSnapshot, len(accountIDs))
if h.gatewayCache == nil || len(accountIDs) == 0 {
return states
}
for _, accountID := range accountIDs {
account, err := h.adminService.GetAccount(ctx, accountID)
if err != nil || account == nil {
continue
}
states[accountID] = affinityStateSnapshot{
enabled: account.IsAffinityEnabled(),
groupIDs: append([]int64(nil), account.GroupIDs...),
}
}
return states
}
func (h *AccountHandler) clearAffinityCacheIfDisabled(ctx context.Context, accountID int64, oldState affinityStateSnapshot) {
if h.gatewayCache == nil || !oldState.enabled {
return
}
account, err := h.adminService.GetAccount(ctx, accountID)
if err != nil || account == nil || account.IsAffinityEnabled() {
return
}
groupIDs := oldState.groupIDs
if len(account.GroupIDs) > 0 {
groupIDs = mergeGroupIDs(groupIDs, account.GroupIDs)
}
h.clearAccountAffinity(ctx, accountID, groupIDs)
}
func (h *AccountHandler) clearAffinityCacheForBulkIfDisabled(ctx context.Context, accountIDs []int64, oldStates map[int64]affinityStateSnapshot) {
for _, accountID := range accountIDs {
oldState, ok := oldStates[accountID]
if !ok {
continue
}
h.clearAffinityCacheIfDisabled(ctx, accountID, oldState)
}
}
// GetAffinityDetails returns the affinity details grouped by user for an account.
// GET /api/v1/admin/accounts/:id/affinity-details
func (h *AccountHandler) GetAffinityDetails(c *gin.Context) {
accountID, err := strconv.ParseInt(c.Param("id"), 10, 64)
if err != nil {
response.BadRequest(c, "Invalid account ID")
return
}
emptyResp := AffinityDetailsResponse{
Users: []AffinityDetailUser{},
TotalUsers: 0,
TotalClients: 0,
PinnedUsers: []int64{},
}
account, err := h.adminService.GetAccount(c.Request.Context(), accountID)
if err != nil {
response.ErrorFrom(c, err)
return
}
if !account.IsAffinityEnabled() {
response.Success(c, emptyResp)
return
}
pinnedUsers := account.GetPinnedUsers()
if pinnedUsers == nil {
pinnedUsers = []int64{}
}
if h.gatewayCache == nil || len(account.GroupIDs) == 0 {
emptyResp.PinnedUsers = pinnedUsers
response.Success(c, emptyResp)
return
}
clients, err := h.gatewayCache.GetAccountAffinityClientsWithScores(
c.Request.Context(), accountID, account.GroupIDs, service.ClientAffinityTTL,
)
if err != nil {
emptyResp.PinnedUsers = pinnedUsers
response.Success(c, emptyResp)
return
}
pinnedSet := make(map[int64]struct{}, len(pinnedUsers))
for _, uid := range pinnedUsers {
pinnedSet[uid] = struct{}{}
}
// 按 UserID 分组
userMap := make(map[int64]*AffinityDetailUser)
var userOrder []int64
for _, cl := range clients {
u, ok := userMap[cl.UserID]
if !ok {
_, pinned := pinnedSet[cl.UserID]
u = &AffinityDetailUser{
UserID: cl.UserID,
UserEmail: "",
ClientCount: 0,
IsPinned: pinned,
Clients: []AffinityDetailClient{},
}
userMap[cl.UserID] = u
userOrder = append(userOrder, cl.UserID)
}
u.Clients = append(u.Clients, AffinityDetailClient{
ClientID: cl.ClientID,
LastActive: cl.LastActive,
IsPinned: false, // 客户端级别暂无 pinned 概念
})
}
// 关联用户邮箱(查询失败时保持空字符串,不影响主流程)
for _, uid := range userOrder {
user, uErr := h.adminService.GetUser(c.Request.Context(), uid)
if uErr != nil || user == nil {
continue
}
if u, ok := userMap[uid]; ok {
if user.Email != "" {
u.UserEmail = user.Email
} else {
u.UserEmail = user.Username
}
}
}
users := make([]AffinityDetailUser, 0, len(userOrder))
for _, uid := range userOrder {
u := userMap[uid]
u.ClientCount = len(u.Clients)
users = append(users, *u)
}
response.Success(c, AffinityDetailsResponse{
Users: users,
TotalUsers: len(users),
TotalClients: len(clients),
PinnedUsers: pinnedUsers,
})
}
@@ -65,6 +65,7 @@ func setupAccountDataRouter() (*gin.Engine, *stubAdminService) {
nil,
nil,
nil,
nil,
)
router.GET("/api/v1/admin/accounts/data", h.ExportData)
@@ -57,6 +57,7 @@ type AccountHandler struct {
sessionLimitCache service.SessionLimitCache
rpmCache service.RPMCache
tokenCacheInvalidator service.TokenCacheInvalidator
gatewayCache service.GatewayCache
}
// NewAccountHandler creates a new admin account handler
@@ -74,6 +75,7 @@ func NewAccountHandler(
sessionLimitCache service.SessionLimitCache,
rpmCache service.RPMCache,
tokenCacheInvalidator service.TokenCacheInvalidator,
gatewayCache service.GatewayCache,
) *AccountHandler {
return &AccountHandler{
adminService: adminService,
@@ -89,6 +91,7 @@ func NewAccountHandler(
sessionLimitCache: sessionLimitCache,
rpmCache: rpmCache,
tokenCacheInvalidator: tokenCacheInvalidator,
gatewayCache: gatewayCache,
}
}
@@ -206,6 +209,34 @@ func (h *AccountHandler) buildAccountResponseWithRuntime(ctx context.Context, ac
}
}
// 亲和客户端数据(启用亲和的账号始终返回 count,即使为 0)
if account.IsAffinityEnabled() {
if h.gatewayCache != nil && len(account.GroupIDs) > 0 {
accountGroups := map[int64][]int64{account.ID: account.GroupIDs}
if clients, err := h.gatewayCache.GetAccountAffinityClientsBatch(ctx, accountGroups, service.ClientAffinityTTL); err == nil {
if cl, ok := clients[account.ID]; ok && len(cl) > 0 {
count := int64(len(cl))
item.AffinityClientCount = &count
item.AffinityClients = cl
userCount := countUniqueUsersFromAffinityMembers(cl)
item.AffinityUserCount = &userCount
} else {
zero := int64(0)
item.AffinityClientCount = &zero
item.AffinityUserCount = &zero
}
} else {
zero := int64(0)
item.AffinityClientCount = &zero
item.AffinityUserCount = &zero
}
} else {
zero := int64(0)
item.AffinityClientCount = &zero
item.AffinityUserCount = &zero
}
}
return item
}
@@ -318,6 +349,21 @@ func (h *AccountHandler) List(c *gin.Context) {
_ = g.Wait()
}
// 获取亲和客户端数据(Redis Pipeline,低开销)
var affinityClients map[int64][]string
if h.gatewayCache != nil {
accountGroups := make(map[int64][]int64)
for i := range accounts {
acc := &accounts[i]
if acc.IsAffinityEnabled() && len(acc.GroupIDs) > 0 {
accountGroups[acc.ID] = acc.GroupIDs
}
}
if len(accountGroups) > 0 {
affinityClients, _ = h.gatewayCache.GetAccountAffinityClientsBatch(c.Request.Context(), accountGroups, service.ClientAffinityTTL)
}
}
// Build response with concurrency info
result := make([]AccountWithConcurrency, len(accounts))
for i := range accounts {
@@ -348,6 +394,22 @@ func (h *AccountHandler) List(c *gin.Context) {
}
}
// 注入亲和客户端数据到 DTO(启用亲和的账号始终返回 count,即使为 0)
if acc.IsAffinityEnabled() {
if clients, ok := affinityClients[acc.ID]; ok && len(clients) > 0 {
count := int64(len(clients))
item.AffinityClientCount = &count
item.AffinityClients = clients
// 从成员列表中解析唯一用户数
userCount := countUniqueUsersFromAffinityMembers(clients)
item.AffinityUserCount = &userCount
} else {
zero := int64(0)
item.AffinityClientCount = &zero
item.AffinityUserCount = &zero
}
}
result[i] = item
}
@@ -568,6 +630,8 @@ func (h *AccountHandler) Update(c *gin.Context) {
// base_rpm 输入校验:负值归零,超过 10000 截断
sanitizeExtraBaseRPM(req.Extra)
oldStates := h.captureAffinityStates(c.Request.Context(), []int64{accountID})
// 确定是否跳过混合渠道检查
skipCheck := req.ConfirmMixedChannelRisk != nil && *req.ConfirmMixedChannelRisk
@@ -604,6 +668,8 @@ func (h *AccountHandler) Update(c *gin.Context) {
return
}
h.clearAffinityCacheForBulkIfDisabled(c.Request.Context(), []int64{accountID}, oldStates)
response.Success(c, h.buildAccountResponseWithRuntime(c.Request.Context(), account))
}
@@ -1320,6 +1386,8 @@ func (h *AccountHandler) BulkUpdate(c *gin.Context) {
return
}
oldStates := h.captureAffinityStates(c.Request.Context(), req.AccountIDs)
result, err := h.adminService.BulkUpdateAccounts(c.Request.Context(), &service.BulkUpdateAccountsInput{
AccountIDs: req.AccountIDs,
Name: req.Name,
@@ -1341,6 +1409,12 @@ func (h *AccountHandler) BulkUpdate(c *gin.Context) {
c.JSON(409, gin.H{
"error": "mixed_channel_warning",
"message": mixedErr.Error(),
"details": gin.H{
"group_id": mixedErr.GroupID,
"group_name": mixedErr.GroupName,
"current_platform": mixedErr.CurrentPlatform,
"other_platform": mixedErr.OtherPlatform,
},
})
return
}
@@ -1348,6 +1422,8 @@ func (h *AccountHandler) BulkUpdate(c *gin.Context) {
return
}
h.clearAffinityCacheForBulkIfDisabled(c.Request.Context(), req.AccountIDs, oldStates)
response.Success(c, result)
}
@@ -1560,6 +1636,25 @@ func (h *AccountHandler) ResetQuota(c *gin.Context) {
response.Success(c, h.buildAccountResponseWithRuntime(c.Request.Context(), account))
}
// mergeGroupIDs 合并两个 groupID 切片并去重。
func mergeGroupIDs(a, b []int64) []int64 {
seen := make(map[int64]struct{}, len(a)+len(b))
result := make([]int64, 0, len(a)+len(b))
for _, id := range a {
if _, ok := seen[id]; !ok {
seen[id] = struct{}{}
result = append(result, id)
}
}
for _, id := range b {
if _, ok := seen[id]; !ok {
seen[id] = struct{}{}
result = append(result, id)
}
}
return result
}
// GetTempUnschedulable handles getting temporary unschedulable status
// GET /api/v1/admin/accounts/:id/temp-unschedulable
func (h *AccountHandler) GetTempUnschedulable(c *gin.Context) {
@@ -0,0 +1,298 @@
package admin
import (
"bytes"
"context"
"encoding/json"
"net/http"
"net/http/httptest"
"testing"
"time"
"github.com/Wei-Shaw/sub2api/internal/service"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/require"
)
type affinityCacheAdminService struct {
*stubAdminService
accounts map[int64]*service.Account
bulkResultIDs []int64
bulkFailedCount int
}
func (s *affinityCacheAdminService) GetAccount(_ context.Context, id int64) (*service.Account, error) {
if acc, ok := s.accounts[id]; ok {
cloned := *acc
return &cloned, nil
}
return s.stubAdminService.GetAccount(context.Background(), id)
}
func (s *affinityCacheAdminService) UpdateAccount(_ context.Context, id int64, input *service.UpdateAccountInput) (*service.Account, error) {
acc := s.accounts[id]
if acc == nil {
return s.stubAdminService.UpdateAccount(context.Background(), id, input)
}
if input.Extra != nil {
acc.Extra = input.Extra
}
if input.GroupIDs != nil {
acc.GroupIDs = *input.GroupIDs
}
cloned := *acc
return &cloned, nil
}
func (s *affinityCacheAdminService) BulkUpdateAccounts(_ context.Context, input *service.BulkUpdateAccountsInput) (*service.BulkUpdateAccountsResult, error) {
for _, accountID := range input.AccountIDs {
acc := s.accounts[accountID]
if acc == nil {
continue
}
if input.Extra != nil {
acc.Extra = input.Extra
}
if input.GroupIDs != nil {
acc.GroupIDs = *input.GroupIDs
}
}
successIDs := append([]int64(nil), input.AccountIDs...)
failed := 0
if len(s.bulkResultIDs) > 0 {
successIDs = append([]int64(nil), s.bulkResultIDs...)
failed = s.bulkFailedCount
}
return &service.BulkUpdateAccountsResult{
Success: len(successIDs),
Failed: failed,
SuccessIDs: successIDs,
Results: []service.BulkUpdateAccountResult{},
}, nil
}
type affinityCacheClearRecorder struct {
clearCalls []struct {
accountID int64
groupIDs []int64
}
}
func (r *affinityCacheClearRecorder) GetSessionAccountID(_ context.Context, _ int64, _ string) (int64, error) {
return 0, nil
}
func (r *affinityCacheClearRecorder) SetSessionAccountID(_ context.Context, _ int64, _ string, _ int64, _ time.Duration) error {
return nil
}
func (r *affinityCacheClearRecorder) RefreshSessionTTL(_ context.Context, _ int64, _ string, _ time.Duration) error {
return nil
}
func (r *affinityCacheClearRecorder) DeleteSessionAccountID(_ context.Context, _ int64, _ string) error {
return nil
}
func (r *affinityCacheClearRecorder) GetAffinityAccounts(_ context.Context, _ int64, _ int64, _ string, _ time.Duration) ([]int64, error) {
return nil, nil
}
func (r *affinityCacheClearRecorder) UpdateAffinity(_ context.Context, _ int64, _ int64, _ string, _ int64, _ time.Duration) error {
return nil
}
func (r *affinityCacheClearRecorder) GetAccountAffinityCountBatch(_ context.Context, _ int64, _ []int64, _ time.Duration) (map[int64]int64, error) {
return map[int64]int64{}, nil
}
func (r *affinityCacheClearRecorder) GetAccountAffinityClientsBatch(_ context.Context, _ map[int64][]int64, _ time.Duration) (map[int64][]string, error) {
return map[int64][]string{}, nil
}
func (r *affinityCacheClearRecorder) GetAccountAffinityClientsWithScores(_ context.Context, _ int64, _ []int64, _ time.Duration) ([]service.AffinityClient, error) {
return nil, nil
}
func (r *affinityCacheClearRecorder) ClearAccountAffinity(_ context.Context, accountID int64, groupIDs []int64) error {
r.clearCalls = append(r.clearCalls, struct {
accountID int64
groupIDs []int64
}{
accountID: accountID,
groupIDs: append([]int64(nil), groupIDs...),
})
return nil
}
func (r *affinityCacheClearRecorder) GetAffinityMultiCount(_ context.Context, _ int64, _ int64, _ int64, _ time.Duration) (users, clients, perUser int64, err error) {
return 0, 0, 0, nil
}
func setupAffinityCacheRouter(adminSvc service.AdminService, gatewayCache service.GatewayCache) *gin.Engine {
gin.SetMode(gin.TestMode)
router := gin.New()
handler := NewAccountHandler(adminSvc, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, gatewayCache)
router.PUT("/api/v1/admin/accounts/:id", handler.Update)
router.POST("/api/v1/admin/accounts/bulk-update", handler.BulkUpdate)
return router
}
func TestAccountHandlerUpdate_DisablingAffinityClearsCache(t *testing.T) {
adminSvc := &affinityCacheAdminService{
stubAdminService: newStubAdminService(),
accounts: map[int64]*service.Account{
3: {
ID: 3,
Platform: service.PlatformAnthropic,
Type: service.AccountTypeOAuth,
Status: service.StatusActive,
GroupIDs: []int64{1, 2},
Extra: map[string]any{
"client_affinity_enabled": true,
},
},
},
}
cache := &affinityCacheClearRecorder{}
router := setupAffinityCacheRouter(adminSvc, cache)
body, _ := json.Marshal(map[string]any{
"extra": map[string]any{},
})
rec := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodPut, "/api/v1/admin/accounts/3", bytes.NewReader(body))
req.Header.Set("Content-Type", "application/json")
router.ServeHTTP(rec, req)
require.Equal(t, http.StatusOK, rec.Code)
require.Len(t, cache.clearCalls, 1)
require.Equal(t, int64(3), cache.clearCalls[0].accountID)
require.Equal(t, []int64{1, 2}, cache.clearCalls[0].groupIDs)
}
func TestAccountHandlerBulkUpdate_DisablingAffinityClearsCache(t *testing.T) {
adminSvc := &affinityCacheAdminService{
stubAdminService: newStubAdminService(),
accounts: map[int64]*service.Account{
1: {
ID: 1,
Platform: service.PlatformAnthropic,
Type: service.AccountTypeOAuth,
Status: service.StatusActive,
GroupIDs: []int64{11},
Extra: map[string]any{
"client_affinity_enabled": true,
},
},
2: {
ID: 2,
Platform: service.PlatformAnthropic,
Type: service.AccountTypeSetupToken,
Status: service.StatusActive,
GroupIDs: []int64{22},
Extra: map[string]any{
"client_affinity_enabled": true,
},
},
},
}
cache := &affinityCacheClearRecorder{}
router := setupAffinityCacheRouter(adminSvc, cache)
body, _ := json.Marshal(map[string]any{
"account_ids": []int64{1, 2},
"extra": map[string]any{
"client_affinity_enabled": false,
},
})
rec := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodPost, "/api/v1/admin/accounts/bulk-update", bytes.NewReader(body))
req.Header.Set("Content-Type", "application/json")
router.ServeHTTP(rec, req)
require.Equal(t, http.StatusOK, rec.Code)
require.Len(t, cache.clearCalls, 2)
require.Equal(t, int64(1), cache.clearCalls[0].accountID)
require.Equal(t, []int64{11}, cache.clearCalls[0].groupIDs)
require.Equal(t, int64(2), cache.clearCalls[1].accountID)
require.Equal(t, []int64{22}, cache.clearCalls[1].groupIDs)
}
func TestAccountHandlerBulkUpdate_DisablingAffinityWithNewFlagClearsCache(t *testing.T) {
adminSvc := &affinityCacheAdminService{
stubAdminService: newStubAdminService(),
accounts: map[int64]*service.Account{
1: {
ID: 1,
Platform: service.PlatformAnthropic,
Type: service.AccountTypeOAuth,
Status: service.StatusActive,
GroupIDs: []int64{33},
Extra: map[string]any{
"affinity_enabled": true,
},
},
},
}
cache := &affinityCacheClearRecorder{}
router := setupAffinityCacheRouter(adminSvc, cache)
body, _ := json.Marshal(map[string]any{
"account_ids": []int64{1},
"extra": map[string]any{
"affinity_enabled": false,
},
})
rec := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodPost, "/api/v1/admin/accounts/bulk-update", bytes.NewReader(body))
req.Header.Set("Content-Type", "application/json")
router.ServeHTTP(rec, req)
require.Equal(t, http.StatusOK, rec.Code)
require.Len(t, cache.clearCalls, 1)
require.Equal(t, int64(1), cache.clearCalls[0].accountID)
require.Equal(t, []int64{33}, cache.clearCalls[0].groupIDs)
}
func TestAccountHandlerBulkUpdate_PartialFailureStillClearsDisabledAffinityCache(t *testing.T) {
adminSvc := &affinityCacheAdminService{
stubAdminService: newStubAdminService(),
accounts: map[int64]*service.Account{
1: {
ID: 1,
Platform: service.PlatformAnthropic,
Type: service.AccountTypeOAuth,
Status: service.StatusActive,
GroupIDs: []int64{44},
Extra: map[string]any{
"affinity_enabled": true,
},
},
2: {
ID: 2,
Platform: service.PlatformAnthropic,
Type: service.AccountTypeSetupToken,
Status: service.StatusActive,
GroupIDs: []int64{55},
Extra: map[string]any{
"client_affinity_enabled": true,
},
},
},
bulkResultIDs: []int64{1},
bulkFailedCount: 1,
}
cache := &affinityCacheClearRecorder{}
router := setupAffinityCacheRouter(adminSvc, cache)
body, _ := json.Marshal(map[string]any{
"account_ids": []int64{1, 2},
"group_ids": []int64{101},
"extra": map[string]any{
"affinity_enabled": false,
"client_affinity_enabled": false,
},
})
rec := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodPost, "/api/v1/admin/accounts/bulk-update", bytes.NewReader(body))
req.Header.Set("Content-Type", "application/json")
router.ServeHTTP(rec, req)
require.Equal(t, http.StatusOK, rec.Code)
require.Len(t, cache.clearCalls, 2)
require.Equal(t, int64(1), cache.clearCalls[0].accountID)
require.Equal(t, []int64{44, 101}, cache.clearCalls[0].groupIDs)
require.Equal(t, int64(2), cache.clearCalls[1].accountID)
require.Equal(t, []int64{55, 101}, cache.clearCalls[1].groupIDs)
}
@@ -0,0 +1,205 @@
package admin
import (
"context"
"encoding/json"
"net/http"
"net/http/httptest"
"testing"
"time"
"github.com/Wei-Shaw/sub2api/internal/service"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/require"
)
type affinityDetailsAdminService struct {
*stubAdminService
account service.Account
usersByID map[int64]service.User
}
func (s *affinityDetailsAdminService) GetAccount(_ context.Context, id int64) (*service.Account, error) {
if s.account.ID == id {
acc := s.account
return &acc, nil
}
return s.stubAdminService.GetAccount(context.Background(), id)
}
func (s *affinityDetailsAdminService) GetUser(_ context.Context, id int64) (*service.User, error) {
if u, ok := s.usersByID[id]; ok {
user := u
return &user, nil
}
return s.stubAdminService.GetUser(context.Background(), id)
}
type affinityDetailsGatewayCacheStub struct {
clients []service.AffinityClient
}
func (s *affinityDetailsGatewayCacheStub) GetSessionAccountID(_ context.Context, _ int64, _ string) (int64, error) {
return 0, nil
}
func (s *affinityDetailsGatewayCacheStub) SetSessionAccountID(_ context.Context, _ int64, _ string, _ int64, _ time.Duration) error {
return nil
}
func (s *affinityDetailsGatewayCacheStub) RefreshSessionTTL(_ context.Context, _ int64, _ string, _ time.Duration) error {
return nil
}
func (s *affinityDetailsGatewayCacheStub) DeleteSessionAccountID(_ context.Context, _ int64, _ string) error {
return nil
}
func (s *affinityDetailsGatewayCacheStub) GetAffinityAccounts(_ context.Context, _ int64, _ int64, _ string, _ time.Duration) ([]int64, error) {
return nil, nil
}
func (s *affinityDetailsGatewayCacheStub) UpdateAffinity(_ context.Context, _ int64, _ int64, _ string, _ int64, _ time.Duration) error {
return nil
}
func (s *affinityDetailsGatewayCacheStub) GetAccountAffinityCountBatch(_ context.Context, _ int64, _ []int64, _ time.Duration) (map[int64]int64, error) {
return map[int64]int64{}, nil
}
func (s *affinityDetailsGatewayCacheStub) GetAccountAffinityClientsBatch(_ context.Context, _ map[int64][]int64, _ time.Duration) (map[int64][]string, error) {
return map[int64][]string{}, nil
}
func (s *affinityDetailsGatewayCacheStub) GetAccountAffinityClientsWithScores(_ context.Context, _ int64, _ []int64, _ time.Duration) ([]service.AffinityClient, error) {
return s.clients, nil
}
func (s *affinityDetailsGatewayCacheStub) ClearAccountAffinity(_ context.Context, _ int64, _ []int64) error {
return nil
}
func (s *affinityDetailsGatewayCacheStub) GetAffinityMultiCount(_ context.Context, _ int64, _ int64, _ int64, _ time.Duration) (users, clients, perUser int64, err error) {
return 0, 0, 0, nil
}
func setupAffinityDetailsRouter(adminSvc service.AdminService, gatewayCache service.GatewayCache) *gin.Engine {
gin.SetMode(gin.TestMode)
router := gin.New()
handler := NewAccountHandler(adminSvc, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, gatewayCache)
router.GET("/api/v1/admin/accounts/:id/affinity-details", handler.GetAffinityDetails)
return router
}
func TestAccountHandlerGetAffinityDetails_ContractFields(t *testing.T) {
adminSvc := &affinityDetailsAdminService{
stubAdminService: newStubAdminService(),
account: service.Account{
ID: 88,
Platform: service.PlatformAnthropic,
Type: service.AccountTypeOAuth,
Status: service.StatusActive,
GroupIDs: []int64{1},
Extra: map[string]any{
"client_affinity_enabled": true,
"pinned_users": []any{float64(42), float64(99)},
},
},
usersByID: map[int64]service.User{
42: {ID: 42, Email: "user42@example.com"},
7: {ID: 7, Username: "user7"},
},
}
cache := &affinityDetailsGatewayCacheStub{
clients: []service.AffinityClient{
{UserID: 42, ClientID: "client-a", LastActive: time.Now().Add(-time.Minute)},
{UserID: 7, ClientID: "client-b", LastActive: time.Now().Add(-2 * time.Minute)},
{UserID: 42, ClientID: "client-c", LastActive: time.Now().Add(-3 * time.Minute)},
},
}
router := setupAffinityDetailsRouter(adminSvc, cache)
rec := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodGet, "/api/v1/admin/accounts/88/affinity-details", nil)
router.ServeHTTP(rec, req)
require.Equal(t, http.StatusOK, rec.Code)
var resp struct {
Code int `json:"code"`
Data struct {
Users []struct {
UserID int64 `json:"user_id"`
UserEmail string `json:"user_email"`
ClientCount int `json:"client_count"`
IsPinned bool `json:"is_pinned"`
Clients []struct {
ClientID string `json:"client_id"`
} `json:"clients"`
} `json:"users"`
TotalUsers int `json:"total_users"`
TotalClients int `json:"total_clients"`
PinnedUsers []int64 `json:"pinned_users"`
} `json:"data"`
}
require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &resp))
require.Equal(t, 0, resp.Code)
require.Equal(t, 2, resp.Data.TotalUsers)
require.Equal(t, 3, resp.Data.TotalClients)
require.Equal(t, []int64{42, 99}, resp.Data.PinnedUsers)
require.Len(t, resp.Data.Users, 2)
usersByID := make(map[int64]struct {
UserEmail string
ClientCount int
IsPinned bool
ClientLen int
}, len(resp.Data.Users))
for _, u := range resp.Data.Users {
usersByID[u.UserID] = struct {
UserEmail string
ClientCount int
IsPinned bool
ClientLen int
}{
UserEmail: u.UserEmail,
ClientCount: u.ClientCount,
IsPinned: u.IsPinned,
ClientLen: len(u.Clients),
}
}
require.Equal(t, "user42@example.com", usersByID[42].UserEmail)
require.Equal(t, 2, usersByID[42].ClientCount)
require.True(t, usersByID[42].IsPinned)
require.Equal(t, 2, usersByID[42].ClientLen)
require.Equal(t, "user7", usersByID[7].UserEmail)
require.Equal(t, 1, usersByID[7].ClientCount)
require.False(t, usersByID[7].IsPinned)
require.Equal(t, 1, usersByID[7].ClientLen)
}
func TestAccountHandlerGetAffinityDetails_DisabledReturnsEmptyContract(t *testing.T) {
adminSvc := &affinityDetailsAdminService{
stubAdminService: newStubAdminService(),
account: service.Account{
ID: 89,
Platform: service.PlatformAnthropic,
Type: service.AccountTypeOAuth,
Status: service.StatusActive,
GroupIDs: []int64{1},
},
usersByID: map[int64]service.User{},
}
router := setupAffinityDetailsRouter(adminSvc, &affinityDetailsGatewayCacheStub{})
rec := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodGet, "/api/v1/admin/accounts/89/affinity-details", nil)
router.ServeHTTP(rec, req)
require.Equal(t, http.StatusOK, rec.Code)
var resp struct {
Code int `json:"code"`
Data struct {
Users []any `json:"users"`
TotalUsers int `json:"total_users"`
TotalClients int `json:"total_clients"`
PinnedUsers []int64 `json:"pinned_users"`
} `json:"data"`
}
require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &resp))
require.Equal(t, 0, resp.Code)
require.Empty(t, resp.Data.Users)
require.Equal(t, 0, resp.Data.TotalUsers)
require.Equal(t, 0, resp.Data.TotalClients)
require.Empty(t, resp.Data.PinnedUsers)
}
@@ -28,7 +28,7 @@ func (s *availableModelsAdminService) GetAccount(_ context.Context, id int64) (*
func setupAvailableModelsRouter(adminSvc service.AdminService) *gin.Engine {
gin.SetMode(gin.TestMode)
router := gin.New()
handler := NewAccountHandler(adminSvc, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil)
handler := NewAccountHandler(adminSvc, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil)
router.GET("/api/v1/admin/accounts/:id/models", handler.GetAvailableModels)
return router
}
@@ -15,7 +15,7 @@ import (
func setupAccountMixedChannelRouter(adminSvc *stubAdminService) *gin.Engine {
gin.SetMode(gin.TestMode)
router := gin.New()
accountHandler := NewAccountHandler(adminSvc, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil)
accountHandler := NewAccountHandler(adminSvc, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil)
router.POST("/api/v1/admin/accounts/check-mixed-channel", accountHandler.CheckMixedChannel)
router.POST("/api/v1/admin/accounts", accountHandler.Create)
router.PUT("/api/v1/admin/accounts/:id", accountHandler.Update)
@@ -111,7 +111,7 @@ func TestAccountHandlerCreateMixedChannelConflictSimplifiedResponse(t *testing.T
var resp map[string]any
require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &resp))
require.Equal(t, "mixed_channel_warning", resp["error"])
require.Contains(t, resp["message"], "mixed_channel_warning")
require.Contains(t, resp["message"], "claude-max")
_, hasDetails := resp["details"]
_, hasRequireConfirmation := resp["require_confirmation"]
require.False(t, hasDetails)
@@ -140,7 +140,7 @@ func TestAccountHandlerUpdateMixedChannelConflictSimplifiedResponse(t *testing.T
var resp map[string]any
require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &resp))
require.Equal(t, "mixed_channel_warning", resp["error"])
require.Contains(t, resp["message"], "mixed_channel_warning")
require.Contains(t, resp["message"], "claude-max")
_, hasDetails := resp["details"]
_, hasRequireConfirmation := resp["require_confirmation"]
require.False(t, hasDetails)
@@ -29,6 +29,7 @@ func TestAccountHandler_Create_AnthropicAPIKeyPassthroughExtraForwarded(t *testi
nil,
nil,
nil,
nil,
)
router := gin.New()
@@ -36,7 +36,7 @@ func (f *failingAdminService) UpdateAccount(ctx context.Context, id int64, input
func setupAccountHandlerWithService(adminSvc service.AdminService) (*gin.Engine, *AccountHandler) {
gin.SetMode(gin.TestMode)
router := gin.New()
handler := NewAccountHandler(adminSvc, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil)
handler := NewAccountHandler(adminSvc, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil)
router.POST("/api/v1/admin/accounts/batch-update-credentials", handler.BatchUpdateCredentials)
return router, handler
}
@@ -0,0 +1,148 @@
package admin
import (
"log/slog"
"net/http"
"github.com/Wei-Shaw/sub2api/internal/pkg/response"
"github.com/Wei-Shaw/sub2api/internal/service"
"github.com/gin-gonic/gin"
)
// GDriveOAuthHandler 处理 Google Drive OAuth 授权流程。
type GDriveOAuthHandler struct {
settingService *service.SettingService
gdriveOAuth *service.SoraGDriveOAuthService
gdriveStorage *service.SoraGDriveStorage
}
// NewGDriveOAuthHandler 创建 GDrive OAuth Handler。
func NewGDriveOAuthHandler(settingService *service.SettingService, gdriveOAuth *service.SoraGDriveOAuthService, gdriveStorage *service.SoraGDriveStorage) *GDriveOAuthHandler {
return &GDriveOAuthHandler{
settingService: settingService,
gdriveOAuth: gdriveOAuth,
gdriveStorage: gdriveStorage,
}
}
// StartOAuthRequest 启动 OAuth 授权请求。
type StartOAuthRequest struct {
ClientID string `json:"client_id" binding:"required"`
ClientSecret string `json:"client_secret" binding:"required"`
RedirectURI string `json:"redirect_uri" binding:"required"`
}
// StartOAuth 生成 Google OAuth 授权 URL。
// POST /api/v1/admin/settings/sora-storage/gdrive-oauth/start
func (h *GDriveOAuthHandler) StartOAuth(c *gin.Context) {
if h.gdriveOAuth == nil {
response.Error(c, http.StatusInternalServerError, "GDrive OAuth service not initialized")
return
}
var req StartOAuthRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.BadRequest(c, "Invalid request: "+err.Error())
return
}
authURL, state, err := h.gdriveOAuth.GenerateAuthURL(req.ClientID, req.ClientSecret, req.RedirectURI)
if err != nil {
response.Error(c, http.StatusInternalServerError, "生成授权 URL 失败: "+err.Error())
return
}
response.Success(c, gin.H{
"auth_url": authURL,
"state": state,
})
}
// OAuthCallbackRequest OAuth 回调请求。
type OAuthCallbackRequest struct {
ClientID string `json:"client_id" binding:"required"`
ClientSecret string `json:"client_secret" binding:"required"`
RedirectURI string `json:"redirect_uri" binding:"required"`
Code string `json:"code" binding:"required"`
ProfileID string `json:"profile_id"` // 要保存到的 profile ID(可选)
}
// OAuthCallback 用授权码换取 refresh_token 并保存到 profile。
// POST /api/v1/admin/settings/sora-storage/gdrive-oauth/callback
func (h *GDriveOAuthHandler) OAuthCallback(c *gin.Context) {
if h.gdriveOAuth == nil {
response.Error(c, http.StatusInternalServerError, "GDrive OAuth service not initialized")
return
}
var req OAuthCallbackRequest
if err := c.ShouldBindJSON(&req); err != nil {
response.BadRequest(c, "Invalid request: "+err.Error())
return
}
refreshToken, err := h.gdriveOAuth.ExchangeCode(c.Request.Context(), req.ClientID, req.ClientSecret, req.RedirectURI, req.Code)
if err != nil {
slog.Error("[GDriveOAuth] exchange failed",
"client_id_len", len(req.ClientID),
"client_secret_len", len(req.ClientSecret),
"redirect_uri", req.RedirectURI,
"code_len", len(req.Code),
"error", err,
)
response.Error(c, http.StatusBadRequest, "换取 refresh_token 失败: "+err.Error())
return
}
// 如果指定了 profile_id,自动保存 refresh_token 到 profile
if req.ProfileID != "" {
profiles, err := h.settingService.ListSoraS3Profiles(c.Request.Context())
if err == nil {
for _, p := range profiles.Items {
if p.ProfileID == req.ProfileID {
_, _ = h.settingService.UpdateSoraS3Profile(c.Request.Context(), req.ProfileID, &service.SoraS3Profile{
Name: p.Name,
Provider: p.Provider,
AccessMode: p.AccessMode,
Enabled: p.Enabled,
Endpoint: p.Endpoint,
Region: p.Region,
Bucket: p.Bucket,
AccessKeyID: p.AccessKeyID,
Prefix: p.Prefix,
ForcePathStyle: p.ForcePathStyle,
CDNURL: p.CDNURL,
DefaultStorageQuotaBytes: p.DefaultStorageQuotaBytes,
AuthType: p.AuthType,
ClientID: p.ClientID,
FolderID: p.FolderID,
RefreshToken: refreshToken,
})
break
}
}
}
}
response.Success(c, gin.H{
"refresh_token": refreshToken,
"message": "OAuth 授权成功",
})
}
// TestGDriveStorage 测试 GDrive 存储的完整上传→下载→删除流程。
// POST /api/v1/admin/settings/sora-storage/gdrive-test
func (h *GDriveOAuthHandler) TestGDriveStorage(c *gin.Context) {
if h.gdriveStorage == nil {
response.Error(c, http.StatusInternalServerError, "GDrive storage not initialized")
return
}
result, err := h.gdriveStorage.TestFullCycle(c.Request.Context())
if err != nil {
response.Error(c, http.StatusBadRequest, "GDrive 测试失败: "+err.Error())
return
}
response.Success(c, result)
}
@@ -103,9 +103,10 @@ type CreateGroupRequest struct {
FallbackGroupID *int64 `json:"fallback_group_id"`
FallbackGroupIDOnInvalidRequest *int64 `json:"fallback_group_id_on_invalid_request"`
// 模型路由配置(仅 anthropic 平台使用)
ModelRouting map[string][]int64 `json:"model_routing"`
ModelRoutingEnabled bool `json:"model_routing_enabled"`
MCPXMLInject *bool `json:"mcp_xml_inject"`
ModelRouting map[string][]int64 `json:"model_routing"`
ModelRoutingEnabled bool `json:"model_routing_enabled"`
MCPXMLInject *bool `json:"mcp_xml_inject"`
SimulateClaudeMaxEnabled *bool `json:"simulate_claude_max_enabled"`
// 支持的模型系列(仅 antigravity 平台使用)
SupportedModelScopes []string `json:"supported_model_scopes"`
// Sora 存储配额
@@ -141,9 +142,10 @@ type UpdateGroupRequest struct {
FallbackGroupID *int64 `json:"fallback_group_id"`
FallbackGroupIDOnInvalidRequest *int64 `json:"fallback_group_id_on_invalid_request"`
// 模型路由配置(仅 anthropic 平台使用)
ModelRouting map[string][]int64 `json:"model_routing"`
ModelRoutingEnabled *bool `json:"model_routing_enabled"`
MCPXMLInject *bool `json:"mcp_xml_inject"`
ModelRouting map[string][]int64 `json:"model_routing"`
ModelRoutingEnabled *bool `json:"model_routing_enabled"`
MCPXMLInject *bool `json:"mcp_xml_inject"`
SimulateClaudeMaxEnabled *bool `json:"simulate_claude_max_enabled"`
// 支持的模型系列(仅 antigravity 平台使用)
SupportedModelScopes *[]string `json:"supported_model_scopes"`
// Sora 存储配额
@@ -264,6 +266,7 @@ func (h *GroupHandler) Create(c *gin.Context) {
ModelRouting: req.ModelRouting,
ModelRoutingEnabled: req.ModelRoutingEnabled,
MCPXMLInject: req.MCPXMLInject,
SimulateClaudeMaxEnabled: req.SimulateClaudeMaxEnabled,
SupportedModelScopes: req.SupportedModelScopes,
SoraStorageQuotaBytes: req.SoraStorageQuotaBytes,
AllowMessagesDispatch: req.AllowMessagesDispatch,
@@ -317,6 +320,7 @@ func (h *GroupHandler) Update(c *gin.Context) {
ModelRouting: req.ModelRouting,
ModelRoutingEnabled: req.ModelRoutingEnabled,
MCPXMLInject: req.MCPXMLInject,
SimulateClaudeMaxEnabled: req.SimulateClaudeMaxEnabled,
SupportedModelScopes: req.SupportedModelScopes,
SoraStorageQuotaBytes: req.SoraStorageQuotaBytes,
AllowMessagesDispatch: req.AllowMessagesDispatch,
+124 -17
View File
@@ -37,21 +37,25 @@ func generateMenuItemID() (string, error) {
// SettingHandler 系统设置处理器
type SettingHandler struct {
settingService *service.SettingService
emailService *service.EmailService
turnstileService *service.TurnstileService
opsService *service.OpsService
soraS3Storage *service.SoraS3Storage
settingService *service.SettingService
emailService *service.EmailService
turnstileService *service.TurnstileService
opsService *service.OpsService
soraS3Storage *service.SoraS3Storage
soraGDriveStorage *service.SoraGDriveStorage
soraGenerationService *service.SoraGenerationService
}
// NewSettingHandler 创建系统设置处理器
func NewSettingHandler(settingService *service.SettingService, emailService *service.EmailService, turnstileService *service.TurnstileService, opsService *service.OpsService, soraS3Storage *service.SoraS3Storage) *SettingHandler {
func NewSettingHandler(settingService *service.SettingService, emailService *service.EmailService, turnstileService *service.TurnstileService, opsService *service.OpsService, soraS3Storage *service.SoraS3Storage, soraGDriveStorage *service.SoraGDriveStorage, soraGenerationService *service.SoraGenerationService) *SettingHandler {
return &SettingHandler{
settingService: settingService,
emailService: emailService,
turnstileService: turnstileService,
opsService: opsService,
soraS3Storage: soraS3Storage,
settingService: settingService,
emailService: emailService,
turnstileService: turnstileService,
opsService: opsService,
soraS3Storage: soraS3Storage,
soraGDriveStorage: soraGDriveStorage,
soraGenerationService: soraGenerationService,
}
}
@@ -1070,6 +1074,8 @@ func toSoraS3ProfileDTO(profile service.SoraS3Profile) dto.SoraS3Profile {
ProfileID: profile.ProfileID,
Name: profile.Name,
IsActive: profile.IsActive,
Provider: profile.GetProvider(),
AccessMode: profile.AccessMode,
Enabled: profile.Enabled,
Endpoint: profile.Endpoint,
Region: profile.Region,
@@ -1081,6 +1087,13 @@ func toSoraS3ProfileDTO(profile service.SoraS3Profile) dto.SoraS3Profile {
CDNURL: profile.CDNURL,
DefaultStorageQuotaBytes: profile.DefaultStorageQuotaBytes,
UpdatedAt: profile.UpdatedAt,
// Google Drive 专属
AuthType: profile.AuthType,
ClientID: profile.ClientID,
ClientSecretConfigured: profile.ClientSecretConfigured,
RefreshTokenConfigured: profile.RefreshTokenConfigured,
ServiceAccountConfigured: profile.ServiceAccountConfigured,
FolderID: profile.FolderID,
}
}
@@ -1160,6 +1173,8 @@ type CreateSoraS3ProfileRequest struct {
ProfileID string `json:"profile_id"`
Name string `json:"name"`
SetActive bool `json:"set_active"`
Provider string `json:"provider"` // "s3" / "gdrive"
AccessMode string `json:"access_mode"` // "direct" / "proxy"
Enabled bool `json:"enabled"`
Endpoint string `json:"endpoint"`
Region string `json:"region"`
@@ -1170,10 +1185,19 @@ type CreateSoraS3ProfileRequest struct {
ForcePathStyle bool `json:"force_path_style"`
CDNURL string `json:"cdn_url"`
DefaultStorageQuotaBytes int64 `json:"default_storage_quota_bytes"`
// Google Drive 专属
AuthType string `json:"auth_type,omitempty"`
ClientID string `json:"client_id,omitempty"`
ClientSecret string `json:"client_secret,omitempty"`
RefreshToken string `json:"refresh_token,omitempty"`
ServiceAccountJSON string `json:"service_account_json,omitempty"`
FolderID string `json:"folder_id,omitempty"`
}
type UpdateSoraS3ProfileRequest struct {
Name string `json:"name"`
Provider string `json:"provider"`
AccessMode string `json:"access_mode"`
Enabled bool `json:"enabled"`
Endpoint string `json:"endpoint"`
Region string `json:"region"`
@@ -1184,6 +1208,13 @@ type UpdateSoraS3ProfileRequest struct {
ForcePathStyle bool `json:"force_path_style"`
CDNURL string `json:"cdn_url"`
DefaultStorageQuotaBytes int64 `json:"default_storage_quota_bytes"`
// Google Drive 专属
AuthType string `json:"auth_type,omitempty"`
ClientID string `json:"client_id,omitempty"`
ClientSecret string `json:"client_secret,omitempty"`
RefreshToken string `json:"refresh_token,omitempty"`
ServiceAccountJSON string `json:"service_account_json,omitempty"`
FolderID string `json:"folder_id,omitempty"`
}
// CreateSoraS3Profile 创建 Sora S3 配置
@@ -1206,14 +1237,23 @@ func (h *SettingHandler) CreateSoraS3Profile(c *gin.Context) {
response.BadRequest(c, "Profile ID is required")
return
}
if err := validateSoraS3RequiredWhenEnabled(req.Enabled, req.Endpoint, req.Bucket, req.AccessKeyID, req.SecretAccessKey, false); err != nil {
response.BadRequest(c, err.Error())
return
// S3 专属字段验证:仅当 provider 为 s3(或未指定)时校验
provider := req.Provider
if provider == "" {
provider = "s3"
}
if provider == "s3" {
if err := validateSoraS3RequiredWhenEnabled(req.Enabled, req.Endpoint, req.Bucket, req.AccessKeyID, req.SecretAccessKey, false); err != nil {
response.BadRequest(c, err.Error())
return
}
}
created, err := h.settingService.CreateSoraS3Profile(c.Request.Context(), &service.SoraS3Profile{
ProfileID: req.ProfileID,
Name: req.Name,
Provider: req.Provider,
AccessMode: req.AccessMode,
Enabled: req.Enabled,
Endpoint: req.Endpoint,
Region: req.Region,
@@ -1224,6 +1264,13 @@ func (h *SettingHandler) CreateSoraS3Profile(c *gin.Context) {
ForcePathStyle: req.ForcePathStyle,
CDNURL: req.CDNURL,
DefaultStorageQuotaBytes: req.DefaultStorageQuotaBytes,
// Google Drive 专属
AuthType: req.AuthType,
ClientID: req.ClientID,
ClientSecret: req.ClientSecret,
RefreshToken: req.RefreshToken,
ServiceAccountJSON: req.ServiceAccountJSON,
FolderID: req.FolderID,
}, req.SetActive)
if err != nil {
response.ErrorFrom(c, err)
@@ -1266,13 +1313,25 @@ func (h *SettingHandler) UpdateSoraS3Profile(c *gin.Context) {
response.ErrorFrom(c, service.ErrSoraS3ProfileNotFound)
return
}
if err := validateSoraS3RequiredWhenEnabled(req.Enabled, req.Endpoint, req.Bucket, req.AccessKeyID, req.SecretAccessKey, existing.SecretAccessKeyConfigured); err != nil {
response.BadRequest(c, err.Error())
return
// S3 专属字段验证
provider := req.Provider
if provider == "" && existing != nil {
provider = existing.GetProvider()
}
if provider == "" {
provider = "s3"
}
if provider == "s3" {
if err := validateSoraS3RequiredWhenEnabled(req.Enabled, req.Endpoint, req.Bucket, req.AccessKeyID, req.SecretAccessKey, existing.SecretAccessKeyConfigured); err != nil {
response.BadRequest(c, err.Error())
return
}
}
updated, updateErr := h.settingService.UpdateSoraS3Profile(c.Request.Context(), profileID, &service.SoraS3Profile{
Name: req.Name,
Provider: req.Provider,
AccessMode: req.AccessMode,
Enabled: req.Enabled,
Endpoint: req.Endpoint,
Region: req.Region,
@@ -1283,6 +1342,13 @@ func (h *SettingHandler) UpdateSoraS3Profile(c *gin.Context) {
ForcePathStyle: req.ForcePathStyle,
CDNURL: req.CDNURL,
DefaultStorageQuotaBytes: req.DefaultStorageQuotaBytes,
// Google Drive 专属
AuthType: req.AuthType,
ClientID: req.ClientID,
ClientSecret: req.ClientSecret,
RefreshToken: req.RefreshToken,
ServiceAccountJSON: req.ServiceAccountJSON,
FolderID: req.FolderID,
})
if updateErr != nil {
response.ErrorFrom(c, updateErr)
@@ -1583,3 +1649,44 @@ func (h *SettingHandler) UpdateStreamTimeoutSettings(c *gin.Context) {
ThresholdWindowMinutes: updatedSettings.ThresholdWindowMinutes,
})
}
// GetGDriveQuota 获取 Google Drive 配额信息。
// GET /api/v1/admin/settings/sora-storage/gdrive-quota
func (h *SettingHandler) GetGDriveQuota(c *gin.Context) {
if h.soraGDriveStorage == nil {
response.Error(c, http.StatusServiceUnavailable, "GDrive storage not configured")
return
}
quota, err := h.soraGDriveStorage.GetQuotaInfo(c.Request.Context())
if err != nil {
response.Error(c, http.StatusInternalServerError, fmt.Sprintf("failed to get GDrive quota: %v", err))
return
}
response.Success(c, quota)
}
// GetStorageVideoStats 获取各存储类型的视频统计信息。
// GET /api/v1/admin/settings/sora-storage/video-stats
func (h *SettingHandler) GetStorageVideoStats(c *gin.Context) {
if h.soraGenerationService == nil {
response.Error(c, http.StatusServiceUnavailable, "generation service not configured")
return
}
storageTypes := []string{service.SoraStorageTypeS3, service.SoraStorageTypeGDrive}
result := make(map[string]*service.StorageVideoStats, len(storageTypes))
for _, st := range storageTypes {
completed, inProgress, err := h.soraGenerationService.CountByStorageType(c.Request.Context(), st)
if err != nil {
log.Printf("[SettingHandler] CountByStorageType(%s) error: %v", st, err)
continue
}
result[st] = &service.StorageVideoStats{
Completed: completed,
InProgress: inProgress,
}
}
response.Success(c, result)
}
+22 -10
View File
@@ -135,16 +135,17 @@ func GroupFromServiceAdmin(g *service.Group) *AdminGroup {
return nil
}
out := &AdminGroup{
Group: groupFromServiceBase(g),
ModelRouting: g.ModelRouting,
ModelRoutingEnabled: g.ModelRoutingEnabled,
MCPXMLInject: g.MCPXMLInject,
DefaultMappedModel: g.DefaultMappedModel,
SupportedModelScopes: g.SupportedModelScopes,
AccountCount: g.AccountCount,
ActiveAccountCount: g.ActiveAccountCount,
RateLimitedAccountCount: g.RateLimitedAccountCount,
SortOrder: g.SortOrder,
Group: groupFromServiceBase(g),
ModelRouting: g.ModelRouting,
ModelRoutingEnabled: g.ModelRoutingEnabled,
MCPXMLInject: g.MCPXMLInject,
DefaultMappedModel: g.DefaultMappedModel,
SimulateClaudeMaxEnabled: g.SimulateClaudeMaxEnabled,
SupportedModelScopes: g.SupportedModelScopes,
AccountCount: g.AccountCount,
ActiveAccountCount: g.ActiveAccountCount,
RateLimitedAccountCount: g.RateLimitedAccountCount,
SortOrder: g.SortOrder,
}
if len(g.AccountGroups) > 0 {
out.AccountGroups = make([]AccountGroup, 0, len(g.AccountGroups))
@@ -266,6 +267,17 @@ func AccountFromServiceShallow(a *service.Account) *Account {
}
}
// 客户端亲和调度(Anthropic 和 Antigravity 账号)
if a.IsAffinityEnabled() {
enabled := true
out.ClientAffinityEnabled = &enabled
allow := a.IsAffinityAllowSwitch()
out.AffinityAllowSwitch = &allow
if pinnedUsers := a.GetPinnedUsers(); len(pinnedUsers) > 0 {
out.PinnedUserIDs = pinnedUsers
}
}
// 提取账号配额限制(apikey / bedrock 类型有效)
if a.IsAPIKeyOrBedrock() {
if limit := a.GetQuotaLimit(); limit > 0 {
+11 -1
View File
@@ -133,11 +133,13 @@ type SoraS3Settings struct {
DefaultStorageQuotaBytes int64 `json:"default_storage_quota_bytes"`
}
// SoraS3Profile Sora S3 存储配置项 DTO(响应用,不含敏感字段)
// SoraS3Profile Sora 存储配置项 DTO(响应用,不含敏感字段)
type SoraS3Profile struct {
ProfileID string `json:"profile_id"`
Name string `json:"name"`
IsActive bool `json:"is_active"`
Provider string `json:"provider"` // "s3" / "gdrive"
AccessMode string `json:"access_mode"` // "direct" / "proxy"
Enabled bool `json:"enabled"`
Endpoint string `json:"endpoint"`
Region string `json:"region"`
@@ -149,6 +151,14 @@ type SoraS3Profile struct {
CDNURL string `json:"cdn_url"`
DefaultStorageQuotaBytes int64 `json:"default_storage_quota_bytes"`
UpdatedAt string `json:"updated_at"`
// --- Google Drive 专属 ---
AuthType string `json:"auth_type,omitempty"`
ClientID string `json:"client_id,omitempty"`
ClientSecretConfigured bool `json:"client_secret_configured"`
RefreshTokenConfigured bool `json:"refresh_token_configured"`
ServiceAccountConfigured bool `json:"service_account_configured"`
FolderID string `json:"folder_id,omitempty"`
}
// ListSoraS3ProfilesResponse Sora S3 配置列表响应
+19
View File
@@ -117,6 +117,8 @@ type AdminGroup struct {
// MCP XML 协议注入(仅 antigravity 平台使用)
MCPXMLInject bool `json:"mcp_xml_inject"`
// Claude usage 模拟开关(仅管理员可见)
SimulateClaudeMaxEnabled bool `json:"simulate_claude_max_enabled"`
// OpenAI Messages 调度配置(仅 openai 平台使用)
DefaultMappedModel string `json:"default_mapped_model"`
@@ -197,6 +199,23 @@ type Account struct {
CacheTTLOverrideEnabled *bool `json:"cache_ttl_override_enabled,omitempty"`
CacheTTLOverrideTarget *string `json:"cache_ttl_override_target,omitempty"`
// 客户端亲和调度(Anthropic 和 Antigravity 账号有效)
// 启用后新会话会优先调度到客户端之前使用过的账号
ClientAffinityEnabled *bool `json:"client_affinity_enabled,omitempty"`
// 亲和允许切换(默认 true)
AffinityAllowSwitch *bool `json:"affinity_allow_switch,omitempty"`
// 亲和用户数量(admin 列表端点注入)
AffinityUserCount *int64 `json:"affinity_user_count,omitempty"`
// 指定亲和用户 ID 列表
PinnedUserIDs []int64 `json:"pinned_user_ids,omitempty"`
// 亲和客户端数据(仅 admin 列表端点注入,不由 mapper 填充)
AffinityClientCount *int64 `json:"affinity_client_count,omitempty"`
AffinityClients []string `json:"affinity_clients,omitempty"`
// API Key 账号配额限制
QuotaLimit *float64 `json:"quota_limit,omitempty"`
QuotaUsed *float64 `json:"quota_used,omitempty"`
+13 -2
View File
@@ -291,7 +291,7 @@ func (h *GatewayHandler) Messages(c *gin.Context) {
}
for {
selection, err := h.gatewayService.SelectAccountWithLoadAwareness(c.Request.Context(), apiKey.GroupID, sessionKey, reqModel, fs.FailedAccountIDs, "") // Gemini 不使用会话限制
selection, err := h.gatewayService.SelectAccountWithLoadAwareness(c.Request.Context(), apiKey.GroupID, sessionKey, reqModel, fs.FailedAccountIDs, "", int64(0)) // Gemini 不使用会话限制
if err != nil {
if len(fs.FailedAccountIDs) == 0 {
h.handleStreamingAwareError(c, http.StatusServiceUnavailable, "api_error", "No available accounts: "+err.Error(), streamStarted)
@@ -453,6 +453,7 @@ func (h *GatewayHandler) Messages(c *gin.Context) {
h.submitUsageRecordTask(func(ctx context.Context) {
if err := h.gatewayService.RecordUsage(ctx, &service.RecordUsageInput{
Result: result,
ParsedRequest: parsedReq,
APIKey: apiKey,
User: apiKey.User,
Account: account,
@@ -500,8 +501,16 @@ func (h *GatewayHandler) Messages(c *gin.Context) {
for {
// 选择支持该模型的账号
selection, err := h.gatewayService.SelectAccountWithLoadAwareness(c.Request.Context(), currentAPIKey.GroupID, sessionKey, reqModel, fs.FailedAccountIDs, parsedReq.MetadataUserID)
selection, err := h.gatewayService.SelectAccountWithLoadAwareness(c.Request.Context(), currentAPIKey.GroupID, sessionKey, reqModel, fs.FailedAccountIDs, parsedReq.MetadataUserID, subject.UserID)
if err != nil {
if errors.Is(err, service.ErrAffinityNoSwitch) {
h.handleStreamingAwareError(c, http.StatusServiceUnavailable, "api_error", "Affinity account unavailable and switching is disabled", streamStarted)
return
}
if errors.Is(err, service.ErrAffinityLimitExceeded) {
h.handleStreamingAwareError(c, http.StatusTooManyRequests, "api_error", "Affinity client limit exceeded", streamStarted)
return
}
if len(fs.FailedAccountIDs) == 0 {
h.handleStreamingAwareError(c, http.StatusServiceUnavailable, "api_error", "No available accounts: "+err.Error(), streamStarted)
return
@@ -647,6 +656,7 @@ func (h *GatewayHandler) Messages(c *gin.Context) {
// ===== 用户消息串行队列 END =====
// 转发请求 - 根据账号平台分流
c.Set("parsed_request", parsedReq)
var result *service.ForwardResult
requestCtx := c.Request.Context()
if fs.SwitchCount > 0 {
@@ -772,6 +782,7 @@ func (h *GatewayHandler) Messages(c *gin.Context) {
h.submitUsageRecordTask(func(ctx context.Context) {
if err := h.gatewayService.RecordUsage(ctx, &service.RecordUsageInput{
Result: result,
ParsedRequest: parsedReq,
APIKey: currentAPIKey,
User: currentAPIKey.User,
Account: account,
@@ -352,7 +352,7 @@ func (h *GatewayHandler) GeminiV1BetaModels(c *gin.Context) {
}
for {
selection, err := h.gatewayService.SelectAccountWithLoadAwareness(c.Request.Context(), apiKey.GroupID, sessionKey, modelName, fs.FailedAccountIDs, "") // Gemini 不使用会话限制
selection, err := h.gatewayService.SelectAccountWithLoadAwareness(c.Request.Context(), apiKey.GroupID, sessionKey, modelName, fs.FailedAccountIDs, "", int64(0)) // Gemini 不使用会话限制
if err != nil {
if len(fs.FailedAccountIDs) == 0 {
googleError(c, http.StatusServiceUnavailable, "No available Gemini accounts: "+err.Error())
+2
View File
@@ -29,6 +29,7 @@ type AdminHandlers struct {
ErrorPassthrough *admin.ErrorPassthroughHandler
APIKey *admin.AdminAPIKeyHandler
ScheduledTest *admin.ScheduledTestHandler
GDriveOAuth *admin.GDriveOAuthHandler
}
// Handlers contains all HTTP handlers
@@ -45,6 +46,7 @@ type Handlers struct {
OpenAIGateway *OpenAIGatewayHandler
SoraGateway *SoraGatewayHandler
SoraClient *SoraClientHandler
SoraVideos *SoraVideosHandler
Setting *SettingHandler
Totp *TotpHandler
}
+41 -37
View File
@@ -31,7 +31,7 @@ const (
type SoraClientHandler struct {
genService *service.SoraGenerationService
quotaService *service.SoraQuotaService
s3Storage *service.SoraS3Storage
objectStorage service.SoraObjectStorage
soraGatewayService *service.SoraGatewayService
gatewayService *service.GatewayService
mediaStorage *service.SoraMediaStorage
@@ -48,7 +48,7 @@ type SoraClientHandler struct {
func NewSoraClientHandler(
genService *service.SoraGenerationService,
quotaService *service.SoraQuotaService,
s3Storage *service.SoraS3Storage,
objectStorage service.SoraObjectStorage,
soraGatewayService *service.SoraGatewayService,
gatewayService *service.GatewayService,
mediaStorage *service.SoraMediaStorage,
@@ -57,7 +57,7 @@ func NewSoraClientHandler(
return &SoraClientHandler{
genService: genService,
quotaService: quotaService,
s3Storage: s3Storage,
objectStorage: objectStorage,
soraGatewayService: soraGatewayService,
gatewayService: gatewayService,
mediaStorage: mediaStorage,
@@ -291,11 +291,11 @@ func (h *SoraClientHandler) processGeneration(genID int64, userID int64, groupID
return
}
// 三层降级存储:S3 → 本地 → 上游临时 URL
// 三层降级存储:对象存储 → 本地 → 上游临时 URL
storedURL, storedURLs, storageType, s3Keys, fileSize := h.storeMediaWithDegradation(ctx, userID, mediaType, mediaURL, mediaURLs)
usageAdded := false
if (storageType == service.SoraStorageTypeS3 || storageType == service.SoraStorageTypeLocal) && fileSize > 0 && h.quotaService != nil {
if (service.IsObjectStorageType(storageType) || storageType == service.SoraStorageTypeLocal) && fileSize > 0 && h.quotaService != nil {
if err := h.quotaService.AddUsage(ctx, userID, fileSize); err != nil {
h.cleanupStoredMedia(ctx, storageType, s3Keys, storedURLs)
var quotaErr *service.QuotaExceededError
@@ -346,39 +346,41 @@ func (h *SoraClientHandler) storeMediaWithDegradation(
urls = []string{mediaURL}
}
// 第一层:尝试 S3
if h.s3Storage != nil && h.s3Storage.Enabled(ctx) {
// 第一层:尝试对象存储(S3 / Google Drive)
if h.objectStorage != nil && h.objectStorage.Enabled(ctx) {
keys := make([]string, 0, len(urls))
var totalSize int64
var actualStorageType string
allOK := true
for _, u := range urls {
key, size, err := h.s3Storage.UploadFromURL(ctx, userID, u)
key, size, st, err := h.objectStorage.UploadFromURL(ctx, userID, u)
if err != nil {
logger.LegacyPrintf("handler.sora_client", "[SoraClient] S3 上传失败 err=%v", err)
logger.LegacyPrintf("handler.sora_client", "[SoraClient] 对象存储上传失败 type=%s err=%v", h.objectStorage.StorageType(), err)
allOK = false
// 清理已上传的文件
if len(keys) > 0 {
_ = h.s3Storage.DeleteObjects(ctx, keys)
_ = h.objectStorage.DeleteObjects(ctx, keys)
}
break
}
keys = append(keys, key)
totalSize += size
actualStorageType = st
}
if allOK && len(keys) > 0 {
accessURLs := make([]string, 0, len(keys))
for _, key := range keys {
accessURL, err := h.s3Storage.GetAccessURL(ctx, key)
accessURL, err := h.objectStorage.GetAccessURL(ctx, key)
if err != nil {
logger.LegacyPrintf("handler.sora_client", "[SoraClient] 生成 S3 访问 URL 失败 err=%v", err)
_ = h.s3Storage.DeleteObjects(ctx, keys)
logger.LegacyPrintf("handler.sora_client", "[SoraClient] 生成访问 URL 失败 type=%s err=%v", h.objectStorage.StorageType(), err)
_ = h.objectStorage.DeleteObjects(ctx, keys)
allOK = false
break
}
accessURLs = append(accessURLs, accessURL)
}
if allOK && len(accessURLs) > 0 {
return accessURLs[0], accessURLs, service.SoraStorageTypeS3, keys, totalSize
return accessURLs[0], accessURLs, actualStorageType, keys, totalSize
}
}
}
@@ -678,7 +680,7 @@ func (h *SoraClientHandler) SaveToStorage(c *gin.Context) {
return
}
if h.s3Storage == nil || !h.s3Storage.Enabled(c.Request.Context()) {
if h.objectStorage == nil || !h.objectStorage.Enabled(c.Request.Context()) {
response.Error(c, http.StatusServiceUnavailable, "云存储未配置,请联系管理员")
return
}
@@ -697,24 +699,24 @@ func (h *SoraClientHandler) SaveToStorage(c *gin.Context) {
var totalSize int64
for _, sourceURL := range sourceURLs {
objectKey, fileSize, uploadErr := h.s3Storage.UploadFromURL(c.Request.Context(), userID, sourceURL)
objectKey, fileSize, _, uploadErr := h.objectStorage.UploadFromURL(c.Request.Context(), userID, sourceURL)
if uploadErr != nil {
if len(uploadedKeys) > 0 {
_ = h.s3Storage.DeleteObjects(c.Request.Context(), uploadedKeys)
_ = h.objectStorage.DeleteObjects(c.Request.Context(), uploadedKeys)
}
var upstreamErr *service.UpstreamDownloadError
if errors.As(uploadErr, &upstreamErr) && (upstreamErr.StatusCode == http.StatusForbidden || upstreamErr.StatusCode == http.StatusNotFound) {
response.Error(c, http.StatusGone, "媒体链接已过期,无法保存")
return
}
response.Error(c, http.StatusInternalServerError, "上传到 S3 失败: "+uploadErr.Error())
response.Error(c, http.StatusInternalServerError, "上传到存储失败: "+uploadErr.Error())
return
}
accessURL, err := h.s3Storage.GetAccessURL(c.Request.Context(), objectKey)
accessURL, err := h.objectStorage.GetAccessURL(c.Request.Context(), objectKey)
if err != nil {
uploadedKeys = append(uploadedKeys, objectKey)
_ = h.s3Storage.DeleteObjects(c.Request.Context(), uploadedKeys)
response.Error(c, http.StatusInternalServerError, "生成 S3 访问链接失败: "+err.Error())
_ = h.objectStorage.DeleteObjects(c.Request.Context(), uploadedKeys)
response.Error(c, http.StatusInternalServerError, "生成访问链接失败: "+err.Error())
return
}
uploadedKeys = append(uploadedKeys, objectKey)
@@ -725,7 +727,7 @@ func (h *SoraClientHandler) SaveToStorage(c *gin.Context) {
usageAdded := false
if totalSize > 0 && h.quotaService != nil {
if err := h.quotaService.AddUsage(c.Request.Context(), userID, totalSize); err != nil {
_ = h.s3Storage.DeleteObjects(c.Request.Context(), uploadedKeys)
_ = h.objectStorage.DeleteObjects(c.Request.Context(), uploadedKeys)
var quotaErr *service.QuotaExceededError
if errors.As(err, &quotaErr) {
response.Error(c, http.StatusTooManyRequests, "存储配额已满,请删除不需要的作品释放空间")
@@ -742,11 +744,11 @@ func (h *SoraClientHandler) SaveToStorage(c *gin.Context) {
id,
accessURLs[0],
accessURLs,
service.SoraStorageTypeS3,
h.objectStorage.StorageType(),
uploadedKeys,
totalSize,
); err != nil {
_ = h.s3Storage.DeleteObjects(c.Request.Context(), uploadedKeys)
_ = h.objectStorage.DeleteObjects(c.Request.Context(), uploadedKeys)
if usageAdded && h.quotaService != nil {
_ = h.quotaService.ReleaseUsage(c.Request.Context(), userID, totalSize)
}
@@ -755,7 +757,7 @@ func (h *SoraClientHandler) SaveToStorage(c *gin.Context) {
}
response.Success(c, gin.H{
"message": "已保存到 S3",
"message": "已保存到云存储",
"object_key": uploadedKeys[0],
"object_keys": uploadedKeys,
})
@@ -764,28 +766,30 @@ func (h *SoraClientHandler) SaveToStorage(c *gin.Context) {
// GetStorageStatus 返回存储状态。
// GET /api/v1/sora/storage-status
func (h *SoraClientHandler) GetStorageStatus(c *gin.Context) {
s3Enabled := h.s3Storage != nil && h.s3Storage.Enabled(c.Request.Context())
s3Healthy := false
if s3Enabled {
s3Healthy = h.s3Storage.IsHealthy(c.Request.Context())
objectStorageEnabled := h.objectStorage != nil && h.objectStorage.Enabled(c.Request.Context())
objectStorageHealthy := false
storageType := ""
if objectStorageEnabled {
objectStorageHealthy = h.objectStorage.IsHealthy(c.Request.Context())
storageType = h.objectStorage.StorageType()
}
localEnabled := h.mediaStorage != nil && h.mediaStorage.Enabled()
response.Success(c, gin.H{
"s3_enabled": s3Enabled,
"s3_healthy": s3Healthy,
"s3_enabled": objectStorageEnabled, // 保留字段名向后兼容
"s3_healthy": objectStorageHealthy,
"storage_type": storageType,
"local_enabled": localEnabled,
})
}
func (h *SoraClientHandler) cleanupStoredMedia(ctx context.Context, storageType string, s3Keys []string, localPaths []string) {
switch storageType {
case service.SoraStorageTypeS3:
if h.s3Storage != nil && len(s3Keys) > 0 {
if err := h.s3Storage.DeleteObjects(ctx, s3Keys); err != nil {
logger.LegacyPrintf("handler.sora_client", "[SoraClient] 清理 S3 文件失败 keys=%v err=%v", s3Keys, err)
if service.IsObjectStorageType(storageType) {
if h.objectStorage != nil && len(s3Keys) > 0 {
if err := h.objectStorage.DeleteObjects(ctx, s3Keys); err != nil {
logger.LegacyPrintf("handler.sora_client", "[SoraClient] 清理存储文件失败 type=%s keys=%v err=%v", storageType, s3Keys, err)
}
}
case service.SoraStorageTypeLocal:
} else if storageType == service.SoraStorageTypeLocal {
if h.mediaStorage != nil && len(localPaths) > 0 {
if err := h.mediaStorage.DeleteByRelativePaths(localPaths); err != nil {
logger.LegacyPrintf("handler.sora_client", "[SoraClient] 清理本地文件失败 paths=%v err=%v", localPaths, err)
@@ -124,6 +124,9 @@ func (r *stubSoraGenRepo) CountByUserAndStatus(_ context.Context, _ int64, _ []s
}
return r.countValue, nil
}
func (r *stubSoraGenRepo) CountByStorageType(_ context.Context, _ string, _ []string) (int64, error) {
return 0, nil
}
// ==================== 辅助函数 ====================
@@ -1641,7 +1644,7 @@ func TestStoreMediaWithDegradation_S3SuccessSingleURL(t *testing.T) {
defer fakeS3.Close()
s3Storage := newS3StorageForHandler(fakeS3.URL)
h := &SoraClientHandler{s3Storage: s3Storage}
h := &SoraClientHandler{objectStorage: s3Storage}
storedURL, storedURLs, storageType, s3Keys, fileSize := h.storeMediaWithDegradation(
context.Background(), 1, "video", sourceServer.URL+"/v.mp4", nil,
@@ -1663,7 +1666,7 @@ func TestStoreMediaWithDegradation_S3SuccessMultiURL(t *testing.T) {
defer fakeS3.Close()
s3Storage := newS3StorageForHandler(fakeS3.URL)
h := &SoraClientHandler{s3Storage: s3Storage}
h := &SoraClientHandler{objectStorage: s3Storage}
urls := []string{sourceServer.URL + "/a.mp4", sourceServer.URL + "/b.mp4"}
storedURL, storedURLs, storageType, s3Keys, fileSize := h.storeMediaWithDegradation(
@@ -1688,7 +1691,7 @@ func TestStoreMediaWithDegradation_S3DownloadFails(t *testing.T) {
defer badSource.Close()
s3Storage := newS3StorageForHandler(fakeS3.URL)
h := &SoraClientHandler{s3Storage: s3Storage}
h := &SoraClientHandler{objectStorage: s3Storage}
_, _, storageType, _, _ := h.storeMediaWithDegradation(
context.Background(), 1, "video", badSource.URL+"/missing.mp4", nil,
@@ -1703,7 +1706,7 @@ func TestStoreMediaWithDegradation_S3FailsSingleURL(t *testing.T) {
defer fakeS3.Close()
s3Storage := newS3StorageForHandler(fakeS3.URL)
h := &SoraClientHandler{s3Storage: s3Storage}
h := &SoraClientHandler{objectStorage: s3Storage}
_, _, storageType, s3Keys, _ := h.storeMediaWithDegradation(
context.Background(), 1, "video", sourceServer.URL+"/v.mp4", nil,
@@ -1720,7 +1723,7 @@ func TestStoreMediaWithDegradation_S3PartialFailureCleanup(t *testing.T) {
defer fakeS3.Close()
s3Storage := newS3StorageForHandler(fakeS3.URL)
h := &SoraClientHandler{s3Storage: s3Storage}
h := &SoraClientHandler{objectStorage: s3Storage}
urls := []string{sourceServer.URL + "/a.mp4", sourceServer.URL + "/b.mp4"}
_, _, storageType, s3Keys, _ := h.storeMediaWithDegradation(
@@ -1804,8 +1807,8 @@ func TestStoreMediaWithDegradation_S3FailsFallbackToLocal(t *testing.T) {
}
mediaStorage := service.NewSoraMediaStorage(cfg)
h := &SoraClientHandler{
s3Storage: s3Storage,
mediaStorage: mediaStorage,
objectStorage: s3Storage,
mediaStorage: mediaStorage,
}
_, _, storageType, _, _ := h.storeMediaWithDegradation(
@@ -1831,14 +1834,14 @@ func TestSaveToStorage_S3EnabledButUploadFails(t *testing.T) {
}
s3Storage := newS3StorageForHandler(fakeS3.URL)
genService := service.NewSoraGenerationService(repo, nil, nil)
h := &SoraClientHandler{genService: genService, s3Storage: s3Storage}
h := &SoraClientHandler{genService: genService, objectStorage: s3Storage}
c, rec := makeGinContext("POST", "/api/v1/sora/generations/1/save", "", 1)
c.Params = gin.Params{{Key: "id", Value: "1"}}
h.SaveToStorage(c)
require.Equal(t, http.StatusInternalServerError, rec.Code)
resp := parseResponse(t, rec)
require.Contains(t, resp["message"], "S3")
require.Contains(t, resp["message"], "上传到存储失败")
}
func TestSaveToStorage_UpstreamURLExpired(t *testing.T) {
@@ -1857,7 +1860,7 @@ func TestSaveToStorage_UpstreamURLExpired(t *testing.T) {
}
s3Storage := newS3StorageForHandler(fakeS3.URL)
genService := service.NewSoraGenerationService(repo, nil, nil)
h := &SoraClientHandler{genService: genService, s3Storage: s3Storage}
h := &SoraClientHandler{genService: genService, objectStorage: s3Storage}
c, rec := makeGinContext("POST", "/api/v1/sora/generations/1/save", "", 1)
c.Params = gin.Params{{Key: "id", Value: "1"}}
@@ -1881,7 +1884,7 @@ func TestSaveToStorage_S3EnabledUploadSuccess(t *testing.T) {
}
s3Storage := newS3StorageForHandler(fakeS3.URL)
genService := service.NewSoraGenerationService(repo, nil, nil)
h := &SoraClientHandler{genService: genService, s3Storage: s3Storage}
h := &SoraClientHandler{genService: genService, objectStorage: s3Storage}
c, rec := makeGinContext("POST", "/api/v1/sora/generations/1/save", "", 1)
c.Params = gin.Params{{Key: "id", Value: "1"}}
@@ -1889,7 +1892,7 @@ func TestSaveToStorage_S3EnabledUploadSuccess(t *testing.T) {
require.Equal(t, http.StatusOK, rec.Code)
resp := parseResponse(t, rec)
data := resp["data"].(map[string]any)
require.Contains(t, data["message"], "S3")
require.Contains(t, data["message"], "已保存到云存储")
require.NotEmpty(t, data["object_key"])
// 验证记录已更新为 S3 存储
require.Equal(t, service.SoraStorageTypeS3, repo.gens[1].StorageType)
@@ -1913,7 +1916,7 @@ func TestSaveToStorage_S3EnabledUploadSuccess_MultiMediaURLs(t *testing.T) {
}
s3Storage := newS3StorageForHandler(fakeS3.URL)
genService := service.NewSoraGenerationService(repo, nil, nil)
h := &SoraClientHandler{genService: genService, s3Storage: s3Storage}
h := &SoraClientHandler{genService: genService, objectStorage: s3Storage}
c, rec := makeGinContext("POST", "/api/v1/sora/generations/1/save", "", 1)
c.Params = gin.Params{{Key: "id", Value: "1"}}
@@ -1949,7 +1952,7 @@ func TestSaveToStorage_S3EnabledUploadSuccessWithQuota(t *testing.T) {
SoraStorageUsedBytes: 0,
}
quotaService := service.NewSoraQuotaService(userRepo, nil, nil)
h := &SoraClientHandler{genService: genService, s3Storage: s3Storage, quotaService: quotaService}
h := &SoraClientHandler{genService: genService, objectStorage: s3Storage, quotaService: quotaService}
c, rec := makeGinContext("POST", "/api/v1/sora/generations/1/save", "", 1)
c.Params = gin.Params{{Key: "id", Value: "1"}}
@@ -1975,7 +1978,7 @@ func TestSaveToStorage_S3UploadSuccessMarkCompletedFails(t *testing.T) {
repo.updateErr = fmt.Errorf("db error")
s3Storage := newS3StorageForHandler(fakeS3.URL)
genService := service.NewSoraGenerationService(repo, nil, nil)
h := &SoraClientHandler{genService: genService, s3Storage: s3Storage}
h := &SoraClientHandler{genService: genService, objectStorage: s3Storage}
c, rec := makeGinContext("POST", "/api/v1/sora/generations/1/save", "", 1)
c.Params = gin.Params{{Key: "id", Value: "1"}}
@@ -1991,7 +1994,7 @@ func TestGetStorageStatus_S3EnabledNotHealthy(t *testing.T) {
defer fakeS3.Close()
s3Storage := newS3StorageForHandler(fakeS3.URL)
h := &SoraClientHandler{s3Storage: s3Storage}
h := &SoraClientHandler{objectStorage: s3Storage}
c, rec := makeGinContext("GET", "/api/v1/sora/storage-status", "", 0)
h.GetStorageStatus(c)
@@ -2007,7 +2010,7 @@ func TestGetStorageStatus_S3EnabledHealthy(t *testing.T) {
defer fakeS3.Close()
s3Storage := newS3StorageForHandler(fakeS3.URL)
h := &SoraClientHandler{s3Storage: s3Storage}
h := &SoraClientHandler{objectStorage: s3Storage}
c, rec := makeGinContext("GET", "/api/v1/sora/storage-status", "", 0)
h.GetStorageStatus(c)
@@ -2447,7 +2450,7 @@ func TestProcessGeneration_FullSuccessWithS3(t *testing.T) {
genService: genService,
gatewayService: gatewayService,
soraGatewayService: soraGatewayService,
s3Storage: s3Storage,
objectStorage: s3Storage,
quotaService: quotaService,
}
@@ -2497,7 +2500,7 @@ func TestProcessGeneration_MarkCompletedFails(t *testing.T) {
// ==================== cleanupStoredMedia 直接测试 ====================
func TestCleanupStoredMedia_S3Path(t *testing.T) {
// S3 清理路径:s3Storage 为 nil 时不 panic
// S3 清理路径:objectStorage 为 nil 时不 panic
h := &SoraClientHandler{}
// 不应 panic
h.cleanupStoredMedia(context.Background(), service.SoraStorageTypeS3, []string{"key1"}, nil)
@@ -2955,7 +2958,7 @@ func TestSaveToStorage_QuotaExceeded(t *testing.T) {
SoraStorageUsedBytes: 10,
}
quotaService := service.NewSoraQuotaService(userRepo, nil, nil)
h := &SoraClientHandler{genService: genService, s3Storage: s3Storage, quotaService: quotaService}
h := &SoraClientHandler{genService: genService, objectStorage: s3Storage, quotaService: quotaService}
c, rec := makeGinContext("POST", "/api/v1/sora/generations/1/save", "", 1)
c.Params = gin.Params{{Key: "id", Value: "1"}}
@@ -2983,7 +2986,7 @@ func TestSaveToStorage_QuotaNonQuotaError(t *testing.T) {
// 用户不存在 → GetByID 失败 → AddUsage 返回普通 error
userRepo := newStubUserRepoForHandler()
quotaService := service.NewSoraQuotaService(userRepo, nil, nil)
h := &SoraClientHandler{genService: genService, s3Storage: s3Storage, quotaService: quotaService}
h := &SoraClientHandler{genService: genService, objectStorage: s3Storage, quotaService: quotaService}
c, rec := makeGinContext("POST", "/api/v1/sora/generations/1/save", "", 1)
c.Params = gin.Params{{Key: "id", Value: "1"}}
@@ -3006,7 +3009,7 @@ func TestSaveToStorage_EmptyMediaURLs(t *testing.T) {
}
s3Storage := newS3StorageForHandler(fakeS3.URL)
genService := service.NewSoraGenerationService(repo, nil, nil)
h := &SoraClientHandler{genService: genService, s3Storage: s3Storage}
h := &SoraClientHandler{genService: genService, objectStorage: s3Storage}
c, rec := makeGinContext("POST", "/api/v1/sora/generations/1/save", "", 1)
c.Params = gin.Params{{Key: "id", Value: "1"}}
@@ -3033,7 +3036,7 @@ func TestSaveToStorage_MultiURL_SecondUploadFails(t *testing.T) {
}
s3Storage := newS3StorageForHandler(fakeS3.URL)
genService := service.NewSoraGenerationService(repo, nil, nil)
h := &SoraClientHandler{genService: genService, s3Storage: s3Storage}
h := &SoraClientHandler{genService: genService, objectStorage: s3Storage}
c, rec := makeGinContext("POST", "/api/v1/sora/generations/1/save", "", 1)
c.Params = gin.Params{{Key: "id", Value: "1"}}
@@ -3066,7 +3069,7 @@ func TestSaveToStorage_MarkCompletedFailsWithQuotaRollback(t *testing.T) {
SoraStorageUsedBytes: 0,
}
quotaService := service.NewSoraQuotaService(userRepo, nil, nil)
h := &SoraClientHandler{genService: genService, s3Storage: s3Storage, quotaService: quotaService}
h := &SoraClientHandler{genService: genService, objectStorage: s3Storage, quotaService: quotaService}
c, rec := makeGinContext("POST", "/api/v1/sora/generations/1/save", "", 1)
c.Params = gin.Params{{Key: "id", Value: "1"}}
@@ -3080,7 +3083,7 @@ func TestCleanupStoredMedia_WithS3Storage_ActualDelete(t *testing.T) {
fakeS3 := newFakeS3Server("ok")
defer fakeS3.Close()
s3Storage := newS3StorageForHandler(fakeS3.URL)
h := &SoraClientHandler{s3Storage: s3Storage}
h := &SoraClientHandler{objectStorage: s3Storage}
h.cleanupStoredMedia(context.Background(), service.SoraStorageTypeS3, []string{"key1", "key2"}, nil)
}
@@ -3089,7 +3092,7 @@ func TestCleanupStoredMedia_S3DeleteFails_LogOnly(t *testing.T) {
fakeS3 := newFakeS3Server("fail")
defer fakeS3.Close()
s3Storage := newS3StorageForHandler(fakeS3.URL)
h := &SoraClientHandler{s3Storage: s3Storage}
h := &SoraClientHandler{objectStorage: s3Storage}
h.cleanupStoredMedia(context.Background(), service.SoraStorageTypeS3, []string{"key1"}, nil)
}
@@ -225,7 +225,7 @@ func (h *SoraGatewayHandler) ChatCompletions(c *gin.Context) {
var lastFailoverHeaders http.Header
for {
selection, err := h.gatewayService.SelectAccountWithLoadAwareness(c.Request.Context(), apiKey.GroupID, sessionHash, reqModel, failedAccountIDs, "")
selection, err := h.gatewayService.SelectAccountWithLoadAwareness(c.Request.Context(), apiKey.GroupID, sessionHash, reqModel, failedAccountIDs, "", int64(0))
if err != nil {
reqLog.Warn("sora.account_select_failed",
zap.Error(err),
@@ -0,0 +1,391 @@
package handler
import (
"encoding/json"
"errors"
"fmt"
"net/http"
"unicode/utf8"
"github.com/Wei-Shaw/sub2api/internal/pkg/logger"
middleware2 "github.com/Wei-Shaw/sub2api/internal/server/middleware"
"github.com/Wei-Shaw/sub2api/internal/service"
"github.com/gin-gonic/gin"
)
// SoraVideosHandler handles Sora video/image async task API.
type SoraVideosHandler struct {
taskService *service.SoraTaskService
gatewayService *service.GatewayService
objectStorage service.SoraObjectStorage
mediaStorage *service.SoraMediaStorage
soraGatewayService *service.SoraGatewayService
}
func NewSoraVideosHandler(
taskService *service.SoraTaskService,
gatewayService *service.GatewayService,
objectStorage service.SoraObjectStorage,
mediaStorage *service.SoraMediaStorage,
soraGatewayService *service.SoraGatewayService,
) *SoraVideosHandler {
if taskService == nil {
return nil
}
return &SoraVideosHandler{
taskService: taskService,
gatewayService: gatewayService,
objectStorage: objectStorage,
mediaStorage: mediaStorage,
soraGatewayService: soraGatewayService,
}
}
func (h *SoraVideosHandler) CreateVideo(c *gin.Context) {
apiKey, account, release, ok := h.selectAccount(c, "")
if !ok {
return
}
defer release()
body, err := readBody(c)
if err != nil {
return
}
var req service.CreateVideoRequest
if err := json.Unmarshal(body, &req); err != nil {
soraErrorResponse(c, http.StatusBadRequest, "invalid_request_error", "Invalid request body")
return
}
if req.Model == "" {
soraErrorResponse(c, http.StatusBadRequest, "parameter_missing", "model is required")
return
}
task, err := h.taskService.CreateVideoTask(c.Request.Context(), apiKey.ID, account, &req, body)
if err != nil {
handleTaskCreateError(c, "CreateVideo", err)
return
}
c.JSON(http.StatusOK, service.TaskToResponse(task))
}
func (h *SoraVideosHandler) GetVideo(c *gin.Context) {
apiKey, ok := h.getAPIKey(c)
if !ok {
return
}
taskID := c.Param("id")
if taskID == "" {
soraErrorResponse(c, http.StatusBadRequest, "parameter_missing", "video_id is required")
return
}
task, err := h.taskService.GetTask(c.Request.Context(), taskID, apiKey.ID)
if err != nil {
soraErrorResponse(c, http.StatusNotFound, "not_found", "Task not found")
return
}
c.JSON(http.StatusOK, service.TaskToResponse(task))
}
func (h *SoraVideosHandler) RemixVideo(c *gin.Context) {
apiKey, ok := h.getAPIKey(c)
if !ok {
return
}
taskID := c.Param("id")
if taskID == "" {
soraErrorResponse(c, http.StatusBadRequest, "parameter_missing", "video_id is required")
return
}
body, err := readBody(c)
if err != nil {
return
}
var req service.RemixRequest
if err := json.Unmarshal(body, &req); err != nil || req.Prompt == "" {
soraErrorResponse(c, http.StatusBadRequest, "parameter_missing", "prompt is required")
return
}
originalTask, err := h.taskService.GetTask(c.Request.Context(), taskID, apiKey.ID)
if err != nil {
soraErrorResponse(c, http.StatusNotFound, "not_found", "Original video task not found")
return
}
if originalTask.Status != service.SoraTaskCompleted {
soraErrorResponse(c, http.StatusBadRequest, "invalid_request_error", "Original video must be completed before remix")
return
}
remixTargetID := originalTask.ShareID
if remixTargetID == "" {
remixTargetID = originalTask.UpstreamTaskID
}
account, err := h.selectAccountByID(c, originalTask.AccountID)
if err != nil {
soraErrorResponse(c, http.StatusServiceUnavailable, "server_error", "Failed to get account")
return
}
videoReq := &service.CreateVideoRequest{
Model: originalTask.Model,
Prompt: req.Prompt,
RemixTargetID: remixTargetID,
}
reqBody, _ := json.Marshal(videoReq)
task, err := h.taskService.CreateVideoTask(c.Request.Context(), apiKey.ID, account, videoReq, reqBody)
if err != nil {
handleTaskCreateError(c, "RemixVideo", err)
return
}
c.JSON(http.StatusOK, service.TaskToResponse(task))
}
// GetVideoContent returns video content based on storage configuration.
func (h *SoraVideosHandler) GetVideoContent(c *gin.Context) {
apiKey, ok := h.getAPIKey(c)
if !ok {
return
}
taskID := c.Param("id")
if taskID == "" {
soraErrorResponse(c, http.StatusBadRequest, "parameter_missing", "video_id is required")
return
}
task, err := h.taskService.GetTask(c.Request.Context(), taskID, apiKey.ID)
if err != nil {
soraErrorResponse(c, http.StatusNotFound, "not_found", "Task not found")
return
}
switch task.Status {
case service.SoraTaskCompleted:
contentURL := h.resolveContentURL(c, task)
if contentURL == "" {
soraErrorResponse(c, http.StatusNotFound, "not_found", "Video URL not available")
return
}
c.Redirect(http.StatusFound, contentURL)
case service.SoraTaskFailed:
c.JSON(http.StatusGone, gin.H{
"id": task.ID,
"object": task.ObjectType,
"status": task.Status,
"error": gin.H{
"message": task.ErrorMessage,
"type": task.ErrorType,
},
})
default:
c.JSON(http.StatusAccepted, service.TaskToResponse(task))
}
}
func (h *SoraVideosHandler) resolveContentURL(c *gin.Context, task *service.SoraTask) string {
if task.StoredKey == "" {
return task.VideoURL
}
switch task.StorageType {
case "s3", "gdrive":
if h.objectStorage != nil {
accessURL, err := h.objectStorage.GetAccessURL(c.Request.Context(), task.StoredKey)
if err != nil {
logger.LegacyPrintf("handler.sora_videos",
"[GetVideoContent] task=%s get access URL error: %v, fallback to upstream", task.ID, err)
return task.VideoURL
}
return accessURL
}
return task.VideoURL
case "local":
return "/sora/media" + task.StoredKey
default:
return task.VideoURL
}
}
func (h *SoraVideosHandler) CreateImage(c *gin.Context) {
apiKey, account, release, ok := h.selectAccount(c, "")
if !ok {
return
}
defer release()
body, err := readBody(c)
if err != nil {
return
}
var req service.CreateImageRequest
if err := json.Unmarshal(body, &req); err != nil {
soraErrorResponse(c, http.StatusBadRequest, "invalid_request_error", "Invalid request body")
return
}
if req.Prompt == "" {
soraErrorResponse(c, http.StatusBadRequest, "parameter_missing", "prompt is required")
return
}
if req.Model == "" {
req.Model = inferImageModel(req.Size)
}
task, err := h.taskService.CreateImageGeneration(c.Request.Context(), apiKey.ID, account, &req, body)
if err != nil {
handleTaskCreateError(c, "CreateImage", err)
return
}
c.JSON(http.StatusOK, service.TaskToResponse(task))
}
func (h *SoraVideosHandler) EditImage(c *gin.Context) {
apiKey, account, release, ok := h.selectAccount(c, "")
if !ok {
return
}
defer release()
body, err := readBody(c)
if err != nil {
return
}
var req service.EditImageRequest
if err := json.Unmarshal(body, &req); err != nil {
soraErrorResponse(c, http.StatusBadRequest, "invalid_request_error", "Invalid request body")
return
}
if req.Image == "" {
soraErrorResponse(c, http.StatusBadRequest, "parameter_missing", "image is required")
return
}
if req.Prompt == "" {
soraErrorResponse(c, http.StatusBadRequest, "parameter_missing", "prompt is required")
return
}
if req.Model == "" {
req.Model = inferImageModel(req.Size)
}
imageReq := &service.CreateImageRequest{
Model: req.Model,
Prompt: req.Prompt,
Image: req.Image,
Size: req.Size,
ResponseFormat: req.ResponseFormat,
N: 1,
}
task, err := h.taskService.CreateImageGeneration(c.Request.Context(), apiKey.ID, account, imageReq, body)
if err != nil {
handleTaskCreateError(c, "EditImage", err)
return
}
c.JSON(http.StatusOK, service.TaskToResponse(task))
}
// ── Internal helpers ──
func (h *SoraVideosHandler) getAPIKey(c *gin.Context) (*service.APIKey, bool) {
apiKey, ok := middleware2.GetAPIKeyFromContext(c)
if !ok {
soraErrorResponse(c, http.StatusUnauthorized, "authentication_error", "Invalid API key")
return nil, false
}
return apiKey, true
}
func (h *SoraVideosHandler) selectAccount(c *gin.Context, model string) (*service.APIKey, *service.Account, func(), bool) {
apiKey, ok := h.getAPIKey(c)
if !ok {
return nil, nil, nil, false
}
selection, err := h.gatewayService.SelectAccountWithLoadAwareness(
c.Request.Context(), apiKey.GroupID, "", model, nil, "", int64(0),
)
if err != nil {
soraErrorResponse(c, http.StatusServiceUnavailable, "server_error", "No available accounts")
return nil, nil, nil, false
}
releaseFunc := func() {}
if selection.ReleaseFunc != nil {
releaseFunc = selection.ReleaseFunc
}
return apiKey, selection.Account, releaseFunc, true
}
func (h *SoraVideosHandler) selectAccountByID(c *gin.Context, accountID int64) (*service.Account, error) {
return h.taskService.GetAccountByID(c.Request.Context(), accountID)
}
func readBody(c *gin.Context) ([]byte, error) {
body, err := c.GetRawData()
if err != nil {
soraErrorResponse(c, http.StatusBadRequest, "invalid_request_error", "Failed to read request body")
return nil, err
}
if len(body) == 0 {
soraErrorResponse(c, http.StatusBadRequest, "invalid_request_error", "Request body is empty")
return nil, fmt.Errorf("empty body")
}
if !utf8.Valid(body) {
soraErrorResponse(c, http.StatusBadRequest, "invalid_request_error", "Request body must be valid UTF-8")
return nil, fmt.Errorf("invalid utf-8")
}
return body, nil
}
func soraErrorResponse(c *gin.Context, status int, errType, message string) {
c.JSON(status, gin.H{
"error": gin.H{
"message": message,
"type": errType,
},
})
}
// handleTaskCreateError writes the appropriate error response, transparently
// forwarding upstream HTTP status codes when available.
func handleTaskCreateError(c *gin.Context, logTag string, err error) {
logger.LegacyPrintf("handler.sora_videos", "[%s] error: %v", logTag, err)
var ue *service.SoraUpstreamError
if errors.As(err, &ue) {
c.Data(ue.StatusCode, "application/json", ue.Body)
return
}
soraErrorResponse(c, http.StatusInternalServerError, "server_error", "Failed to create task")
}
func inferImageModel(size string) string {
switch size {
case "540x360":
return "gpt-image-landscape"
case "360x540":
return "gpt-image-portrait"
default:
return "gpt-image"
}
}
+5
View File
@@ -20,6 +20,7 @@ func ProvideAdminHandlers(
openaiOAuthHandler *admin.OpenAIOAuthHandler,
geminiOAuthHandler *admin.GeminiOAuthHandler,
antigravityOAuthHandler *admin.AntigravityOAuthHandler,
gdriveOAuthHandler *admin.GDriveOAuthHandler,
proxyHandler *admin.ProxyHandler,
redeemHandler *admin.RedeemHandler,
promoHandler *admin.PromoHandler,
@@ -57,6 +58,7 @@ func ProvideAdminHandlers(
ErrorPassthrough: errorPassthroughHandler,
APIKey: apiKeyHandler,
ScheduledTest: scheduledTestHandler,
GDriveOAuth: gdriveOAuthHandler,
}
}
@@ -84,6 +86,7 @@ func ProvideHandlers(
openaiGatewayHandler *OpenAIGatewayHandler,
soraGatewayHandler *SoraGatewayHandler,
soraClientHandler *SoraClientHandler,
soraVideosHandler *SoraVideosHandler,
settingHandler *SettingHandler,
totpHandler *TotpHandler,
_ *service.IdempotencyCoordinator,
@@ -102,6 +105,7 @@ func ProvideHandlers(
OpenAIGateway: openaiGatewayHandler,
SoraGateway: soraGatewayHandler,
SoraClient: soraClientHandler,
SoraVideos: soraVideosHandler,
Setting: settingHandler,
Totp: totpHandler,
}
@@ -147,6 +151,7 @@ var ProviderSet = wire.NewSet(
admin.NewErrorPassthroughHandler,
admin.NewAdminAPIKeyHandler,
admin.NewScheduledTestHandler,
admin.NewGDriveOAuthHandler,
// AdminHandlers and Handlers constructors
ProvideAdminHandlers,
@@ -18,6 +18,9 @@ const (
BlockTypeFunction
)
// UsageMapHook is a callback that can modify usage data before it's emitted in SSE events.
type UsageMapHook func(usageMap map[string]any)
// StreamingProcessor 流式响应处理器
type StreamingProcessor struct {
blockType BlockType
@@ -30,6 +33,7 @@ type StreamingProcessor struct {
originalModel string
webSearchQueries []string
groundingChunks []GeminiGroundingChunk
usageMapHook UsageMapHook
// 累计 usage
inputTokens int
@@ -45,6 +49,25 @@ func NewStreamingProcessor(originalModel string) *StreamingProcessor {
}
}
// SetUsageMapHook sets an optional hook that modifies usage maps before they are emitted.
func (p *StreamingProcessor) SetUsageMapHook(fn UsageMapHook) {
p.usageMapHook = fn
}
func usageToMap(u ClaudeUsage) map[string]any {
m := map[string]any{
"input_tokens": u.InputTokens,
"output_tokens": u.OutputTokens,
}
if u.CacheCreationInputTokens > 0 {
m["cache_creation_input_tokens"] = u.CacheCreationInputTokens
}
if u.CacheReadInputTokens > 0 {
m["cache_read_input_tokens"] = u.CacheReadInputTokens
}
return m
}
// ProcessLine 处理 SSE 行,返回 Claude SSE 事件
func (p *StreamingProcessor) ProcessLine(line string) []byte {
line = strings.TrimSpace(line)
@@ -168,6 +191,13 @@ func (p *StreamingProcessor) emitMessageStart(v1Resp *V1InternalResponse) []byte
responseID = "msg_" + generateRandomID()
}
var usageValue any = usage
if p.usageMapHook != nil {
usageMap := usageToMap(usage)
p.usageMapHook(usageMap)
usageValue = usageMap
}
message := map[string]any{
"id": responseID,
"type": "message",
@@ -176,7 +206,7 @@ func (p *StreamingProcessor) emitMessageStart(v1Resp *V1InternalResponse) []byte
"model": p.originalModel,
"stop_reason": nil,
"stop_sequence": nil,
"usage": usage,
"usage": usageValue,
}
event := map[string]any{
@@ -487,13 +517,20 @@ func (p *StreamingProcessor) emitFinish(finishReason string) []byte {
CacheReadInputTokens: p.cacheReadTokens,
}
var usageValue any = usage
if p.usageMapHook != nil {
usageMap := usageToMap(usage)
p.usageMapHook(usageMap)
usageValue = usageMap
}
deltaEvent := map[string]any{
"type": "message_delta",
"delta": map[string]any{
"stop_reason": stopReason,
"stop_sequence": nil,
},
"usage": usage,
"usage": usageValue,
}
_, _ = result.Write(p.formatSSE("message_delta", deltaEvent))
+13 -3
View File
@@ -10,7 +10,13 @@ import (
)
func TestInit_DualOutput(t *testing.T) {
tmpDir := t.TempDir()
// Use os.MkdirTemp instead of t.TempDir to avoid cleanup failures
// when lumberjack holds file handles on Windows.
tmpDir, err := os.MkdirTemp("", "logger-test-*")
if err != nil {
t.Fatalf("create temp dir: %v", err)
}
t.Cleanup(func() { _ = os.RemoveAll(tmpDir) })
logPath := filepath.Join(tmpDir, "logs", "sub2api.log")
origStdout := os.Stdout
@@ -57,7 +63,9 @@ func TestInit_DualOutput(t *testing.T) {
L().Info("dual-output-info")
L().Warn("dual-output-warn")
Sync()
// Skip Sync() — on Windows, fsync on pipes deadlocks (FlushFileBuffers).
// The log data is already in the pipe buffer; closing writers is sufficient.
_ = stdoutW.Close()
_ = stderrW.Close()
@@ -166,7 +174,9 @@ func TestInit_CallerShouldPointToCallsite(t *testing.T) {
}
L().Info("caller-check")
Sync()
// Skip Sync() — on Windows, fsync on pipes deadlocks (FlushFileBuffers).
os.Stdout = origStdout
os.Stderr = origStderr
_ = stdoutW.Close()
logBytes, _ := io.ReadAll(stdoutR)
@@ -77,7 +77,7 @@ func TestStdLogBridgeRoutesLevels(t *testing.T) {
log.Printf("service started")
log.Printf("Warning: queue full")
log.Printf("Forward request failed: timeout")
Sync()
// Skip Sync() — on Windows, fsync on pipes deadlocks (FlushFileBuffers).
_ = stdoutW.Close()
_ = stderrW.Close()
@@ -139,7 +139,7 @@ func TestLegacyPrintfRoutesLevels(t *testing.T) {
LegacyPrintf("service.test", "request started")
LegacyPrintf("service.test", "Warning: queue full")
LegacyPrintf("service.test", "forward failed: timeout")
Sync()
// Skip Sync() — on Windows, fsync on pipes deadlocks (FlushFileBuffers).
_ = stdoutW.Close()
_ = stderrW.Close()
@@ -164,6 +164,7 @@ func (r *apiKeyRepository) GetByKeyForAuth(ctx context.Context, key string) (*se
group.FieldModelRoutingEnabled,
group.FieldModelRouting,
group.FieldMcpXMLInject,
group.FieldSimulateClaudeMaxEnabled,
group.FieldSupportedModelScopes,
group.FieldAllowMessagesDispatch,
group.FieldDefaultMappedModel,
@@ -645,6 +646,7 @@ func groupEntityToService(g *dbent.Group) *service.Group {
ModelRouting: g.ModelRouting,
ModelRoutingEnabled: g.ModelRoutingEnabled,
MCPXMLInject: g.McpXMLInject,
SimulateClaudeMaxEnabled: g.SimulateClaudeMaxEnabled,
SupportedModelScopes: g.SupportedModelScopes,
SortOrder: g.SortOrder,
AllowMessagesDispatch: g.AllowMessagesDispatch,
+317 -7
View File
@@ -2,14 +2,46 @@ package repository
import (
"context"
_ "embed"
"fmt"
"strconv"
"strings"
"time"
"github.com/Wei-Shaw/sub2api/internal/service"
"github.com/redis/go-redis/v9"
)
const stickySessionPrefix = "sticky_session:"
const (
stickySessionPrefix = "sticky_session:"
affinityKeyPrefix = "affinity:"
affinityRevKeyPrefix = "affinity_rev:"
)
var (
//go:embed lua/get_affinity.lua
getAffinityLua string
//go:embed lua/update_affinity.lua
updateAffinityLua string
//go:embed lua/get_affinity_count.lua
getAffinityCountLua string
//go:embed lua/get_affinity_clients.lua
getAffinityClientsLua string
//go:embed lua/get_affinity_clients_with_scores.lua
getAffinityClientsWithScoresLua string
//go:embed lua/clear_account_affinity.lua
clearAccountAffinityLua string
//go:embed lua/get_affinity_multi_count.lua
getAffinityMultiCountLua string
getAffinityScript = redis.NewScript(getAffinityLua)
updateAffinityScript = redis.NewScript(updateAffinityLua)
getAffinityCountScript = redis.NewScript(getAffinityCountLua)
getAffinityClientsScript = redis.NewScript(getAffinityClientsLua)
getAffinityClientsWithScoresScript = redis.NewScript(getAffinityClientsWithScoresLua)
clearAccountAffinityScript = redis.NewScript(clearAccountAffinityLua)
getAffinityMultiCountScript = redis.NewScript(getAffinityMultiCountLua)
)
type gatewayCache struct {
rdb *redis.Client
@@ -19,6 +51,16 @@ func NewGatewayCache(rdb *redis.Client) service.GatewayCache {
return &gatewayCache{rdb: rdb}
}
// ensureScriptLoaded 确保 Lua 脚本已加载到 Redis 服务器的脚本缓存中。
// Pipeline 中的 Script.Run 只发送 EVALSHA,如果 Redis 重启过导致脚本缓存丢失,
// EVALSHA 会返回 NOSCRIPT 错误。此方法提前加载脚本以避免该问题。
func ensureScriptLoaded(ctx context.Context, rdb *redis.Client, script *redis.Script) {
exists, err := script.Exists(ctx, rdb).Result()
if err != nil || len(exists) == 0 || !exists[0] {
_ = script.Load(ctx, rdb).Err()
}
}
// buildSessionKey 构建 session key,包含 groupID 实现分组隔离
// 格式: sticky_session:{groupID}:{sessionHash}
func buildSessionKey(groupID int64, sessionHash string) string {
@@ -41,13 +83,281 @@ func (c *gatewayCache) RefreshSessionTTL(ctx context.Context, groupID int64, ses
}
// DeleteSessionAccountID 删除粘性会话与账号的绑定关系。
// 当检测到绑定的账号不可用(如状态错误、禁用、不可调度等)时调用,
// 以便下次请求能够重新选择可用账号。
//
// DeleteSessionAccountID removes the sticky session binding for the given session.
// Called when the bound account becomes unavailable (e.g., error status, disabled,
// or unschedulable), allowing subsequent requests to select a new available account.
func (c *gatewayCache) DeleteSessionAccountID(ctx context.Context, groupID int64, sessionHash string) error {
key := buildSessionKey(groupID, sessionHash)
return c.rdb.Del(ctx, key).Err()
}
// buildAffinityKey 构建正向亲和 key(member → accounts)
// 格式: affinity:{groupID}:{userID}/{clientID}
func buildAffinityKey(groupID int64, userID int64, clientID string) string {
return fmt.Sprintf("%s%d:%s", affinityKeyPrefix, groupID, buildAffinityMember(userID, clientID))
}
// buildAffinityReverseKey 构建反向亲和 key(account → members)
// 格式: affinity_rev:{groupID}:{accountID}
func buildAffinityReverseKey(groupID int64, accountID int64) string {
return fmt.Sprintf("%s%d:%d", affinityRevKeyPrefix, groupID, accountID)
}
// buildAffinityMember 构建亲和成员标识
// 格式: {userID}/{clientID}
func buildAffinityMember(userID int64, clientID string) string {
return fmt.Sprintf("%d/%s", userID, clientID)
}
// parseAffinityMember 解析亲和成员标识为 userID 和 clientID
func parseAffinityMember(member string) (userID int64, clientID string) {
idx := strings.IndexByte(member, '/')
if idx < 0 {
// 兼容旧格式(纯 clientID,无 userID 前缀)
return 0, member
}
userID, _ = strconv.ParseInt(member[:idx], 10, 64)
clientID = member[idx+1:]
return userID, clientID
}
// GetAffinityAccounts 获取亲和账号列表(按最近使用降序),同时清理过期成员
func (c *gatewayCache) GetAffinityAccounts(ctx context.Context, groupID int64, userID int64, clientID string, ttl time.Duration) ([]int64, error) {
key := buildAffinityKey(groupID, userID, clientID)
now := time.Now().Unix()
expireThreshold := now - int64(ttl.Seconds())
result, err := getAffinityScript.Run(ctx, c.rdb, []string{key}, expireThreshold).StringSlice()
if err != nil {
if err == redis.Nil {
return nil, nil
}
return nil, err
}
accountIDs := make([]int64, 0, len(result))
for _, s := range result {
id, err := strconv.ParseInt(s, 10, 64)
if err != nil {
continue
}
accountIDs = append(accountIDs, id)
}
return accountIDs, nil
}
// UpdateAffinity 添加/更新亲和关系(更新 score 为当前时间戳,刷新 key TTL)
func (c *gatewayCache) UpdateAffinity(ctx context.Context, groupID int64, userID int64, clientID string, accountID int64, ttl time.Duration) error {
fwdKey := buildAffinityKey(groupID, userID, clientID)
revKey := buildAffinityReverseKey(groupID, accountID)
now := time.Now().Unix()
ttlSeconds := int64(ttl.Seconds())
expireThreshold := now - ttlSeconds
member := buildAffinityMember(userID, clientID)
return updateAffinityScript.Run(ctx, c.rdb, []string{fwdKey, revKey},
now, ttlSeconds, accountID, expireThreshold, member,
).Err()
}
// GetAccountAffinityCountBatch 批量获取账号的亲和成员数量(惰性清理过期成员)
func (c *gatewayCache) GetAccountAffinityCountBatch(ctx context.Context, groupID int64, accountIDs []int64, ttl time.Duration) (map[int64]int64, error) {
if len(accountIDs) == 0 {
return map[int64]int64{}, nil
}
now := time.Now().Unix()
expireThreshold := now - int64(ttl.Seconds())
ensureScriptLoaded(ctx, c.rdb, getAffinityCountScript)
pipe := c.rdb.Pipeline()
cmds := make([]*redis.Cmd, len(accountIDs))
for i, accID := range accountIDs {
key := buildAffinityReverseKey(groupID, accID)
cmds[i] = getAffinityCountScript.Run(ctx, pipe, []string{key}, expireThreshold)
}
_, err := pipe.Exec(ctx)
if err != nil && err != redis.Nil {
return nil, err
}
result := make(map[int64]int64, len(accountIDs))
for i, accID := range accountIDs {
count, _ := cmds[i].Int64()
result[accID] = count
}
return result, nil
}
// GetAccountAffinityClientsBatch 批量获取每个账号跨所有分组的亲和成员列表(去重)。
// accountGroups: map[accountID][]groupID,对每个 (groupID, accountID) 组合查询反向索引。
// 返回值成员格式为 {userID}/{clientID}。
func (c *gatewayCache) GetAccountAffinityClientsBatch(ctx context.Context, accountGroups map[int64][]int64, ttl time.Duration) (map[int64][]string, error) {
if len(accountGroups) == 0 {
return map[int64][]string{}, nil
}
now := time.Now().Unix()
expireThreshold := now - int64(ttl.Seconds())
// 构建所有 (accountID, groupID) 组合的查询
type queryItem struct {
accountID int64
groupID int64
}
var queries []queryItem
for accID, groupIDs := range accountGroups {
for _, gID := range groupIDs {
queries = append(queries, queryItem{accountID: accID, groupID: gID})
}
}
ensureScriptLoaded(ctx, c.rdb, getAffinityClientsScript)
pipe := c.rdb.Pipeline()
cmds := make([]*redis.Cmd, len(queries))
for i, q := range queries {
key := buildAffinityReverseKey(q.groupID, q.accountID)
cmds[i] = getAffinityClientsScript.Run(ctx, pipe, []string{key}, expireThreshold)
}
_, err := pipe.Exec(ctx)
if err != nil && err != redis.Nil {
return nil, err
}
// 合并结果:同一个 accountID 跨多个 group 的成员去重
result := make(map[int64][]string, len(accountGroups))
seen := make(map[int64]map[string]struct{}, len(accountGroups))
for i, q := range queries {
members, _ := cmds[i].StringSlice()
if len(members) == 0 {
continue
}
if seen[q.accountID] == nil {
seen[q.accountID] = make(map[string]struct{})
}
for _, member := range members {
if _, exists := seen[q.accountID][member]; !exists {
seen[q.accountID][member] = struct{}{}
result[q.accountID] = append(result[q.accountID], member)
}
}
}
return result, nil
}
// GetAccountAffinityClientsWithScores 获取单个账号跨所有分组的亲和客户端列表(含最后活跃时间戳,去重取最近)。
func (c *gatewayCache) GetAccountAffinityClientsWithScores(
ctx context.Context,
accountID int64,
groupIDs []int64,
ttl time.Duration,
) ([]service.AffinityClient, error) {
if len(groupIDs) == 0 {
return nil, nil
}
now := time.Now().Unix()
expireThreshold := now - int64(ttl.Seconds())
ensureScriptLoaded(ctx, c.rdb, getAffinityClientsWithScoresScript)
pipe := c.rdb.Pipeline()
cmds := make([]*redis.Cmd, len(groupIDs))
for i, gID := range groupIDs {
key := buildAffinityReverseKey(gID, accountID)
cmds[i] = getAffinityClientsWithScoresScript.Run(ctx, pipe, []string{key}, expireThreshold)
}
_, err := pipe.Exec(ctx)
if err != nil && err != redis.Nil {
return nil, err
}
// 合并跨组结果,同一 member 取最近的 lastActive
type memberInfo struct {
userID int64
clientID string
ts int64
}
seen := make(map[string]*memberInfo) // member string → info
for _, cmd := range cmds {
vals, _ := cmd.StringSlice()
// vals 格式: [member1, score1, member2, score2, ...]
for j := 0; j+1 < len(vals); j += 2 {
member := vals[j]
ts, _ := strconv.ParseInt(vals[j+1], 10, 64)
if existing, ok := seen[member]; !ok || ts > existing.ts {
uid, cid := parseAffinityMember(member)
seen[member] = &memberInfo{userID: uid, clientID: cid, ts: ts}
}
}
}
result := make([]service.AffinityClient, 0, len(seen))
for _, info := range seen {
result = append(result, service.AffinityClient{
UserID: info.userID,
ClientID: info.clientID,
LastActive: time.Unix(info.ts, 0),
})
}
// 按最后活跃时间降序排序
service.SortAffinityClients(result)
return result, nil
}
// ClearAccountAffinity 清除指定账号在所有分组的亲和记录(正向+反向索引)。
// 对每个 groupID 执行 Lua 脚本:读取反向索引获取所有成员,
// 从每个成员的正向索引中移除该账号,然后删除反向索引。
func (c *gatewayCache) ClearAccountAffinity(ctx context.Context, accountID int64, groupIDs []int64) error {
if len(groupIDs) == 0 {
return nil
}
ensureScriptLoaded(ctx, c.rdb, clearAccountAffinityScript)
pipe := c.rdb.Pipeline()
for _, gID := range groupIDs {
revKey := buildAffinityReverseKey(gID, accountID)
clearAccountAffinityScript.Run(ctx, pipe, []string{revKey}, gID, accountID)
}
_, err := pipe.Exec(ctx)
if err != nil && err != redis.Nil {
return err
}
return nil
}
// GetAffinityMultiCount 获取账号的多维度亲和计数。
// 返回: uniqueUsers(独立用户数), uniqueClients(独立客户端数), perUserClients(目标用户的客户端数)
func (c *gatewayCache) GetAffinityMultiCount(
ctx context.Context,
groupID int64,
accountID int64,
targetUserID int64,
ttl time.Duration,
) (users, clients, perUser int64, err error) {
key := buildAffinityReverseKey(groupID, accountID)
now := time.Now().Unix()
expireThreshold := now - int64(ttl.Seconds())
targetUserStr := ""
if targetUserID > 0 {
targetUserStr = strconv.FormatInt(targetUserID, 10)
}
result, err := getAffinityMultiCountScript.Run(ctx, c.rdb, []string{key}, expireThreshold, targetUserStr).Int64Slice()
if err != nil {
if err == redis.Nil {
return 0, 0, 0, nil
}
return 0, 0, 0, err
}
if len(result) < 4 {
return 0, 0, 0, nil
}
// result: {totalMembers, uniqueUsers, uniqueClients, perUserClients}
return result[1], result[2], result[3], nil
}
+4 -2
View File
@@ -61,7 +61,8 @@ func (r *groupRepository) Create(ctx context.Context, groupIn *service.Group) er
SetMcpXMLInject(groupIn.MCPXMLInject).
SetSoraStorageQuotaBytes(groupIn.SoraStorageQuotaBytes).
SetAllowMessagesDispatch(groupIn.AllowMessagesDispatch).
SetDefaultMappedModel(groupIn.DefaultMappedModel)
SetDefaultMappedModel(groupIn.DefaultMappedModel).
SetSimulateClaudeMaxEnabled(groupIn.SimulateClaudeMaxEnabled)
// 设置模型路由配置
if groupIn.ModelRouting != nil {
@@ -130,7 +131,8 @@ func (r *groupRepository) Update(ctx context.Context, groupIn *service.Group) er
SetMcpXMLInject(groupIn.MCPXMLInject).
SetSoraStorageQuotaBytes(groupIn.SoraStorageQuotaBytes).
SetAllowMessagesDispatch(groupIn.AllowMessagesDispatch).
SetDefaultMappedModel(groupIn.DefaultMappedModel)
SetDefaultMappedModel(groupIn.DefaultMappedModel).
SetSimulateClaudeMaxEnabled(groupIn.SimulateClaudeMaxEnabled)
// 显式处理可空字段:nil 需要 clear,非 nil 需要 set。
if groupIn.DailyLimitUSD != nil {
@@ -0,0 +1,29 @@
-- 清除单个账号在指定分组的所有亲和记录(正向+反向)
-- KEYS[1] = affinity_rev:{groupID}:{accountID} (反向索引)
-- ARGV[1] = groupID (用于构建正向 key)
-- ARGV[2] = accountID (正向索引中要移除的成员)
-- 返回: 清理的成员数量
local rev_key = KEYS[1]
local group_id = ARGV[1]
local account_id = ARGV[2]
-- 获取反向索引中所有成员 ({userID}/{clientID})
local members = redis.call('ZRANGE', rev_key, 0, -1)
if #members == 0 then
return 0
end
-- 从每个成员的正向索引中移除该账号
for _, member in ipairs(members) do
local fwd_key = 'affinity:' .. group_id .. ':' .. member
redis.call('ZREM', fwd_key, account_id)
-- 如果正向索引为空,删除 key
if redis.call('ZCARD', fwd_key) == 0 then
redis.call('DEL', fwd_key)
end
end
-- 删除反向索引
redis.call('DEL', rev_key)
return #members
@@ -0,0 +1,5 @@
-- 清理过期成员后返回亲和账号列表(按最近使用降序)
-- KEYS[1] = affinity:{groupID}:{userID}/{clientID}
-- ARGV[1] = 过期阈值时间戳 (now - ttl)
redis.call('ZREMRANGEBYSCORE', KEYS[1], '-inf', ARGV[1])
return redis.call('ZREVRANGE', KEYS[1], 0, -1)
@@ -0,0 +1,6 @@
-- 清理过期成员后返回反向索引的成员列表(按最近使用降序)
-- 成员格式: {userID}/{clientID}
-- KEYS[1] = affinity_rev:{groupID}:{accountID}
-- ARGV[1] = 过期阈值时间戳 (now - ttl)
redis.call('ZREMRANGEBYSCORE', KEYS[1], '-inf', ARGV[1])
return redis.call('ZREVRANGE', KEYS[1], 0, -1)
@@ -0,0 +1,7 @@
-- 清理过期成员后返回反向索引的成员列表及其 score(最后活跃时间戳)
-- 成员格式: {userID}/{clientID}
-- KEYS[1] = affinity_rev:{groupID}:{accountID}
-- ARGV[1] = 过期阈值时间戳 (now - ttl)
-- 返回: {member1, score1, member2, score2, ...}(按最近使用降序)
redis.call('ZREMRANGEBYSCORE', KEYS[1], '-inf', ARGV[1])
return redis.call('ZREVRANGEBYSCORE', KEYS[1], '+inf', '-inf', 'WITHSCORES')
@@ -0,0 +1,5 @@
-- 清理过期成员后返回反向索引的成员数量
-- KEYS[1] = affinity_rev:{groupID}:{accountID}
-- ARGV[1] = 过期阈值时间戳 (now - ttl)
redis.call('ZREMRANGEBYSCORE', KEYS[1], '-inf', ARGV[1])
return redis.call('ZCARD', KEYS[1])
@@ -0,0 +1,40 @@
-- 从反向索引解析多维度计数(用户/客户端/每用户客户端)
-- KEYS[1] = affinity_rev:{groupID}:{accountID}
-- ARGV[1] = 过期阈值时间戳 (now - ttl)
-- ARGV[2] = 目标 userID(传 "" 则不计算 perUserClients)
-- 返回: {totalMembers, uniqueUsers, uniqueClients, perUserClients}
redis.call('ZREMRANGEBYSCORE', KEYS[1], '-inf', ARGV[1])
local members = redis.call('ZRANGE', KEYS[1], 0, -1)
local total = #members
if total == 0 then
return {0, 0, 0, 0}
end
local target_user = ARGV[2]
local users = {}
local clients = {}
local user_count = 0
local client_count = 0
local per_user_count = 0
for _, member in ipairs(members) do
local sep = string.find(member, '/', 1, true)
if sep then
local uid = string.sub(member, 1, sep - 1)
local cid = string.sub(member, sep + 1)
if not users[uid] then
users[uid] = true
user_count = user_count + 1
end
if not clients[cid] then
clients[cid] = true
client_count = client_count + 1
end
if target_user ~= '' and uid == target_user then
per_user_count = per_user_count + 1
end
end
end
return {total, user_count, client_count, per_user_count}
@@ -0,0 +1,15 @@
-- 原子双写正向+反向索引
-- KEYS[1] = affinity:{groupID}:{userID}/{clientID} (正向: member → accounts)
-- KEYS[2] = affinity_rev:{groupID}:{accountID} (反向: account → members)
-- ARGV[1] = 当前时间戳 (score)
-- ARGV[2] = TTL 秒数
-- ARGV[3] = accountID (正向索引的成员)
-- ARGV[4] = 过期阈值时间戳 (now - ttl)
-- ARGV[5] = {userID}/{clientID} (反向索引的成员)
redis.call('ZREMRANGEBYSCORE', KEYS[1], '-inf', ARGV[4])
redis.call('ZADD', KEYS[1], ARGV[1], ARGV[3])
redis.call('EXPIRE', KEYS[1], ARGV[2])
redis.call('ZREMRANGEBYSCORE', KEYS[2], '-inf', ARGV[4])
redis.call('ZADD', KEYS[2], ARGV[1], ARGV[5])
redis.call('EXPIRE', KEYS[2], ARGV[2])
return 1
@@ -417,3 +417,22 @@ func (r *soraGenerationRepository) CountByUserAndStatus(ctx context.Context, use
err := r.sql.QueryRowContext(ctx, query, args...).Scan(&count)
return count, err
}
// CountByStorageType 按存储类型和状态统计生成记录数。
func (r *soraGenerationRepository) CountByStorageType(ctx context.Context, storageType string, statuses []string) (int64, error) {
if len(statuses) == 0 {
return 0, nil
}
placeholders := make([]string, len(statuses))
args := []any{storageType}
for i, s := range statuses {
placeholders[i] = fmt.Sprintf("$%d", i+2)
args = append(args, s)
}
var count int64
query := fmt.Sprintf("SELECT COUNT(*) FROM sora_generations WHERE storage_type = $1 AND status IN (%s)", strings.Join(placeholders, ","))
err := r.sql.QueryRowContext(ctx, query, args...).Scan(&count)
return count, err
}
@@ -0,0 +1,136 @@
package repository
import (
"context"
"database/sql"
"encoding/json"
"unicode/utf8"
"github.com/Wei-Shaw/sub2api/internal/service"
)
type SoraTaskRepository struct {
db *sql.DB
}
func NewSoraTaskRepository(sqlDB *sql.DB) service.SoraTaskRepository {
return &SoraTaskRepository{db: sqlDB}
}
func (r *SoraTaskRepository) Create(ctx context.Context, task *service.SoraTask) error {
var charStr, reqStr *string
if task.CharacterInfo != nil {
b, _ := json.Marshal(task.CharacterInfo)
s := string(b)
charStr = &s
}
if len(task.RequestBody) > 0 && utf8.Valid(task.RequestBody) {
s := string(task.RequestBody)
reqStr = &s
}
_, err := r.db.ExecContext(ctx, `
INSERT INTO sora_tasks (
id, account_id, api_key_id, upstream_task_id, object_type,
model, prompt, status, progress, video_url, stored_key, storage_type,
share_id, character_info, error_message, error_type,
request_body, seconds, size, created_at, completed_at
) VALUES ($1,$2,$3,$4,$5,$6,$7,$8,$9,$10,$11,$12,$13,$14,$15,$16,$17,$18,$19,$20,$21)`,
task.ID, task.AccountID, task.APIKeyID, task.UpstreamTaskID, task.ObjectType,
task.Model, task.Prompt, task.Status, task.Progress, task.VideoURL,
task.StoredKey, task.StorageType,
task.ShareID, charStr, task.ErrorMessage, task.ErrorType,
reqStr, task.Seconds, task.Size, task.CreatedAt, task.CompletedAt,
)
return err
}
func (r *SoraTaskRepository) GetByID(ctx context.Context, id string) (*service.SoraTask, error) {
row := r.db.QueryRowContext(ctx, `SELECT `+soraTaskColumns+` FROM sora_tasks WHERE id = $1`, id)
return scanTask(row)
}
func (r *SoraTaskRepository) GetByIDAndAPIKey(ctx context.Context, id string, apiKeyID int64) (*service.SoraTask, error) {
row := r.db.QueryRowContext(ctx,
`SELECT `+soraTaskColumns+` FROM sora_tasks WHERE id = $1 AND api_key_id = $2`,
id, apiKeyID,
)
return scanTask(row)
}
func (r *SoraTaskRepository) Update(ctx context.Context, task *service.SoraTask) error {
var charStr *string
if task.CharacterInfo != nil {
b, _ := json.Marshal(task.CharacterInfo)
s := string(b)
charStr = &s
}
_, err := r.db.ExecContext(ctx, `
UPDATE sora_tasks SET
upstream_task_id = $2, status = $3, progress = $4,
video_url = $5, stored_key = $6, storage_type = $7,
share_id = $8, character_info = $9,
error_message = $10, error_type = $11,
completed_at = $12, seconds = $13, size = $14, object_type = $15
WHERE id = $1`,
task.ID, task.UpstreamTaskID, task.Status, task.Progress,
task.VideoURL, task.StoredKey, task.StorageType,
task.ShareID, charStr,
task.ErrorMessage, task.ErrorType,
task.CompletedAt, task.Seconds, task.Size, task.ObjectType,
)
return err
}
func (r *SoraTaskRepository) ListPending(ctx context.Context) ([]*service.SoraTask, error) {
rows, err := r.db.QueryContext(ctx,
`SELECT `+soraTaskColumns+` FROM sora_tasks WHERE status IN ('queued', 'in_progress') ORDER BY created_at ASC LIMIT 200`,
)
if err != nil {
return nil, err
}
defer func() { _ = rows.Close() }()
var tasks []*service.SoraTask
for rows.Next() {
t, err := scanTaskFromRow(rows)
if err != nil {
return nil, err
}
tasks = append(tasks, t)
}
return tasks, rows.Err()
}
const soraTaskColumns = `id, account_id, api_key_id, upstream_task_id, object_type,
model, prompt, status, progress, video_url, stored_key, storage_type,
share_id, character_info, error_message, error_type,
request_body, seconds, size, created_at, completed_at`
func scanTask(s scannable) (*service.SoraTask, error) {
var t service.SoraTask
var charJSON, reqBody []byte
err := s.Scan(
&t.ID, &t.AccountID, &t.APIKeyID, &t.UpstreamTaskID, &t.ObjectType,
&t.Model, &t.Prompt, &t.Status, &t.Progress, &t.VideoURL,
&t.StoredKey, &t.StorageType,
&t.ShareID, &charJSON, &t.ErrorMessage, &t.ErrorType,
&reqBody, &t.Seconds, &t.Size, &t.CreatedAt, &t.CompletedAt,
)
if err != nil {
return nil, err
}
if len(charJSON) > 0 {
var ch service.SoraCharacter
if json.Unmarshal(charJSON, &ch) == nil && ch.Username != "" {
t.CharacterInfo = &ch
}
}
t.RequestBody = reqBody
return &t, nil
}
func scanTaskFromRow(rows *sql.Rows) (*service.SoraTask, error) {
return scanTask(rows)
}
@@ -2956,7 +2956,7 @@ func (r *usageLogRepository) GetGroupStatsWithFilters(ctx context.Context, start
query := `
SELECT
COALESCE(ul.group_id, 0) as group_id,
COALESCE(g.name, '') as group_name,
COALESCE(g.name, '(无分组)') as group_name,
COUNT(*) as requests,
COALESCE(SUM(ul.input_tokens + ul.output_tokens + ul.cache_creation_tokens + ul.cache_read_tokens), 0) as total_tokens,
COALESCE(SUM(ul.total_cost), 0) as cost,
@@ -100,7 +100,7 @@ func (r *userGroupRateRepository) GetByGroupID(ctx context.Context, groupID int6
query := `
SELECT ugr.user_id, u.username, u.email, COALESCE(u.notes, ''), u.status, ugr.rate_multiplier
FROM user_group_rate_multipliers ugr
JOIN users u ON u.id = ugr.user_id
JOIN users u ON u.id = ugr.user_id AND u.deleted_at IS NULL
WHERE ugr.group_id = $1
ORDER BY ugr.user_id
`
+5 -6
View File
@@ -210,10 +210,9 @@ func TestAPIContracts(t *testing.T) {
"sora_video_price_per_request": null,
"sora_video_price_per_request_hd": null,
"claude_code_only": false,
"allow_messages_dispatch": false,
"fallback_group_id": null,
"fallback_group_id_on_invalid_request": null,
"allow_messages_dispatch": false,
"allow_messages_dispatch": false,
"fallback_group_id": null,
"fallback_group_id_on_invalid_request": null,
"created_at": "2025-01-02T03:04:05Z",
"updated_at": "2025-01-02T03:04:05Z"
}
@@ -651,8 +650,8 @@ func newContractDeps(t *testing.T) *contractDeps {
authHandler := handler.NewAuthHandler(cfg, nil, userService, settingService, nil, redeemService, nil)
apiKeyHandler := handler.NewAPIKeyHandler(apiKeyService)
usageHandler := handler.NewUsageHandler(usageService, apiKeyService)
adminSettingHandler := adminhandler.NewSettingHandler(settingService, nil, nil, nil, nil)
adminAccountHandler := adminhandler.NewAccountHandler(adminService, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil)
adminSettingHandler := adminhandler.NewSettingHandler(settingService, nil, nil, nil, nil, nil, nil)
adminAccountHandler := adminhandler.NewAccountHandler(adminService, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil)
jwtAuth := func(c *gin.Context) {
c.Set(string(middleware.ContextKeyUser), middleware.AuthSubject{
+19 -1
View File
@@ -264,6 +264,8 @@ func registerAccountRoutes(admin *gin.RouterGroup, h *handler.Handlers) {
accounts.POST("/today-stats/batch", h.Admin.Account.GetBatchTodayStats)
accounts.POST("/:id/clear-rate-limit", h.Admin.Account.ClearRateLimit)
accounts.POST("/:id/reset-quota", h.Admin.Account.ResetQuota)
accounts.GET("/:id/affinity-clients", h.Admin.Account.GetAffinityClients)
accounts.GET("/:id/affinity-details", h.Admin.Account.GetAffinityDetails)
accounts.GET("/:id/temp-unschedulable", h.Admin.Account.GetTempUnschedulable)
accounts.DELETE("/:id/temp-unschedulable", h.Admin.Account.ClearTempUnschedulable)
accounts.POST("/:id/schedulable", h.Admin.Account.SetSchedulable)
@@ -414,7 +416,7 @@ func registerSettingsRoutes(admin *gin.RouterGroup, h *handler.Handlers) {
// Beta 策略配置
adminSettings.GET("/beta-policy", h.Admin.Setting.GetBetaPolicySettings)
adminSettings.PUT("/beta-policy", h.Admin.Setting.UpdateBetaPolicySettings)
// Sora S3 存储配置
// Sora S3 存储配置(旧路由,保留兼容)
adminSettings.GET("/sora-s3", h.Admin.Setting.GetSoraS3Settings)
adminSettings.PUT("/sora-s3", h.Admin.Setting.UpdateSoraS3Settings)
adminSettings.POST("/sora-s3/test", h.Admin.Setting.TestSoraS3Connection)
@@ -423,6 +425,22 @@ func registerSettingsRoutes(admin *gin.RouterGroup, h *handler.Handlers) {
adminSettings.PUT("/sora-s3/profiles/:profile_id", h.Admin.Setting.UpdateSoraS3Profile)
adminSettings.DELETE("/sora-s3/profiles/:profile_id", h.Admin.Setting.DeleteSoraS3Profile)
adminSettings.POST("/sora-s3/profiles/:profile_id/activate", h.Admin.Setting.SetActiveSoraS3Profile)
// Sora 统一存储配置(新路由,指向相同 handler)
adminSettings.GET("/sora-storage", h.Admin.Setting.GetSoraS3Settings)
adminSettings.PUT("/sora-storage", h.Admin.Setting.UpdateSoraS3Settings)
adminSettings.POST("/sora-storage/test", h.Admin.Setting.TestSoraS3Connection)
adminSettings.GET("/sora-storage/profiles", h.Admin.Setting.ListSoraS3Profiles)
adminSettings.POST("/sora-storage/profiles", h.Admin.Setting.CreateSoraS3Profile)
adminSettings.PUT("/sora-storage/profiles/:profile_id", h.Admin.Setting.UpdateSoraS3Profile)
adminSettings.DELETE("/sora-storage/profiles/:profile_id", h.Admin.Setting.DeleteSoraS3Profile)
adminSettings.POST("/sora-storage/profiles/:profile_id/activate", h.Admin.Setting.SetActiveSoraS3Profile)
// Google Drive OAuth
adminSettings.POST("/sora-storage/gdrive-oauth/start", h.Admin.GDriveOAuth.StartOAuth)
adminSettings.POST("/sora-storage/gdrive-oauth/callback", h.Admin.GDriveOAuth.OAuthCallback)
adminSettings.POST("/sora-storage/gdrive-test", h.Admin.GDriveOAuth.TestGDriveStorage)
// Sora 存储统计
adminSettings.GET("/sora-storage/gdrive-quota", h.Admin.Setting.GetGDriveQuota)
adminSettings.GET("/sora-storage/video-stats", h.Admin.Setting.GetStorageVideoStats)
}
}
+10
View File
@@ -144,6 +144,16 @@ func RegisterGatewayRoutes(
{
soraV1.POST("/chat/completions", h.SoraGateway.ChatCompletions)
soraV1.GET("/models", h.Gateway.Models)
// Sora Videos/Images async task API
if h.SoraVideos != nil {
soraV1.POST("/videos", h.SoraVideos.CreateVideo)
soraV1.GET("/videos/:id", h.SoraVideos.GetVideo)
soraV1.POST("/videos/:id/remix", h.SoraVideos.RemixVideo)
soraV1.GET("/videos/:id/content", h.SoraVideos.GetVideoContent)
soraV1.POST("/images/generations", h.SoraVideos.CreateImage)
soraV1.POST("/images/edits", h.SoraVideos.EditImage)
}
}
// Sora 媒体代理(可选 API Key 验证)
+238
View File
@@ -1204,6 +1204,244 @@ func (a *Account) IsSessionIDMaskingEnabled() bool {
return false
}
// IsAffinityEnabled 检查是否启用亲和调度(统一入口)
// 仅适用于 Anthropic 账号,同时检查新字段 affinity_enabled 和旧字段 client_affinity_enabled(向后兼容)
func (a *Account) IsAffinityEnabled() bool {
if a.Platform != PlatformAnthropic {
return false
}
if a.Extra == nil {
return false
}
if v, ok := a.Extra["affinity_enabled"]; ok {
if enabled, ok := v.(bool); ok {
return enabled
}
return false
}
if v, ok := a.Extra["client_affinity_enabled"]; ok {
if enabled, ok := v.(bool); ok {
return enabled
}
}
return false
}
// IsClientAffinityEnabled 向后兼容别名,内部调用 IsAffinityEnabled
func (a *Account) IsClientAffinityEnabled() bool {
return a.IsAffinityEnabled()
}
// AffinityZone 表示账号的客户端亲和分区
type AffinityZone int
const (
AffinityZoneGreen AffinityZone = iota // 绿区:允许绑定新客户端,优先调度
AffinityZoneYellow // 黄区:允许绑定,仅在无绿区账号时降级调度
AffinityZoneRed // 红区:禁止调度
)
// GetAffinityBase 获取亲和基础限制(绿区上限),0 表示未配置
func (a *Account) GetAffinityBase() int {
if a.Extra == nil {
return 0
}
if v, ok := a.Extra["affinity_base"]; ok {
return parseExtraInt(v)
}
return 0
}
// GetAffinityBuffer 获取亲和缓冲区大小(黄区范围)
// 返回 (value, configured):
// - (0, false): 未配置 → 无限黄区(超过 base 永远黄区,永不红区)
// - (0, true): 显式设为 0 → 无黄区,超过 base 直接红区
// - (n, true): n > 0 → 黄区范围为 base+1 到 base+n
func (a *Account) GetAffinityBuffer() (int, bool) {
if a.Extra == nil {
return 0, false
}
v, ok := a.Extra["affinity_buffer"]
if !ok {
return 0, false
}
// 显式设为 null/nil → 视为未配置
if v == nil {
return 0, false
}
return parseExtraInt(v), true
}
// GetAffinityZone 根据当前绑定的客户端数量计算账号的亲和分区。
// 未开启亲和或未配置 base 的账号永远返回绿区。
func (a *Account) GetAffinityZone(clientCount int64) AffinityZone {
if !a.IsClientAffinityEnabled() {
return AffinityZoneGreen
}
base := a.GetAffinityBase()
if base <= 0 {
return AffinityZoneGreen
}
if clientCount <= int64(base) {
return AffinityZoneGreen
}
buffer, configured := a.GetAffinityBuffer()
if !configured {
return AffinityZoneYellow // 未配置 buffer → 无限黄区
}
if buffer == 0 {
return AffinityZoneRed // buffer=0 → 无黄区,直接红区
}
if clientCount <= int64(base+buffer) {
return AffinityZoneYellow
}
return AffinityZoneRed
}
// IsAffinityAllowSwitch 检查亲和账号是否允许在无绿区时切换到其他账号
// 默认 true(允许切换),设为 false 时即使全部红区也不切换
func (a *Account) IsAffinityAllowSwitch() bool {
if a.Extra == nil {
return true
}
if v, ok := a.Extra["affinity_allow_switch"]; ok {
if allow, ok := v.(bool); ok {
return allow
}
}
return true
}
// GetPinnedUsers 获取指定亲和用户列表
func (a *Account) GetPinnedUsers() []int64 {
if a.Extra == nil {
return nil
}
v, ok := a.Extra["pinned_users"]
if !ok || v == nil {
return nil
}
arr, ok := v.([]any)
if !ok || len(arr) == 0 {
return nil
}
result := make([]int64, 0, len(arr))
for _, item := range arr {
switch id := item.(type) {
case float64:
result = append(result, int64(id))
case int64:
result = append(result, id)
case int:
result = append(result, int64(id))
case json.Number:
if i, err := id.Int64(); err == nil {
result = append(result, i)
}
}
}
return result
}
// IsPinnedUser 检查 userID 是否在指定亲和用户列表中
func (a *Account) IsPinnedUser(userID int64) bool {
for _, id := range a.GetPinnedUsers() {
if id == userID {
return true
}
}
return false
}
// GetAffinityUserBase 获取用户维度亲和基础限制(绿区上限),0 表示未配置
func (a *Account) GetAffinityUserBase() int {
if a.Extra == nil {
return 0
}
if v, ok := a.Extra["affinity_user_base"]; ok {
return parseExtraInt(v)
}
return 0
}
// GetAffinityUserBuffer 获取用户维度亲和缓冲区大小(黄区范围)
// 返回 (value, configured),语义与 GetAffinityBuffer 相同
func (a *Account) GetAffinityUserBuffer() (int, bool) {
if a.Extra == nil {
return 0, false
}
v, ok := a.Extra["affinity_user_buffer"]
if !ok {
return 0, false
}
if v == nil {
return 0, false
}
return parseExtraInt(v), true
}
// GetPerUserClientLimit 获取每用户客户端限制,0 表示不限制
func (a *Account) GetPerUserClientLimit() int {
if a.Extra == nil {
return 0
}
if v, ok := a.Extra["per_user_client_limit"]; ok {
return parseExtraInt(v)
}
return 0
}
// getAffinityZoneForDim 根据单一维度的计数和三区参数计算亲和分区
func getAffinityZoneForDim(count int64, base int, buffer int, bufferConfigured bool) AffinityZone {
if base <= 0 {
return AffinityZoneGreen
}
if count <= int64(base) {
return AffinityZoneGreen
}
if !bufferConfigured {
return AffinityZoneYellow
}
if buffer == 0 {
return AffinityZoneRed
}
if count <= int64(base+buffer) {
return AffinityZoneYellow
}
return AffinityZoneRed
}
// GetMultiDimAffinityZone 根据多维度计数计算亲和分区,取所有维度中最严格的区域。
// 未开启亲和的账号永远返回绿区。
func (a *Account) GetMultiDimAffinityZone(userCount, clientCount, perUserCount int64) AffinityZone {
if !a.IsAffinityEnabled() {
return AffinityZoneGreen
}
worst := AffinityZoneGreen
// 客户端维度
clientBase := a.GetAffinityBase()
clientBuffer, clientBufCfg := a.GetAffinityBuffer()
if z := getAffinityZoneForDim(clientCount, clientBase, clientBuffer, clientBufCfg); z > worst {
worst = z
}
// 用户维度
userBase := a.GetAffinityUserBase()
userBuffer, userBufCfg := a.GetAffinityUserBuffer()
if z := getAffinityZoneForDim(userCount, userBase, userBuffer, userBufCfg); z > worst {
worst = z
}
// 每用户客户端维度
perUserLimit := a.GetPerUserClientLimit()
if perUserLimit > 0 && perUserCount > int64(perUserLimit) {
worst = AffinityZoneRed
}
return worst
}
// IsCacheTTLOverrideEnabled 检查是否启用缓存 TTL 强制替换
// 仅适用于 Anthropic OAuth/SetupToken 类型账号
// 启用后将所有 cache creation tokens 归入指定的 TTL 类型(5m 或 1h)
+25 -6
View File
@@ -143,9 +143,10 @@ type CreateGroupInput struct {
// 无效请求兜底分组 ID(仅 anthropic 平台使用)
FallbackGroupIDOnInvalidRequest *int64
// 模型路由配置(仅 anthropic 平台使用)
ModelRouting map[string][]int64
ModelRoutingEnabled bool // 是否启用模型路由
MCPXMLInject *bool
ModelRouting map[string][]int64
ModelRoutingEnabled bool // 是否启用模型路由
MCPXMLInject *bool
SimulateClaudeMaxEnabled *bool
// 支持的模型系列(仅 antigravity 平台使用)
SupportedModelScopes []string
// Sora 存储配额
@@ -182,9 +183,10 @@ type UpdateGroupInput struct {
// 无效请求兜底分组 ID(仅 anthropic 平台使用)
FallbackGroupIDOnInvalidRequest *int64
// 模型路由配置(仅 anthropic 平台使用)
ModelRouting map[string][]int64
ModelRoutingEnabled *bool // 是否启用模型路由
MCPXMLInject *bool
ModelRouting map[string][]int64
ModelRoutingEnabled *bool // 是否启用模型路由
MCPXMLInject *bool
SimulateClaudeMaxEnabled *bool
// 支持的模型系列(仅 antigravity 平台使用)
SupportedModelScopes *[]string
// Sora 存储配额
@@ -868,6 +870,13 @@ func (s *adminServiceImpl) CreateGroup(ctx context.Context, input *CreateGroupIn
if input.MCPXMLInject != nil {
mcpXMLInject = *input.MCPXMLInject
}
simulateClaudeMaxEnabled := false
if input.SimulateClaudeMaxEnabled != nil {
if platform != PlatformAnthropic && *input.SimulateClaudeMaxEnabled {
return nil, fmt.Errorf("simulate_claude_max_enabled only supported for anthropic groups")
}
simulateClaudeMaxEnabled = *input.SimulateClaudeMaxEnabled
}
// 如果指定了复制账号的源分组,先获取账号 ID 列表
var accountIDsToCopy []int64
@@ -924,6 +933,7 @@ func (s *adminServiceImpl) CreateGroup(ctx context.Context, input *CreateGroupIn
FallbackGroupIDOnInvalidRequest: fallbackOnInvalidRequest,
ModelRouting: input.ModelRouting,
MCPXMLInject: mcpXMLInject,
SimulateClaudeMaxEnabled: simulateClaudeMaxEnabled,
SupportedModelScopes: input.SupportedModelScopes,
SoraStorageQuotaBytes: input.SoraStorageQuotaBytes,
AllowMessagesDispatch: input.AllowMessagesDispatch,
@@ -1130,6 +1140,15 @@ func (s *adminServiceImpl) UpdateGroup(ctx context.Context, id int64, input *Upd
if input.MCPXMLInject != nil {
group.MCPXMLInject = *input.MCPXMLInject
}
if input.SimulateClaudeMaxEnabled != nil {
if group.Platform != PlatformAnthropic && *input.SimulateClaudeMaxEnabled {
return nil, fmt.Errorf("simulate_claude_max_enabled only supported for anthropic groups")
}
group.SimulateClaudeMaxEnabled = *input.SimulateClaudeMaxEnabled
}
if group.Platform != PlatformAnthropic {
group.SimulateClaudeMaxEnabled = false
}
// 支持的模型系列(仅 antigravity 平台使用)
if input.SupportedModelScopes != nil {
@@ -43,6 +43,16 @@ func (s *accountRepoStubForBulkUpdate) BindGroups(_ context.Context, accountID i
return nil
}
func (s *accountRepoStubForBulkUpdate) ListByGroup(_ context.Context, groupID int64) ([]Account, error) {
if err, ok := s.listByGroupErr[groupID]; ok {
return nil, err
}
if rows, ok := s.listByGroupData[groupID]; ok {
return rows, nil
}
return nil, nil
}
func (s *accountRepoStubForBulkUpdate) GetByIDs(_ context.Context, ids []int64) ([]*Account, error) {
s.getByIDsCalled = true
s.getByIDsIDs = append([]int64{}, ids...)
@@ -63,16 +73,6 @@ func (s *accountRepoStubForBulkUpdate) GetByID(_ context.Context, id int64) (*Ac
return nil, errors.New("account not found")
}
func (s *accountRepoStubForBulkUpdate) ListByGroup(_ context.Context, groupID int64) ([]Account, error) {
if err, ok := s.listByGroupErr[groupID]; ok {
return nil, err
}
if rows, ok := s.listByGroupData[groupID]; ok {
return rows, nil
}
return nil, nil
}
// TestAdminService_BulkUpdateAccounts_AllSuccessIDs 验证批量更新成功时返回 success_ids/failed_ids。
func TestAdminService_BulkUpdateAccounts_AllSuccessIDs(t *testing.T) {
repo := &accountRepoStubForBulkUpdate{}
@@ -785,3 +785,57 @@ func TestAdminService_UpdateGroup_InvalidRequestFallbackAllowsAntigravity(t *tes
require.NotNil(t, repo.updated)
require.Equal(t, fallbackID, *repo.updated.FallbackGroupIDOnInvalidRequest)
}
func TestAdminService_CreateGroup_SimulateClaudeMaxRequiresAnthropic(t *testing.T) {
repo := &groupRepoStubForAdmin{}
svc := &adminServiceImpl{groupRepo: repo}
enabled := true
_, err := svc.CreateGroup(context.Background(), &CreateGroupInput{
Name: "openai-group",
Platform: PlatformOpenAI,
SimulateClaudeMaxEnabled: &enabled,
})
require.Error(t, err)
require.Contains(t, err.Error(), "simulate_claude_max_enabled only supported for anthropic groups")
require.Nil(t, repo.created)
}
func TestAdminService_UpdateGroup_SimulateClaudeMaxRequiresAnthropic(t *testing.T) {
existingGroup := &Group{
ID: 1,
Name: "openai-group",
Platform: PlatformOpenAI,
Status: StatusActive,
}
repo := &groupRepoStubForAdmin{getByID: existingGroup}
svc := &adminServiceImpl{groupRepo: repo}
enabled := true
_, err := svc.UpdateGroup(context.Background(), 1, &UpdateGroupInput{
SimulateClaudeMaxEnabled: &enabled,
})
require.Error(t, err)
require.Contains(t, err.Error(), "simulate_claude_max_enabled only supported for anthropic groups")
require.Nil(t, repo.updated)
}
func TestAdminService_UpdateGroup_ClearsSimulateClaudeMaxWhenPlatformChanges(t *testing.T) {
existingGroup := &Group{
ID: 1,
Name: "anthropic-group",
Platform: PlatformAnthropic,
Status: StatusActive,
SimulateClaudeMaxEnabled: true,
}
repo := &groupRepoStubForAdmin{getByID: existingGroup}
svc := &adminServiceImpl{groupRepo: repo}
group, err := svc.UpdateGroup(context.Background(), 1, &UpdateGroupInput{
Platform: PlatformOpenAI,
})
require.NoError(t, err)
require.NotNil(t, group)
require.NotNil(t, repo.updated)
require.False(t, repo.updated.SimulateClaudeMaxEnabled)
}
@@ -1717,7 +1717,7 @@ func (s *AntigravityGatewayService) Forward(ctx context.Context, c *gin.Context,
var clientDisconnect bool
if claudeReq.Stream {
// 客户端要求流式,直接透传转换
streamRes, err := s.handleClaudeStreamingResponse(c, resp, startTime, originalModel)
streamRes, err := s.handleClaudeStreamingResponse(c, resp, startTime, originalModel, account.ID)
if err != nil {
logger.LegacyPrintf("service.antigravity_gateway", "%s status=stream_error error=%v", prefix, err)
return nil, err
@@ -1727,7 +1727,7 @@ func (s *AntigravityGatewayService) Forward(ctx context.Context, c *gin.Context,
clientDisconnect = streamRes.clientDisconnect
} else {
// 客户端要求非流式,收集流式响应后转换返回
streamRes, err := s.handleClaudeStreamToNonStreaming(c, resp, startTime, originalModel)
streamRes, err := s.handleClaudeStreamToNonStreaming(c, resp, startTime, originalModel, account.ID)
if err != nil {
logger.LegacyPrintf("service.antigravity_gateway", "%s status=stream_collect_error error=%v", prefix, err)
return nil, err
@@ -1736,6 +1736,9 @@ func (s *AntigravityGatewayService) Forward(ctx context.Context, c *gin.Context,
firstTokenMs = streamRes.firstTokenMs
}
// Claude Max cache billing: 同步 ForwardResult.Usage 与客户端响应体一致
applyClaudeMaxCacheBillingPolicyToUsage(usage, parsedRequestFromGinContext(c), claudeMaxGroupFromGinContext(c), originalModel, account.ID)
return &ForwardResult{
RequestID: requestID,
Usage: *usage,
@@ -3670,7 +3673,7 @@ func (s *AntigravityGatewayService) writeGoogleError(c *gin.Context, status int,
// handleClaudeStreamToNonStreaming 收集上游流式响应,转换为 Claude 非流式格式返回
// 用于处理客户端非流式请求但上游只支持流式的情况
func (s *AntigravityGatewayService) handleClaudeStreamToNonStreaming(c *gin.Context, resp *http.Response, startTime time.Time, originalModel string) (*antigravityStreamResult, error) {
func (s *AntigravityGatewayService) handleClaudeStreamToNonStreaming(c *gin.Context, resp *http.Response, startTime time.Time, originalModel string, accountID int64) (*antigravityStreamResult, error) {
scanner := bufio.NewScanner(resp.Body)
maxLineSize := defaultMaxLineSize
if s.settingService.cfg != nil && s.settingService.cfg.Gateway.MaxLineSize > 0 {
@@ -3828,6 +3831,9 @@ returnResponse:
return nil, s.writeClaudeError(c, http.StatusBadGateway, "upstream_error", "Failed to parse upstream response")
}
// Claude Max cache billing simulation (non-streaming)
claudeResp = applyClaudeMaxNonStreamingRewrite(c, claudeResp, agUsage, originalModel, accountID)
c.Data(http.StatusOK, "application/json", claudeResp)
// 转换为 service.ClaudeUsage
@@ -3842,7 +3848,7 @@ returnResponse:
}
// handleClaudeStreamingResponse 处理 Claude 流式响应(Gemini SSE → Claude SSE 转换)
func (s *AntigravityGatewayService) handleClaudeStreamingResponse(c *gin.Context, resp *http.Response, startTime time.Time, originalModel string) (*antigravityStreamResult, error) {
func (s *AntigravityGatewayService) handleClaudeStreamingResponse(c *gin.Context, resp *http.Response, startTime time.Time, originalModel string, accountID int64) (*antigravityStreamResult, error) {
c.Header("Content-Type", "text/event-stream")
c.Header("Cache-Control", "no-cache")
c.Header("Connection", "keep-alive")
@@ -3855,6 +3861,8 @@ func (s *AntigravityGatewayService) handleClaudeStreamingResponse(c *gin.Context
}
processor := antigravity.NewStreamingProcessor(originalModel)
setupClaudeMaxStreamingHook(c, processor, originalModel, accountID)
var firstTokenMs *int
// 使用 Scanner 并限制单行大小,避免 ReadString 无上限导致 OOM
scanner := bufio.NewScanner(resp.Body)
@@ -922,7 +922,7 @@ func TestHandleClaudeStreamingResponse_NormalComplete(t *testing.T) {
fmt.Fprintln(pw, "")
}()
result, err := svc.handleClaudeStreamingResponse(c, resp, time.Now(), "claude-sonnet-4-5")
result, err := svc.handleClaudeStreamingResponse(c, resp, time.Now(), "claude-sonnet-4-5", 0)
_ = pr.Close()
require.NoError(t, err)
@@ -999,7 +999,7 @@ func TestHandleClaudeStreamingResponse_ThoughtsTokenCount(t *testing.T) {
fmt.Fprintln(pw, "")
}()
result, err := svc.handleClaudeStreamingResponse(c, resp, time.Now(), "gemini-2.5-pro")
result, err := svc.handleClaudeStreamingResponse(c, resp, time.Now(), "gemini-2.5-pro", 0)
_ = pr.Close()
require.NoError(t, err)
@@ -1202,7 +1202,7 @@ func TestHandleClaudeStreamingResponse_ClientDisconnect(t *testing.T) {
fmt.Fprintln(pw, "")
}()
result, err := svc.handleClaudeStreamingResponse(c, resp, time.Now(), "claude-sonnet-4-5")
result, err := svc.handleClaudeStreamingResponse(c, resp, time.Now(), "claude-sonnet-4-5", 0)
_ = pr.Close()
require.NoError(t, err)
@@ -1234,7 +1234,7 @@ func TestHandleClaudeStreamingResponse_EmptyStream(t *testing.T) {
fmt.Fprintln(pw, "")
}()
_, err := svc.handleClaudeStreamingResponse(c, resp, time.Now(), "claude-sonnet-4-5")
_, err := svc.handleClaudeStreamingResponse(c, resp, time.Now(), "claude-sonnet-4-5", 0)
_ = pr.Close()
// 应当返回 UpstreamFailoverError 而非 nil,以便上层触发 failover
@@ -1266,7 +1266,7 @@ func TestHandleClaudeStreamingResponse_ContextCanceled(t *testing.T) {
resp := &http.Response{StatusCode: http.StatusOK, Body: cancelReadCloser{}, Header: http.Header{}}
result, err := svc.handleClaudeStreamingResponse(c, resp, time.Now(), "claude-sonnet-4-5")
result, err := svc.handleClaudeStreamingResponse(c, resp, time.Now(), "claude-sonnet-4-5", 0)
require.NoError(t, err)
require.NotNil(t, result)
@@ -9,6 +9,7 @@ import (
"net/http"
"strings"
"testing"
"time"
"github.com/stretchr/testify/require"
)
@@ -29,6 +30,27 @@ func (c *stubSmartRetryCache) DeleteSessionAccountID(_ context.Context, groupID
c.deleteCalls = append(c.deleteCalls, deleteSessionCall{groupID: groupID, sessionHash: sessionHash})
return nil
}
func (c *stubSmartRetryCache) GetAffinityAccounts(_ context.Context, _ int64, _ int64, _ string, _ time.Duration) ([]int64, error) {
return nil, nil
}
func (c *stubSmartRetryCache) UpdateAffinity(_ context.Context, _ int64, _ int64, _ string, _ int64, _ time.Duration) error {
return nil
}
func (c *stubSmartRetryCache) GetAffinityMultiCount(_ context.Context, _ int64, _ int64, _ int64, _ time.Duration) (int64, int64, int64, error) {
return 0, 0, 0, nil
}
func (c *stubSmartRetryCache) GetAccountAffinityCountBatch(_ context.Context, _ int64, _ []int64, _ time.Duration) (map[int64]int64, error) {
return map[int64]int64{}, nil
}
func (c *stubSmartRetryCache) GetAccountAffinityClientsBatch(_ context.Context, _ map[int64][]int64, _ time.Duration) (map[int64][]string, error) {
return map[int64][]string{}, nil
}
func (c *stubSmartRetryCache) GetAccountAffinityClientsWithScores(_ context.Context, _ int64, _ []int64, _ time.Duration) ([]AffinityClient, error) {
return nil, nil
}
func (c *stubSmartRetryCache) ClearAccountAffinity(_ context.Context, _ int64, _ []int64) error {
return nil
}
// mockSmartRetryUpstream 用于 handleSmartRetry 测试的 mock upstream
type mockSmartRetryUpstream struct {
@@ -80,17 +102,12 @@ func (m *mockSmartRetryUpstream) Do(req *http.Request, proxyURL string, accountI
m.responseBodies[respIdx] = bodyBytes
}
// 用缓存的 body 字节重建新的 reader
var body io.ReadCloser
// 用缓存的 body 重建 reader(支持重试场景多次读取)
cloned := *resp
if m.responseBodies[respIdx] != nil {
body = io.NopCloser(bytes.NewReader(m.responseBodies[respIdx]))
cloned.Body = io.NopCloser(bytes.NewReader(m.responseBodies[respIdx]))
}
return &http.Response{
StatusCode: resp.StatusCode,
Header: resp.Header.Clone(),
Body: body,
}, respErr
return &cloned, respErr
}
func (m *mockSmartRetryUpstream) DoWithTLS(req *http.Request, proxyURL string, accountID int64, accountConcurrency int, enableTLSFingerprint bool) (*http.Response, error) {
@@ -59,9 +59,10 @@ type APIKeyAuthGroupSnapshot struct {
// Model routing is used by gateway account selection, so it must be part of auth cache snapshot.
// Only anthropic groups use these fields; others may leave them empty.
ModelRouting map[string][]int64 `json:"model_routing,omitempty"`
ModelRoutingEnabled bool `json:"model_routing_enabled"`
MCPXMLInject bool `json:"mcp_xml_inject"`
ModelRouting map[string][]int64 `json:"model_routing,omitempty"`
ModelRoutingEnabled bool `json:"model_routing_enabled"`
MCPXMLInject bool `json:"mcp_xml_inject"`
SimulateClaudeMaxEnabled bool `json:"simulate_claude_max_enabled"`
// 支持的模型系列(仅 antigravity 平台使用)
SupportedModelScopes []string `json:"supported_model_scopes,omitempty"`
@@ -244,6 +244,7 @@ func (s *APIKeyService) snapshotFromAPIKey(apiKey *APIKey) *APIKeyAuthSnapshot {
ModelRouting: apiKey.Group.ModelRouting,
ModelRoutingEnabled: apiKey.Group.ModelRoutingEnabled,
MCPXMLInject: apiKey.Group.MCPXMLInject,
SimulateClaudeMaxEnabled: apiKey.Group.SimulateClaudeMaxEnabled,
SupportedModelScopes: apiKey.Group.SupportedModelScopes,
AllowMessagesDispatch: apiKey.Group.AllowMessagesDispatch,
DefaultMappedModel: apiKey.Group.DefaultMappedModel,
@@ -303,6 +304,7 @@ func (s *APIKeyService) snapshotToAPIKey(key string, snapshot *APIKeyAuthSnapsho
ModelRouting: snapshot.Group.ModelRouting,
ModelRoutingEnabled: snapshot.Group.ModelRoutingEnabled,
MCPXMLInject: snapshot.Group.MCPXMLInject,
SimulateClaudeMaxEnabled: snapshot.Group.SimulateClaudeMaxEnabled,
SupportedModelScopes: snapshot.Group.SupportedModelScopes,
AllowMessagesDispatch: snapshot.Group.AllowMessagesDispatch,
DefaultMappedModel: snapshot.Group.DefaultMappedModel,
@@ -0,0 +1,450 @@
package service
import (
"encoding/json"
"strings"
"github.com/Wei-Shaw/sub2api/internal/pkg/claude"
"github.com/Wei-Shaw/sub2api/internal/pkg/logger"
"github.com/tidwall/gjson"
)
type claudeMaxCacheBillingOutcome struct {
Simulated bool
}
func applyClaudeMaxCacheBillingPolicyToUsage(usage *ClaudeUsage, parsed *ParsedRequest, group *Group, model string, accountID int64) claudeMaxCacheBillingOutcome {
var out claudeMaxCacheBillingOutcome
if usage == nil || !shouldApplyClaudeMaxBillingRulesForUsage(group, model, parsed) {
return out
}
resolvedModel := strings.TrimSpace(model)
if resolvedModel == "" && parsed != nil {
resolvedModel = strings.TrimSpace(parsed.Model)
}
if hasCacheCreationTokens(*usage) {
// Upstream already returned cache creation usage; keep original usage.
return out
}
if !shouldSimulateClaudeMaxUsageForUsage(*usage, parsed) {
return out
}
beforeInputTokens := usage.InputTokens
out.Simulated = safelyProjectUsageToClaudeMax1H(usage, parsed)
if out.Simulated {
logger.LegacyPrintf("service.gateway", "simulate_claude_max_usage: model=%s account=%d input_tokens:%d->%d cache_creation_1h=%d",
resolvedModel,
accountID,
beforeInputTokens,
usage.InputTokens,
usage.CacheCreation1hTokens,
)
}
return out
}
func isClaudeFamilyModel(model string) bool {
normalized := strings.ToLower(strings.TrimSpace(claude.NormalizeModelID(model)))
if normalized == "" {
return false
}
return strings.Contains(normalized, "claude-")
}
func shouldApplyClaudeMaxBillingRules(input *RecordUsageInput) bool {
if input == nil || input.Result == nil || input.APIKey == nil || input.APIKey.Group == nil {
return false
}
return shouldApplyClaudeMaxBillingRulesForUsage(input.APIKey.Group, input.Result.Model, input.ParsedRequest)
}
func shouldApplyClaudeMaxBillingRulesForUsage(group *Group, model string, parsed *ParsedRequest) bool {
if group == nil {
return false
}
if !group.SimulateClaudeMaxEnabled || group.Platform != PlatformAnthropic {
return false
}
resolvedModel := model
if resolvedModel == "" && parsed != nil {
resolvedModel = parsed.Model
}
if !isClaudeFamilyModel(resolvedModel) {
return false
}
return true
}
func hasCacheCreationTokens(usage ClaudeUsage) bool {
return usage.CacheCreationInputTokens > 0 || usage.CacheCreation5mTokens > 0 || usage.CacheCreation1hTokens > 0
}
func shouldSimulateClaudeMaxUsage(input *RecordUsageInput) bool {
if input == nil || input.Result == nil {
return false
}
if !shouldApplyClaudeMaxBillingRules(input) {
return false
}
return shouldSimulateClaudeMaxUsageForUsage(input.Result.Usage, input.ParsedRequest)
}
func shouldSimulateClaudeMaxUsageForUsage(usage ClaudeUsage, parsed *ParsedRequest) bool {
if usage.InputTokens <= 0 {
return false
}
if hasCacheCreationTokens(usage) {
return false
}
if !hasClaudeCacheSignals(parsed) {
return false
}
return true
}
func safelyProjectUsageToClaudeMax1H(usage *ClaudeUsage, parsed *ParsedRequest) (changed bool) {
defer func() {
if r := recover(); r != nil {
logger.LegacyPrintf("service.gateway", "simulate_claude_max_usage skipped: panic=%v", r)
changed = false
}
}()
return projectUsageToClaudeMax1H(usage, parsed)
}
func projectUsageToClaudeMax1H(usage *ClaudeUsage, parsed *ParsedRequest) bool {
if usage == nil {
return false
}
totalWindowTokens := usage.InputTokens + usage.CacheCreation5mTokens + usage.CacheCreation1hTokens
if totalWindowTokens <= 1 {
return false
}
simulatedInputTokens := computeClaudeMaxProjectedInputTokens(totalWindowTokens, parsed)
if simulatedInputTokens <= 0 {
simulatedInputTokens = 1
}
if simulatedInputTokens >= totalWindowTokens {
simulatedInputTokens = totalWindowTokens - 1
}
cacheCreation1hTokens := totalWindowTokens - simulatedInputTokens
if usage.InputTokens == simulatedInputTokens &&
usage.CacheCreation5mTokens == 0 &&
usage.CacheCreation1hTokens == cacheCreation1hTokens &&
usage.CacheCreationInputTokens == cacheCreation1hTokens {
return false
}
usage.InputTokens = simulatedInputTokens
usage.CacheCreation5mTokens = 0
usage.CacheCreation1hTokens = cacheCreation1hTokens
usage.CacheCreationInputTokens = cacheCreation1hTokens
return true
}
type claudeCacheProjection struct {
HasBreakpoint bool
BreakpointCount int
TotalEstimatedTokens int
TailEstimatedTokens int
}
func computeClaudeMaxProjectedInputTokens(totalWindowTokens int, parsed *ParsedRequest) int {
if totalWindowTokens <= 1 {
return totalWindowTokens
}
projection := analyzeClaudeCacheProjection(parsed)
if !projection.HasBreakpoint || projection.TotalEstimatedTokens <= 0 || projection.TailEstimatedTokens <= 0 {
return totalWindowTokens
}
totalEstimate := int64(projection.TotalEstimatedTokens)
tailEstimate := int64(projection.TailEstimatedTokens)
if tailEstimate > totalEstimate {
tailEstimate = totalEstimate
}
scaled := (int64(totalWindowTokens)*tailEstimate + totalEstimate/2) / totalEstimate
if scaled <= 0 {
scaled = 1
}
if scaled >= int64(totalWindowTokens) {
scaled = int64(totalWindowTokens - 1)
}
return int(scaled)
}
func hasClaudeCacheSignals(parsed *ParsedRequest) bool {
if parsed == nil {
return false
}
if hasTopLevelEphemeralCacheControl(parsed) {
return true
}
return countExplicitCacheBreakpoints(parsed) > 0
}
func hasTopLevelEphemeralCacheControl(parsed *ParsedRequest) bool {
if parsed == nil || len(parsed.Body) == 0 {
return false
}
cacheType := strings.TrimSpace(gjson.GetBytes(parsed.Body, "cache_control.type").String())
return strings.EqualFold(cacheType, "ephemeral")
}
func analyzeClaudeCacheProjection(parsed *ParsedRequest) claudeCacheProjection {
var projection claudeCacheProjection
if parsed == nil {
return projection
}
total := 0
lastBreakpointAt := -1
switch system := parsed.System.(type) {
case string:
total += claudeMaxMessageOverheadTokens + estimateClaudeTextTokens(system)
case []any:
for _, raw := range system {
block, ok := raw.(map[string]any)
if !ok {
total += claudeMaxUnknownContentTokens
continue
}
total += estimateClaudeBlockTokens(block)
if hasEphemeralCacheControl(block) {
lastBreakpointAt = total
projection.BreakpointCount++
projection.HasBreakpoint = true
}
}
}
for _, rawMsg := range parsed.Messages {
total += claudeMaxMessageOverheadTokens
msg, ok := rawMsg.(map[string]any)
if !ok {
total += claudeMaxUnknownContentTokens
continue
}
content, exists := msg["content"]
if !exists {
continue
}
msgTokens, msgLastBreak, msgBreakCount := estimateClaudeContentTokens(content)
total += msgTokens
if msgBreakCount > 0 {
lastBreakpointAt = total - msgTokens + msgLastBreak
projection.BreakpointCount += msgBreakCount
projection.HasBreakpoint = true
}
}
if total <= 0 {
total = 1
}
projection.TotalEstimatedTokens = total
if projection.HasBreakpoint && lastBreakpointAt >= 0 {
tail := total - lastBreakpointAt
if tail <= 0 {
tail = 1
}
projection.TailEstimatedTokens = tail
return projection
}
if hasTopLevelEphemeralCacheControl(parsed) {
tail := estimateLastUserMessageTokens(parsed)
if tail <= 0 {
tail = 1
}
projection.HasBreakpoint = true
projection.BreakpointCount = 1
projection.TailEstimatedTokens = tail
}
return projection
}
func countExplicitCacheBreakpoints(parsed *ParsedRequest) int {
if parsed == nil {
return 0
}
total := 0
if system, ok := parsed.System.([]any); ok {
for _, raw := range system {
if block, ok := raw.(map[string]any); ok && hasEphemeralCacheControl(block) {
total++
}
}
}
for _, rawMsg := range parsed.Messages {
msg, ok := rawMsg.(map[string]any)
if !ok {
continue
}
content, ok := msg["content"].([]any)
if !ok {
continue
}
for _, raw := range content {
if block, ok := raw.(map[string]any); ok && hasEphemeralCacheControl(block) {
total++
}
}
}
return total
}
func hasEphemeralCacheControl(block map[string]any) bool {
if block == nil {
return false
}
raw, ok := block["cache_control"]
if !ok || raw == nil {
return false
}
switch cc := raw.(type) {
case map[string]any:
cacheType, _ := cc["type"].(string)
return strings.EqualFold(strings.TrimSpace(cacheType), "ephemeral")
case map[string]string:
return strings.EqualFold(strings.TrimSpace(cc["type"]), "ephemeral")
default:
return false
}
}
func estimateClaudeContentTokens(content any) (tokens int, lastBreakAt int, breakpointCount int) {
switch value := content.(type) {
case string:
return estimateClaudeTextTokens(value), -1, 0
case []any:
total := 0
lastBreak := -1
breaks := 0
for _, raw := range value {
block, ok := raw.(map[string]any)
if !ok {
total += claudeMaxUnknownContentTokens
continue
}
total += estimateClaudeBlockTokens(block)
if hasEphemeralCacheControl(block) {
lastBreak = total
breaks++
}
}
return total, lastBreak, breaks
default:
return estimateStructuredTokens(value), -1, 0
}
}
func estimateClaudeBlockTokens(block map[string]any) int {
if block == nil {
return claudeMaxUnknownContentTokens
}
tokens := claudeMaxBlockOverheadTokens
blockType, _ := block["type"].(string)
switch blockType {
case "text":
if text, ok := block["text"].(string); ok {
tokens += estimateClaudeTextTokens(text)
}
case "tool_result":
if content, ok := block["content"]; ok {
nested, _, _ := estimateClaudeContentTokens(content)
tokens += nested
}
case "tool_use":
if name, ok := block["name"].(string); ok {
tokens += estimateClaudeTextTokens(name)
}
if input, ok := block["input"]; ok {
tokens += estimateStructuredTokens(input)
}
default:
if text, ok := block["text"].(string); ok {
tokens += estimateClaudeTextTokens(text)
} else if content, ok := block["content"]; ok {
nested, _, _ := estimateClaudeContentTokens(content)
tokens += nested
}
}
if tokens <= claudeMaxBlockOverheadTokens {
tokens += claudeMaxUnknownContentTokens
}
return tokens
}
func estimateLastUserMessageTokens(parsed *ParsedRequest) int {
if parsed == nil || len(parsed.Messages) == 0 {
return 0
}
for i := len(parsed.Messages) - 1; i >= 0; i-- {
msg, ok := parsed.Messages[i].(map[string]any)
if !ok {
continue
}
role, _ := msg["role"].(string)
if !strings.EqualFold(strings.TrimSpace(role), "user") {
continue
}
tokens, _, _ := estimateClaudeContentTokens(msg["content"])
return claudeMaxMessageOverheadTokens + tokens
}
return 0
}
func estimateStructuredTokens(v any) int {
if v == nil {
return 0
}
raw, err := json.Marshal(v)
if err != nil {
return claudeMaxUnknownContentTokens
}
return estimateClaudeTextTokens(string(raw))
}
func estimateClaudeTextTokens(text string) int {
if tokens, ok := estimateTokensByThirdPartyTokenizer(text); ok {
return tokens
}
return estimateClaudeTextTokensHeuristic(text)
}
func estimateClaudeTextTokensHeuristic(text string) int {
normalized := strings.Join(strings.Fields(strings.TrimSpace(text)), " ")
if normalized == "" {
return 0
}
asciiChars := 0
nonASCIIChars := 0
for _, r := range normalized {
if r <= 127 {
asciiChars++
} else {
nonASCIIChars++
}
}
tokens := nonASCIIChars
if asciiChars > 0 {
tokens += (asciiChars + 3) / 4
}
if words := len(strings.Fields(normalized)); words > tokens {
tokens = words
}
if tokens <= 0 {
return 1
}
return tokens
}
@@ -0,0 +1,156 @@
package service
import (
"strings"
"testing"
)
func TestProjectUsageToClaudeMax1H_Conservation(t *testing.T) {
usage := &ClaudeUsage{
InputTokens: 1200,
CacheCreationInputTokens: 0,
CacheCreation5mTokens: 0,
CacheCreation1hTokens: 0,
}
parsed := &ParsedRequest{
Model: "claude-sonnet-4-5",
Messages: []any{
map[string]any{
"role": "user",
"content": []any{
map[string]any{
"type": "text",
"text": strings.Repeat("cached context ", 200),
"cache_control": map[string]any{"type": "ephemeral"},
},
map[string]any{
"type": "text",
"text": "summarize quickly",
},
},
},
},
}
changed := projectUsageToClaudeMax1H(usage, parsed)
if !changed {
t.Fatalf("expected usage to be projected")
}
total := usage.InputTokens + usage.CacheCreation5mTokens + usage.CacheCreation1hTokens
if total != 1200 {
t.Fatalf("total tokens changed: got=%d want=%d", total, 1200)
}
if usage.CacheCreation5mTokens != 0 {
t.Fatalf("cache_creation_5m should be 0, got=%d", usage.CacheCreation5mTokens)
}
if usage.InputTokens <= 0 || usage.InputTokens >= 1200 {
t.Fatalf("simulated input out of range, got=%d", usage.InputTokens)
}
if usage.InputTokens > 100 {
t.Fatalf("simulated input should stay near cache breakpoint tail, got=%d", usage.InputTokens)
}
if usage.CacheCreation1hTokens <= 0 {
t.Fatalf("cache_creation_1h should be > 0, got=%d", usage.CacheCreation1hTokens)
}
if usage.CacheCreationInputTokens != usage.CacheCreation1hTokens {
t.Fatalf("cache_creation_input_tokens mismatch: got=%d want=%d", usage.CacheCreationInputTokens, usage.CacheCreation1hTokens)
}
}
func TestComputeClaudeMaxProjectedInputTokens_Deterministic(t *testing.T) {
parsed := &ParsedRequest{
Model: "claude-opus-4-5",
Messages: []any{
map[string]any{
"role": "user",
"content": []any{
map[string]any{
"type": "text",
"text": "build context",
"cache_control": map[string]any{"type": "ephemeral"},
},
map[string]any{
"type": "text",
"text": "what is failing now",
},
},
},
},
}
got1 := computeClaudeMaxProjectedInputTokens(4096, parsed)
got2 := computeClaudeMaxProjectedInputTokens(4096, parsed)
if got1 != got2 {
t.Fatalf("non-deterministic input tokens: %d != %d", got1, got2)
}
}
func TestShouldSimulateClaudeMaxUsage(t *testing.T) {
group := &Group{
Platform: PlatformAnthropic,
SimulateClaudeMaxEnabled: true,
}
input := &RecordUsageInput{
Result: &ForwardResult{
Model: "claude-sonnet-4-5",
Usage: ClaudeUsage{
InputTokens: 3000,
CacheCreationInputTokens: 0,
CacheCreation5mTokens: 0,
CacheCreation1hTokens: 0,
},
},
ParsedRequest: &ParsedRequest{
Messages: []any{
map[string]any{
"role": "user",
"content": []any{
map[string]any{
"type": "text",
"text": "cached",
"cache_control": map[string]any{"type": "ephemeral"},
},
map[string]any{
"type": "text",
"text": "tail",
},
},
},
},
},
APIKey: &APIKey{Group: group},
}
if !shouldSimulateClaudeMaxUsage(input) {
t.Fatalf("expected simulate=true for claude group with cache signal")
}
input.ParsedRequest = &ParsedRequest{
Messages: []any{
map[string]any{"role": "user", "content": "no cache signal"},
},
}
if shouldSimulateClaudeMaxUsage(input) {
t.Fatalf("expected simulate=false when request has no cache signal")
}
input.ParsedRequest = &ParsedRequest{
Messages: []any{
map[string]any{
"role": "user",
"content": []any{
map[string]any{
"type": "text",
"text": "cached",
"cache_control": map[string]any{"type": "ephemeral"},
},
},
},
},
}
input.Result.Usage.CacheCreationInputTokens = 100
if shouldSimulateClaudeMaxUsage(input) {
t.Fatalf("expected simulate=false when cache creation already exists")
}
}
@@ -0,0 +1,41 @@
package service
import (
"sync"
tiktoken "github.com/pkoukk/tiktoken-go"
tiktokenloader "github.com/pkoukk/tiktoken-go-loader"
)
var (
claudeTokenizerOnce sync.Once
claudeTokenizer *tiktoken.Tiktoken
)
func getClaudeTokenizer() *tiktoken.Tiktoken {
claudeTokenizerOnce.Do(func() {
// Use offline loader to avoid runtime dictionary download.
tiktoken.SetBpeLoader(tiktokenloader.NewOfflineLoader())
// Use a high-capacity tokenizer as the default approximation for Claude payloads.
enc, err := tiktoken.GetEncoding(tiktoken.MODEL_O200K_BASE)
if err != nil {
enc, err = tiktoken.GetEncoding(tiktoken.MODEL_CL100K_BASE)
}
if err == nil {
claudeTokenizer = enc
}
})
return claudeTokenizer
}
func estimateTokensByThirdPartyTokenizer(text string) (int, bool) {
enc := getClaudeTokenizer()
if enc == nil {
return 0, false
}
tokens := len(enc.EncodeOrdinary(text))
if tokens <= 0 {
return 0, false
}
return tokens, true
}
@@ -343,8 +343,9 @@ func (s *ConcurrencyService) StartSlotCleanupWorker(accountRepo AccountRepositor
}()
}
// GetAccountConcurrencyBatch gets current concurrency counts for multiple accounts
// Returns a map of accountID -> current concurrency count
// GetAccountConcurrencyBatch gets current concurrency counts for multiple accounts.
// Uses a detached context with timeout to prevent HTTP request cancellation from
// causing the entire batch to fail (which would show all concurrency as 0).
func (s *ConcurrencyService) GetAccountConcurrencyBatch(ctx context.Context, accountIDs []int64) (map[int64]int, error) {
if len(accountIDs) == 0 {
return map[int64]int{}, nil
@@ -356,5 +357,11 @@ func (s *ConcurrencyService) GetAccountConcurrencyBatch(ctx context.Context, acc
}
return result, nil
}
return s.cache.GetAccountConcurrencyBatch(ctx, accountIDs)
// Use a detached context so that a cancelled HTTP request doesn't cause
// the Redis pipeline to fail and return all-zero concurrency counts.
redisCtx, cancel := context.WithTimeout(context.Background(), 3*time.Second)
defer cancel()
return s.cache.GetAccountConcurrencyBatch(redisCtx, accountIDs)
}
@@ -87,6 +87,7 @@ func (c *stubConcurrencyCacheForTest) GetAccountsLoadBatch(_ context.Context, _
func (c *stubConcurrencyCacheForTest) GetUsersLoadBatch(_ context.Context, _ []UserWithConcurrency) (map[int64]*UserLoadInfo, error) {
return c.usersLoadBatch, c.usersLoadErr
}
func (c *stubConcurrencyCacheForTest) CleanupExpiredAccountSlots(_ context.Context, _ int64) error {
return c.cleanupErr
}
@@ -220,7 +220,7 @@ func TestApplyErrorPassthroughRule_SkipMonitoringSetsContextKey(t *testing.T) {
v, exists := c.Get(OpsSkipPassthroughKey)
assert.True(t, exists, "OpsSkipPassthroughKey should be set when skip_monitoring=true")
boolVal, ok := v.(bool)
assert.True(t, ok, "value should be bool")
assert.True(t, ok, "value should be a bool")
assert.True(t, boolVal)
}
+26 -5
View File
@@ -110,13 +110,12 @@ func TestCheckErrorPolicy(t *testing.T) {
expected: ErrorPolicyTempUnscheduled,
},
{
// Antigravity 401 不走升级逻辑(由 applyErrorPolicy 的 temp_unschedulable_rules 自行控制),
// second hit 仍然返回 TempUnscheduled。
name: "temp_unschedulable_401_second_hit_antigravity_stays_temp",
// Gemini OAuth 401 second hit 会升级为 error(返回 None,交由默认错误逻辑处理)。
name: "temp_unschedulable_401_second_hit_gemini_escalates",
account: &Account{
ID: 15,
Type: AccountTypeOAuth,
Platform: PlatformAntigravity,
Platform: PlatformGemini, // 非 Antigravity 平台 401 second hit 升级
TempUnschedulableReason: `{"status_code":401,"until_unix":1735689600}`,
Credentials: map[string]any{
"temp_unschedulable_enabled": true,
@@ -131,7 +130,29 @@ func TestCheckErrorPolicy(t *testing.T) {
},
statusCode: 401,
body: []byte(`unauthorized`),
expected: ErrorPolicyTempUnscheduled,
expected: ErrorPolicyNone, // Gemini 401 second hit 升级为 error
},
{
name: "temp_unschedulable_401_antigravity_no_escalation",
account: &Account{
ID: 16,
Type: AccountTypeOAuth,
Platform: PlatformAntigravity, // Antigravity 跳过 401 升级,由 rules 正常处理
TempUnschedulableReason: `{"status_code":401,"until_unix":1735689600}`,
Credentials: map[string]any{
"temp_unschedulable_enabled": true,
"temp_unschedulable_rules": []any{
map[string]any{
"error_code": float64(401),
"keywords": []any{"unauthorized"},
"duration_minutes": float64(10),
},
},
},
},
statusCode: 401,
body: []byte(`unauthorized`),
expected: ErrorPolicyTempUnscheduled, // Antigravity 不升级,继续走规则匹配
},
{
name: "temp_unschedulable_body_miss_returns_none",
@@ -0,0 +1,239 @@
package service
import "context"
// gatewayAffinityFlow encapsulates affinity-specific scheduling steps so the
// main account selection flow can stay focused on generic scheduling.
type gatewayAffinityFlow struct {
svc *GatewayService
ctx context.Context
groupID *int64
sessionHash string
requestedModel string
affinityClientID string
affinityUserID int64
platform string
useMixed bool
accountByID map[int64]*Account
isExcluded func(int64) bool
}
type affinityWaitCandidate struct {
account *Account
}
func newGatewayAffinityFlow(
svc *GatewayService,
ctx context.Context,
groupID *int64,
sessionHash string,
requestedModel string,
affinityClientID string,
affinityUserID int64,
platform string,
useMixed bool,
accountByID map[int64]*Account,
isExcluded func(int64) bool,
) *gatewayAffinityFlow {
return &gatewayAffinityFlow{
svc: svc,
ctx: ctx,
groupID: groupID,
sessionHash: sessionHash,
requestedModel: requestedModel,
affinityClientID: affinityClientID,
affinityUserID: affinityUserID,
platform: platform,
useMixed: useMixed,
accountByID: accountByID,
isExcluded: isExcluded,
}
}
// shouldFilterAccountWithoutClientID excludes affinity-enabled Anthropic accounts
// when metadata.user_id does not provide a usable client_id.
func shouldFilterAccountWithoutClientID(account *Account, affinityClientID string) bool {
if account == nil || affinityClientID != "" {
return false
}
if account.Platform != PlatformAnthropic {
return false
}
return account.IsAffinityEnabled()
}
func filterAccountsWithoutClientID(accounts []Account, affinityClientID string) []Account {
if affinityClientID != "" {
return accounts
}
filtered := make([]Account, 0, len(accounts))
for _, acc := range accounts {
if shouldFilterAccountWithoutClientID(&acc, affinityClientID) {
continue
}
filtered = append(filtered, acc)
}
return filtered
}
func (f *gatewayAffinityFlow) preprocessPinnedUsers(accounts []Account) {
if f.affinityUserID <= 0 || f.affinityClientID == "" || f.svc.cache == nil {
return
}
for i := range accounts {
if accounts[i].IsPinnedUser(f.affinityUserID) && accounts[i].IsAffinityEnabled() {
_ = f.svc.cache.UpdateAffinity(
f.ctx,
derefGroupID(f.groupID),
f.affinityUserID,
f.affinityClientID,
accounts[i].ID,
ClientAffinityTTL,
)
}
}
}
// trySelectAffinityAccount runs Layer 1.4 and returns:
// - result != nil: affinity path selected an account or wait plan
// - affinityHit == true: an effective affinity-enabled record was considered and should suppress sticky fallback
func (f *gatewayAffinityFlow) trySelectAffinityAccount() (*AccountSelectionResult, bool, error) {
if f.affinityClientID == "" || f.affinityUserID <= 0 || f.svc.cache == nil {
return nil, false, nil
}
gid := derefGroupID(f.groupID)
affinityAccountIDs, err := f.svc.cache.GetAffinityAccounts(f.ctx, gid, f.affinityUserID, f.affinityClientID, ClientAffinityTTL)
if err != nil || len(affinityAccountIDs) == 0 {
return nil, false, nil
}
noSwitchBlocked := false
anyAllowSwitch := false
effectiveAffinityHit := false
waitCandidates := make(map[int64]*affinityWaitCandidate)
for _, affinityAccID := range affinityAccountIDs {
account, ok := f.accountByID[affinityAccID]
if f.isExcluded != nil && f.isExcluded(affinityAccID) {
checkAcc := account
if !ok && f.svc.accountRepo != nil {
if acc, repoErr := f.svc.accountRepo.GetByID(f.ctx, affinityAccID); repoErr == nil && acc != nil {
checkAcc = acc
}
}
if checkAcc != nil && checkAcc.IsAffinityEnabled() {
effectiveAffinityHit = true
if !checkAcc.IsAffinityAllowSwitch() {
noSwitchBlocked = true
} else {
anyAllowSwitch = true
}
}
continue
}
if !ok || !f.svc.isAccountSchedulableForSelection(account) {
checkAcc := account
if !ok && f.svc.accountRepo != nil {
if acc, repoErr := f.svc.accountRepo.GetByID(f.ctx, affinityAccID); repoErr == nil && acc != nil {
checkAcc = acc
}
}
if checkAcc != nil && checkAcc.IsAffinityEnabled() {
effectiveAffinityHit = true
if !checkAcc.IsAffinityAllowSwitch() {
noSwitchBlocked = true
} else {
anyAllowSwitch = true
}
}
continue
}
if !account.IsAffinityEnabled() {
continue
}
effectiveAffinityHit = true
if account.IsAffinityAllowSwitch() {
anyAllowSwitch = true
}
if !f.svc.isAccountAllowedForPlatform(account, f.platform, f.useMixed) ||
(f.requestedModel != "" && !f.svc.isModelSupportedByAccountWithContext(f.ctx, account, f.requestedModel)) ||
!f.svc.isAccountSchedulableForModelSelection(f.ctx, account, f.requestedModel) ||
!f.svc.isAccountSchedulableForQuota(account) ||
!f.svc.isAccountSchedulableForWindowCost(f.ctx, account, false) ||
!f.svc.isAccountSchedulableForRPM(f.ctx, account, false) {
if !account.IsAffinityAllowSwitch() {
noSwitchBlocked = true
}
continue
}
userCount, clientCount, perUserCount, multiErr := f.svc.cache.GetAffinityMultiCount(
f.ctx, gid, affinityAccID, f.affinityUserID, ClientAffinityTTL,
)
if multiErr == nil {
zone := account.GetMultiDimAffinityZone(userCount, clientCount, perUserCount)
if zone == AffinityZoneRed {
if !account.IsAffinityAllowSwitch() {
noSwitchBlocked = true
}
continue
}
}
result, acquireErr := f.svc.tryAcquireAccountSlot(f.ctx, affinityAccID, account.Concurrency)
if acquireErr == nil && result.Acquired {
if !f.svc.checkAndRegisterSession(f.ctx, account, f.sessionHash) {
result.ReleaseFunc()
continue
}
_ = f.svc.cache.UpdateAffinity(f.ctx, gid, f.affinityUserID, f.affinityClientID, affinityAccID, ClientAffinityTTL)
if f.sessionHash != "" {
_ = f.svc.cache.SetSessionAccountID(f.ctx, gid, f.sessionHash, affinityAccID, stickySessionTTL)
}
return &AccountSelectionResult{
Account: account,
Acquired: true,
ReleaseFunc: result.ReleaseFunc,
}, true, nil
}
if acquireErr == nil && !result.Acquired && !account.IsAffinityAllowSwitch() {
noSwitchBlocked = true
waitCandidates[affinityAccID] = &affinityWaitCandidate{account: account}
}
}
if noSwitchBlocked && !anyAllowSwitch && f.svc.concurrencyService != nil {
for _, waitAccID := range affinityAccountIDs {
candidate, ok := waitCandidates[waitAccID]
if !ok || candidate == nil || candidate.account == nil {
continue
}
acc := candidate.account
waitingCount, _ := f.svc.concurrencyService.GetAccountWaitingCount(f.ctx, waitAccID)
if waitingCount >= f.svc.schedulingConfig().StickySessionMaxWaiting {
continue
}
if !f.svc.checkAndRegisterSession(f.ctx, acc, f.sessionHash) {
continue
}
cfg := f.svc.schedulingConfig()
return &AccountSelectionResult{
Account: acc,
WaitPlan: &AccountWaitPlan{
AccountID: waitAccID,
MaxConcurrency: acc.Concurrency,
Timeout: cfg.StickySessionWaitTimeout,
MaxWaiting: cfg.StickySessionMaxWaiting,
},
}, true, nil
}
return nil, effectiveAffinityHit, ErrAffinityNoSwitch
}
return nil, effectiveAffinityHit, nil
}
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,196 @@
package service
import (
"context"
"encoding/json"
"github.com/Wei-Shaw/sub2api/internal/pkg/antigravity"
"github.com/gin-gonic/gin"
"github.com/tidwall/sjson"
)
type claudeMaxResponseRewriteContext struct {
Parsed *ParsedRequest
Group *Group
}
type claudeMaxResponseRewriteContextKeyType struct{}
var claudeMaxResponseRewriteContextKey = claudeMaxResponseRewriteContextKeyType{}
func withClaudeMaxResponseRewriteContext(ctx context.Context, c *gin.Context, parsed *ParsedRequest) context.Context {
if ctx == nil {
ctx = context.Background()
}
value := claudeMaxResponseRewriteContext{
Parsed: parsed,
Group: claudeMaxGroupFromGinContext(c),
}
return context.WithValue(ctx, claudeMaxResponseRewriteContextKey, value)
}
func claudeMaxResponseRewriteContextFromContext(ctx context.Context) claudeMaxResponseRewriteContext {
if ctx == nil {
return claudeMaxResponseRewriteContext{}
}
value, _ := ctx.Value(claudeMaxResponseRewriteContextKey).(claudeMaxResponseRewriteContext)
return value
}
func claudeMaxGroupFromGinContext(c *gin.Context) *Group {
if c == nil {
return nil
}
raw, exists := c.Get("api_key")
if !exists {
return nil
}
apiKey, ok := raw.(*APIKey)
if !ok || apiKey == nil {
return nil
}
return apiKey.Group
}
func parsedRequestFromGinContext(c *gin.Context) *ParsedRequest {
if c == nil {
return nil
}
raw, exists := c.Get("parsed_request")
if !exists {
return nil
}
parsed, _ := raw.(*ParsedRequest)
return parsed
}
func applyClaudeMaxSimulationToUsage(ctx context.Context, usage *ClaudeUsage, model string, accountID int64) claudeMaxCacheBillingOutcome {
var out claudeMaxCacheBillingOutcome
if usage == nil {
return out
}
rewriteCtx := claudeMaxResponseRewriteContextFromContext(ctx)
return applyClaudeMaxCacheBillingPolicyToUsage(usage, rewriteCtx.Parsed, rewriteCtx.Group, model, accountID)
}
func applyClaudeMaxSimulationToUsageJSONMap(ctx context.Context, usageObj map[string]any, model string, accountID int64) claudeMaxCacheBillingOutcome {
var out claudeMaxCacheBillingOutcome
if usageObj == nil {
return out
}
usage := claudeUsageFromJSONMap(usageObj)
out = applyClaudeMaxSimulationToUsage(ctx, &usage, model, accountID)
if out.Simulated {
rewriteClaudeUsageJSONMap(usageObj, usage)
}
return out
}
func rewriteClaudeUsageJSONBytes(body []byte, usage ClaudeUsage) []byte {
updated := body
var err error
updated, err = sjson.SetBytes(updated, "usage.input_tokens", usage.InputTokens)
if err != nil {
return body
}
updated, err = sjson.SetBytes(updated, "usage.cache_creation_input_tokens", usage.CacheCreationInputTokens)
if err != nil {
return body
}
updated, err = sjson.SetBytes(updated, "usage.cache_creation.ephemeral_5m_input_tokens", usage.CacheCreation5mTokens)
if err != nil {
return body
}
updated, err = sjson.SetBytes(updated, "usage.cache_creation.ephemeral_1h_input_tokens", usage.CacheCreation1hTokens)
if err != nil {
return body
}
return updated
}
func claudeUsageFromJSONMap(usageObj map[string]any) ClaudeUsage {
var usage ClaudeUsage
if usageObj == nil {
return usage
}
usage.InputTokens = usageIntFromAny(usageObj["input_tokens"])
usage.OutputTokens = usageIntFromAny(usageObj["output_tokens"])
usage.CacheCreationInputTokens = usageIntFromAny(usageObj["cache_creation_input_tokens"])
usage.CacheReadInputTokens = usageIntFromAny(usageObj["cache_read_input_tokens"])
if ccObj, ok := usageObj["cache_creation"].(map[string]any); ok {
usage.CacheCreation5mTokens = usageIntFromAny(ccObj["ephemeral_5m_input_tokens"])
usage.CacheCreation1hTokens = usageIntFromAny(ccObj["ephemeral_1h_input_tokens"])
}
return usage
}
func rewriteClaudeUsageJSONMap(usageObj map[string]any, usage ClaudeUsage) {
if usageObj == nil {
return
}
usageObj["input_tokens"] = usage.InputTokens
usageObj["cache_creation_input_tokens"] = usage.CacheCreationInputTokens
ccObj, _ := usageObj["cache_creation"].(map[string]any)
if ccObj == nil {
ccObj = make(map[string]any, 2)
usageObj["cache_creation"] = ccObj
}
ccObj["ephemeral_5m_input_tokens"] = usage.CacheCreation5mTokens
ccObj["ephemeral_1h_input_tokens"] = usage.CacheCreation1hTokens
}
func usageIntFromAny(v any) int {
switch value := v.(type) {
case int:
return value
case int64:
return int(value)
case float64:
return int(value)
case json.Number:
if n, err := value.Int64(); err == nil {
return int(n)
}
}
return 0
}
// setupClaudeMaxStreamingHook 为 Antigravity 流式路径设置 SSE usage 改写 hook。
func setupClaudeMaxStreamingHook(c *gin.Context, processor *antigravity.StreamingProcessor, originalModel string, accountID int64) {
group := claudeMaxGroupFromGinContext(c)
parsed := parsedRequestFromGinContext(c)
if !shouldApplyClaudeMaxBillingRulesForUsage(group, originalModel, parsed) {
return
}
processor.SetUsageMapHook(func(usageMap map[string]any) {
svcUsage := claudeUsageFromJSONMap(usageMap)
outcome := applyClaudeMaxCacheBillingPolicyToUsage(&svcUsage, parsed, group, originalModel, accountID)
if outcome.Simulated {
rewriteClaudeUsageJSONMap(usageMap, svcUsage)
}
})
}
// applyClaudeMaxNonStreamingRewrite 为 Antigravity 非流式路径改写响应体中的 usage。
func applyClaudeMaxNonStreamingRewrite(c *gin.Context, claudeResp []byte, agUsage *antigravity.ClaudeUsage, originalModel string, accountID int64) []byte {
group := claudeMaxGroupFromGinContext(c)
parsed := parsedRequestFromGinContext(c)
if !shouldApplyClaudeMaxBillingRulesForUsage(group, originalModel, parsed) {
return claudeResp
}
svcUsage := &ClaudeUsage{
InputTokens: agUsage.InputTokens,
OutputTokens: agUsage.OutputTokens,
CacheCreationInputTokens: agUsage.CacheCreationInputTokens,
CacheReadInputTokens: agUsage.CacheReadInputTokens,
}
outcome := applyClaudeMaxCacheBillingPolicyToUsage(svcUsage, parsed, group, originalModel, accountID)
if outcome.Simulated {
return rewriteClaudeUsageJSONBytes(claudeResp, *svcUsage)
}
return claudeResp
}
@@ -143,6 +143,27 @@ func (s *stickyGatewayCacheHotpathStub) RefreshSessionTTL(ctx context.Context, g
func (s *stickyGatewayCacheHotpathStub) DeleteSessionAccountID(ctx context.Context, groupID int64, sessionHash string) error {
return nil
}
func (s *stickyGatewayCacheHotpathStub) GetAffinityAccounts(_ context.Context, _ int64, _ int64, _ string, _ time.Duration) ([]int64, error) {
return nil, nil
}
func (s *stickyGatewayCacheHotpathStub) UpdateAffinity(_ context.Context, _ int64, _ int64, _ string, _ int64, _ time.Duration) error {
return nil
}
func (s *stickyGatewayCacheHotpathStub) GetAffinityMultiCount(_ context.Context, _ int64, _ int64, _ int64, _ time.Duration) (int64, int64, int64, error) {
return 0, 0, 0, nil
}
func (s *stickyGatewayCacheHotpathStub) GetAccountAffinityCountBatch(_ context.Context, _ int64, _ []int64, _ time.Duration) (map[int64]int64, error) {
return map[int64]int64{}, nil
}
func (s *stickyGatewayCacheHotpathStub) GetAccountAffinityClientsBatch(_ context.Context, _ map[int64][]int64, _ time.Duration) (map[int64][]string, error) {
return map[int64][]string{}, nil
}
func (s *stickyGatewayCacheHotpathStub) GetAccountAffinityClientsWithScores(_ context.Context, _ int64, _ []int64, _ time.Duration) ([]AffinityClient, error) {
return nil, nil
}
func (s *stickyGatewayCacheHotpathStub) ClearAccountAffinity(_ context.Context, _ int64, _ []int64) error {
return nil
}
func (s *modelsListAccountRepoStub) ListSchedulableByGroupID(ctx context.Context, groupID int64) ([]Account, error) {
s.listByGroupCalls.Add(1)
@@ -732,7 +753,7 @@ func TestSelectAccountWithLoadAwareness_StickyReadReuse(t *testing.T) {
modelsListCacheTTL: time.Minute,
}
result, err := svc.SelectAccountWithLoadAwareness(baseCtx, nil, "sess-hash", "", nil, "")
result, err := svc.SelectAccountWithLoadAwareness(baseCtx, nil, "sess-hash", "", nil, "", int64(0))
require.NoError(t, err)
require.NotNil(t, result)
require.NotNil(t, result.Account)
@@ -754,7 +775,7 @@ func TestSelectAccountWithLoadAwareness_StickyReadReuse(t *testing.T) {
ctx := context.WithValue(baseCtx, ctxkey.PrefetchedStickyAccountID, account.ID)
ctx = context.WithValue(ctx, ctxkey.PrefetchedStickyGroupID, int64(0))
result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "sess-hash", "", nil, "")
result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "sess-hash", "", nil, "", int64(0))
require.NoError(t, err)
require.NotNil(t, result)
require.NotNil(t, result.Account)
@@ -776,7 +797,7 @@ func TestSelectAccountWithLoadAwareness_StickyReadReuse(t *testing.T) {
ctx := context.WithValue(baseCtx, ctxkey.PrefetchedStickyAccountID, int64(999))
ctx = context.WithValue(ctx, ctxkey.PrefetchedStickyGroupID, int64(77))
result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "sess-hash", "", nil, "")
result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "sess-hash", "", nil, "", int64(0))
require.NoError(t, err)
require.NotNil(t, result)
require.NotNil(t, result.Account)
@@ -235,6 +235,28 @@ func (m *mockGatewayCacheForPlatform) DeleteSessionAccountID(ctx context.Context
return nil
}
func (m *mockGatewayCacheForPlatform) GetAffinityAccounts(_ context.Context, _ int64, _ int64, _ string, _ time.Duration) ([]int64, error) {
return nil, nil
}
func (m *mockGatewayCacheForPlatform) UpdateAffinity(_ context.Context, _ int64, _ int64, _ string, _ int64, _ time.Duration) error {
return nil
}
func (m *mockGatewayCacheForPlatform) GetAffinityMultiCount(_ context.Context, _ int64, _ int64, _ int64, _ time.Duration) (int64, int64, int64, error) {
return 0, 0, 0, nil
}
func (m *mockGatewayCacheForPlatform) GetAccountAffinityCountBatch(_ context.Context, _ int64, _ []int64, _ time.Duration) (map[int64]int64, error) {
return map[int64]int64{}, nil
}
func (m *mockGatewayCacheForPlatform) GetAccountAffinityClientsBatch(_ context.Context, _ map[int64][]int64, _ time.Duration) (map[int64][]string, error) {
return map[int64][]string{}, nil
}
func (m *mockGatewayCacheForPlatform) GetAccountAffinityClientsWithScores(_ context.Context, _ int64, _ []int64, _ time.Duration) ([]AffinityClient, error) {
return nil, nil
}
func (m *mockGatewayCacheForPlatform) ClearAccountAffinity(_ context.Context, _ int64, _ []int64) error {
return nil
}
type mockGroupRepoForGateway struct {
groups map[int64]*Group
getByIDCalls int
@@ -2031,7 +2053,7 @@ func TestGatewayService_SelectAccountWithLoadAwareness(t *testing.T) {
concurrencyService: nil, // No concurrency service
}
result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", nil, "")
result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", nil, "", int64(0))
require.NoError(t, err)
require.NotNil(t, result)
require.NotNil(t, result.Account)
@@ -2084,7 +2106,7 @@ func TestGatewayService_SelectAccountWithLoadAwareness(t *testing.T) {
concurrencyService: nil, // legacy path
}
result, err := svc.SelectAccountWithLoadAwareness(ctx, &groupID, sessionHash, "claude-b", nil, "")
result, err := svc.SelectAccountWithLoadAwareness(ctx, &groupID, sessionHash, "claude-b", nil, "", int64(0))
require.NoError(t, err)
require.NotNil(t, result)
require.NotNil(t, result.Account)
@@ -2116,7 +2138,7 @@ func TestGatewayService_SelectAccountWithLoadAwareness(t *testing.T) {
concurrencyService: nil,
}
result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", nil, "")
result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", nil, "", int64(0))
require.NoError(t, err)
require.NotNil(t, result)
require.NotNil(t, result.Account)
@@ -2148,13 +2170,314 @@ func TestGatewayService_SelectAccountWithLoadAwareness(t *testing.T) {
}
excludedIDs := map[int64]struct{}{1: {}}
result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", excludedIDs, "")
result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", excludedIDs, "", int64(0))
require.NoError(t, err)
require.NotNil(t, result)
require.NotNil(t, result.Account)
require.Equal(t, int64(2), result.Account.ID, "不应选择被排除的账号")
})
t.Run("无客户端ID时过滤Anthropic OAuth亲和账号-load-aware路径", func(t *testing.T) {
repo := &mockAccountRepoForPlatform{
accounts: []Account{
{
ID: 1,
Platform: PlatformAnthropic,
Type: AccountTypeOAuth,
Priority: 1,
Status: StatusActive,
Schedulable: true,
Concurrency: 5,
Extra: map[string]any{
"client_affinity_enabled": true,
},
},
{
ID: 2,
Platform: PlatformAnthropic,
Type: AccountTypeOAuth,
Priority: 2,
Status: StatusActive,
Schedulable: true,
Concurrency: 5,
},
},
accountsByID: map[int64]*Account{},
}
for i := range repo.accounts {
repo.accountsByID[repo.accounts[i].ID] = &repo.accounts[i]
}
cfg := testConfig()
cfg.Gateway.Scheduling.LoadBatchEnabled = true
svc := &GatewayService{
accountRepo: repo,
cache: &mockGatewayCacheForPlatform{},
cfg: cfg,
concurrencyService: NewConcurrencyService(&mockConcurrencyCache{}),
}
result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", nil, "", 0)
require.NoError(t, err)
require.NotNil(t, result)
require.NotNil(t, result.Account)
require.Equal(t, int64(2), result.Account.ID, "无客户端ID时应过滤亲和OAuth账号")
})
t.Run("无客户端ID且亲和关闭时-load-aware路径恢复旧逻辑", func(t *testing.T) {
repo := &mockAccountRepoForPlatform{
accounts: []Account{
{
ID: 1,
Platform: PlatformAnthropic,
Type: AccountTypeOAuth,
Priority: 1,
Status: StatusActive,
Schedulable: true,
Concurrency: 5,
Extra: map[string]any{
"affinity_enabled": false,
},
},
{
ID: 2,
Platform: PlatformAnthropic,
Type: AccountTypeOAuth,
Priority: 2,
Status: StatusActive,
Schedulable: true,
Concurrency: 5,
},
},
accountsByID: map[int64]*Account{},
}
for i := range repo.accounts {
repo.accountsByID[repo.accounts[i].ID] = &repo.accounts[i]
}
cfg := testConfig()
cfg.Gateway.Scheduling.LoadBatchEnabled = true
svc := &GatewayService{
accountRepo: repo,
cache: &mockGatewayCacheForPlatform{},
cfg: cfg,
concurrencyService: NewConcurrencyService(&mockConcurrencyCache{}),
}
result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", nil, "", 0)
require.NoError(t, err)
require.NotNil(t, result)
require.NotNil(t, result.Account)
require.Equal(t, int64(1), result.Account.ID, "关闭亲和后不应再被无clientID过滤")
})
t.Run("有客户端ID时不过滤Anthropic OAuth亲和账号", func(t *testing.T) {
repo := &mockAccountRepoForPlatform{
accounts: []Account{
{
ID: 1,
Platform: PlatformAnthropic,
Type: AccountTypeOAuth,
Priority: 1,
Status: StatusActive,
Schedulable: true,
Concurrency: 5,
Extra: map[string]any{
"client_affinity_enabled": true,
},
},
{
ID: 2,
Platform: PlatformAnthropic,
Type: AccountTypeOAuth,
Priority: 2,
Status: StatusActive,
Schedulable: true,
Concurrency: 5,
},
},
accountsByID: map[int64]*Account{},
}
for i := range repo.accounts {
repo.accountsByID[repo.accounts[i].ID] = &repo.accounts[i]
}
cfg := testConfig()
cfg.Gateway.Scheduling.LoadBatchEnabled = true
svc := &GatewayService{
accountRepo: repo,
cache: &mockGatewayCacheForPlatform{},
cfg: cfg,
concurrencyService: NewConcurrencyService(&mockConcurrencyCache{}),
}
metadataUserID := "user_0123456789abcdef0123456789abcdef0123456789abcdef0123456789abcdef_account_test_session_test"
result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", nil, metadataUserID, 123)
require.NoError(t, err)
require.NotNil(t, result)
require.NotNil(t, result.Account)
require.Equal(t, int64(1), result.Account.ID, "有客户端ID时应允许亲和OAuth账号参与调度")
})
t.Run("过滤对Anthropic全类型生效-SetupToken也受影响", func(t *testing.T) {
repo := &mockAccountRepoForPlatform{
accounts: []Account{
{
ID: 1,
Platform: PlatformAnthropic,
Type: AccountTypeSetupToken,
Priority: 1,
Status: StatusActive,
Schedulable: true,
Concurrency: 5,
Extra: map[string]any{
"client_affinity_enabled": true,
},
},
{
ID: 2,
Platform: PlatformAnthropic,
Type: AccountTypeOAuth,
Priority: 2,
Status: StatusActive,
Schedulable: true,
Concurrency: 5,
Extra: map[string]any{
"client_affinity_enabled": true,
},
},
{
ID: 3,
Platform: PlatformAnthropic,
Type: AccountTypeOAuth,
Priority: 3,
Status: StatusActive,
Schedulable: true,
Concurrency: 5,
},
},
accountsByID: map[int64]*Account{},
}
for i := range repo.accounts {
repo.accountsByID[repo.accounts[i].ID] = &repo.accounts[i]
}
cfg := testConfig()
cfg.Gateway.Scheduling.LoadBatchEnabled = true
svc := &GatewayService{
accountRepo: repo,
cache: &mockGatewayCacheForPlatform{},
cfg: cfg,
concurrencyService: NewConcurrencyService(&mockConcurrencyCache{}),
}
result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", nil, "", 0)
require.NoError(t, err)
require.NotNil(t, result)
require.NotNil(t, result.Account)
require.Equal(t, int64(3), result.Account.ID, "无客户端ID时,Anthropic SetupToken/OAuth(开启亲和)都应被过滤")
})
t.Run("无客户端ID时过滤Anthropic OAuth亲和账号-legacy路径", func(t *testing.T) {
repo := &mockAccountRepoForPlatform{
accounts: []Account{
{
ID: 1,
Platform: PlatformAnthropic,
Type: AccountTypeOAuth,
Priority: 1,
Status: StatusActive,
Schedulable: true,
Concurrency: 5,
Extra: map[string]any{
"client_affinity_enabled": true,
},
},
{
ID: 2,
Platform: PlatformAnthropic,
Type: AccountTypeOAuth,
Priority: 2,
Status: StatusActive,
Schedulable: true,
Concurrency: 5,
},
},
accountsByID: map[int64]*Account{},
}
for i := range repo.accounts {
repo.accountsByID[repo.accounts[i].ID] = &repo.accounts[i]
}
cfg := testConfig()
cfg.Gateway.Scheduling.LoadBatchEnabled = false
svc := &GatewayService{
accountRepo: repo,
cache: &mockGatewayCacheForPlatform{},
cfg: cfg,
concurrencyService: nil,
}
result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", nil, "", 0)
require.NoError(t, err)
require.NotNil(t, result)
require.NotNil(t, result.Account)
require.Equal(t, int64(2), result.Account.ID, "legacy路径也应过滤亲和OAuth账号")
})
t.Run("无客户端ID且亲和关闭时-legacy路径恢复旧逻辑", func(t *testing.T) {
repo := &mockAccountRepoForPlatform{
accounts: []Account{
{
ID: 1,
Platform: PlatformAnthropic,
Type: AccountTypeOAuth,
Priority: 1,
Status: StatusActive,
Schedulable: true,
Concurrency: 5,
Extra: map[string]any{
"client_affinity_enabled": false,
},
},
{
ID: 2,
Platform: PlatformAnthropic,
Type: AccountTypeOAuth,
Priority: 2,
Status: StatusActive,
Schedulable: true,
Concurrency: 5,
},
},
accountsByID: map[int64]*Account{},
}
for i := range repo.accounts {
repo.accountsByID[repo.accounts[i].ID] = &repo.accounts[i]
}
cfg := testConfig()
cfg.Gateway.Scheduling.LoadBatchEnabled = false
svc := &GatewayService{
accountRepo: repo,
cache: &mockGatewayCacheForPlatform{},
cfg: cfg,
concurrencyService: nil,
}
result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", nil, "", 0)
require.NoError(t, err)
require.NotNil(t, result)
require.NotNil(t, result.Account)
require.Equal(t, int64(1), result.Account.ID, "关闭亲和后 legacy 路径不应再被无clientID过滤")
})
t.Run("粘性命中-不调用GetByID", func(t *testing.T) {
repo := &mockAccountRepoForPlatform{
accounts: []Account{
@@ -2182,7 +2505,7 @@ func TestGatewayService_SelectAccountWithLoadAwareness(t *testing.T) {
concurrencyService: NewConcurrencyService(concurrencyCache),
}
result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "sticky", "claude-3-5-sonnet-20241022", nil, "")
result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "sticky", "claude-3-5-sonnet-20241022", nil, "", int64(0))
require.NoError(t, err)
require.NotNil(t, result)
require.NotNil(t, result.Account)
@@ -2218,7 +2541,7 @@ func TestGatewayService_SelectAccountWithLoadAwareness(t *testing.T) {
concurrencyService: NewConcurrencyService(concurrencyCache),
}
result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "sticky", "claude-3-5-sonnet-20241022", nil, "")
result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "sticky", "claude-3-5-sonnet-20241022", nil, "", int64(0))
require.NoError(t, err)
require.NotNil(t, result)
require.NotNil(t, result.Account)
@@ -2259,7 +2582,7 @@ func TestGatewayService_SelectAccountWithLoadAwareness(t *testing.T) {
concurrencyService: NewConcurrencyService(concurrencyCache),
}
result, err := svc.SelectAccountWithLoadAwareness(testCtx, nil, "sticky", "claude-3-5-sonnet-20241022", nil, "")
result, err := svc.SelectAccountWithLoadAwareness(testCtx, nil, "sticky", "claude-3-5-sonnet-20241022", nil, "", int64(0))
require.NoError(t, err)
require.NotNil(t, result)
require.NotNil(t, result.Account)
@@ -2287,7 +2610,7 @@ func TestGatewayService_SelectAccountWithLoadAwareness(t *testing.T) {
concurrencyService: nil,
}
result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", nil, "")
result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", nil, "", int64(0))
require.Error(t, err)
require.Nil(t, result)
require.ErrorIs(t, err, ErrNoAvailableAccounts)
@@ -2319,7 +2642,7 @@ func TestGatewayService_SelectAccountWithLoadAwareness(t *testing.T) {
concurrencyService: nil,
}
result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", nil, "")
result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", nil, "", int64(0))
require.NoError(t, err)
require.NotNil(t, result)
require.NotNil(t, result.Account)
@@ -2352,7 +2675,7 @@ func TestGatewayService_SelectAccountWithLoadAwareness(t *testing.T) {
concurrencyService: nil,
}
result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", nil, "")
result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", nil, "", int64(0))
require.NoError(t, err)
require.NotNil(t, result)
require.NotNil(t, result.Account)
@@ -2390,7 +2713,7 @@ func TestGatewayService_SelectAccountWithLoadAwareness(t *testing.T) {
concurrencyService: NewConcurrencyService(concurrencyCache),
}
result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "sticky", "claude-3-5-sonnet-20241022", nil, "")
result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "sticky", "claude-3-5-sonnet-20241022", nil, "", int64(0))
require.NoError(t, err)
require.NotNil(t, result)
require.NotNil(t, result.WaitPlan)
@@ -2426,7 +2749,7 @@ func TestGatewayService_SelectAccountWithLoadAwareness(t *testing.T) {
concurrencyService: NewConcurrencyService(concurrencyCache),
}
result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "legacy", "claude-3-5-sonnet-20241022", nil, "")
result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "legacy", "claude-3-5-sonnet-20241022", nil, "", int64(0))
require.NoError(t, err)
require.NotNil(t, result)
require.NotNil(t, result.Account)
@@ -2485,7 +2808,7 @@ func TestGatewayService_SelectAccountWithLoadAwareness(t *testing.T) {
concurrencyService: NewConcurrencyService(concurrencyCache),
}
result, err := svc.SelectAccountWithLoadAwareness(ctx, &groupID, sessionHash, "claude-3-5-sonnet-20241022", nil, "")
result, err := svc.SelectAccountWithLoadAwareness(ctx, &groupID, sessionHash, "claude-3-5-sonnet-20241022", nil, "", int64(0))
require.NoError(t, err)
require.NotNil(t, result)
require.NotNil(t, result.WaitPlan)
@@ -2539,7 +2862,7 @@ func TestGatewayService_SelectAccountWithLoadAwareness(t *testing.T) {
concurrencyService: NewConcurrencyService(concurrencyCache),
}
result, err := svc.SelectAccountWithLoadAwareness(ctx, &groupID, sessionHash, "claude-3-5-sonnet-20241022", nil, "")
result, err := svc.SelectAccountWithLoadAwareness(ctx, &groupID, sessionHash, "claude-3-5-sonnet-20241022", nil, "", int64(0))
require.NoError(t, err)
require.NotNil(t, result)
require.NotNil(t, result.Account)
@@ -2593,7 +2916,7 @@ func TestGatewayService_SelectAccountWithLoadAwareness(t *testing.T) {
concurrencyService: NewConcurrencyService(concurrencyCache),
}
result, err := svc.SelectAccountWithLoadAwareness(ctx, &groupID, sessionHash, "claude-3-5-sonnet-20241022", nil, "")
result, err := svc.SelectAccountWithLoadAwareness(ctx, &groupID, sessionHash, "claude-3-5-sonnet-20241022", nil, "", int64(0))
require.NoError(t, err)
require.NotNil(t, result)
require.NotNil(t, result.Account)
@@ -2651,7 +2974,7 @@ func TestGatewayService_SelectAccountWithLoadAwareness(t *testing.T) {
concurrencyService: NewConcurrencyService(concurrencyCache),
}
result, err := svc.SelectAccountWithLoadAwareness(ctx, &groupID, "route", "claude-3-5-sonnet-20241022", nil, "")
result, err := svc.SelectAccountWithLoadAwareness(ctx, &groupID, "route", "claude-3-5-sonnet-20241022", nil, "", int64(0))
require.NoError(t, err)
require.NotNil(t, result)
require.NotNil(t, result.Account)
@@ -2709,7 +3032,7 @@ func TestGatewayService_SelectAccountWithLoadAwareness(t *testing.T) {
concurrencyService: NewConcurrencyService(concurrencyCache),
}
result, err := svc.SelectAccountWithLoadAwareness(ctx, &groupID, "route-full", "claude-3-5-sonnet-20241022", nil, "")
result, err := svc.SelectAccountWithLoadAwareness(ctx, &groupID, "route-full", "claude-3-5-sonnet-20241022", nil, "", int64(0))
require.NoError(t, err)
require.NotNil(t, result)
require.NotNil(t, result.WaitPlan)
@@ -2767,7 +3090,7 @@ func TestGatewayService_SelectAccountWithLoadAwareness(t *testing.T) {
concurrencyService: NewConcurrencyService(concurrencyCache),
}
result, err := svc.SelectAccountWithLoadAwareness(ctx, &groupID, "fallback", "claude-3-5-sonnet-20241022", nil, "")
result, err := svc.SelectAccountWithLoadAwareness(ctx, &groupID, "fallback", "claude-3-5-sonnet-20241022", nil, "", int64(0))
require.NoError(t, err)
require.NotNil(t, result)
require.NotNil(t, result.Account)
@@ -2804,7 +3127,7 @@ func TestGatewayService_SelectAccountWithLoadAwareness(t *testing.T) {
concurrencyService: NewConcurrencyService(concurrencyCache),
}
result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", nil, "")
result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "", "claude-3-5-sonnet-20241022", nil, "", int64(0))
require.NoError(t, err)
require.NotNil(t, result)
require.NotNil(t, result.WaitPlan)
@@ -2856,7 +3179,7 @@ func TestGatewayService_SelectAccountWithLoadAwareness(t *testing.T) {
concurrencyService: NewConcurrencyService(concurrencyCache),
}
result, err := svc.SelectAccountWithLoadAwareness(ctx, &groupID, "gemini", "gemini-2.5-pro", nil, "")
result, err := svc.SelectAccountWithLoadAwareness(ctx, &groupID, "gemini", "gemini-2.5-pro", nil, "", int64(0))
require.NoError(t, err)
require.NotNil(t, result)
require.NotNil(t, result.Account)
@@ -2934,7 +3257,7 @@ func TestGatewayService_SelectAccountWithLoadAwareness(t *testing.T) {
}
excluded := map[int64]struct{}{1: {}}
result, err := svc.SelectAccountWithLoadAwareness(ctx, &groupID, "", "claude-3-5-sonnet-20241022", excluded, "")
result, err := svc.SelectAccountWithLoadAwareness(ctx, &groupID, "", "claude-3-5-sonnet-20241022", excluded, "", int64(0))
require.NoError(t, err)
require.NotNil(t, result)
require.NotNil(t, result.Account)
@@ -2988,7 +3311,7 @@ func TestGatewayService_SelectAccountWithLoadAwareness(t *testing.T) {
concurrencyService: nil,
}
result, err := svc.SelectAccountWithLoadAwareness(ctx, &groupID, "", "gemini-2.5-pro", nil, "")
result, err := svc.SelectAccountWithLoadAwareness(ctx, &groupID, "", "gemini-2.5-pro", nil, "", int64(0))
require.NoError(t, err)
require.NotNil(t, result)
require.NotNil(t, result.Account)
@@ -3021,7 +3344,7 @@ func TestGatewayService_SelectAccountWithLoadAwareness(t *testing.T) {
concurrencyService: nil,
}
result, err := svc.SelectAccountWithLoadAwareness(ctx, &groupID, "", "claude-3-5-sonnet-20241022", nil, "")
result, err := svc.SelectAccountWithLoadAwareness(ctx, &groupID, "", "claude-3-5-sonnet-20241022", nil, "", int64(0))
require.Error(t, err)
require.Nil(t, result)
require.ErrorIs(t, err, ErrClaudeCodeOnly)
@@ -3059,7 +3382,7 @@ func TestGatewayService_SelectAccountWithLoadAwareness(t *testing.T) {
concurrencyService: NewConcurrencyService(concurrencyCache),
}
result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "wait", "claude-3-5-sonnet-20241022", nil, "")
result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "wait", "claude-3-5-sonnet-20241022", nil, "", int64(0))
require.NoError(t, err)
require.NotNil(t, result)
require.NotNil(t, result.WaitPlan)
@@ -3097,7 +3420,7 @@ func TestGatewayService_SelectAccountWithLoadAwareness(t *testing.T) {
concurrencyService: NewConcurrencyService(concurrencyCache),
}
result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "missing-load", "claude-3-5-sonnet-20241022", nil, "")
result, err := svc.SelectAccountWithLoadAwareness(ctx, nil, "missing-load", "claude-3-5-sonnet-20241022", nil, "", int64(0))
require.NoError(t, err)
require.NotNil(t, result)
require.NotNil(t, result.Account)
@@ -0,0 +1,199 @@
package service
import (
"context"
"testing"
"time"
"github.com/Wei-Shaw/sub2api/internal/config"
"github.com/stretchr/testify/require"
)
type usageLogRepoRecordUsageStub struct {
UsageLogRepository
last *UsageLog
inserted bool
err error
}
func (s *usageLogRepoRecordUsageStub) Create(_ context.Context, log *UsageLog) (bool, error) {
copied := *log
s.last = &copied
return s.inserted, s.err
}
func newGatewayServiceForRecordUsageTest(repo UsageLogRepository) *GatewayService {
return &GatewayService{
usageLogRepo: repo,
billingService: NewBillingService(&config.Config{}, nil),
cfg: &config.Config{RunMode: config.RunModeSimple},
deferredService: &DeferredService{},
}
}
func TestRecordUsage_SimulateClaudeMaxEnabled_ProjectsUsageAndSkipsTTLOverride(t *testing.T) {
repo := &usageLogRepoRecordUsageStub{inserted: true}
svc := newGatewayServiceForRecordUsageTest(repo)
groupID := int64(11)
input := &RecordUsageInput{
Result: &ForwardResult{
RequestID: "req-sim-1",
Model: "claude-sonnet-4",
Duration: time.Second,
Usage: ClaudeUsage{
InputTokens: 160,
},
},
ParsedRequest: &ParsedRequest{
Model: "claude-sonnet-4",
Messages: []any{
map[string]any{
"role": "user",
"content": []any{
map[string]any{
"type": "text",
"text": "long cached context for prior turns",
"cache_control": map[string]any{"type": "ephemeral"},
},
map[string]any{
"type": "text",
"text": "please summarize the logs and provide root cause analysis",
},
},
},
},
},
APIKey: &APIKey{
ID: 1,
GroupID: &groupID,
Group: &Group{
ID: groupID,
Platform: PlatformAnthropic,
RateMultiplier: 1,
SimulateClaudeMaxEnabled: true,
},
},
User: &User{ID: 2},
Account: &Account{
ID: 3,
Platform: PlatformAnthropic,
Type: AccountTypeOAuth,
Extra: map[string]any{
"cache_ttl_override_enabled": true,
"cache_ttl_override_target": "5m",
},
},
}
err := svc.RecordUsage(context.Background(), input)
require.NoError(t, err)
require.NotNil(t, repo.last)
log := repo.last
require.Equal(t, 80, log.InputTokens)
require.Equal(t, 80, log.CacheCreationTokens)
require.Equal(t, 0, log.CacheCreation5mTokens)
require.Equal(t, 80, log.CacheCreation1hTokens)
require.False(t, log.CacheTTLOverridden, "simulate outcome should skip account ttl override")
}
func TestRecordUsage_SimulateClaudeMaxDisabled_AppliesTTLOverride(t *testing.T) {
repo := &usageLogRepoRecordUsageStub{inserted: true}
svc := newGatewayServiceForRecordUsageTest(repo)
groupID := int64(12)
input := &RecordUsageInput{
Result: &ForwardResult{
RequestID: "req-sim-2",
Model: "claude-sonnet-4",
Duration: time.Second,
Usage: ClaudeUsage{
InputTokens: 40,
CacheCreationInputTokens: 120,
CacheCreation1hTokens: 120,
},
},
APIKey: &APIKey{
ID: 2,
GroupID: &groupID,
Group: &Group{
ID: groupID,
Platform: PlatformAnthropic,
RateMultiplier: 1,
SimulateClaudeMaxEnabled: false,
},
},
User: &User{ID: 3},
Account: &Account{
ID: 4,
Platform: PlatformAnthropic,
Type: AccountTypeOAuth,
Extra: map[string]any{
"cache_ttl_override_enabled": true,
"cache_ttl_override_target": "5m",
},
},
}
err := svc.RecordUsage(context.Background(), input)
require.NoError(t, err)
require.NotNil(t, repo.last)
log := repo.last
require.Equal(t, 120, log.CacheCreationTokens)
require.Equal(t, 120, log.CacheCreation5mTokens)
require.Equal(t, 0, log.CacheCreation1hTokens)
require.True(t, log.CacheTTLOverridden)
}
func TestRecordUsage_SimulateClaudeMaxEnabled_ExistingCacheCreationBypassesSimulation(t *testing.T) {
repo := &usageLogRepoRecordUsageStub{inserted: true}
svc := newGatewayServiceForRecordUsageTest(repo)
groupID := int64(13)
input := &RecordUsageInput{
Result: &ForwardResult{
RequestID: "req-sim-3",
Model: "claude-sonnet-4",
Duration: time.Second,
Usage: ClaudeUsage{
InputTokens: 20,
CacheCreationInputTokens: 120,
CacheCreation5mTokens: 120,
},
},
APIKey: &APIKey{
ID: 3,
GroupID: &groupID,
Group: &Group{
ID: groupID,
Platform: PlatformAnthropic,
RateMultiplier: 1,
SimulateClaudeMaxEnabled: true,
},
},
User: &User{ID: 4},
Account: &Account{
ID: 5,
Platform: PlatformAnthropic,
Type: AccountTypeOAuth,
Extra: map[string]any{
"cache_ttl_override_enabled": true,
"cache_ttl_override_target": "5m",
},
},
}
err := svc.RecordUsage(context.Background(), input)
require.NoError(t, err)
require.NotNil(t, repo.last)
log := repo.last
require.Equal(t, 20, log.InputTokens)
require.Equal(t, 120, log.CacheCreation5mTokens)
require.Equal(t, 0, log.CacheCreation1hTokens)
require.Equal(t, 120, log.CacheCreationTokens)
require.False(t, log.CacheTTLOverridden, "existing cache_creation with SimulateClaudeMax enabled should skip account ttl override")
}
@@ -0,0 +1,170 @@
package service
import (
"bytes"
"context"
"encoding/json"
"net/http"
"net/http/httptest"
"testing"
"github.com/Wei-Shaw/sub2api/internal/config"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/require"
"github.com/tidwall/gjson"
)
func TestHandleNonStreamingResponse_UsageAlignedWithClaudeMaxSimulation(t *testing.T) {
gin.SetMode(gin.TestMode)
svc := &GatewayService{
cfg: &config.Config{},
rateLimitService: &RateLimitService{},
}
account := &Account{
ID: 11,
Platform: PlatformAnthropic,
Type: AccountTypeOAuth,
Extra: map[string]any{
"cache_ttl_override_enabled": true,
"cache_ttl_override_target": "5m",
},
}
group := &Group{
ID: 99,
Platform: PlatformAnthropic,
SimulateClaudeMaxEnabled: true,
}
parsed := &ParsedRequest{
Model: "claude-sonnet-4",
Messages: []any{
map[string]any{
"role": "user",
"content": []any{
map[string]any{
"type": "text",
"text": "long cached context",
"cache_control": map[string]any{"type": "ephemeral"},
},
map[string]any{
"type": "text",
"text": "new user question",
},
},
},
},
}
upstreamBody := []byte(`{"id":"msg_1","model":"claude-sonnet-4","usage":{"input_tokens":120,"output_tokens":8}}`)
resp := &http.Response{
StatusCode: http.StatusOK,
Header: http.Header{"Content-Type": []string{"application/json"}},
Body: ioNopCloserBytes(upstreamBody),
}
rec := httptest.NewRecorder()
c, _ := gin.CreateTestContext(rec)
c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages", bytes.NewReader(nil))
c.Set("api_key", &APIKey{Group: group})
requestCtx := withClaudeMaxResponseRewriteContext(context.Background(), c, parsed)
usage, err := svc.handleNonStreamingResponse(requestCtx, resp, c, account, "claude-sonnet-4", "claude-sonnet-4")
require.NoError(t, err)
require.NotNil(t, usage)
var rendered struct {
Usage ClaudeUsage `json:"usage"`
}
require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &rendered))
rendered.Usage.CacheCreation5mTokens = int(gjson.GetBytes(rec.Body.Bytes(), "usage.cache_creation.ephemeral_5m_input_tokens").Int())
rendered.Usage.CacheCreation1hTokens = int(gjson.GetBytes(rec.Body.Bytes(), "usage.cache_creation.ephemeral_1h_input_tokens").Int())
require.Equal(t, rendered.Usage.InputTokens, usage.InputTokens)
require.Equal(t, rendered.Usage.OutputTokens, usage.OutputTokens)
require.Equal(t, rendered.Usage.CacheCreationInputTokens, usage.CacheCreationInputTokens)
require.Equal(t, rendered.Usage.CacheCreation5mTokens, usage.CacheCreation5mTokens)
require.Equal(t, rendered.Usage.CacheCreation1hTokens, usage.CacheCreation1hTokens)
require.Equal(t, rendered.Usage.CacheReadInputTokens, usage.CacheReadInputTokens)
require.Greater(t, usage.CacheCreation1hTokens, 0)
require.Equal(t, 0, usage.CacheCreation5mTokens)
require.Less(t, usage.InputTokens, 120)
}
func TestHandleNonStreamingResponse_ClaudeMaxDisabled_NoSimulationIntercept(t *testing.T) {
gin.SetMode(gin.TestMode)
svc := &GatewayService{
cfg: &config.Config{},
rateLimitService: &RateLimitService{},
}
account := &Account{
ID: 12,
Platform: PlatformAnthropic,
Type: AccountTypeOAuth,
Extra: map[string]any{
"cache_ttl_override_enabled": true,
"cache_ttl_override_target": "5m",
},
}
group := &Group{
ID: 100,
Platform: PlatformAnthropic,
SimulateClaudeMaxEnabled: false,
}
parsed := &ParsedRequest{
Model: "claude-sonnet-4",
Messages: []any{
map[string]any{
"role": "user",
"content": []any{
map[string]any{
"type": "text",
"text": "long cached context",
"cache_control": map[string]any{"type": "ephemeral"},
},
map[string]any{
"type": "text",
"text": "new user question",
},
},
},
},
}
upstreamBody := []byte(`{"id":"msg_2","model":"claude-sonnet-4","usage":{"input_tokens":120,"output_tokens":8}}`)
resp := &http.Response{
StatusCode: http.StatusOK,
Header: http.Header{"Content-Type": []string{"application/json"}},
Body: ioNopCloserBytes(upstreamBody),
}
rec := httptest.NewRecorder()
c, _ := gin.CreateTestContext(rec)
c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages", bytes.NewReader(nil))
c.Set("api_key", &APIKey{Group: group})
requestCtx := withClaudeMaxResponseRewriteContext(context.Background(), c, parsed)
usage, err := svc.handleNonStreamingResponse(requestCtx, resp, c, account, "claude-sonnet-4", "claude-sonnet-4")
require.NoError(t, err)
require.NotNil(t, usage)
require.Equal(t, 120, usage.InputTokens)
require.Equal(t, 0, usage.CacheCreationInputTokens)
require.Equal(t, 0, usage.CacheCreation5mTokens)
require.Equal(t, 0, usage.CacheCreation1hTokens)
}
func ioNopCloserBytes(b []byte) *readCloserFromBytes {
return &readCloserFromBytes{Reader: bytes.NewReader(b)}
}
type readCloserFromBytes struct {
*bytes.Reader
}
func (r *readCloserFromBytes) Close() error {
return nil
}
+280 -24
View File
@@ -40,7 +40,8 @@ import (
const (
claudeAPIURL = "https://api.anthropic.com/v1/messages?beta=true"
claudeAPICountTokensURL = "https://api.anthropic.com/v1/messages/count_tokens?beta=true"
stickySessionTTL = time.Hour // 粘性会话TTL
stickySessionTTL = time.Hour // 粘性会话TTL
ClientAffinityTTL = 24 * time.Hour // 客户端亲和TTL
defaultMaxLineSize = 500 * 1024 * 1024
// Canonical Claude Code banner. Keep it EXACT (no trailing whitespace/newlines)
// to match real Claude CLI traffic as closely as possible. When we need a visual
@@ -57,14 +58,21 @@ const (
claudeMimicDebugInfoKey = "claude_mimic_debug_info"
)
const (
claudeMaxMessageOverheadTokens = 3
claudeMaxBlockOverheadTokens = 1
claudeMaxUnknownContentTokens = 4
)
// ForceCacheBillingContextKey 强制缓存计费上下文键
// 用于粘性会话切换时,将 input_tokens 转为 cache_read_input_tokens 计费
type forceCacheBillingKeyType struct{}
// accountWithLoad 账号与负载信息的组合,用于负载感知调度
type accountWithLoad struct {
account *Account
loadInfo *AccountLoadInfo
account *Account
loadInfo *AccountLoadInfo
affinityCount int64 // 亲和客户端数量(反向索引),越少越优先
}
var ForceCacheBillingContextKey = forceCacheBillingKeyType{}
@@ -328,6 +336,10 @@ var (
sseDataRe = regexp.MustCompile(`^data:\s*`)
claudeCliUserAgentRe = regexp.MustCompile(`^claude-cli/\d+\.\d+\.\d+`)
// clientIDFromMetadataRegex 从 metadata.user_id 中提取客户端 ID(64位 hex)
// 格式: user_{64位hex}_account_...
clientIDFromMetadataRegex = regexp.MustCompile(`^user_([a-f0-9]{64})_account_`)
// claudeCodePromptPrefixes 用于检测 Claude Code 系统提示词的前缀列表
// 支持多种变体:标准版、Agent SDK 版、Explore Agent 版、Compact 版等
// 注意:前缀之间不应存在包含关系,否则会导致冗余匹配
@@ -351,6 +363,12 @@ var ErrNoAvailableAccounts = errors.New("no available accounts")
// ErrClaudeCodeOnly 表示分组仅允许 Claude Code 客户端访问
var ErrClaudeCodeOnly = errors.New("this group only allows Claude Code clients")
// ErrAffinityNoSwitch 表示亲和账号不可用且不允许切换到其他账号
var ErrAffinityNoSwitch = errors.New("affinity account unavailable and switching is disabled")
// ErrAffinityLimitExceeded 表示亲和客户端限制已达上限
var ErrAffinityLimitExceeded = errors.New("affinity client limit exceeded")
// allowedHeaders 白名单headers(参考CRS项目)
var allowedHeaders = map[string]bool{
"accept": true,
@@ -391,6 +409,39 @@ type GatewayCache interface {
// DeleteSessionAccountID 删除粘性会话绑定,用于账号不可用时主动清理
// Delete sticky session binding, used to proactively clean up when account becomes unavailable
DeleteSessionAccountID(ctx context.Context, groupID int64, sessionHash string) error
// GetAffinityAccounts 获取亲和账号列表(按最近使用降序),同时清理过期成员
GetAffinityAccounts(ctx context.Context, groupID int64, userID int64, clientID string, ttl time.Duration) ([]int64, error)
// UpdateAffinity 添加/更新亲和关系(更新 score 为当前时间戳,刷新 key TTL)
UpdateAffinity(ctx context.Context, groupID int64, userID int64, clientID string, accountID int64, ttl time.Duration) error
// GetAccountAffinityCountBatch 批量获取账号的亲和成员数量(惰性清理过期成员)
GetAccountAffinityCountBatch(ctx context.Context, groupID int64, accountIDs []int64, ttl time.Duration) (map[int64]int64, error)
// GetAccountAffinityClientsBatch 批量获取每个账号跨所有分组的亲和成员列表(去重)
// accountGroups: map[accountID][]groupID
// 返回值成员格式为 {userID}/{clientID}
GetAccountAffinityClientsBatch(ctx context.Context, accountGroups map[int64][]int64, ttl time.Duration) (map[int64][]string, error)
// GetAccountAffinityClientsWithScores 获取单个账号跨所有分组的亲和客户端列表(含最后活跃时间)
GetAccountAffinityClientsWithScores(ctx context.Context, accountID int64, groupIDs []int64, ttl time.Duration) ([]AffinityClient, error)
// ClearAccountAffinity 清除指定账号在所有分组的亲和记录(正向+反向索引)
// 用于账号关闭亲和时立即清理旧绑定
ClearAccountAffinity(ctx context.Context, accountID int64, groupIDs []int64) error
// GetAffinityMultiCount 获取账号的多维度亲和计数
// 返回: uniqueUsers, uniqueClients, perUserClients
GetAffinityMultiCount(ctx context.Context, groupID int64, accountID int64, targetUserID int64, ttl time.Duration) (users, clients, perUser int64, err error)
}
// AffinityClient 亲和客户端信息(含用户 ID 和最后活跃时间)
type AffinityClient struct {
UserID int64 `json:"user_id"`
ClientID string `json:"client_id"`
LastActive time.Time `json:"last_active"`
}
// SortAffinityClients 按最后活跃时间降序排序
func SortAffinityClients(clients []AffinityClient) {
sort.Slice(clients, func(i, j int) bool {
return clients[i].LastActive.After(clients[j].LastActive)
})
}
// derefGroupID safely dereferences *int64 to int64, returning 0 if nil
@@ -461,6 +512,20 @@ func shouldClearStickySession(account *Account, requestedModel string) bool {
return false
}
// extractClientIDFromMetadata 从 metadata.user_id 中提取客户端 ID(64位 hex)。
// 格式: user_{64位hex}_account_..._session_...
// 返回空字符串表示无法提取(非 Claude Code/Console 客户端)。
func extractClientIDFromMetadata(metadataUserID string) string {
if metadataUserID == "" {
return ""
}
matches := clientIDFromMetadataRegex.FindStringSubmatch(metadataUserID)
if matches == nil {
return ""
}
return matches[1]
}
type AccountWaitPlan struct {
AccountID int64
MaxConcurrency int
@@ -1094,8 +1159,10 @@ func (s *GatewayService) SelectAccountForModelWithExclusions(ctx context.Context
}
// SelectAccountWithLoadAwareness selects account with load-awareness and wait plan.
// metadataUserID: 已废弃参数,会话限制现在统一使用 sessionHash
func (s *GatewayService) SelectAccountWithLoadAwareness(ctx context.Context, groupID *int64, sessionHash string, requestedModel string, excludedIDs map[int64]struct{}, metadataUserID string) (*AccountSelectionResult, error) {
// 调度流程文档见 docs/ACCOUNT_SCHEDULING_FLOW.md 。
// metadataUserID: 用于客户端亲和调度,从中提取客户端 ID
// sub2apiUserID: 系统用户 ID,用于二维亲和调度
func (s *GatewayService) SelectAccountWithLoadAwareness(ctx context.Context, groupID *int64, sessionHash string, requestedModel string, excludedIDs map[int64]struct{}, metadataUserID string, sub2apiUserID int64) (*AccountSelectionResult, error) {
// 调试日志:记录调度入口参数
excludedIDsList := make([]int64, 0, len(excludedIDs))
for id := range excludedIDs {
@@ -1125,6 +1192,10 @@ func (s *GatewayService) SelectAccountWithLoadAwareness(ctx context.Context, gro
}
}
// 提取客户端 ID(用于客户端亲和调度)
affinityClientID := extractClientIDFromMetadata(metadataUserID)
affinityUserID := sub2apiUserID
if s.debugModelRoutingEnabled() && requestedModel != "" {
groupPlatform := ""
if group != nil {
@@ -1146,6 +1217,10 @@ func (s *GatewayService) SelectAccountWithLoadAwareness(ctx context.Context, gro
if err != nil {
return nil, err
}
if shouldFilterAccountWithoutClientID(account, affinityClientID) {
localExcluded[account.ID] = struct{}{}
continue
}
result, err := s.tryAcquireAccountSlot(ctx, account.ID, account.Concurrency)
if err == nil && result.Acquired {
@@ -1207,12 +1282,18 @@ func (s *GatewayService) SelectAccountWithLoadAwareness(ctx context.Context, gro
if err != nil {
return nil, err
}
accounts = filterAccountsWithoutClientID(accounts, affinityClientID)
if len(accounts) == 0 {
return nil, ErrNoAvailableAccounts
}
ctx = s.withWindowCostPrefetch(ctx, accounts)
ctx = s.withRPMPrefetch(ctx, accounts)
// 提前构建 accountByID(供 Layer 1 和 Layer 1.5 使用)
accountByID := make(map[int64]*Account, len(accounts))
for i := range accounts {
accountByID[accounts[i].ID] = &accounts[i]
}
isExcluded := func(accountID int64) bool {
if excludedIDs == nil {
return false
@@ -1220,12 +1301,19 @@ func (s *GatewayService) SelectAccountWithLoadAwareness(ctx context.Context, gro
_, excluded := excludedIDs[accountID]
return excluded
}
// 提前构建 accountByID(供 Layer 1 和 Layer 1.5 使用)
accountByID := make(map[int64]*Account, len(accounts))
for i := range accounts {
accountByID[accounts[i].ID] = &accounts[i]
}
affinityFlow := newGatewayAffinityFlow(
s,
ctx,
groupID,
sessionHash,
requestedModel,
affinityClientID,
affinityUserID,
platform,
useMixed,
accountByID,
isExcluded,
)
// 获取模型路由配置(仅 anthropic 平台)
var routingAccountIDs []int64
@@ -1388,7 +1476,10 @@ func (s *GatewayService) SelectAccountWithLoadAwareness(ctx context.Context, gro
}
if len(routingAvailable) > 0 {
// 排序:优先级 > 负载率 > 最后使用时间
// 批量获取亲和客户端数量
s.populateAffinityCounts(ctx, routingAvailable, derefGroupID(groupID))
// 排序:优先级 > 负载率 > 亲和客户端数 > 最后使用时间
sort.SliceStable(routingAvailable, func(i, j int) bool {
a, b := routingAvailable[i], routingAvailable[j]
if a.account.Priority != b.account.Priority {
@@ -1397,6 +1488,9 @@ func (s *GatewayService) SelectAccountWithLoadAwareness(ctx context.Context, gro
if a.loadInfo.LoadRate != b.loadInfo.LoadRate {
return a.loadInfo.LoadRate < b.loadInfo.LoadRate
}
if a.affinityCount != b.affinityCount {
return a.affinityCount < b.affinityCount
}
switch {
case a.account.LastUsedAt == nil && b.account.LastUsedAt != nil:
return true
@@ -1422,6 +1516,9 @@ func (s *GatewayService) SelectAccountWithLoadAwareness(ctx context.Context, gro
if sessionHash != "" && s.cache != nil {
_ = s.cache.SetSessionAccountID(ctx, derefGroupID(groupID), sessionHash, item.account.ID, stickySessionTTL)
}
if affinityClientID != "" && affinityUserID > 0 && s.cache != nil && item.account.IsAffinityEnabled() {
_ = s.cache.UpdateAffinity(ctx, derefGroupID(groupID), affinityUserID, affinityClientID, item.account.ID, ClientAffinityTTL)
}
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)
}
@@ -1459,14 +1556,27 @@ func (s *GatewayService) SelectAccountWithLoadAwareness(ctx context.Context, gro
}
}
// ============ Layer 1.5: 粘性会话(仅在无模型路由配置时生效) ============
if len(routingAccountIDs) == 0 && sessionHash != "" && stickyAccountID > 0 && !isExcluded(stickyAccountID) {
// ============ Layer 1.3: 用户亲和预处理(pinned_users 自动注入) ============
affinityFlow.preprocessPinnedUsers(accounts)
// ============ Layer 1.4: 客户端亲和调度(优先于粘性会话) ============
affinityHit := false
if affinityResult, hit, err := affinityFlow.trySelectAffinityAccount(); err != nil {
return nil, err
} else {
affinityHit = hit
if affinityResult != nil {
return affinityResult, nil
}
}
// ============ Layer 1.5: 粘性会话(仅在无模型路由配置 且 亲和未命中时生效) ============
if !affinityHit && len(routingAccountIDs) == 0 && sessionHash != "" && stickyAccountID > 0 && !isExcluded(stickyAccountID) {
accountID := stickyAccountID
if accountID > 0 && !isExcluded(accountID) {
account, ok := accountByID[accountID]
if ok {
// 检查账户是否需要清理粘性会话绑定
// Check if the account needs sticky session cleanup
clearSticky := shouldClearStickySession(account, requestedModel)
if clearSticky {
_ = s.cache.DeleteSessionAccountID(ctx, derefGroupID(groupID), sessionHash)
@@ -1482,7 +1592,6 @@ func (s *GatewayService) SelectAccountWithLoadAwareness(ctx context.Context, gro
result, err := s.tryAcquireAccountSlot(ctx, accountID, account.Concurrency)
if err == nil && result.Acquired {
// 会话数量限制检查
// Session count limit check
if !s.checkAndRegisterSession(ctx, account, sessionHash) {
result.ReleaseFunc() // 释放槽位,继续到 Layer 2
} else {
@@ -1497,10 +1606,8 @@ func (s *GatewayService) SelectAccountWithLoadAwareness(ctx context.Context, gro
waitingCount, _ := s.concurrencyService.GetAccountWaitingCount(ctx, accountID)
if waitingCount < cfg.StickySessionMaxWaiting {
// 会话数量限制检查(等待计划也需要占用会话配额)
// Session count limit check (wait plan also requires session quota)
if !s.checkAndRegisterSession(ctx, account, sessionHash) {
// 会话限制已满,继续到 Layer 2
// Session limit full, continue to Layer 2
} else {
return &AccountSelectionResult{
Account: account,
@@ -1570,6 +1677,9 @@ func (s *GatewayService) SelectAccountWithLoadAwareness(ctx context.Context, gro
loadMap, err := s.concurrencyService.GetAccountsLoadBatch(ctx, accountLoads)
if err != nil {
if result, ok := s.tryAcquireByLegacyOrder(ctx, candidates, groupID, sessionHash, preferOAuth); ok {
if affinityClientID != "" && affinityUserID > 0 && s.cache != nil && result.Account != nil && result.Account.IsAffinityEnabled() {
_ = s.cache.UpdateAffinity(ctx, derefGroupID(groupID), affinityUserID, affinityClientID, result.Account.ID, ClientAffinityTTL)
}
return result, nil
}
} else {
@@ -1587,13 +1697,37 @@ func (s *GatewayService) SelectAccountWithLoadAwareness(ctx context.Context, gro
}
}
// 分层过滤选择:优先级 → 负载率 → LRU
// 批量获取亲和客户端数量(用于均衡分配新客户端)
s.populateAffinityCounts(ctx, available, derefGroupID(groupID))
// 分层过滤选择:优先级 → 亲和三区 → 负载率 → 亲和客户端数 → LRU
for len(available) > 0 {
// 1. 取优先级最小的集合
candidates := filterByMinPriority(available)
// 2. 取负载率最低的集合
// 2. 按亲和三区过滤:绿区优先 → 黄区降级 → 红区移除(在同优先级内)
candidates = classifyByAffinityZone(candidates)
if len(candidates) == 0 {
// 当前优先级组全部在红区,移除后回退到下一优先级组
minPri := available[0].account.Priority
for _, a := range available[1:] {
if a.account.Priority < minPri {
minPri = a.account.Priority
}
}
newAvailable := make([]accountWithLoad, 0, len(available))
for _, a := range available {
if a.account.Priority != minPri {
newAvailable = append(newAvailable, a)
}
}
available = newAvailable
continue
}
// 3. 取负载率最低的集合
candidates = filterByMinLoadRate(candidates)
// 3. LRU 选择最久未用的账号
// 3. 取亲和客户端数最少的集合
candidates = filterByMinAffinityCount(candidates)
// 4. LRU 选择最久未用的账号
selected := selectByLRU(candidates, preferOAuth)
if selected == nil {
break
@@ -1608,6 +1742,10 @@ func (s *GatewayService) SelectAccountWithLoadAwareness(ctx context.Context, gro
if sessionHash != "" && s.cache != nil {
_ = s.cache.SetSessionAccountID(ctx, derefGroupID(groupID), sessionHash, selected.account.ID, stickySessionTTL)
}
// 更新亲和关系
if affinityClientID != "" && affinityUserID > 0 && s.cache != nil && selected.account.IsAffinityEnabled() {
_ = s.cache.UpdateAffinity(ctx, derefGroupID(groupID), affinityUserID, affinityClientID, selected.account.ID, ClientAffinityTTL)
}
return &AccountSelectionResult{
Account: selected.account,
Acquired: true,
@@ -2365,6 +2503,36 @@ func (s *GatewayService) getSchedulableAccount(ctx context.Context, accountID in
return s.accountRepo.GetByID(ctx, accountID)
}
// populateAffinityCounts 批量获取账号的亲和客户端数量并填入 accountWithLoad 切片。
// 仅当存在开启了客户端亲和的账号时才查询 Redis,否则跳过。
func (s *GatewayService) populateAffinityCounts(ctx context.Context, accounts []accountWithLoad, groupID int64) {
if s.cache == nil || len(accounts) == 0 {
return
}
// 快速检查:是否有任何账号开启了亲和
hasAffinity := false
for _, acc := range accounts {
if acc.account.IsAffinityEnabled() {
hasAffinity = true
break
}
}
if !hasAffinity {
return
}
accountIDs := make([]int64, len(accounts))
for i, acc := range accounts {
accountIDs[i] = acc.account.ID
}
countMap, err := s.cache.GetAccountAffinityCountBatch(ctx, groupID, accountIDs, ClientAffinityTTL)
if err != nil {
return // 查询失败不影响调度,affinityCount 保持 0
}
for i := range accounts {
accounts[i].affinityCount = countMap[accounts[i].account.ID]
}
}
// filterByMinPriority 过滤出优先级最小的账号集合
func filterByMinPriority(accounts []accountWithLoad) []accountWithLoad {
if len(accounts) == 0 {
@@ -2405,6 +2573,64 @@ func filterByMinLoadRate(accounts []accountWithLoad) []accountWithLoad {
return result
}
// filterByMinAffinityCount 过滤出亲和客户端数最少的账号集合
func filterByMinAffinityCount(accounts []accountWithLoad) []accountWithLoad {
if len(accounts) == 0 {
return accounts
}
minCount := accounts[0].affinityCount
for _, acc := range accounts[1:] {
if acc.affinityCount < minCount {
minCount = acc.affinityCount
}
}
result := make([]accountWithLoad, 0, len(accounts))
for _, acc := range accounts {
if acc.affinityCount == minCount {
result = append(result, acc)
}
}
return result
}
// classifyByAffinityZone 按亲和分区对候选账号进行分类。
// 返回值:仅绿区账号(有绿区时),否则返回黄区账号。红区账号被移除。
// 如果没有任何账号开启了亲和三区配置(即 affinity_base <= 0),则原样返回所有账号。
func classifyByAffinityZone(accounts []accountWithLoad) []accountWithLoad {
if len(accounts) == 0 {
return accounts
}
// 快速检查:是否有任何账号配置了 affinity_base
hasZoneConfig := false
for _, acc := range accounts {
if acc.account.IsAffinityEnabled() && acc.account.GetAffinityBase() > 0 {
hasZoneConfig = true
break
}
}
if !hasZoneConfig {
return accounts
}
greens := make([]accountWithLoad, 0, len(accounts))
yellows := make([]accountWithLoad, 0, len(accounts))
for _, acc := range accounts {
zone := acc.account.GetAffinityZone(acc.affinityCount)
switch zone {
case AffinityZoneGreen:
greens = append(greens, acc)
case AffinityZoneYellow:
yellows = append(yellows, acc)
case AffinityZoneRed:
// 红区:移除,不参与调度
}
}
if len(greens) > 0 {
return greens
}
return yellows
}
// selectByLRU 从集合中选择最久未用的账号
// 如果有多个账号具有相同的最小 LastUsedAt,则随机选择一个
func selectByLRU(accounts []accountWithLoad, preferOAuth bool) *accountWithLoad {
@@ -3378,6 +3604,10 @@ func (s *GatewayService) isModelSupportedByAccount(account *Account, requestedMo
_, ok := ResolveBedrockModelID(account, requestedModel)
return ok
}
// OpenAI 透传模式:仅替换认证,允许所有模型
if account.Platform == PlatformOpenAI && account.IsOpenAIPassthroughEnabled() {
return true
}
// OAuth/SetupToken 账号使用 Anthropic 标准映射(短ID → 长ID)
if account.Platform == PlatformAnthropic && account.Type != AccountTypeAPIKey {
requestedModel = claude.NormalizeModelID(requestedModel)
@@ -4486,6 +4716,7 @@ func (s *GatewayService) Forward(ctx context.Context, c *gin.Context, account *A
}
// 处理正常响应
ctx = withClaudeMaxResponseRewriteContext(ctx, c, parsed)
// 触发上游接受回调(提前释放串行锁,不等流完成)
if parsed.OnUpstreamAccepted != nil {
@@ -6591,6 +6822,7 @@ func (s *GatewayService) handleStreamingResponse(ctx context.Context, resp *http
needModelReplace := originalModel != mappedModel
clientDisconnected := false // 客户端断开标志,断开后继续读取上游以获取完整usage
sawTerminalEvent := false
skipAccountTTLOverride := false
pendingEventLines := make([]string, 0, 4)
@@ -6652,17 +6884,25 @@ func (s *GatewayService) handleStreamingResponse(ctx context.Context, resp *http
if msg, ok := event["message"].(map[string]any); ok {
if u, ok := msg["usage"].(map[string]any); ok {
eventChanged = reconcileCachedTokens(u) || eventChanged
claudeMaxOutcome := applyClaudeMaxSimulationToUsageJSONMap(ctx, u, originalModel, account.ID)
if claudeMaxOutcome.Simulated {
skipAccountTTLOverride = true
}
}
}
}
if eventType == "message_delta" {
if u, ok := event["usage"].(map[string]any); ok {
eventChanged = reconcileCachedTokens(u) || eventChanged
claudeMaxOutcome := applyClaudeMaxSimulationToUsageJSONMap(ctx, u, originalModel, account.ID)
if claudeMaxOutcome.Simulated {
skipAccountTTLOverride = true
}
}
}
// Cache TTL Override: 重写 SSE 事件中的 cache_creation 分类
if account.IsCacheTTLOverrideEnabled() {
if account.IsCacheTTLOverrideEnabled() && !skipAccountTTLOverride {
overrideTarget := account.GetCacheTTLOverrideTarget()
if eventType == "message_start" {
if msg, ok := event["message"].(map[string]any); ok {
@@ -7094,8 +7334,13 @@ func (s *GatewayService) handleNonStreamingResponse(ctx context.Context, resp *h
}
}
claudeMaxOutcome := applyClaudeMaxSimulationToUsage(ctx, &response.Usage, originalModel, account.ID)
if claudeMaxOutcome.Simulated {
body = rewriteClaudeUsageJSONBytes(body, response.Usage)
}
// Cache TTL Override: 重写 non-streaming 响应中的 cache_creation 分类
if account.IsCacheTTLOverrideEnabled() {
if account.IsCacheTTLOverrideEnabled() && !claudeMaxOutcome.Simulated {
overrideTarget := account.GetCacheTTLOverrideTarget()
if applyCacheTTLOverride(&response.Usage, overrideTarget) {
// 同步更新 body JSON 中的嵌套 cache_creation 对象
@@ -7161,6 +7406,7 @@ func (s *GatewayService) getUserGroupRateMultiplier(ctx context.Context, userID,
// RecordUsageInput 记录使用量的输入参数
type RecordUsageInput struct {
Result *ForwardResult
ParsedRequest *ParsedRequest
APIKey *APIKey
User *User
Account *Account
@@ -7473,9 +7719,19 @@ func (s *GatewayService) RecordUsage(ctx context.Context, input *RecordUsageInpu
result.Usage.InputTokens = 0
}
// Claude Max cache billing policy (group-level):
// - GatewayService 路径: Forward 已改写 usage(含 cache tokens)→ apply 见到 cache tokens 跳过 → simulatedClaudeMax=true(通过第二条件)
// - Antigravity 路径: Forward 中 hook 改写了客户端 SSE,但 ForwardResult.Usage 是原始值 → apply 实际执行模拟 → simulatedClaudeMax=true
var apiKeyGroup *Group
if apiKey != nil {
apiKeyGroup = apiKey.Group
}
claudeMaxOutcome := applyClaudeMaxCacheBillingPolicyToUsage(&result.Usage, input.ParsedRequest, apiKeyGroup, result.Model, account.ID)
simulatedClaudeMax := claudeMaxOutcome.Simulated ||
(shouldApplyClaudeMaxBillingRulesForUsage(apiKeyGroup, result.Model, input.ParsedRequest) && hasCacheCreationTokens(result.Usage))
// Cache TTL Override: 确保计费时 token 分类与账号设置一致
cacheTTLOverridden := false
if account.IsCacheTTLOverrideEnabled() {
if account.IsCacheTTLOverrideEnabled() && !simulatedClaudeMax {
applyCacheTTLOverride(&result.Usage, account.GetCacheTTLOverrideTarget())
cacheTTLOverridden = (result.Usage.CacheCreation5mTokens + result.Usage.CacheCreation1hTokens) > 0
}
@@ -0,0 +1,122 @@
//go:build unit
// Package service 提供 API 网关核心服务。
// 本文件包含 SortAffinityClients 函数的单元测试,
// 验证 AffinityClient 切片排序逻辑在各种输入条件下的正确行为。
//
// This file contains unit tests for the SortAffinityClients function,
// verifying correct sorting behavior for AffinityClient slices under
// various input conditions including empty, single, sorted, reverse,
// and duplicate-timestamp scenarios.
package service
import (
"testing"
"time"
"github.com/stretchr/testify/require"
)
func TestSortAffinityClients_Empty(t *testing.T) {
var clients []AffinityClient
SortAffinityClients(clients)
require.Empty(t, clients)
clients = []AffinityClient{}
SortAffinityClients(clients)
require.Empty(t, clients)
}
func TestSortAffinityClients_SingleElement(t *testing.T) {
now := time.Now()
clients := []AffinityClient{
{ClientID: "client-1", LastActive: now},
}
SortAffinityClients(clients)
require.Len(t, clients, 1)
require.Equal(t, "client-1", clients[0].ClientID)
require.Equal(t, now, clients[0].LastActive)
}
func TestSortAffinityClients_AlreadySorted(t *testing.T) {
now := time.Now()
clients := []AffinityClient{
{ClientID: "newest", LastActive: now},
{ClientID: "middle", LastActive: now.Add(-1 * time.Hour)},
{ClientID: "oldest", LastActive: now.Add(-2 * time.Hour)},
}
SortAffinityClients(clients)
require.Equal(t, "newest", clients[0].ClientID)
require.Equal(t, "middle", clients[1].ClientID)
require.Equal(t, "oldest", clients[2].ClientID)
}
func TestSortAffinityClients_ReverseOrder(t *testing.T) {
now := time.Now()
clients := []AffinityClient{
{ClientID: "oldest", LastActive: now.Add(-2 * time.Hour)},
{ClientID: "middle", LastActive: now.Add(-1 * time.Hour)},
{ClientID: "newest", LastActive: now},
}
SortAffinityClients(clients)
require.Equal(t, "newest", clients[0].ClientID)
require.Equal(t, "middle", clients[1].ClientID)
require.Equal(t, "oldest", clients[2].ClientID)
}
func TestSortAffinityClients_SameTimestamps(t *testing.T) {
now := time.Now()
clients := []AffinityClient{
{ClientID: "c1", LastActive: now},
{ClientID: "c2", LastActive: now},
{ClientID: "c3", LastActive: now},
}
SortAffinityClients(clients)
// 所有时间戳相同时,排序结果应保持稳定(sort.Slice 不保证稳定性,
// 但只要结果是某种确定的顺序即可)。
// 验证所有元素仍然存在且时间相同。
require.Len(t, clients, 3)
ids := map[string]bool{}
for _, c := range clients {
ids[c.ClientID] = true
require.Equal(t, now, c.LastActive)
}
require.True(t, ids["c1"])
require.True(t, ids["c2"])
require.True(t, ids["c3"])
}
func TestSortAffinityClients_MixedOrder(t *testing.T) {
now := time.Now()
clients := []AffinityClient{
{ClientID: "c3", LastActive: now.Add(-30 * time.Minute)},
{ClientID: "c1", LastActive: now},
{ClientID: "c5", LastActive: now.Add(-2 * time.Hour)},
{ClientID: "c2", LastActive: now.Add(-10 * time.Minute)},
{ClientID: "c4", LastActive: now.Add(-1 * time.Hour)},
}
SortAffinityClients(clients)
// 按 LastActive 降序排列
require.Equal(t, "c1", clients[0].ClientID) // now
require.Equal(t, "c2", clients[1].ClientID) // -10m
require.Equal(t, "c3", clients[2].ClientID) // -30m
require.Equal(t, "c4", clients[3].ClientID) // -1h
require.Equal(t, "c5", clients[4].ClientID) // -2h
}
func TestSortAffinityClients_SubSecondDifferences(t *testing.T) {
base := time.Now()
clients := []AffinityClient{
{ClientID: "early", LastActive: base},
{ClientID: "late", LastActive: base.Add(500 * time.Millisecond)},
}
SortAffinityClients(clients)
// 500ms 差异也应正确排序(更晚的在前)
require.Equal(t, "late", clients[0].ClientID)
require.Equal(t, "early", clients[1].ClientID)
}
@@ -288,6 +288,28 @@ func (m *mockGatewayCacheForGemini) DeleteSessionAccountID(ctx context.Context,
return nil
}
func (m *mockGatewayCacheForGemini) GetAffinityAccounts(_ context.Context, _ int64, _ int64, _ string, _ time.Duration) ([]int64, error) {
return nil, nil
}
func (m *mockGatewayCacheForGemini) UpdateAffinity(_ context.Context, _ int64, _ int64, _ string, _ int64, _ time.Duration) error {
return nil
}
func (m *mockGatewayCacheForGemini) GetAffinityMultiCount(_ context.Context, _ int64, _ int64, _ int64, _ time.Duration) (int64, int64, int64, error) {
return 0, 0, 0, nil
}
func (m *mockGatewayCacheForGemini) GetAccountAffinityCountBatch(_ context.Context, _ int64, _ []int64, _ time.Duration) (map[int64]int64, error) {
return map[int64]int64{}, nil
}
func (m *mockGatewayCacheForGemini) GetAccountAffinityClientsBatch(_ context.Context, _ map[int64][]int64, _ time.Duration) (map[int64][]string, error) {
return map[int64][]string{}, nil
}
func (m *mockGatewayCacheForGemini) GetAccountAffinityClientsWithScores(_ context.Context, _ int64, _ []int64, _ time.Duration) ([]AffinityClient, error) {
return nil, nil
}
func (m *mockGatewayCacheForGemini) ClearAccountAffinity(_ context.Context, _ int64, _ []int64) error {
return nil
}
// TestGeminiMessagesCompatService_SelectAccountForModelWithExclusions_GeminiPlatform 测试 Gemini 单平台选择
func TestGeminiMessagesCompatService_SelectAccountForModelWithExclusions_GeminiPlatform(t *testing.T) {
ctx := context.Background()
+3
View File
@@ -50,6 +50,9 @@ type Group struct {
// MCP XML 协议注入开关(仅 antigravity 平台使用)
MCPXMLInject bool
// Claude usage 模拟开关:将无写缓存 usage 模拟为 claude-max 风格
SimulateClaudeMaxEnabled bool
// 支持的模型系列(仅 antigravity 平台使用)
// 可选值: claude, gemini_text, gemini_image
SupportedModelScopes []string
@@ -323,7 +323,7 @@ func (s *defaultOpenAIAccountScheduler) selectBySessionHash(
_ = s.service.deleteStickySessionAccountID(ctx, req.GroupID, sessionHash)
return nil, nil
}
if req.RequestedModel != "" && !account.IsModelSupported(req.RequestedModel) {
if req.RequestedModel != "" && !account.IsOpenAIPassthroughEnabled() && !account.IsModelSupported(req.RequestedModel) {
return nil, nil
}
if !s.isAccountTransportCompatible(account, req.RequiredTransport) {
@@ -582,7 +582,7 @@ func (s *defaultOpenAIAccountScheduler) selectByLoadBalance(
if !account.IsSchedulable() || !account.IsOpenAI() {
continue
}
if req.RequestedModel != "" && !account.IsModelSupported(req.RequestedModel) {
if req.RequestedModel != "" && !account.IsOpenAIPassthroughEnabled() && !account.IsModelSupported(req.RequestedModel) {
continue
}
if !s.isAccountTransportCompatible(account, req.RequiredTransport) {
@@ -1393,7 +1393,7 @@ func (s *OpenAIGatewayService) SelectAccountWithLoadAwareness(ctx context.Contex
if !acc.IsSchedulable() {
continue
}
if requestedModel != "" && !acc.IsModelSupported(requestedModel) {
if requestedModel != "" && !acc.IsOpenAIPassthroughEnabled() && !acc.IsModelSupported(requestedModel) {
continue
}
candidates = append(candidates, acc)
@@ -1554,7 +1554,7 @@ func (s *OpenAIGatewayService) resolveFreshSchedulableOpenAIAccount(ctx context.
if !fresh.IsSchedulable() || !fresh.IsOpenAI() {
return nil
}
if requestedModel != "" && !fresh.IsModelSupported(requestedModel) {
if requestedModel != "" && !fresh.IsOpenAIPassthroughEnabled() && !fresh.IsModelSupported(requestedModel) {
return nil
}
return fresh
@@ -282,6 +282,28 @@ func (c *stubGatewayCache) DeleteSessionAccountID(ctx context.Context, groupID i
return nil
}
func (c *stubGatewayCache) GetAffinityAccounts(_ context.Context, _ int64, _ int64, _ string, _ time.Duration) ([]int64, error) {
return nil, nil
}
func (c *stubGatewayCache) UpdateAffinity(_ context.Context, _ int64, _ int64, _ string, _ int64, _ time.Duration) error {
return nil
}
func (c *stubGatewayCache) GetAffinityMultiCount(_ context.Context, _ int64, _ int64, _ int64, _ time.Duration) (int64, int64, int64, error) {
return 0, 0, 0, nil
}
func (c *stubGatewayCache) GetAccountAffinityCountBatch(_ context.Context, _ int64, _ []int64, _ time.Duration) (map[int64]int64, error) {
return map[int64]int64{}, nil
}
func (c *stubGatewayCache) GetAccountAffinityClientsBatch(_ context.Context, _ map[int64][]int64, _ time.Duration) (map[int64][]string, error) {
return map[int64][]string{}, nil
}
func (c *stubGatewayCache) GetAccountAffinityClientsWithScores(_ context.Context, _ int64, _ []int64, _ time.Duration) ([]AffinityClient, error) {
return nil, nil
}
func (c *stubGatewayCache) ClearAccountAffinity(_ context.Context, _ int64, _ []int64) error {
return nil
}
func TestOpenAISelectAccountWithLoadAwareness_FiltersUnschedulable(t *testing.T) {
now := time.Now()
resetAt := now.Add(10 * time.Minute)
@@ -193,6 +193,28 @@ func (c *openAIWSStateStoreTimeoutProbeCache) DeleteSessionAccountID(ctx context
return nil
}
func (c *openAIWSStateStoreTimeoutProbeCache) GetAffinityAccounts(_ context.Context, _ int64, _ int64, _ string, _ time.Duration) ([]int64, error) {
return nil, nil
}
func (c *openAIWSStateStoreTimeoutProbeCache) UpdateAffinity(_ context.Context, _ int64, _ int64, _ string, _ int64, _ time.Duration) error {
return nil
}
func (c *openAIWSStateStoreTimeoutProbeCache) GetAffinityMultiCount(_ context.Context, _ int64, _ int64, _ int64, _ time.Duration) (int64, int64, int64, error) {
return 0, 0, 0, nil
}
func (c *openAIWSStateStoreTimeoutProbeCache) GetAccountAffinityCountBatch(_ context.Context, _ int64, _ []int64, _ time.Duration) (map[int64]int64, error) {
return map[int64]int64{}, nil
}
func (c *openAIWSStateStoreTimeoutProbeCache) GetAccountAffinityClientsBatch(_ context.Context, _ map[int64][]int64, _ time.Duration) (map[int64][]string, error) {
return map[int64][]string{}, nil
}
func (c *openAIWSStateStoreTimeoutProbeCache) GetAccountAffinityClientsWithScores(_ context.Context, _ int64, _ []int64, _ time.Duration) ([]AffinityClient, error) {
return nil, nil
}
func (c *openAIWSStateStoreTimeoutProbeCache) ClearAccountAffinity(_ context.Context, _ int64, _ []int64) error {
return nil
}
func TestOpenAIWSStateStore_RedisOpsUseShortTimeout(t *testing.T) {
probe := &openAIWSStateStoreTimeoutProbeCache{}
store := NewOpenAIWSStateStore(probe)
+3 -6
View File
@@ -64,12 +64,9 @@ func (s *OpsService) getAccountsLoadMapBestEffort(ctx context.Context, accounts
if acc.ID <= 0 {
continue
}
c := acc.Concurrency
if c <= 0 {
c = 1
}
if prev, ok := unique[acc.ID]; !ok || c > prev {
unique[acc.ID] = c
lf := acc.EffectiveLoadFactor()
if prev, ok := unique[acc.ID]; !ok || lf > prev {
unique[acc.ID] = lf
}
}
@@ -391,7 +391,7 @@ func (c *OpsMetricsCollector) collectConcurrencyQueueDepth(parentCtx context.Con
}
batch = append(batch, AccountWithConcurrency{
ID: acc.ID,
MaxConcurrency: acc.Concurrency,
MaxConcurrency: acc.EffectiveLoadFactor(),
})
}
if len(batch) == 0 {
+1 -1
View File
@@ -519,7 +519,7 @@ func (s *OpsService) selectAccountForRetry(ctx context.Context, reqType opsRetry
if s.gatewayService == nil {
return nil, fmt.Errorf("gateway service not available")
}
return s.gatewayService.SelectAccountWithLoadAwareness(ctx, groupID, "", model, excludedIDs, "") // 重试不使用会话限制
return s.gatewayService.SelectAccountWithLoadAwareness(ctx, groupID, "", model, excludedIDs, "", int64(0)) // 重试不使用会话限制
default:
return nil, fmt.Errorf("unsupported retry type: %s", reqType)
}
@@ -183,6 +183,15 @@ func TestOpsSystemLogSink_StartStopAndFlushSuccess(t *testing.T) {
if strings.TrimSpace(item.Message) == "" {
t.Fatalf("message should not be empty")
}
// writtenCount is incremented after BatchInsertSystemLogsFn returns,
// so poll briefly to avoid a race between the done signal and the atomic add.
deadline := time.Now().Add(time.Second)
for time.Now().Before(deadline) {
if sink.Health().WrittenCount > 0 {
break
}
time.Sleep(time.Millisecond)
}
health := sink.Health()
if health.WrittenCount == 0 {
t.Fatalf("written_count should be >0")
@@ -94,6 +94,8 @@ func TestRateLimitService_HandleUpstreamError_OAuth401SetsTempUnschedulable(t *t
})
}
// TestRateLimitService_HandleUpstreamError_OAuth401InvalidatorError
// OpenAI OAuth 401 缓存失效出错时仍走 temp_unschedulable
func TestRateLimitService_HandleUpstreamError_OAuth401InvalidatorError(t *testing.T) {
repo := &rateLimitAccountRepoStub{}
invalidator := &tokenCacheInvalidatorRecorder{err: errors.New("boom")}
@@ -101,7 +103,7 @@ func TestRateLimitService_HandleUpstreamError_OAuth401InvalidatorError(t *testin
service.SetTokenCacheInvalidator(invalidator)
account := &Account{
ID: 101,
Platform: PlatformGemini,
Platform: PlatformOpenAI,
Type: AccountTypeOAuth,
}
@@ -210,6 +210,49 @@ func (s *SettingService) SetOnS3UpdateCallback(callback func()) {
s.onS3Update = callback
}
// SetOnStorageUpdateCallback 设置存储配置变更时的回调函数(用于刷新所有存储客户端缓存)。
// 替代 SetOnS3UpdateCallback,支持 S3 + GDrive 统一刷新。
func (s *SettingService) SetOnStorageUpdateCallback(callback func()) {
s.onS3Update = callback
}
// --- 统一存储 Profile 方法别名 ---
// ListSoraStorageProfiles 获取 Sora 存储多配置列表(统一方法名)。
func (s *SettingService) ListSoraStorageProfiles(ctx context.Context) (*SoraS3ProfileList, error) {
return s.ListSoraS3Profiles(ctx)
}
// CreateSoraStorageProfile 创建 Sora 存储配置(统一方法名)。
func (s *SettingService) CreateSoraStorageProfile(ctx context.Context, profile *SoraS3Profile, setActive bool) (*SoraS3Profile, error) {
return s.CreateSoraS3Profile(ctx, profile, setActive)
}
// UpdateSoraStorageProfile 更新 Sora 存储配置(统一方法名)。
func (s *SettingService) UpdateSoraStorageProfile(ctx context.Context, profileID string, profile *SoraS3Profile) (*SoraS3Profile, error) {
return s.UpdateSoraS3Profile(ctx, profileID, profile)
}
// DeleteSoraStorageProfile 删除 Sora 存储配置(统一方法名)。
func (s *SettingService) DeleteSoraStorageProfile(ctx context.Context, profileID string) error {
return s.DeleteSoraS3Profile(ctx, profileID)
}
// SetActiveSoraStorageProfile 设置激活的 Sora 存储配置(统一方法名)。
func (s *SettingService) SetActiveSoraStorageProfile(ctx context.Context, profileID string) (*SoraS3Profile, error) {
return s.SetActiveSoraS3Profile(ctx, profileID)
}
// GetActiveStorageProfile 获取当前激活的存储配置 profile。
func (s *SettingService) GetActiveStorageProfile(ctx context.Context) (*SoraS3Profile, error) {
profiles, err := s.ListSoraS3Profiles(ctx)
if err != nil {
return nil, err
}
active := pickActiveSoraS3Profile(profiles.Items, profiles.ActiveProfileID)
return active, nil
}
// SetVersion sets the application version for injection into public settings
func (s *SettingService) SetVersion(version string) {
s.version = version
@@ -1475,6 +1518,8 @@ type soraS3ProfilesStore struct {
type soraS3ProfileStoreItem struct {
ProfileID string `json:"profile_id"`
Name string `json:"name"`
Provider string `json:"provider,omitempty"` // "s3" / "gdrive",空值视为 "s3"
AccessMode string `json:"access_mode,omitempty"` // "direct" / "proxy",空值视为 "direct"
Enabled bool `json:"enabled"`
Endpoint string `json:"endpoint"`
Region string `json:"region"`
@@ -1486,6 +1531,14 @@ type soraS3ProfileStoreItem struct {
CDNURL string `json:"cdn_url"`
DefaultStorageQuotaBytes int64 `json:"default_storage_quota_bytes"`
UpdatedAt string `json:"updated_at"`
// --- Google Drive 专属 ---
AuthType string `json:"auth_type,omitempty"`
ClientID string `json:"client_id,omitempty"`
ClientSecret string `json:"client_secret,omitempty"`
RefreshToken string `json:"refresh_token,omitempty"`
ServiceAccountJSON string `json:"service_account_json,omitempty"`
FolderID string `json:"folder_id,omitempty"`
}
// GetSoraS3Settings 获取 Sora S3 存储配置(兼容旧单配置语义:返回当前激活配置)
@@ -1597,6 +1650,8 @@ func (s *SettingService) CreateSoraS3Profile(ctx context.Context, profile *SoraS
store.Items = append(store.Items, soraS3ProfileStoreItem{
ProfileID: profileID,
Name: name,
Provider: profile.Provider,
AccessMode: profile.AccessMode,
Enabled: profile.Enabled,
Endpoint: strings.TrimSpace(profile.Endpoint),
Region: strings.TrimSpace(profile.Region),
@@ -1608,6 +1663,13 @@ func (s *SettingService) CreateSoraS3Profile(ctx context.Context, profile *SoraS
CDNURL: strings.TrimSpace(profile.CDNURL),
DefaultStorageQuotaBytes: maxInt64(profile.DefaultStorageQuotaBytes, 0),
UpdatedAt: now,
// Google Drive 专属
AuthType: profile.AuthType,
ClientID: strings.TrimSpace(profile.ClientID),
ClientSecret: profile.ClientSecret,
RefreshToken: profile.RefreshToken,
ServiceAccountJSON: profile.ServiceAccountJSON,
FolderID: strings.TrimSpace(profile.FolderID),
})
if setActive || store.ActiveProfileID == "" {
@@ -1653,6 +1715,8 @@ func (s *SettingService) UpdateSoraS3Profile(ctx context.Context, profileID stri
return nil, infraerrors.BadRequest("SORA_S3_PROFILE_NAME_REQUIRED", "name is required")
}
target.Name = name
target.Provider = profile.Provider
target.AccessMode = profile.AccessMode
target.Enabled = profile.Enabled
target.Endpoint = strings.TrimSpace(profile.Endpoint)
target.Region = strings.TrimSpace(profile.Region)
@@ -1665,6 +1729,19 @@ func (s *SettingService) UpdateSoraS3Profile(ctx context.Context, profileID stri
if profile.SecretAccessKey != "" {
target.SecretAccessKey = profile.SecretAccessKey
}
// Google Drive 专属
target.AuthType = profile.AuthType
target.ClientID = strings.TrimSpace(profile.ClientID)
if profile.ClientSecret != "" {
target.ClientSecret = profile.ClientSecret
}
if profile.RefreshToken != "" {
target.RefreshToken = profile.RefreshToken
}
if profile.ServiceAccountJSON != "" {
target.ServiceAccountJSON = profile.ServiceAccountJSON
}
target.FolderID = strings.TrimSpace(profile.FolderID)
target.UpdatedAt = time.Now().UTC().Format(time.RFC3339)
store.Items[targetIndex] = target
@@ -1967,6 +2044,8 @@ func convertSoraS3ProfilesStore(store *soraS3ProfilesStore) *SoraS3ProfileList {
ProfileID: item.ProfileID,
Name: item.Name,
IsActive: item.ProfileID == store.ActiveProfileID,
Provider: item.Provider,
AccessMode: item.AccessMode,
Enabled: item.Enabled,
Endpoint: item.Endpoint,
Region: item.Region,
@@ -1979,6 +2058,16 @@ func convertSoraS3ProfilesStore(store *soraS3ProfilesStore) *SoraS3ProfileList {
CDNURL: item.CDNURL,
DefaultStorageQuotaBytes: item.DefaultStorageQuotaBytes,
UpdatedAt: item.UpdatedAt,
// Google Drive 专属
AuthType: item.AuthType,
ClientID: item.ClientID,
ClientSecret: item.ClientSecret,
ClientSecretConfigured: item.ClientSecret != "",
RefreshToken: item.RefreshToken,
RefreshTokenConfigured: item.RefreshToken != "",
ServiceAccountJSON: item.ServiceAccountJSON,
ServiceAccountConfigured: item.ServiceAccountJSON != "",
FolderID: item.FolderID,
})
}
return &SoraS3ProfileList{
+21
View File
@@ -129,6 +129,8 @@ type SoraS3Profile struct {
ProfileID string `json:"profile_id"`
Name string `json:"name"`
IsActive bool `json:"is_active"`
Provider string `json:"provider"` // "s3" / "gdrive",空值视为 "s3"
AccessMode string `json:"access_mode"` // "direct" / "proxy",空值视为 "direct"
Enabled bool `json:"enabled"`
Endpoint string `json:"endpoint"`
Region string `json:"region"`
@@ -141,6 +143,25 @@ type SoraS3Profile struct {
CDNURL string `json:"cdn_url"`
DefaultStorageQuotaBytes int64 `json:"default_storage_quota_bytes"`
UpdatedAt string `json:"updated_at"`
// --- Google Drive 专属 ---
AuthType string `json:"auth_type,omitempty"` // "oauth2" / "service_account"
ClientID string `json:"client_id,omitempty"`
ClientSecret string `json:"-"`
ClientSecretConfigured bool `json:"client_secret_configured"`
RefreshToken string `json:"-"`
RefreshTokenConfigured bool `json:"refresh_token_configured"`
ServiceAccountJSON string `json:"-"`
ServiceAccountConfigured bool `json:"service_account_configured"`
FolderID string `json:"folder_id,omitempty"`
}
// GetProvider 返回 Provider,空值视为 "s3"。
func (p *SoraS3Profile) GetProvider() string {
if p.Provider == "" {
return SoraStorageTypeS3
}
return p.Provider
}
// SoraS3ProfileList Sora S3 多配置列表
@@ -0,0 +1,71 @@
package service
import (
"context"
"crypto/rand"
"encoding/hex"
"fmt"
"golang.org/x/oauth2"
"golang.org/x/oauth2/google"
"google.golang.org/api/drive/v3"
)
// SoraGDriveOAuthService 处理 Google Drive OAuth2 授权流程。
type SoraGDriveOAuthService struct{}
// NewSoraGDriveOAuthService 创建 GDrive OAuth 服务。
func NewSoraGDriveOAuthService(_ *SettingService) *SoraGDriveOAuthService {
return &SoraGDriveOAuthService{}
}
// GenerateAuthURL 生成 Google OAuth 授权 URL。
func (s *SoraGDriveOAuthService) GenerateAuthURL(clientID, clientSecret, redirectURI string) (authURL, state string, err error) {
if clientID == "" || clientSecret == "" || redirectURI == "" {
return "", "", fmt.Errorf("client_id, client_secret, redirect_uri are required")
}
config := &oauth2.Config{
ClientID: clientID,
ClientSecret: clientSecret,
Endpoint: google.Endpoint,
Scopes: []string{drive.DriveFileScope},
RedirectURL: redirectURI,
}
// 生成随机 state
stateBytes := make([]byte, 16)
if _, err := rand.Read(stateBytes); err != nil {
return "", "", fmt.Errorf("generate state: %w", err)
}
state = hex.EncodeToString(stateBytes)
authURL = config.AuthCodeURL(state, oauth2.AccessTypeOffline, oauth2.ApprovalForce)
return authURL, state, nil
}
// ExchangeCode 用授权码换取 refresh_token。
func (s *SoraGDriveOAuthService) ExchangeCode(ctx context.Context, clientID, clientSecret, redirectURI, code string) (string, error) {
if code == "" {
return "", fmt.Errorf("authorization code is required")
}
config := &oauth2.Config{
ClientID: clientID,
ClientSecret: clientSecret,
Endpoint: google.Endpoint,
Scopes: []string{drive.DriveFileScope},
RedirectURL: redirectURI,
}
token, err := config.Exchange(ctx, code)
if err != nil {
return "", fmt.Errorf("exchange code: %w", err)
}
if token.RefreshToken == "" {
return "", fmt.Errorf("no refresh_token received, please revoke app access and try again")
}
return token.RefreshToken, nil
}

Some files were not shown because too many files have changed in this diff Show More