Merge branch 'main' into release/custom-0.1.99

# Conflicts:
#	backend/cmd/server/wire_gen.go
#	backend/cmd/server/wire_gen_test.go
#	backend/internal/handler/dto/mappers.go
#	backend/internal/service/ratelimit_service.go
#	backend/internal/service/ratelimit_service_401_db_fallback_test.go
#	backend/internal/service/ratelimit_service_401_test.go
#	frontend/src/components/account/AccountCapacityCell.vue
#	frontend/src/components/account/CreateAccountModal.vue
#	frontend/src/components/account/EditAccountModal.vue
#	frontend/src/views/admin/DataManagementView.vue
This commit is contained in:
erio
2026-03-14 21:03:09 +08:00
145 changed files with 11863 additions and 768 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'
+9 -1
View File
@@ -78,6 +78,7 @@ Desktop.ini
# ===================
tmp/
temp/
logs/
*.tmp
*.temp
*.log
@@ -128,8 +129,15 @@ deploy/docker-compose.override.yml
vite.config.js
docs/*
.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.98.1
+24 -8
View File
@@ -142,7 +142,7 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) {
crsSyncService := service.NewCRSSyncService(accountRepository, proxyRepository, oAuthService, openAIOAuthService, geminiOAuthService, configConfig)
sessionLimitCache := repository.ProvideSessionLimitCache(redisClient, configConfig)
rpmCache := repository.NewRPMCache(redisClient)
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)
@@ -175,11 +175,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)
@@ -205,7 +209,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)
@@ -214,13 +218,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)
@@ -236,7 +245,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,
@@ -290,6 +299,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)
@@ -431,6 +441,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
+109 -15
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=
@@ -93,15 +97,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/clipperhouse/stringish v0.1.1 h1:+NSqMOr3GR6k1FdRhhnXrLfztGzuG+VuFDfatpWHKCs=
github.com/clipperhouse/stringish v0.1.1/go.mod h1:v/WhFtE1q0ovMta2+m+UbpZ+2/HEXNWYXQgCt4hdOzA=
github.com/clipperhouse/uax29/v2 v2.5.0 h1:x7T0T4eTHDONxFJsL94uKNKPHrclyFI0lm7+w94cO8U=
github.com/clipperhouse/uax29/v2 v2.5.0/go.mod h1:Wn1g7MK6OoeDT0vL+Q0SQLDz/KpfsVRgg6W7ihQeh4g=
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=
@@ -128,6 +131,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=
@@ -138,6 +143,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=
@@ -175,7 +184,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=
@@ -184,10 +215,18 @@ 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/subcommands v1.2.0/go.mod h1:ZjhPrFU+Olkh9WazFPsl27BQ4UPiG37m3yTrtFlrHVk=
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=
@@ -238,8 +277,6 @@ github.com/mattn/go-colorable v0.1.13/go.mod h1:7S9/ev0klgBDR4GtXTXX8a3vIGJpMovk
github.com/mattn/go-isatty v0.0.16/go.mod h1:kYGgaQfpe5nmfYZH+SKPsOc2e4SrIfOl2e/yFXSvRLM=
github.com/mattn/go-isatty v0.0.20 h1:xfD0iDuEKnDkl03q4limB+vH+GxLEtL/jb4xVJSWWEY=
github.com/mattn/go-isatty v0.0.20/go.mod h1:W+V8PltTTMOvKvAeJH7IuucS94S2C6jfK/D7dTCTo3Y=
github.com/mattn/go-runewidth v0.0.19 h1:v++JhqYnZuu5jSKrk9RbgF5v4CGUjqRfBm05byFGLdw=
github.com/mattn/go-runewidth v0.0.19/go.mod h1:XBkDxAl56ILZc9knddidhrOlY5R/pDhgLpndooCuJAs=
github.com/mattn/go-sqlite3 v1.14.17 h1:mCRHCLDUBXgpKAqIKsaAaAsrAlbkeomtRFKXh2L6YIM=
github.com/mattn/go-sqlite3 v1.14.17/go.mod h1:2eHXhiwb8IkHr+BDWZGa96P6+rkvnG63S2DGjv9HUNg=
github.com/mdelapenya/tlscert v0.2.0 h1:7H81W6Z/4weDvZBNOfQte5GpIMo0lGYEeWbkGp5LJHI=
@@ -273,8 +310,6 @@ github.com/morikuni/aec v1.0.0 h1:nP9CBfwrvYnBRgY6qfDQkygYDmYwOilePFkwzv4dU8A=
github.com/morikuni/aec v1.0.0/go.mod h1:BbKIizmSmc5MMPqRYbxO4ZU0S0+P200+tUnFx7PXmsc=
github.com/ncruces/go-strftime v1.0.0 h1:HMFp8mLCTPp341M/ZnA4qaf7ZlsbTc+miZjCLOFAw7w=
github.com/ncruces/go-strftime v1.0.0/go.mod h1:Fwc5htZGVVkseilnfgOVb9mKy6w1naJmn9CehxcKcls=
github.com/olekukonko/tablewriter v0.0.5 h1:P2Ga83D34wi1o9J6Wh1mRuqd4mF/x/lgBS7N7AbDhec=
github.com/olekukonko/tablewriter v0.0.5/go.mod h1:hPp6KlRPjbx+hW8ykQs1w3UBbZlj6HuIJcUGPhkA7kY=
github.com/opencontainers/go-digest v1.0.0 h1:apOUWs51W5PlhuyGyz9FCeeBIOUDA/6nW8Oi/yOhh5U=
github.com/opencontainers/go-digest v1.0.0/go.mod h1:0JzlMkj0TRzQZfJkVvzbP0HBR3IKzErnv2BNG4W4MAM=
github.com/opencontainers/image-spec v1.1.1 h1:y0fUlFfIZhPF1W537XOLg0/fcx6zcHCJwooC2xJA040=
@@ -285,6 +320,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=
@@ -294,6 +333,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=
@@ -326,8 +366,6 @@ github.com/spf13/afero v1.11.0 h1:WJQKhtpdm3v2IzqG8VMqrr6Rf3UYpEF239Jy9wNepM8=
github.com/spf13/afero v1.11.0/go.mod h1:GH9Y3pIexgf1MTIWtNGyogA5MwRIDXGUr+hbWNoBjkY=
github.com/spf13/cast v1.6.0 h1:GEiTHELF+vaR5dhz3VqZfFSzZjYbgeKDpBxQVS4GYJ0=
github.com/spf13/cast v1.6.0/go.mod h1:ancEpBxwJDODSW/UG4rDrAqiKolqNNh2DX3mk86cAdo=
github.com/spf13/cobra v1.7.0 h1:hyqWnYt1ZQShIddO5kBpj3vu05/++x6tJ6dg8EC572I=
github.com/spf13/cobra v1.7.0/go.mod h1:uLxZILRyS/50WlhOIKD7W6V5bgeIt+4sICxh6uRMrb0=
github.com/spf13/pflag v1.0.5 h1:iy+VFUOCP1a+8yFto/drg2CJ5u0yRoB7fZw3DKv/JXA=
github.com/spf13/pflag v1.0.5/go.mod h1:McXfInJRrz4CZXVZOBLb0bTZqETkiAhM9Iw0y3An2Bg=
github.com/spf13/viper v1.18.2 h1:LUXCnvUvSM6FXAsj6nnfc8Q2tp1dIgUfY9Kc8GsSOiQ=
@@ -337,6 +375,8 @@ github.com/stretchr/objx v0.4.0/go.mod h1:YvHI0jy2hoMjB+UWwv71VJQ9isScKT/TqJzVSS
github.com/stretchr/objx v0.5.0/go.mod h1:Yh+to48EsGEfYuaHDzXPcE3xhTkx73EhmCGUpEOglKo=
github.com/stretchr/objx v0.5.2 h1:xuMeJ0Sdp5ZMRXx/aWO6RZxdr3beISkG5/G/aIRr3pY=
github.com/stretchr/objx v0.5.2/go.mod h1:FRsXN1f5AsAjCGJKqEizvkpNtU+EGNCLh3NxZ/8L+MA=
github.com/stretchr/testify v1.11.1 h1:7s2iGBzp5EwR7/aIZr8ao5+dra3wiQyKjjFuvgVKu7U=
github.com/stretchr/testify v1.11.1/go.mod h1:wZwfW3scLgRK+23gO65QZefKpKQRnfz6sD981Nm4B6U=
github.com/stretchr/testify v1.3.0/go.mod h1:M5WIy9Dh21IEIfnGCwXGc5bZfKNJtfHm1UVUgZn+9EI=
github.com/stretchr/testify v1.7.0/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg=
github.com/stretchr/testify v1.7.1/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg=
@@ -345,8 +385,6 @@ github.com/stretchr/testify v1.8.1/go.mod h1:w2LPCIKwWwSfY2zedu0+kehJoqGctiVI29o
github.com/stretchr/testify v1.8.2/go.mod h1:w2LPCIKwWwSfY2zedu0+kehJoqGctiVI29o6fzry7u4=
github.com/stretchr/testify v1.8.4/go.mod h1:sz/lmYIOXD/1dqDmKjjqLyZ2RngseejIcXlSw2iwfAo=
github.com/stretchr/testify v1.9.0/go.mod h1:r2ic/lqez/lEtzL7wO/rwa5dbSLXVDPFyf8C91i36aY=
github.com/stretchr/testify v1.11.1 h1:7s2iGBzp5EwR7/aIZr8ao5+dra3wiQyKjjFuvgVKu7U=
github.com/stretchr/testify v1.11.1/go.mod h1:wZwfW3scLgRK+23gO65QZefKpKQRnfz6sD981Nm4B6U=
github.com/subosito/gotenv v1.6.0 h1:9NlTDc1FTs4qu0DDq7AEtTPNw6SVm7uBMsUCUjABIf8=
github.com/subosito/gotenv v1.6.0/go.mod h1:Dk4QP5c2W3ibzajGcXpNraDfq2IrhjMIvMSWPKKo0FU=
github.com/tam7t/hpkp v0.0.0-20160821193359-2b70b4024ed5 h1:YqAladjX7xpA6BM04leXMWAEjS0mTZ5kUU9KRBriQJc=
@@ -384,6 +422,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=
@@ -398,6 +438,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=
@@ -417,18 +459,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=
@@ -436,30 +500,58 @@ golang.org/x/sys v0.0.0-20210616094352-59db8d763f22/go.mod h1:oPkhp1MJrh7nUepCBc
golang.org/x/sys v0.0.0-20220704084225-05e143d24a9e/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.0.0-20220715151400-c0bba94af5f8/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.0.0-20220811171246-fbc7d0a398ab/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.6.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.8.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.11.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.41.0 h1:Ivj+2Cp/ylzLiEU89QhWblYnOE9zerudt9Ftecq2C6k=
golang.org/x/sys v0.41.0/go.mod h1:OgkHotnGiDImocRcuBABYBEXf8A9a87e/uXjp9XT3ks=
golang.org/x/sys v0.6.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.8.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
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=
@@ -474,6 +566,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=
@@ -65,6 +65,7 @@ func setupAccountDataRouter() (*gin.Engine, *stubAdminService) {
nil,
nil,
nil,
nil,
)
router.GET("/api/v1/admin/accounts/data", h.ExportData)
@@ -9,6 +9,7 @@ import (
"errors"
"fmt"
"log"
"log/slog"
"net/http"
"strconv"
"strings"
@@ -57,6 +58,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 +76,7 @@ func NewAccountHandler(
sessionLimitCache service.SessionLimitCache,
rpmCache service.RPMCache,
tokenCacheInvalidator service.TokenCacheInvalidator,
gatewayCache service.GatewayCache,
) *AccountHandler {
return &AccountHandler{
adminService: adminService,
@@ -89,6 +92,7 @@ func NewAccountHandler(
sessionLimitCache: sessionLimitCache,
rpmCache: rpmCache,
tokenCacheInvalidator: tokenCacheInvalidator,
gatewayCache: gatewayCache,
}
}
@@ -206,6 +210,29 @@ func (h *AccountHandler) buildAccountResponseWithRuntime(ctx context.Context, ac
}
}
// 亲和客户端数据(启用亲和的账号始终返回 count,即使为 0)
if account.IsClientAffinityEnabled() {
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
} else {
zero := int64(0)
item.AffinityClientCount = &zero
}
} else {
zero := int64(0)
item.AffinityClientCount = &zero
}
} else {
zero := int64(0)
item.AffinityClientCount = &zero
}
}
return item
}
@@ -318,6 +345,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.IsClientAffinityEnabled() && 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 +390,18 @@ func (h *AccountHandler) List(c *gin.Context) {
}
}
// 注入亲和客户端数据到 DTO(启用亲和的账号始终返回 count,即使为 0)
if acc.IsClientAffinityEnabled() {
if clients, ok := affinityClients[acc.ID]; ok && len(clients) > 0 {
count := int64(len(clients))
item.AffinityClientCount = &count
item.AffinityClients = clients
} else {
zero := int64(0)
item.AffinityClientCount = &zero
}
}
result[i] = item
}
@@ -568,6 +622,16 @@ func (h *AccountHandler) Update(c *gin.Context) {
// base_rpm 输入校验:负值归零,超过 10000 截断
sanitizeExtraBaseRPM(req.Extra)
// 记录更新前的亲和状态,用于检测亲和关闭时清理 Redis 记录
oldAffinityEnabled := false
var oldGroupIDs []int64
if len(req.Extra) > 0 && h.gatewayCache != nil {
if oldAccount, err := h.adminService.GetAccount(c.Request.Context(), accountID); err == nil {
oldAffinityEnabled = oldAccount.IsClientAffinityEnabled()
oldGroupIDs = oldAccount.GroupIDs
}
}
// 确定是否跳过混合渠道检查
skipCheck := req.ConfirmMixedChannelRisk != nil && *req.ConfirmMixedChannelRisk
@@ -604,6 +668,15 @@ func (h *AccountHandler) Update(c *gin.Context) {
return
}
// 亲和关闭时清理 Redis 中的亲和记录
if oldAffinityEnabled && !account.IsClientAffinityEnabled() {
groupIDs := oldGroupIDs
if len(account.GroupIDs) > 0 {
groupIDs = mergeGroupIDs(oldGroupIDs, account.GroupIDs)
}
h.clearAccountAffinity(c.Request.Context(), accountID, groupIDs)
}
response.Success(c, h.buildAccountResponseWithRuntime(c.Request.Context(), account))
}
@@ -1341,6 +1414,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
}
@@ -1560,6 +1639,75 @@ func (h *AccountHandler) ResetQuota(c *gin.Context) {
response.Success(c, h.buildAccountResponseWithRuntime(c.Request.Context(), account))
}
// 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.IsClientAffinityEnabled() {
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,
)
}
}
// 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) {
@@ -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)
}
@@ -46,9 +46,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 存储配额
@@ -84,9 +85,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 存储配额
@@ -207,6 +209,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,
@@ -260,6 +263,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,
}
}
@@ -1002,6 +1006,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,
@@ -1013,6 +1019,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,
}
}
@@ -1092,6 +1105,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"`
@@ -1102,10 +1117,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"`
@@ -1116,6 +1140,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 配置
@@ -1138,14 +1169,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,
@@ -1156,6 +1196,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)
@@ -1198,13 +1245,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,
@@ -1215,6 +1274,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)
@@ -1515,3 +1581,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)
}
+15 -8
View File
@@ -135,14 +135,15 @@ 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,
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,
SortOrder: g.SortOrder,
}
if len(g.AccountGroups) > 0 {
out.AccountGroups = make([]AccountGroup, 0, len(g.AccountGroups))
@@ -264,6 +265,12 @@ func AccountFromServiceShallow(a *service.Account) *Account {
}
}
// 客户端亲和调度(Anthropic 和 Antigravity 账号)
if a.IsClientAffinityEnabled() {
enabled := true
out.ClientAffinityEnabled = &enabled
}
// 提取账号配额限制(apikey / bedrock 类型有效)
if a.IsAPIKeyOrBedrock() {
if limit := a.GetQuotaLimit(); limit > 0 {
+11 -1
View File
@@ -132,11 +132,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"`
@@ -148,6 +150,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 配置列表响应
+10
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"`
@@ -195,6 +197,14 @@ 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"`
// 亲和客户端数据(仅 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"`
@@ -440,6 +440,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,
@@ -632,6 +633,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 {
@@ -744,6 +746,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,
+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)
}
@@ -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, "",
)
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))
@@ -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,
+250 -7
View File
@@ -2,14 +2,42 @@ package repository
import (
"context"
_ "embed"
"fmt"
"strconv"
"time"
"github.com/Wei-Shaw/sub2api/internal/service"
"github.com/redis/go-redis/v9"
)
const stickySessionPrefix = "sticky_session:"
const (
stickySessionPrefix = "sticky_session:"
clientAffinityPrefix = "client_affinity:"
clientAffinityReversePrefix = "client_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
getAffinityScript = redis.NewScript(getAffinityLua)
updateAffinityScript = redis.NewScript(updateAffinityLua)
getAffinityCountScript = redis.NewScript(getAffinityCountLua)
getAffinityClientsScript = redis.NewScript(getAffinityClientsLua)
getAffinityClientsWithScoresScript = redis.NewScript(getAffinityClientsWithScoresLua)
clearAccountAffinityScript = redis.NewScript(clearAccountAffinityLua)
)
type gatewayCache struct {
rdb *redis.Client
@@ -19,6 +47,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 +79,218 @@ 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(client → accounts)
// 格式: client_affinity:{groupID}:{clientID}
func buildAffinityKey(groupID int64, clientID string) string {
return fmt.Sprintf("%s%d:%s", clientAffinityPrefix, groupID, clientID)
}
// buildAffinityReverseKey 构建反向亲和 key(account → clients)
// 格式: client_affinity_rev:{groupID}:{accountID}
func buildAffinityReverseKey(groupID int64, accountID int64) string {
return fmt.Sprintf("%s%d:%d", clientAffinityReversePrefix, groupID, accountID)
}
func (c *gatewayCache) GetClientAffinityAccounts(ctx context.Context, groupID int64, clientID string, ttl time.Duration) ([]int64, error) {
key := buildAffinityKey(groupID, 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
}
func (c *gatewayCache) UpdateClientAffinity(ctx context.Context, groupID int64, clientID string, accountID int64, ttl time.Duration) error {
fwdKey := buildAffinityKey(groupID, clientID)
revKey := buildAffinityReverseKey(groupID, accountID)
now := time.Now().Unix()
ttlSeconds := int64(ttl.Seconds())
expireThreshold := now - ttlSeconds
return updateAffinityScript.Run(ctx, c.rdb, []string{fwdKey, revKey},
now, ttlSeconds, accountID, expireThreshold, clientID,
).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) 组合查询反向索引。
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 的 clientID 去重
result := make(map[int64][]string, len(accountGroups))
seen := make(map[int64]map[string]struct{}, len(accountGroups))
for i, q := range queries {
clients, _ := cmds[i].StringSlice()
if len(clients) == 0 {
continue
}
if seen[q.accountID] == nil {
seen[q.accountID] = make(map[string]struct{})
}
for _, clientID := range clients {
if _, exists := seen[q.accountID][clientID]; !exists {
seen[q.accountID][clientID] = struct{}{}
result[q.accountID] = append(result[q.accountID], clientID)
}
}
}
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
}
// 合并跨组结果,同一 clientID 取最近的 lastActive
seen := make(map[string]int64) // clientID → max timestamp
for _, cmd := range cmds {
vals, _ := cmd.StringSlice()
// vals 格式: [clientID1, score1, clientID2, score2, ...]
for j := 0; j+1 < len(vals); j += 2 {
clientID := vals[j]
ts, _ := strconv.ParseInt(vals[j+1], 10, 64)
if existing, ok := seen[clientID]; !ok || ts > existing {
seen[clientID] = ts
}
}
}
result := make([]service.AffinityClient, 0, len(seen))
for clientID, ts := range seen {
result = append(result, service.AffinityClient{
ClientID: clientID,
LastActive: time.Unix(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
}
+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 {
@@ -129,7 +130,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] = client_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]
-- 获取反向索引中所有客户端 ID
local clients = redis.call('ZRANGE', rev_key, 0, -1)
if #clients == 0 then
return 0
end
-- 从每个客户端的正向索引中移除该账号
for _, client_id in ipairs(clients) do
local fwd_key = 'client_affinity:' .. group_id .. ':' .. client_id
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 #clients
@@ -0,0 +1,5 @@
-- 清理过期成员后返回亲和账号列表(按最近使用降序)
-- KEYS[1] = client_affinity:{groupID}:{clientID}
-- ARGV[1] = 过期阈值时间戳 (now - ttl)
redis.call('ZREMRANGEBYSCORE', KEYS[1], '-inf', ARGV[1])
return redis.call('ZREVRANGE', KEYS[1], 0, -1)
@@ -0,0 +1,5 @@
-- 清理过期成员后返回反向索引的 clientID 列表(按最近使用降序)
-- KEYS[1] = client_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,6 @@
-- 清理过期成员后返回反向索引的 clientID 列表及其 score(最后活跃时间戳)
-- KEYS[1] = client_affinity_rev:{groupID}:{accountID}
-- ARGV[1] = 过期阈值时间戳 (now - ttl)
-- 返回: {clientID1, score1, clientID2, score2, ...}(按最近使用降序)
redis.call('ZREMRANGEBYSCORE', KEYS[1], '-inf', ARGV[1])
return redis.call('ZREVRANGEBYSCORE', KEYS[1], '+inf', '-inf', 'WITHSCORES')
@@ -0,0 +1,5 @@
-- 清理过期成员后返回反向索引的成员数量
-- KEYS[1] = client_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,15 @@
-- 原子双写正向+反向索引
-- KEYS[1] = client_affinity:{groupID}:{clientID} (正向: client → accounts)
-- KEYS[2] = client_affinity_rev:{groupID}:{accountID} (反向: account → clients)
-- ARGV[1] = 当前时间戳 (score)
-- ARGV[2] = TTL 秒数
-- ARGV[3] = accountID (正向索引的成员)
-- ARGV[4] = 过期阈值时间戳 (now - ttl)
-- ARGV[5] = 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)
}
@@ -2904,7 +2904,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"
}
@@ -650,8 +649,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{
+18 -1
View File
@@ -261,6 +261,7 @@ 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/temp-unschedulable", h.Admin.Account.GetTempUnschedulable)
accounts.DELETE("/:id/temp-unschedulable", h.Admin.Account.ClearTempUnschedulable)
accounts.POST("/:id/schedulable", h.Admin.Account.SetSchedulable)
@@ -408,7 +409,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)
@@ -417,6 +418,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
@@ -138,6 +138,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 验证)
+84
View File
@@ -1188,6 +1188,90 @@ func (a *Account) IsSessionIDMaskingEnabled() bool {
return false
}
// IsClientAffinityEnabled 检查是否启用客户端亲和调度
// 仅适用于 Anthropic 账号(OAuth/SetupToken/APIKey)
// 启用后,新会话会优先调度到之前使用过的账号
func (a *Account) IsClientAffinityEnabled() bool {
if a.Platform != PlatformAnthropic {
return false
}
if a.Extra == nil {
return false
}
if v, ok := a.Extra["client_affinity_enabled"]; ok {
if enabled, ok := v.(bool); ok {
return enabled
}
}
return false
}
// 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
}
// IsCacheTTLOverrideEnabled 检查是否启用缓存 TTL 强制替换
// 仅适用于 Anthropic OAuth/SetupToken 类型账号
// 启用后将所有 cache creation tokens 归入指定的 TTL 类型(5m 或 1h)
+29 -10
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 存储配额
@@ -368,6 +370,10 @@ type ProxyExitInfoProber interface {
ProbeProxy(ctx context.Context, proxyURL string) (*ProxyExitInfo, int64, error)
}
type groupExistenceBatchReader interface {
ExistsByIDs(ctx context.Context, ids []int64) (map[int64]bool, error)
}
type proxyQualityTarget struct {
Target string
URL string
@@ -445,10 +451,6 @@ type userGroupRateBatchReader interface {
GetByUserIDs(ctx context.Context, userIDs []int64) (map[int64]map[int64]float64, error)
}
type groupExistenceBatchReader interface {
ExistsByIDs(ctx context.Context, ids []int64) (map[int64]bool, error)
}
// NewAdminService creates a new AdminService
func NewAdminService(
userRepo UserRepository,
@@ -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,
@@ -1135,6 +1145,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)
}
@@ -1673,7 +1673,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
@@ -1683,7 +1683,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
@@ -1692,6 +1692,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,
@@ -3595,7 +3598,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 {
@@ -3753,6 +3756,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
@@ -3767,7 +3773,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")
@@ -3780,6 +3786,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,24 @@ func (c *stubSmartRetryCache) DeleteSessionAccountID(_ context.Context, groupID
c.deleteCalls = append(c.deleteCalls, deleteSessionCall{groupID: groupID, sessionHash: sessionHash})
return nil
}
func (c *stubSmartRetryCache) GetClientAffinityAccounts(_ context.Context, _ int64, _ string, _ time.Duration) ([]int64, error) {
return nil, nil
}
func (c *stubSmartRetryCache) UpdateClientAffinity(_ context.Context, _ int64, _ string, _ int64, _ time.Duration) error {
return 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 {
@@ -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)
}
+23 -1
View File
@@ -116,7 +116,7 @@ func TestCheckErrorPolicy(t *testing.T) {
account: &Account{
ID: 15,
Type: AccountTypeOAuth,
Platform: PlatformAntigravity,
Platform: PlatformGemini, // 非 Antigravity 平台才有 401 升级逻辑
TempUnschedulableReason: `{"status_code":401,"until_unix":1735689600}`,
Credentials: map[string]any{
"temp_unschedulable_enabled": true,
@@ -133,6 +133,28 @@ func TestCheckErrorPolicy(t *testing.T) {
body: []byte(`unauthorized`),
expected: ErrorPolicyTempUnscheduled,
},
{
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",
account: &Account{
@@ -0,0 +1,800 @@
//go:build unit
package service
import (
"context"
"errors"
"sort"
"testing"
"time"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
// ---------------------------------------------------------------------------
// Mock: GatewayCache for affinity tests
// ---------------------------------------------------------------------------
// mockAffinityCache 为亲和调度测试提供可控的 GatewayCache mock。
// 通过 getCountBatchFunc 可以自定义 GetAccountAffinityCountBatch 的行为。
type mockAffinityCache struct {
getCountBatchFunc func(ctx context.Context, groupID int64, accountIDs []int64, ttl time.Duration) (map[int64]int64, error)
getCountBatchCalls int // 记录 GetAccountAffinityCountBatch 被调用次数
}
func (m *mockAffinityCache) GetSessionAccountID(_ context.Context, _ int64, _ string) (int64, error) {
return 0, errors.New("not found")
}
func (m *mockAffinityCache) SetSessionAccountID(_ context.Context, _ int64, _ string, _ int64, _ time.Duration) error {
return nil
}
func (m *mockAffinityCache) RefreshSessionTTL(_ context.Context, _ int64, _ string, _ time.Duration) error {
return nil
}
func (m *mockAffinityCache) DeleteSessionAccountID(_ context.Context, _ int64, _ string) error {
return nil
}
func (m *mockAffinityCache) GetClientAffinityAccounts(_ context.Context, _ int64, _ string, _ time.Duration) ([]int64, error) {
return nil, nil
}
func (m *mockAffinityCache) UpdateClientAffinity(_ context.Context, _ int64, _ string, _ int64, _ time.Duration) error {
return nil
}
func (m *mockAffinityCache) GetAccountAffinityCountBatch(ctx context.Context, groupID int64, accountIDs []int64, ttl time.Duration) (map[int64]int64, error) {
m.getCountBatchCalls++
if m.getCountBatchFunc != nil {
return m.getCountBatchFunc(ctx, groupID, accountIDs, ttl)
}
return map[int64]int64{}, nil
}
func (m *mockAffinityCache) GetAccountAffinityClientsBatch(_ context.Context, _ map[int64][]int64, _ time.Duration) (map[int64][]string, error) {
return map[int64][]string{}, nil
}
func (m *mockAffinityCache) GetAccountAffinityClientsWithScores(_ context.Context, _ int64, _ []int64, _ time.Duration) ([]AffinityClient, error) {
return nil, nil
}
func (m *mockAffinityCache) ClearAccountAffinity(_ context.Context, _ int64, _ []int64) error {
return nil
}
// ---------------------------------------------------------------------------
// Helper: 构造启用了客户端亲和的 Anthropic 账号
// ---------------------------------------------------------------------------
func newAffinityAccount(id int64, priority int, affinityEnabled bool) *Account {
acc := &Account{
ID: id,
Platform: PlatformAnthropic,
Priority: priority,
Status: StatusActive,
}
if affinityEnabled {
acc.Extra = map[string]any{"client_affinity_enabled": true}
}
return acc
}
func newAffinityAccountWithLoad(id int64, priority int, loadRate int, affinityCount int64, lastUsedAt *time.Time) accountWithLoad {
return accountWithLoad{
account: newAffinityAccount(id, priority, true),
loadInfo: &AccountLoadInfo{AccountID: id, LoadRate: loadRate},
affinityCount: affinityCount,
}
}
// ===========================================================================
// 1. filterByMinAffinityCount 测试
// ===========================================================================
func TestAffinityFilterByMinAffinityCount(t *testing.T) {
t.Run("empty slice returns empty", func(t *testing.T) {
result := filterByMinAffinityCount(nil)
require.Empty(t, result)
})
t.Run("single element returned as-is", func(t *testing.T) {
accounts := []accountWithLoad{
{account: &Account{ID: 1}, loadInfo: &AccountLoadInfo{}, affinityCount: 5},
}
result := filterByMinAffinityCount(accounts)
require.Len(t, result, 1)
require.Equal(t, int64(1), result[0].account.ID)
})
t.Run("all same affinityCount returns all", func(t *testing.T) {
accounts := []accountWithLoad{
{account: &Account{ID: 1}, loadInfo: &AccountLoadInfo{}, affinityCount: 3},
{account: &Account{ID: 2}, loadInfo: &AccountLoadInfo{}, affinityCount: 3},
{account: &Account{ID: 3}, loadInfo: &AccountLoadInfo{}, affinityCount: 3},
}
result := filterByMinAffinityCount(accounts)
require.Len(t, result, 3)
})
t.Run("filters to min affinityCount only", func(t *testing.T) {
accounts := []accountWithLoad{
{account: &Account{ID: 1}, loadInfo: &AccountLoadInfo{}, affinityCount: 10},
{account: &Account{ID: 2}, loadInfo: &AccountLoadInfo{}, affinityCount: 2},
{account: &Account{ID: 3}, loadInfo: &AccountLoadInfo{}, affinityCount: 5},
{account: &Account{ID: 4}, loadInfo: &AccountLoadInfo{}, affinityCount: 2},
}
result := filterByMinAffinityCount(accounts)
require.Len(t, result, 2)
require.Equal(t, int64(2), result[0].account.ID)
require.Equal(t, int64(4), result[1].account.ID)
})
t.Run("zero affinityCount is smallest", func(t *testing.T) {
accounts := []accountWithLoad{
{account: &Account{ID: 1}, loadInfo: &AccountLoadInfo{}, affinityCount: 5},
{account: &Account{ID: 2}, loadInfo: &AccountLoadInfo{}, affinityCount: 0},
{account: &Account{ID: 3}, loadInfo: &AccountLoadInfo{}, affinityCount: 3},
{account: &Account{ID: 4}, loadInfo: &AccountLoadInfo{}, affinityCount: 0},
}
result := filterByMinAffinityCount(accounts)
require.Len(t, result, 2)
require.Equal(t, int64(2), result[0].account.ID)
require.Equal(t, int64(4), result[1].account.ID)
})
t.Run("preserves order within same affinityCount", func(t *testing.T) {
accounts := []accountWithLoad{
{account: &Account{ID: 5}, loadInfo: &AccountLoadInfo{}, affinityCount: 1},
{account: &Account{ID: 3}, loadInfo: &AccountLoadInfo{}, affinityCount: 1},
{account: &Account{ID: 7}, loadInfo: &AccountLoadInfo{}, affinityCount: 2},
{account: &Account{ID: 1}, loadInfo: &AccountLoadInfo{}, affinityCount: 1},
}
result := filterByMinAffinityCount(accounts)
require.Len(t, result, 3)
// 验证保持原始顺序
require.Equal(t, int64(5), result[0].account.ID)
require.Equal(t, int64(3), result[1].account.ID)
require.Equal(t, int64(1), result[2].account.ID)
})
}
// ===========================================================================
// 2. populateAffinityCounts 测试
// ===========================================================================
func TestAffinityPopulateAffinityCounts(t *testing.T) {
ctx := context.Background()
t.Run("nil cache does not panic", func(t *testing.T) {
svc := &GatewayService{cache: nil}
accounts := []accountWithLoad{
{account: newAffinityAccount(1, 1, true), loadInfo: &AccountLoadInfo{}},
}
// 不应 panic
svc.populateAffinityCounts(ctx, accounts, 0)
// affinityCount 保持零值
require.Equal(t, int64(0), accounts[0].affinityCount)
})
t.Run("empty accounts returns immediately", func(t *testing.T) {
cache := &mockAffinityCache{}
svc := &GatewayService{cache: cache}
svc.populateAffinityCounts(ctx, nil, 0)
require.Equal(t, 0, cache.getCountBatchCalls, "should not call Redis for empty accounts")
})
t.Run("no affinity-enabled accounts skips Redis call", func(t *testing.T) {
cache := &mockAffinityCache{}
svc := &GatewayService{cache: cache}
accounts := []accountWithLoad{
// Anthropic 但未启用亲和
{account: newAffinityAccount(1, 1, false), loadInfo: &AccountLoadInfo{}},
// 非 Anthropic 平台
{account: &Account{ID: 2, Platform: PlatformOpenAI}, loadInfo: &AccountLoadInfo{}},
}
svc.populateAffinityCounts(ctx, accounts, 0)
require.Equal(t, 0, cache.getCountBatchCalls, "should skip Redis when no affinity-enabled accounts")
})
t.Run("correctly populates affinityCount from Redis", func(t *testing.T) {
cache := &mockAffinityCache{
getCountBatchFunc: func(_ context.Context, _ int64, accountIDs []int64, _ time.Duration) (map[int64]int64, error) {
result := map[int64]int64{}
for _, id := range accountIDs {
switch id {
case 1:
result[1] = 5
case 2:
result[2] = 0
case 3:
result[3] = 12
}
}
return result, nil
},
}
svc := &GatewayService{cache: cache}
accounts := []accountWithLoad{
{account: newAffinityAccount(1, 1, true), loadInfo: &AccountLoadInfo{}},
{account: newAffinityAccount(2, 1, false), loadInfo: &AccountLoadInfo{}}, // 未启用,但仍在列表中
{account: newAffinityAccount(3, 1, true), loadInfo: &AccountLoadInfo{}},
}
svc.populateAffinityCounts(ctx, accounts, 100)
require.Equal(t, 1, cache.getCountBatchCalls, "should call Redis exactly once")
require.Equal(t, int64(5), accounts[0].affinityCount, "account 1 should have count 5")
require.Equal(t, int64(0), accounts[1].affinityCount, "account 2 should have count 0")
require.Equal(t, int64(12), accounts[2].affinityCount, "account 3 should have count 12")
})
t.Run("Redis error degrades gracefully with counts at 0", func(t *testing.T) {
cache := &mockAffinityCache{
getCountBatchFunc: func(_ context.Context, _ int64, _ []int64, _ time.Duration) (map[int64]int64, error) {
return nil, errors.New("redis connection refused")
},
}
svc := &GatewayService{cache: cache}
accounts := []accountWithLoad{
{account: newAffinityAccount(1, 1, true), loadInfo: &AccountLoadInfo{}},
{account: newAffinityAccount(2, 1, true), loadInfo: &AccountLoadInfo{}},
}
svc.populateAffinityCounts(ctx, accounts, 0)
require.Equal(t, 1, cache.getCountBatchCalls)
require.Equal(t, int64(0), accounts[0].affinityCount, "should remain 0 on error")
require.Equal(t, int64(0), accounts[1].affinityCount, "should remain 0 on error")
})
t.Run("partial Redis result fills only known accounts", func(t *testing.T) {
cache := &mockAffinityCache{
getCountBatchFunc: func(_ context.Context, _ int64, _ []int64, _ time.Duration) (map[int64]int64, error) {
// 只返回部分账号的计数
return map[int64]int64{1: 7}, nil
},
}
svc := &GatewayService{cache: cache}
accounts := []accountWithLoad{
{account: newAffinityAccount(1, 1, true), loadInfo: &AccountLoadInfo{}},
{account: newAffinityAccount(2, 1, true), loadInfo: &AccountLoadInfo{}},
}
svc.populateAffinityCounts(ctx, accounts, 0)
require.Equal(t, int64(7), accounts[0].affinityCount)
require.Equal(t, int64(0), accounts[1].affinityCount, "missing account should default to 0")
})
t.Run("queries all account IDs regardless of affinity status", func(t *testing.T) {
// 验证:只要有至少一个 affinity-enabled 账号,就查询 ALL 账号的计数
var queriedIDs []int64
cache := &mockAffinityCache{
getCountBatchFunc: func(_ context.Context, _ int64, accountIDs []int64, _ time.Duration) (map[int64]int64, error) {
queriedIDs = accountIDs
return map[int64]int64{}, nil
},
}
svc := &GatewayService{cache: cache}
accounts := []accountWithLoad{
{account: newAffinityAccount(1, 1, true), loadInfo: &AccountLoadInfo{}},
{account: newAffinityAccount(2, 1, false), loadInfo: &AccountLoadInfo{}},
{account: &Account{ID: 3, Platform: PlatformOpenAI}, loadInfo: &AccountLoadInfo{}},
}
svc.populateAffinityCounts(ctx, accounts, 0)
require.Equal(t, 1, cache.getCountBatchCalls)
require.Equal(t, []int64{1, 2, 3}, queriedIDs, "should query ALL account IDs, not just affinity-enabled ones")
})
}
// ===========================================================================
// 3. Layer 1 排序链测试(sort.SliceStable 中 affinityCount 的正确性)
// ===========================================================================
func TestAffinityLayer1SortChain(t *testing.T) {
now := time.Now()
earlier := now.Add(-1 * time.Hour)
muchEarlier := now.Add(-2 * time.Hour)
// 复现 Layer 1 的排序逻辑
sortByLayer1 := func(accounts []accountWithLoad) {
sort.SliceStable(accounts, func(i, j int) bool {
a, b := accounts[i], accounts[j]
if a.account.Priority != b.account.Priority {
return a.account.Priority < b.account.Priority
}
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
case a.account.LastUsedAt != nil && b.account.LastUsedAt == nil:
return false
case a.account.LastUsedAt == nil && b.account.LastUsedAt == nil:
return false
default:
return a.account.LastUsedAt.Before(*b.account.LastUsedAt)
}
})
}
t.Run("same priority same loadRate sorts by affinityCount asc", func(t *testing.T) {
accounts := []accountWithLoad{
{account: &Account{ID: 1, Priority: 1, LastUsedAt: &now}, loadInfo: &AccountLoadInfo{LoadRate: 50}, affinityCount: 10},
{account: &Account{ID: 2, Priority: 1, LastUsedAt: &now}, loadInfo: &AccountLoadInfo{LoadRate: 50}, affinityCount: 2},
{account: &Account{ID: 3, Priority: 1, LastUsedAt: &now}, loadInfo: &AccountLoadInfo{LoadRate: 50}, affinityCount: 5},
}
sortByLayer1(accounts)
require.Equal(t, int64(2), accounts[0].account.ID, "lowest affinityCount first")
require.Equal(t, int64(3), accounts[1].account.ID)
require.Equal(t, int64(1), accounts[2].account.ID, "highest affinityCount last")
})
t.Run("priority takes precedence over affinityCount", func(t *testing.T) {
accounts := []accountWithLoad{
{account: &Account{ID: 1, Priority: 2, LastUsedAt: &now}, loadInfo: &AccountLoadInfo{LoadRate: 50}, affinityCount: 0},
{account: &Account{ID: 2, Priority: 1, LastUsedAt: &now}, loadInfo: &AccountLoadInfo{LoadRate: 50}, affinityCount: 100},
}
sortByLayer1(accounts)
require.Equal(t, int64(2), accounts[0].account.ID, "lower priority wins despite higher affinityCount")
})
t.Run("loadRate takes precedence over affinityCount", func(t *testing.T) {
accounts := []accountWithLoad{
{account: &Account{ID: 1, Priority: 1, LastUsedAt: &now}, loadInfo: &AccountLoadInfo{LoadRate: 80}, affinityCount: 0},
{account: &Account{ID: 2, Priority: 1, LastUsedAt: &now}, loadInfo: &AccountLoadInfo{LoadRate: 20}, affinityCount: 100},
}
sortByLayer1(accounts)
require.Equal(t, int64(2), accounts[0].account.ID, "lower loadRate wins despite higher affinityCount")
})
t.Run("affinityCount takes precedence over LRU", func(t *testing.T) {
accounts := []accountWithLoad{
{account: &Account{ID: 1, Priority: 1, LastUsedAt: &muchEarlier}, loadInfo: &AccountLoadInfo{LoadRate: 50}, affinityCount: 5},
{account: &Account{ID: 2, Priority: 1, LastUsedAt: &now}, loadInfo: &AccountLoadInfo{LoadRate: 50}, affinityCount: 1},
}
sortByLayer1(accounts)
require.Equal(t, int64(2), accounts[0].account.ID, "lower affinityCount wins despite older LRU")
})
t.Run("same affinityCount falls through to LRU", func(t *testing.T) {
accounts := []accountWithLoad{
{account: &Account{ID: 1, Priority: 1, LastUsedAt: &now}, loadInfo: &AccountLoadInfo{LoadRate: 50}, affinityCount: 3},
{account: &Account{ID: 2, Priority: 1, LastUsedAt: &earlier}, loadInfo: &AccountLoadInfo{LoadRate: 50}, affinityCount: 3},
{account: &Account{ID: 3, Priority: 1, LastUsedAt: &muchEarlier}, loadInfo: &AccountLoadInfo{LoadRate: 50}, affinityCount: 3},
}
sortByLayer1(accounts)
require.Equal(t, int64(3), accounts[0].account.ID, "LRU: oldest used first")
require.Equal(t, int64(2), accounts[1].account.ID)
require.Equal(t, int64(1), accounts[2].account.ID, "LRU: most recently used last")
})
t.Run("full chain: priority > loadRate > affinityCount > LRU", func(t *testing.T) {
accounts := []accountWithLoad{
// 优先级 2 - 不管其他维度如何,排在后面
{account: &Account{ID: 10, Priority: 2, LastUsedAt: &muchEarlier}, loadInfo: &AccountLoadInfo{LoadRate: 0}, affinityCount: 0},
// 优先级 1,负载 80% - 负载高
{account: &Account{ID: 20, Priority: 1, LastUsedAt: &muchEarlier}, loadInfo: &AccountLoadInfo{LoadRate: 80}, affinityCount: 0},
// 优先级 1,负载 20%,亲和 5
{account: &Account{ID: 30, Priority: 1, LastUsedAt: &muchEarlier}, loadInfo: &AccountLoadInfo{LoadRate: 20}, affinityCount: 5},
// 优先级 1,负载 20%,亲和 1,最近使用
{account: &Account{ID: 40, Priority: 1, LastUsedAt: &now}, loadInfo: &AccountLoadInfo{LoadRate: 20}, affinityCount: 1},
// 优先级 1,负载 20%,亲和 1,更早使用(应排最前)
{account: &Account{ID: 50, Priority: 1, LastUsedAt: &earlier}, loadInfo: &AccountLoadInfo{LoadRate: 20}, affinityCount: 1},
}
sortByLayer1(accounts)
// 期望排序:50 → 40 → 30 → 20 → 10
require.Equal(t, int64(50), accounts[0].account.ID, "best: p1, lr20, ac1, LRU earlier")
require.Equal(t, int64(40), accounts[1].account.ID, "second: p1, lr20, ac1, LRU now")
require.Equal(t, int64(30), accounts[2].account.ID, "third: p1, lr20, ac5")
require.Equal(t, int64(20), accounts[3].account.ID, "fourth: p1, lr80")
require.Equal(t, int64(10), accounts[4].account.ID, "last: p2")
})
}
// ===========================================================================
// 4. Layer 2 分层过滤链完整性测试
// ===========================================================================
func TestAffinityLayer2FilterChain(t *testing.T) {
now := time.Now()
earlier := now.Add(-1 * time.Hour)
muchEarlier := now.Add(-2 * time.Hour)
// 模拟 Layer 2 的完整过滤链:Priority → LoadRate → AffinityCount → LRU
applyLayer2 := func(accounts []accountWithLoad) *accountWithLoad {
candidates := filterByMinPriority(accounts)
candidates = filterByMinLoadRate(candidates)
candidates = filterByMinAffinityCount(candidates)
return selectByLRU(candidates, false)
}
t.Run("priority different - affinityCount does not matter", func(t *testing.T) {
accounts := []accountWithLoad{
{account: &Account{ID: 1, Priority: 2, LastUsedAt: &muchEarlier}, loadInfo: &AccountLoadInfo{LoadRate: 0}, affinityCount: 0},
{account: &Account{ID: 2, Priority: 1, LastUsedAt: &now}, loadInfo: &AccountLoadInfo{LoadRate: 50}, affinityCount: 100},
}
selected := applyLayer2(accounts)
require.NotNil(t, selected)
require.Equal(t, int64(2), selected.account.ID, "higher priority dimension overrides affinityCount")
})
t.Run("same priority same loadRate different affinityCount", func(t *testing.T) {
accounts := []accountWithLoad{
{account: &Account{ID: 1, Priority: 1, LastUsedAt: &now}, loadInfo: &AccountLoadInfo{LoadRate: 30}, affinityCount: 10},
{account: &Account{ID: 2, Priority: 1, LastUsedAt: &now}, loadInfo: &AccountLoadInfo{LoadRate: 30}, affinityCount: 2},
{account: &Account{ID: 3, Priority: 1, LastUsedAt: &now}, loadInfo: &AccountLoadInfo{LoadRate: 30}, affinityCount: 5},
}
selected := applyLayer2(accounts)
require.NotNil(t, selected)
require.Equal(t, int64(2), selected.account.ID, "lowest affinityCount wins")
})
t.Run("same priority same loadRate same affinityCount falls through to LRU", func(t *testing.T) {
accounts := []accountWithLoad{
{account: &Account{ID: 1, Priority: 1, LastUsedAt: &now}, loadInfo: &AccountLoadInfo{LoadRate: 30}, affinityCount: 5},
{account: &Account{ID: 2, Priority: 1, LastUsedAt: &earlier}, loadInfo: &AccountLoadInfo{LoadRate: 30}, affinityCount: 5},
{account: &Account{ID: 3, Priority: 1, LastUsedAt: &muchEarlier}, loadInfo: &AccountLoadInfo{LoadRate: 30}, affinityCount: 5},
}
selected := applyLayer2(accounts)
require.NotNil(t, selected)
require.Equal(t, int64(3), selected.account.ID, "LRU selects oldest")
})
t.Run("loadRate different overrides affinityCount", func(t *testing.T) {
accounts := []accountWithLoad{
{account: &Account{ID: 1, Priority: 1, LastUsedAt: &now}, loadInfo: &AccountLoadInfo{LoadRate: 80}, affinityCount: 0},
{account: &Account{ID: 2, Priority: 1, LastUsedAt: &now}, loadInfo: &AccountLoadInfo{LoadRate: 10}, affinityCount: 100},
}
selected := applyLayer2(accounts)
require.NotNil(t, selected)
require.Equal(t, int64(2), selected.account.ID, "lower loadRate wins over lower affinityCount")
})
t.Run("full chain integration: p → lr → ac → lru", func(t *testing.T) {
accounts := []accountWithLoad{
// p=2 淘汰
{account: &Account{ID: 1, Priority: 2, LastUsedAt: &muchEarlier}, loadInfo: &AccountLoadInfo{LoadRate: 0}, affinityCount: 0},
// p=1, lr=50 淘汰
{account: &Account{ID: 2, Priority: 1, LastUsedAt: &muchEarlier}, loadInfo: &AccountLoadInfo{LoadRate: 50}, affinityCount: 0},
// p=1, lr=10, ac=8 淘汰
{account: &Account{ID: 3, Priority: 1, LastUsedAt: &muchEarlier}, loadInfo: &AccountLoadInfo{LoadRate: 10}, affinityCount: 8},
// p=1, lr=10, ac=2, lru=now 淘汰
{account: &Account{ID: 4, Priority: 1, LastUsedAt: &now}, loadInfo: &AccountLoadInfo{LoadRate: 10}, affinityCount: 2},
// p=1, lr=10, ac=2, lru=muchEarlier → 胜出
{account: &Account{ID: 5, Priority: 1, LastUsedAt: &muchEarlier}, loadInfo: &AccountLoadInfo{LoadRate: 10}, affinityCount: 2},
}
selected := applyLayer2(accounts)
require.NotNil(t, selected)
require.Equal(t, int64(5), selected.account.ID, "full chain selects ID=5")
})
t.Run("empty input returns nil", func(t *testing.T) {
selected := applyLayer2(nil)
require.Nil(t, selected)
})
t.Run("single account always selected", func(t *testing.T) {
accounts := []accountWithLoad{
{account: &Account{ID: 42, Priority: 1, LastUsedAt: &now}, loadInfo: &AccountLoadInfo{LoadRate: 50}, affinityCount: 100},
}
selected := applyLayer2(accounts)
require.NotNil(t, selected)
require.Equal(t, int64(42), selected.account.ID)
})
t.Run("affinityCount zero preferred among same p and lr", func(t *testing.T) {
accounts := []accountWithLoad{
{account: &Account{ID: 1, Priority: 1, LastUsedAt: &now}, loadInfo: &AccountLoadInfo{LoadRate: 0}, affinityCount: 5},
{account: &Account{ID: 2, Priority: 1, LastUsedAt: &now}, loadInfo: &AccountLoadInfo{LoadRate: 0}, affinityCount: 0},
{account: &Account{ID: 3, Priority: 1, LastUsedAt: &now}, loadInfo: &AccountLoadInfo{LoadRate: 0}, affinityCount: 3},
}
selected := applyLayer2(accounts)
require.NotNil(t, selected)
require.Equal(t, int64(2), selected.account.ID, "zero affinityCount preferred")
})
}
// ===========================================================================
// 5. populateAffinityCounts + filterByMinAffinityCount 联合测试
// ===========================================================================
func TestAffinityPopulateAndFilterIntegration(t *testing.T) {
ctx := context.Background()
t.Run("populate then filter selects least-loaded accounts", func(t *testing.T) {
cache := &mockAffinityCache{
getCountBatchFunc: func(_ context.Context, _ int64, _ []int64, _ time.Duration) (map[int64]int64, error) {
return map[int64]int64{
1: 10,
2: 3,
3: 3,
4: 7,
}, nil
},
}
svc := &GatewayService{cache: cache}
accounts := []accountWithLoad{
{account: newAffinityAccount(1, 1, true), loadInfo: &AccountLoadInfo{}},
{account: newAffinityAccount(2, 1, true), loadInfo: &AccountLoadInfo{}},
{account: newAffinityAccount(3, 1, true), loadInfo: &AccountLoadInfo{}},
{account: newAffinityAccount(4, 1, true), loadInfo: &AccountLoadInfo{}},
}
svc.populateAffinityCounts(ctx, accounts, 0)
result := filterByMinAffinityCount(accounts)
require.Len(t, result, 2)
require.Equal(t, int64(2), result[0].account.ID)
require.Equal(t, int64(3), result[1].account.ID)
})
t.Run("Redis failure results in all accounts having 0 affinityCount", func(t *testing.T) {
cache := &mockAffinityCache{
getCountBatchFunc: func(_ context.Context, _ int64, _ []int64, _ time.Duration) (map[int64]int64, error) {
return nil, errors.New("timeout")
},
}
svc := &GatewayService{cache: cache}
accounts := []accountWithLoad{
{account: newAffinityAccount(1, 1, true), loadInfo: &AccountLoadInfo{}},
{account: newAffinityAccount(2, 1, true), loadInfo: &AccountLoadInfo{}},
}
svc.populateAffinityCounts(ctx, accounts, 0)
result := filterByMinAffinityCount(accounts)
// 全部为 0,全部返回
require.Len(t, result, 2, "all accounts should pass filter when Redis fails (all have count 0)")
})
}
// ===========================================================================
// 6. IsClientAffinityEnabled 边界测试
// ===========================================================================
func TestAffinityIsClientAffinityEnabled(t *testing.T) {
t.Run("Anthropic with enabled flag", func(t *testing.T) {
acc := &Account{
Platform: PlatformAnthropic,
Extra: map[string]any{"client_affinity_enabled": true},
}
assert.True(t, acc.IsClientAffinityEnabled())
})
t.Run("Anthropic with disabled flag", func(t *testing.T) {
acc := &Account{
Platform: PlatformAnthropic,
Extra: map[string]any{"client_affinity_enabled": false},
}
assert.False(t, acc.IsClientAffinityEnabled())
})
t.Run("Anthropic with nil Extra", func(t *testing.T) {
acc := &Account{
Platform: PlatformAnthropic,
Extra: nil,
}
assert.False(t, acc.IsClientAffinityEnabled())
})
t.Run("Anthropic without the key", func(t *testing.T) {
acc := &Account{
Platform: PlatformAnthropic,
Extra: map[string]any{"other_key": true},
}
assert.False(t, acc.IsClientAffinityEnabled())
})
t.Run("non-Anthropic platform always false", func(t *testing.T) {
platforms := []string{PlatformOpenAI, PlatformGemini, PlatformAntigravity}
for _, p := range platforms {
acc := &Account{
Platform: p,
Extra: map[string]any{"client_affinity_enabled": true},
}
assert.False(t, acc.IsClientAffinityEnabled(), "platform=%s should not support affinity", p)
}
})
t.Run("wrong type for enabled value", func(t *testing.T) {
acc := &Account{
Platform: PlatformAnthropic,
Extra: map[string]any{"client_affinity_enabled": "true"}, // string 而非 bool
}
assert.False(t, acc.IsClientAffinityEnabled(), "string 'true' should not enable affinity")
})
}
// ===========================================================================
// GetAffinityZone 测试
// ===========================================================================
func TestGetAffinityZone(t *testing.T) {
makeAccount := func(enabled bool, base int, buffer any) *Account {
extra := map[string]any{"client_affinity_enabled": enabled}
if base > 0 {
extra["affinity_base"] = base
}
if buffer != nil {
extra["affinity_buffer"] = buffer
}
return &Account{
Platform: PlatformAnthropic,
Extra: extra,
}
}
t.Run("affinity disabled always green", func(t *testing.T) {
acc := makeAccount(false, 5, 3)
assert.Equal(t, AffinityZoneGreen, acc.GetAffinityZone(100))
})
t.Run("no base configured always green", func(t *testing.T) {
acc := makeAccount(true, 0, nil)
assert.Equal(t, AffinityZoneGreen, acc.GetAffinityZone(100))
})
t.Run("within base limit is green", func(t *testing.T) {
acc := makeAccount(true, 5, 3)
assert.Equal(t, AffinityZoneGreen, acc.GetAffinityZone(0))
assert.Equal(t, AffinityZoneGreen, acc.GetAffinityZone(3))
assert.Equal(t, AffinityZoneGreen, acc.GetAffinityZone(5))
})
t.Run("no buffer configured infinite yellow", func(t *testing.T) {
acc := makeAccount(true, 5, nil) // buffer not set → infinite yellow
assert.Equal(t, AffinityZoneYellow, acc.GetAffinityZone(6))
assert.Equal(t, AffinityZoneYellow, acc.GetAffinityZone(100))
assert.Equal(t, AffinityZoneYellow, acc.GetAffinityZone(9999))
})
t.Run("buffer zero no yellow zone", func(t *testing.T) {
acc := makeAccount(true, 5, 0) // buffer=0 → no yellow, direct red
assert.Equal(t, AffinityZoneGreen, acc.GetAffinityZone(5))
assert.Equal(t, AffinityZoneRed, acc.GetAffinityZone(6))
assert.Equal(t, AffinityZoneRed, acc.GetAffinityZone(100))
})
t.Run("within buffer is yellow", func(t *testing.T) {
acc := makeAccount(true, 5, 3)
assert.Equal(t, AffinityZoneYellow, acc.GetAffinityZone(6))
assert.Equal(t, AffinityZoneYellow, acc.GetAffinityZone(7))
assert.Equal(t, AffinityZoneYellow, acc.GetAffinityZone(8)) // base(5)+buffer(3)=8
})
t.Run("beyond buffer is red", func(t *testing.T) {
acc := makeAccount(true, 5, 3)
assert.Equal(t, AffinityZoneRed, acc.GetAffinityZone(9))
assert.Equal(t, AffinityZoneRed, acc.GetAffinityZone(100))
})
t.Run("boundary exactly at base", func(t *testing.T) {
acc := makeAccount(true, 10, 5)
assert.Equal(t, AffinityZoneGreen, acc.GetAffinityZone(10))
assert.Equal(t, AffinityZoneYellow, acc.GetAffinityZone(11))
})
t.Run("boundary exactly at base plus buffer", func(t *testing.T) {
acc := makeAccount(true, 10, 5)
assert.Equal(t, AffinityZoneYellow, acc.GetAffinityZone(15))
assert.Equal(t, AffinityZoneRed, acc.GetAffinityZone(16))
})
}
// ===========================================================================
// classifyByAffinityZone 测试
// ===========================================================================
func TestClassifyByAffinityZone(t *testing.T) {
makeAWL := func(id int64, base int, buffer any, count int64) accountWithLoad {
extra := map[string]any{"client_affinity_enabled": true}
if base > 0 {
extra["affinity_base"] = base
}
if buffer != nil {
extra["affinity_buffer"] = buffer
}
return accountWithLoad{
account: &Account{ID: id, Platform: PlatformAnthropic, Extra: extra},
loadInfo: &AccountLoadInfo{AccountID: id},
affinityCount: count,
}
}
t.Run("empty input returns empty", func(t *testing.T) {
result := classifyByAffinityZone(nil)
require.Empty(t, result)
})
t.Run("no zone config returns all", func(t *testing.T) {
// 没有账号配置 affinity_base → 原样返回
accs := []accountWithLoad{
{account: newAffinityAccount(1, 50, true), loadInfo: &AccountLoadInfo{AccountID: 1}},
{account: newAffinityAccount(2, 50, true), loadInfo: &AccountLoadInfo{AccountID: 2}},
}
result := classifyByAffinityZone(accs)
require.Len(t, result, 2)
})
t.Run("greens preferred over yellows", func(t *testing.T) {
accs := []accountWithLoad{
makeAWL(1, 5, 3, 3), // green (3 ≤ 5)
makeAWL(2, 5, 3, 7), // yellow (5 < 7 ≤ 8)
makeAWL(3, 5, 3, 2), // green (2 ≤ 5)
}
result := classifyByAffinityZone(accs)
require.Len(t, result, 2)
ids := []int64{result[0].account.ID, result[1].account.ID}
sort.Slice(ids, func(i, j int) bool { return ids[i] < ids[j] })
assert.Equal(t, []int64{1, 3}, ids)
})
t.Run("reds excluded", func(t *testing.T) {
accs := []accountWithLoad{
makeAWL(1, 5, 3, 10), // red (10 > 8)
makeAWL(2, 5, 3, 6), // yellow (5 < 6 ≤ 8)
makeAWL(3, 5, 3, 9), // red (9 > 8)
}
result := classifyByAffinityZone(accs)
require.Len(t, result, 1)
assert.Equal(t, int64(2), result[0].account.ID)
})
t.Run("all red returns empty", func(t *testing.T) {
accs := []accountWithLoad{
makeAWL(1, 5, 0, 6), // buffer=0 → red (6 > 5)
makeAWL(2, 5, 0, 10), // buffer=0 → red (10 > 5)
}
result := classifyByAffinityZone(accs)
require.Empty(t, result)
})
t.Run("mixed with unconfigured accounts", func(t *testing.T) {
// 账号 1: 配置了 base=5,buffer=3 → green(3≤5)
// 账号 2: 未配置 base → 视为 green
// 账号 3: 配置了 base=5,buffer=3 → red(10>8)
accs := []accountWithLoad{
makeAWL(1, 5, 3, 3),
{account: newAffinityAccount(2, 50, true), loadInfo: &AccountLoadInfo{AccountID: 2}, affinityCount: 20},
makeAWL(3, 5, 3, 10),
}
result := classifyByAffinityZone(accs)
require.Len(t, result, 2)
ids := []int64{result[0].account.ID, result[1].account.ID}
sort.Slice(ids, func(i, j int) bool { return ids[i] < ids[j] })
assert.Equal(t, []int64{1, 2}, ids)
})
t.Run("infinite yellow never reds", func(t *testing.T) {
// buffer 未配置 → 无限黄区,永不红区
accs := []accountWithLoad{
makeAWL(1, 5, nil, 100), // yellow (100 > 5, no buffer → infinite yellow)
makeAWL(2, 5, nil, 3), // green (3 ≤ 5)
}
result := classifyByAffinityZone(accs)
// green 优先
require.Len(t, result, 1)
assert.Equal(t, int64(2), result[0].account.ID)
})
t.Run("only yellows when no greens", func(t *testing.T) {
accs := []accountWithLoad{
makeAWL(1, 5, nil, 10), // yellow
makeAWL(2, 5, nil, 20), // yellow
}
result := classifyByAffinityZone(accs)
require.Len(t, result, 2)
})
}
@@ -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,24 @@ 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) GetClientAffinityAccounts(_ context.Context, _ int64, _ string, _ time.Duration) ([]int64, error) {
return nil, nil
}
func (s *stickyGatewayCacheHotpathStub) UpdateClientAffinity(_ context.Context, _ int64, _ string, _ int64, _ time.Duration) error {
return 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)
@@ -235,6 +235,25 @@ func (m *mockGatewayCacheForPlatform) DeleteSessionAccountID(ctx context.Context
return nil
}
func (m *mockGatewayCacheForPlatform) GetClientAffinityAccounts(_ context.Context, _ int64, _ string, _ time.Duration) ([]int64, error) {
return nil, nil
}
func (m *mockGatewayCacheForPlatform) UpdateClientAffinity(_ context.Context, _ int64, _ string, _ int64, _ time.Duration) error {
return 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
@@ -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
}
+296 -11
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{}
@@ -329,6 +337,10 @@ var (
sessionIDRegex = regexp.MustCompile(`session_([a-f0-9-]{36})`)
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 版等
// 注意:前缀之间不应存在包含关系,否则会导致冗余匹配
@@ -389,6 +401,34 @@ 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
// GetClientAffinityAccounts 获取客户端亲和账号列表(按最近使用降序),同时清理过期成员
GetClientAffinityAccounts(ctx context.Context, groupID int64, clientID string, ttl time.Duration) ([]int64, error)
// UpdateClientAffinity 添加/更新客户端亲和关系(更新 score 为当前时间戳,刷新 key TTL)
UpdateClientAffinity(ctx context.Context, groupID 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
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
}
// AffinityClient 亲和客户端信息(含最后活跃时间)
type AffinityClient struct {
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
@@ -459,6 +499,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
@@ -1090,7 +1144,7 @@ func (s *GatewayService) SelectAccountForModelWithExclusions(ctx context.Context
}
// SelectAccountWithLoadAwareness selects account with load-awareness and wait plan.
// metadataUserID: 已废弃参数,会话限制现在统一使用 sessionHash
// metadataUserID: 用于客户端亲和调度,从中提取客户端 ID
func (s *GatewayService) SelectAccountWithLoadAwareness(ctx context.Context, groupID *int64, sessionHash string, requestedModel string, excludedIDs map[int64]struct{}, metadataUserID string) (*AccountSelectionResult, error) {
// 调试日志:记录调度入口参数
excludedIDsList := make([]int64, 0, len(excludedIDs))
@@ -1121,6 +1175,9 @@ func (s *GatewayService) SelectAccountWithLoadAwareness(ctx context.Context, gro
}
}
// 提取客户端 ID(用于客户端亲和调度)
affinityClientID := extractClientIDFromMetadata(metadataUserID)
if s.debugModelRoutingEnabled() && requestedModel != "" {
groupPlatform := ""
if group != nil {
@@ -1384,7 +1441,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 {
@@ -1393,6 +1453,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
@@ -1418,6 +1481,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 != "" && s.cache != nil && item.account.IsClientAffinityEnabled() {
_ = s.cache.UpdateClientAffinity(ctx, derefGroupID(groupID), 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)
}
@@ -1514,6 +1580,76 @@ func (s *GatewayService) SelectAccountWithLoadAwareness(ctx context.Context, gro
}
}
// ============ Layer 1.6: 客户端亲和(仅在粘性会话未命中时生效) ============
if affinityClientID != "" && s.cache != nil && stickyAccountID <= 0 {
affinityAccountIDs, err := s.cache.GetClientAffinityAccounts(ctx, derefGroupID(groupID), affinityClientID, ClientAffinityTTL)
if err == nil && len(affinityAccountIDs) > 0 {
for _, affinityAccID := range affinityAccountIDs {
if isExcluded(affinityAccID) {
continue
}
account, ok := accountByID[affinityAccID]
if !ok || !s.isAccountSchedulableForSelection(account) {
continue
}
if !account.IsClientAffinityEnabled() {
continue
}
if !s.isAccountAllowedForPlatform(account, platform, useMixed) {
continue
}
if requestedModel != "" && !s.isModelSupportedByAccountWithContext(ctx, account, requestedModel) {
continue
}
if !s.isAccountSchedulableForModelSelection(ctx, account, requestedModel) {
continue
}
if !s.isAccountSchedulableForQuota(account) {
continue
}
if !s.isAccountSchedulableForWindowCost(ctx, account, false) {
continue
}
if !s.isAccountSchedulableForRPM(ctx, account, false) {
continue
}
// 亲和三区检查:红区账号不可通过亲和命中调度
if account.GetAffinityBase() > 0 && s.cache != nil {
countMap, err := s.cache.GetAccountAffinityCountBatch(ctx, derefGroupID(groupID), []int64{affinityAccID}, ClientAffinityTTL)
if err == nil {
zone := account.GetAffinityZone(countMap[affinityAccID])
if zone == AffinityZoneRed {
continue
}
}
}
result, err := s.tryAcquireAccountSlot(ctx, affinityAccID, account.Concurrency)
if err == nil && result.Acquired {
if !s.checkAndRegisterSession(ctx, account, sessionHash) {
result.ReleaseFunc()
continue
}
// 亲和命中:更新亲和 score + 绑定粘性会话
_ = s.cache.UpdateClientAffinity(ctx, derefGroupID(groupID), affinityClientID, affinityAccID, ClientAffinityTTL)
if sessionHash != "" {
_ = s.cache.SetSessionAccountID(ctx, derefGroupID(groupID), sessionHash, affinityAccID, stickySessionTTL)
}
slog.Debug("client_affinity_hit",
"group_id", derefGroupID(groupID),
"client_id", affinityClientID[:8]+"...",
"account_id", affinityAccID)
return &AccountSelectionResult{
Account: account,
Acquired: true,
ReleaseFunc: result.ReleaseFunc,
}, nil
}
}
// 所有亲和账号不可用,继续到 Layer 2
}
}
// ============ Layer 2: 负载感知选择 ============
candidates := make([]*Account, 0, len(accounts))
for i := range accounts {
@@ -1566,6 +1702,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 != "" && s.cache != nil && result.Account != nil && result.Account.IsClientAffinityEnabled() {
_ = s.cache.UpdateClientAffinity(ctx, derefGroupID(groupID), affinityClientID, result.Account.ID, ClientAffinityTTL)
}
return result, nil
}
} else {
@@ -1583,13 +1722,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
@@ -1604,6 +1767,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 != "" && s.cache != nil && selected.account.IsClientAffinityEnabled() {
_ = s.cache.UpdateClientAffinity(ctx, derefGroupID(groupID), affinityClientID, selected.account.ID, ClientAffinityTTL)
}
return &AccountSelectionResult{
Account: selected.account,
Acquired: true,
@@ -2361,6 +2528,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.IsClientAffinityEnabled() {
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 {
@@ -2401,6 +2598,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.IsClientAffinityEnabled() && 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 {
@@ -3374,6 +3629,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)
@@ -4476,6 +4735,7 @@ func (s *GatewayService) Forward(ctx context.Context, c *gin.Context, account *A
}
// 处理正常响应
ctx = withClaudeMaxResponseRewriteContext(ctx, c, parsed)
// 触发上游接受回调(提前释放串行锁,不等流完成)
if parsed.OnUpstreamAccepted != nil {
@@ -6552,6 +6812,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)
@@ -6613,17 +6874,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 {
@@ -7055,8 +7324,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 对象
@@ -7122,6 +7396,7 @@ func (s *GatewayService) getUserGroupRateMultiplier(ctx context.Context, userID,
// RecordUsageInput 记录使用量的输入参数
type RecordUsageInput struct {
Result *ForwardResult
ParsedRequest *ParsedRequest
APIKey *APIKey
User *User
Account *Account
@@ -7432,9 +7707,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,25 @@ func (m *mockGatewayCacheForGemini) DeleteSessionAccountID(ctx context.Context,
return nil
}
func (m *mockGatewayCacheForGemini) GetClientAffinityAccounts(_ context.Context, _ int64, _ string, _ time.Duration) ([]int64, error) {
return nil, nil
}
func (m *mockGatewayCacheForGemini) UpdateClientAffinity(_ context.Context, _ int64, _ string, _ int64, _ time.Duration) error {
return 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) {
@@ -1375,7 +1375,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)
@@ -1536,7 +1536,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,25 @@ func (c *stubGatewayCache) DeleteSessionAccountID(ctx context.Context, groupID i
return nil
}
func (c *stubGatewayCache) GetClientAffinityAccounts(_ context.Context, _ int64, _ string, _ time.Duration) ([]int64, error) {
return nil, nil
}
func (c *stubGatewayCache) UpdateClientAffinity(_ context.Context, _ int64, _ string, _ int64, _ time.Duration) error {
return 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,25 @@ func (c *openAIWSStateStoreTimeoutProbeCache) DeleteSessionAccountID(ctx context
return nil
}
func (c *openAIWSStateStoreTimeoutProbeCache) GetClientAffinityAccounts(_ context.Context, _ int64, _ string, _ time.Duration) ([]int64, error) {
return nil, nil
}
func (c *openAIWSStateStoreTimeoutProbeCache) UpdateClientAffinity(_ context.Context, _ int64, _ string, _ int64, _ time.Duration) error {
return 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 {
@@ -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,
}
@@ -201,6 +201,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
@@ -1413,6 +1456,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"`
@@ -1424,6 +1469,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 存储配置(兼容旧单配置语义:返回当前激活配置)
@@ -1535,6 +1588,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),
@@ -1546,6 +1601,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 == "" {
@@ -1591,6 +1653,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)
@@ -1603,6 +1667,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
@@ -1905,6 +1982,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,
@@ -1917,6 +1996,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
@@ -128,6 +128,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"`
@@ -140,6 +142,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
}
@@ -0,0 +1,433 @@
package service
import (
"context"
"fmt"
"io"
"net/http"
"strings"
"sync"
"time"
"github.com/google/uuid"
"golang.org/x/oauth2"
"golang.org/x/oauth2/google"
"google.golang.org/api/drive/v3"
"google.golang.org/api/option"
"github.com/Wei-Shaw/sub2api/internal/pkg/logger"
)
// SoraGDriveStorage 负责 Sora 媒体文件的 Google Drive 存储操作。
type SoraGDriveStorage struct {
settingService *SettingService
mu sync.RWMutex
srv *drive.Service
cfg *SoraS3Profile // 缓存当前 GDrive 配置
healthCheckedAt time.Time
healthErr error
healthTTL time.Duration
}
const defaultGDriveHealthTTL = 30 * time.Second
// NewSoraGDriveStorage 创建 Google Drive 存储服务实例。
func NewSoraGDriveStorage(settingService *SettingService) *SoraGDriveStorage {
return &SoraGDriveStorage{
settingService: settingService,
healthTTL: defaultGDriveHealthTTL,
}
}
// StorageType 返回存储类型标识。
func (s *SoraGDriveStorage) StorageType() string {
return SoraStorageTypeGDrive
}
// Enabled 返回 Google Drive 存储是否已启用。
func (s *SoraGDriveStorage) Enabled(ctx context.Context) bool {
profile := s.getActiveGDriveProfile(ctx)
if profile == nil {
return false
}
return profile.Enabled && s.hasValidCredentials(profile)
}
// getActiveGDriveProfile 获取当前激活的 GDrive 配置。
func (s *SoraGDriveStorage) getActiveGDriveProfile(ctx context.Context) *SoraS3Profile {
if s.settingService == nil {
return nil
}
profile, err := s.settingService.GetActiveStorageProfile(ctx)
if err != nil || profile == nil {
return nil
}
if profile.GetProvider() != SoraStorageTypeGDrive {
return nil
}
return profile
}
// hasValidCredentials 检查 GDrive 配置是否有有效凭证。
func (s *SoraGDriveStorage) hasValidCredentials(profile *SoraS3Profile) bool {
switch profile.AuthType {
case "oauth2":
return profile.ClientID != "" && profile.ClientSecret != "" && profile.RefreshToken != ""
case "service_account":
return profile.ServiceAccountJSON != ""
default:
return false
}
}
// getService 获取或初始化 Drive 服务(带缓存)。
func (s *SoraGDriveStorage) getService(ctx context.Context) (*drive.Service, *SoraS3Profile, error) {
s.mu.RLock()
if s.srv != nil && s.cfg != nil {
srv, cfg := s.srv, s.cfg
s.mu.RUnlock()
return srv, cfg, nil
}
s.mu.RUnlock()
return s.initService(ctx)
}
func (s *SoraGDriveStorage) initService(ctx context.Context) (*drive.Service, *SoraS3Profile, error) {
s.mu.Lock()
defer s.mu.Unlock()
// 双重检查
if s.srv != nil && s.cfg != nil {
return s.srv, s.cfg, nil
}
profile := s.getActiveGDriveProfile(ctx)
if profile == nil {
return nil, nil, fmt.Errorf("no active gdrive profile found")
}
if !profile.Enabled {
return nil, nil, fmt.Errorf("gdrive storage is disabled")
}
srv, err := s.buildDriveService(ctx, profile)
if err != nil {
return nil, nil, fmt.Errorf("build gdrive service: %w", err)
}
s.srv = srv
s.cfg = profile
logger.LegacyPrintf("service.sora_gdrive", "[SoraGDrive] 客户端已初始化 auth_type=%s folder_id=%s", profile.AuthType, profile.FolderID)
return srv, profile, nil
}
// buildDriveService 根据认证类型创建 Google Drive 服务。
func (s *SoraGDriveStorage) buildDriveService(ctx context.Context, profile *SoraS3Profile) (*drive.Service, error) {
switch profile.AuthType {
case "oauth2":
return s.buildOAuth2Service(ctx, profile)
case "service_account":
return s.buildServiceAccountService(ctx, profile)
default:
return nil, fmt.Errorf("unsupported auth_type: %s", profile.AuthType)
}
}
func (s *SoraGDriveStorage) buildOAuth2Service(ctx context.Context, profile *SoraS3Profile) (*drive.Service, error) {
config := &oauth2.Config{
ClientID: profile.ClientID,
ClientSecret: profile.ClientSecret,
Endpoint: google.Endpoint,
Scopes: []string{drive.DriveFileScope},
}
token := &oauth2.Token{
RefreshToken: profile.RefreshToken,
}
tokenSource := config.TokenSource(ctx, token)
srv, err := drive.NewService(ctx, option.WithTokenSource(tokenSource))
if err != nil {
return nil, fmt.Errorf("create gdrive oauth2 service: %w", err)
}
return srv, nil
}
func (s *SoraGDriveStorage) buildServiceAccountService(ctx context.Context, profile *SoraS3Profile) (*drive.Service, error) {
srv, err := drive.NewService(ctx, option.WithCredentialsJSON([]byte(profile.ServiceAccountJSON))) //nolint:staticcheck // SA1019: admin-controlled service account JSON, safe to use
if err != nil {
return nil, fmt.Errorf("create gdrive service account service: %w", err)
}
return srv, nil
}
// RefreshClient 清除缓存的 Drive 客户端。
func (s *SoraGDriveStorage) RefreshClient() {
s.mu.Lock()
defer s.mu.Unlock()
s.srv = nil
s.cfg = nil
s.healthCheckedAt = time.Time{}
s.healthErr = nil
logger.LegacyPrintf("service.sora_gdrive", "[SoraGDrive] 客户端缓存已清除")
}
// GDriveQuotaInfo 包含 Google Drive 配额信息。
type GDriveQuotaInfo struct {
LimitBytes int64 `json:"limit_bytes"`
UsedBytes int64 `json:"used_bytes"`
}
// TestConnection 测试 Google Drive 连接。
func (s *SoraGDriveStorage) TestConnection(ctx context.Context) error {
srv, _, err := s.getService(ctx)
if err != nil {
return err
}
_, err = srv.About.Get().Fields("storageQuota").Context(ctx).Do()
if err != nil {
return fmt.Errorf("gdrive About.Get failed: %w", err)
}
return nil
}
// GetQuotaInfo 获取 Google Drive 配额信息(总量和已用量)。
func (s *SoraGDriveStorage) GetQuotaInfo(ctx context.Context) (*GDriveQuotaInfo, error) {
srv, _, err := s.getService(ctx)
if err != nil {
return nil, err
}
about, err := srv.About.Get().Fields("storageQuota").Context(ctx).Do()
if err != nil {
return nil, fmt.Errorf("gdrive About.Get failed: %w", err)
}
if about.StorageQuota == nil {
return nil, fmt.Errorf("storageQuota not available")
}
return &GDriveQuotaInfo{
LimitBytes: about.StorageQuota.Limit,
UsedBytes: about.StorageQuota.Usage,
}, nil
}
// TestFullCycle 执行完整的上传→获取链接→删除测试。
func (s *SoraGDriveStorage) TestFullCycle(ctx context.Context) (map[string]any, error) {
srv, cfg, err := s.getService(ctx)
if err != nil {
return nil, fmt.Errorf("init client: %w", err)
}
result := map[string]any{}
// 1. 测试 API 连接
about, err := srv.About.Get().Fields("storageQuota").Context(ctx).Do()
if err != nil {
return nil, fmt.Errorf("API connection failed: %w", err)
}
if about.StorageQuota != nil {
result["quota_limit_bytes"] = about.StorageQuota.Limit
result["quota_used_bytes"] = about.StorageQuota.Usage
}
// 2. 上传测试文件
testContent := "sub2api GDrive test file - " + time.Now().Format(time.RFC3339)
fileMeta := &drive.File{
Name: "sub2api_test_" + uuid.NewString()[:8] + ".txt",
MimeType: "text/plain",
}
if cfg.FolderID != "" {
fileMeta.Parents = []string{cfg.FolderID}
}
uploaded, err := srv.Files.Create(fileMeta).
Media(strings.NewReader(testContent)).
Fields("id,name,size,webViewLink").
Context(ctx).Do()
if err != nil {
return nil, fmt.Errorf("upload test file failed: %w", err)
}
result["uploaded_file_id"] = uploaded.Id
result["uploaded_file_name"] = uploaded.Name
result["uploaded_file_size"] = uploaded.Size
result["web_view_link"] = uploaded.WebViewLink
// 3. 获取访问链接
accessURL, err := s.GetAccessURL(ctx, uploaded.Id)
if err != nil {
// 即使获取链接失败,仍尝试清理
_ = srv.Files.Delete(uploaded.Id).Context(ctx).Do()
return nil, fmt.Errorf("get access URL failed: %w", err)
}
result["access_url"] = accessURL
// 4. 删除测试文件
if err := srv.Files.Delete(uploaded.Id).Context(ctx).Do(); err != nil {
result["delete_warning"] = fmt.Sprintf("delete failed (manual cleanup needed): %v", err)
} else {
result["deleted"] = true
}
result["status"] = "ok"
return result, nil
}
// IsHealthy 返回 Google Drive 健康状态(带短缓存)。
func (s *SoraGDriveStorage) IsHealthy(ctx context.Context) bool {
if s == nil {
return false
}
now := time.Now()
s.mu.RLock()
lastCheck := s.healthCheckedAt
lastErr := s.healthErr
ttl := s.healthTTL
s.mu.RUnlock()
if ttl <= 0 {
ttl = defaultGDriveHealthTTL
}
if !lastCheck.IsZero() && now.Sub(lastCheck) < ttl {
return lastErr == nil
}
err := s.TestConnection(ctx)
s.mu.Lock()
s.healthCheckedAt = time.Now()
s.healthErr = err
s.mu.Unlock()
return err == nil
}
// UploadFromURL 从上游 URL 下载并上传到 Google Drive。
// 返回 Google Drive 文件 ID 作为 objectKey、文件大小、存储类型。
func (s *SoraGDriveStorage) UploadFromURL(ctx context.Context, userID int64, sourceURL string) (string, int64, string, error) {
srv, cfg, err := s.getService(ctx)
if err != nil {
return "", 0, "", err
}
// 下载源文件
req, err := http.NewRequestWithContext(ctx, http.MethodGet, sourceURL, nil)
if err != nil {
return "", 0, "", fmt.Errorf("create download request: %w", err)
}
httpClient := &http.Client{Timeout: 5 * time.Minute}
resp, err := httpClient.Do(req)
if err != nil {
return "", 0, "", fmt.Errorf("download from upstream: %w", err)
}
defer func() {
_ = resp.Body.Close()
}()
if resp.StatusCode != http.StatusOK {
return "", 0, "", &UpstreamDownloadError{StatusCode: resp.StatusCode}
}
// 推断文件扩展名和 MIME
ext := fileExtFromURL(sourceURL)
if ext == "" {
ext = fileExtFromContentType(resp.Header.Get("Content-Type"))
}
if ext == "" {
ext = ".bin"
}
contentType := resp.Header.Get("Content-Type")
if contentType == "" {
contentType = "application/octet-stream"
}
// 生成文件名
datePath := time.Now().Format("2006-01-02")
fileName := fmt.Sprintf("sora_%d_%s_%s%s", userID, datePath, uuid.NewString()[:8], ext)
// 创建文件元数据
fileMeta := &drive.File{
Name: fileName,
MimeType: contentType,
}
if cfg.FolderID != "" {
fileMeta.Parents = []string{cfg.FolderID}
}
// 使用 CountingReader 统计大小
cr := &countingReader{Reader: resp.Body}
// 上传到 Google Drive
created, err := srv.Files.Create(fileMeta).
Media(cr).
Fields("id, size").
Context(ctx).
Do()
if err != nil {
return "", 0, "", fmt.Errorf("gdrive upload: %w", err)
}
fileSize := cr.BytesRead
if created.Size > 0 {
fileSize = created.Size
}
// 根据 access_mode 设置权限
if cfg.AccessMode == "" || cfg.AccessMode == "direct" {
// 设为任何人可读
_, permErr := srv.Permissions.Create(created.Id, &drive.Permission{
Type: "anyone",
Role: "reader",
}).Context(ctx).Do()
if permErr != nil {
logger.LegacyPrintf("service.sora_gdrive", "[SoraGDrive] 设置公开权限失败 fileID=%s err=%v", created.Id, permErr)
}
}
logger.LegacyPrintf("service.sora_gdrive", "[SoraGDrive] 上传完成 fileID=%s size=%d", created.Id, fileSize)
return created.Id, fileSize, SoraStorageTypeGDrive, nil
}
// DeleteObjects 删除一组 Google Drive 文件。
func (s *SoraGDriveStorage) DeleteObjects(ctx context.Context, objectKeys []string) error {
if len(objectKeys) == 0 {
return nil
}
srv, _, err := s.getService(ctx)
if err != nil {
return err
}
var lastErr error
for _, fileID := range objectKeys {
if err := srv.Files.Delete(fileID).Context(ctx).Do(); err != nil {
logger.LegacyPrintf("service.sora_gdrive", "[SoraGDrive] 删除失败 fileID=%s err=%v", fileID, err)
lastErr = err
}
}
return lastErr
}
// GetAccessURL 获取 Google Drive 文件的访问 URL。
func (s *SoraGDriveStorage) GetAccessURL(ctx context.Context, objectKey string) (string, error) {
_, cfg, err := s.getService(ctx)
if err != nil {
return "", err
}
// CDN URL 优先
if cfg.CDNURL != "" {
cdnBase := strings.TrimRight(cfg.CDNURL, "/")
return cdnBase + "/" + objectKey, nil
}
// 默认使用 Google Drive 直链
return fmt.Sprintf("https://drive.google.com/uc?export=download&id=%s", objectKey), nil
}
// countingReader 包装 io.Reader 以统计读取的字节数。
type countingReader struct {
Reader io.Reader
BytesRead int64
}
func (r *countingReader) Read(p []byte) (int, error) {
n, err := r.Reader.Read(p)
r.BytesRead += int64(n)
return n, err
}
@@ -0,0 +1,124 @@
//go:build unit
package service
import (
"context"
"testing"
"github.com/stretchr/testify/assert"
)
func TestSoraGDriveStorage_StorageType(t *testing.T) {
s := NewSoraGDriveStorage(nil)
assert.Equal(t, SoraStorageTypeGDrive, s.StorageType())
}
func TestSoraGDriveStorage_EnabledWithNilSettingService(t *testing.T) {
s := NewSoraGDriveStorage(nil)
assert.False(t, s.Enabled(context.Background()))
}
func TestSoraGDriveStorage_IsHealthyWithNilReceiver(t *testing.T) {
var s *SoraGDriveStorage
assert.False(t, s.IsHealthy(context.Background()))
}
func TestSoraGDriveStorage_GetServiceWithoutProfile(t *testing.T) {
s := NewSoraGDriveStorage(nil)
_, _, err := s.getService(context.Background())
assert.Error(t, err)
assert.Contains(t, err.Error(), "no active gdrive profile")
}
func TestSoraGDriveStorage_DeleteObjectsEmpty(t *testing.T) {
s := NewSoraGDriveStorage(nil)
err := s.DeleteObjects(context.Background(), []string{})
assert.NoError(t, err)
}
func TestSoraGDriveStorage_RefreshClient(t *testing.T) {
s := NewSoraGDriveStorage(nil)
// 不应 panic
s.RefreshClient()
assert.Nil(t, s.srv)
assert.Nil(t, s.cfg)
}
func TestSoraGDriveStorage_HasValidCredentials(t *testing.T) {
s := NewSoraGDriveStorage(nil)
tests := []struct {
name string
profile *SoraS3Profile
want bool
}{
{
name: "oauth2 with all fields",
profile: &SoraS3Profile{
AuthType: "oauth2",
ClientID: "id",
ClientSecret: "secret",
RefreshToken: "token",
},
want: true,
},
{
name: "oauth2 missing refresh token",
profile: &SoraS3Profile{
AuthType: "oauth2",
ClientID: "id",
ClientSecret: "secret",
},
want: false,
},
{
name: "service_account with json",
profile: &SoraS3Profile{
AuthType: "service_account",
ServiceAccountJSON: `{"type":"service_account"}`,
},
want: true,
},
{
name: "service_account without json",
profile: &SoraS3Profile{
AuthType: "service_account",
},
want: false,
},
{
name: "unknown auth type",
profile: &SoraS3Profile{
AuthType: "unknown",
},
want: false,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got := s.hasValidCredentials(tt.profile)
assert.Equal(t, tt.want, got)
})
}
}
func TestSoraStorageRouter_DefaultsToS3(t *testing.T) {
s3 := NewSoraS3Storage(nil)
router := NewSoraStorageRouter(nil, s3, nil)
// settingService 为 nil,应返回 s3Storage
backend := router.activeBackend(context.Background())
assert.Equal(t, s3, backend)
}
func TestSoraStorageRouter_StorageType(t *testing.T) {
router := NewSoraStorageRouter(nil, nil, nil)
assert.Equal(t, SoraStorageTypeS3, router.StorageType())
}
func TestSoraStorageRouter_RefreshAllNoPanic(t *testing.T) {
router := NewSoraStorageRouter(nil, nil, nil)
// 不应 panic
router.RefreshAll()
}
@@ -37,6 +37,7 @@ const (
// Sora 存储类型常量
const (
SoraStorageTypeS3 = "s3"
SoraStorageTypeGDrive = "gdrive"
SoraStorageTypeLocal = "local"
SoraStorageTypeUpstream = "upstream"
SoraStorageTypeNone = "none"
@@ -60,4 +61,5 @@ type SoraGenerationRepository interface {
Delete(ctx context.Context, id int64) error
List(ctx context.Context, params SoraGenerationListParams) ([]*SoraGeneration, int64, error)
CountByUserAndStatus(ctx context.Context, userID int64, statuses []string) (int64, error)
CountByStorageType(ctx context.Context, storageType string, statuses []string) (int64, error)
}
@@ -35,21 +35,21 @@ type soraGenerationRepoConditionalUpdater interface {
// SoraGenerationService 管理 Sora 客户端的生成记录 CRUD。
type SoraGenerationService struct {
genRepo SoraGenerationRepository
s3Storage *SoraS3Storage
quotaService *SoraQuotaService
genRepo SoraGenerationRepository
objectStorage SoraObjectStorage
quotaService *SoraQuotaService
}
// NewSoraGenerationService 创建生成记录服务。
func NewSoraGenerationService(
genRepo SoraGenerationRepository,
s3Storage *SoraS3Storage,
objectStorage SoraObjectStorage,
quotaService *SoraQuotaService,
) *SoraGenerationService {
return &SoraGenerationService{
genRepo: genRepo,
s3Storage: s3Storage,
quotaService: quotaService,
genRepo: genRepo,
objectStorage: objectStorage,
quotaService: quotaService,
}
}
@@ -268,15 +268,15 @@ func (s *SoraGenerationService) Delete(ctx context.Context, id, userID int64) er
return fmt.Errorf("无权删除此生成记录")
}
// 清理 S3 文件
if gen.StorageType == SoraStorageTypeS3 && len(gen.S3ObjectKeys) > 0 && s.s3Storage != nil {
if err := s.s3Storage.DeleteObjects(ctx, gen.S3ObjectKeys); err != nil {
logger.LegacyPrintf("service.sora_gen", "[SoraGen] S3 清理失败 id=%d err=%v", id, err)
// 清理存储文件(S3 / Google Drive)
if IsObjectStorageType(gen.StorageType) && len(gen.S3ObjectKeys) > 0 && s.objectStorage != nil {
if err := s.objectStorage.DeleteObjects(ctx, gen.S3ObjectKeys); err != nil {
logger.LegacyPrintf("service.sora_gen", "[SoraGen] 存储清理失败 id=%d type=%s err=%v", id, gen.StorageType, err)
}
}
// 释放配额(S3/本地均释放)
if gen.FileSizeBytes > 0 && (gen.StorageType == SoraStorageTypeS3 || gen.StorageType == SoraStorageTypeLocal) && s.quotaService != nil {
// 释放配额(对象存储/本地均释放)
if gen.FileSizeBytes > 0 && (IsObjectStorageType(gen.StorageType) || gen.StorageType == SoraStorageTypeLocal) && s.quotaService != nil {
if err := s.quotaService.ReleaseUsage(ctx, userID, gen.FileSizeBytes); err != nil {
logger.LegacyPrintf("service.sora_gen", "[SoraGen] 配额释放失败 id=%d err=%v", id, err)
}
@@ -290,9 +290,9 @@ func (s *SoraGenerationService) CountActiveByUser(ctx context.Context, userID in
return s.genRepo.CountByUserAndStatus(ctx, userID, []string{SoraGenStatusPending, SoraGenStatusGenerating})
}
// ResolveMediaURLs 为 S3 记录动态生成预签名 URL。
// ResolveMediaURLs 为对象存储记录动态生成访问 URL。
func (s *SoraGenerationService) ResolveMediaURLs(ctx context.Context, gen *SoraGeneration) error {
if gen == nil || gen.StorageType != SoraStorageTypeS3 || s.s3Storage == nil {
if gen == nil || !IsObjectStorageType(gen.StorageType) || s.objectStorage == nil {
return nil
}
if len(gen.S3ObjectKeys) == 0 {
@@ -308,7 +308,7 @@ func (s *SoraGenerationService) ResolveMediaURLs(ctx context.Context, gen *SoraG
wg.Add(1)
go func(i int, objectKey string) {
defer wg.Done()
url, err := s.s3Storage.GetAccessURL(ctx, objectKey)
url, err := s.objectStorage.GetAccessURL(ctx, objectKey)
if err != nil {
errMu.Lock()
if firstErr == nil {
@@ -330,3 +330,22 @@ func (s *SoraGenerationService) ResolveMediaURLs(ctx context.Context, gen *SoraG
return nil
}
// StorageVideoStats 各存储类型的视频统计。
type StorageVideoStats struct {
Completed int64 `json:"completed"`
InProgress int64 `json:"in_progress"`
}
// CountByStorageType 按存储类型统计视频数量(completed 和 in_progress)。
func (s *SoraGenerationService) CountByStorageType(ctx context.Context, storageType string) (completed, inProgress int64, err error) {
completed, err = s.genRepo.CountByStorageType(ctx, storageType, []string{SoraGenStatusCompleted})
if err != nil {
return 0, 0, fmt.Errorf("count completed: %w", err)
}
inProgress, err = s.genRepo.CountByStorageType(ctx, storageType, []string{SoraGenStatusPending, SoraGenStatusGenerating})
if err != nil {
return 0, 0, fmt.Errorf("count in_progress: %w", err)
}
return completed, inProgress, nil
}
@@ -115,6 +115,25 @@ func (r *stubGenRepo) CountByUserAndStatus(_ context.Context, userID int64, stat
return count, nil
}
func (r *stubGenRepo) CountByStorageType(_ context.Context, storageType string, statuses []string) (int64, error) {
if r.countErr != nil {
return 0, r.countErr
}
var count int64
statusSet := make(map[string]struct{})
for _, s := range statuses {
statusSet[s] = struct{}{}
}
for _, gen := range r.gens {
if gen.StorageType == storageType {
if _, ok := statusSet[gen.Status]; ok {
count++
}
}
}
return count, nil
}
// ==================== Stub: UserRepository (用于 SoraQuotaService) ====================
var _ UserRepository = (*stubUserRepoForQuota)(nil)
@@ -519,7 +538,7 @@ func TestDelete_S3Cleanup_NilS3(t *testing.T) {
svc := NewSoraGenerationService(repo, nil, nil)
err := svc.Delete(context.Background(), 1, 1)
require.NoError(t, err) // s3Storage 为 nil,跳过清理
require.NoError(t, err) // objectStorage 为 nil,跳过清理
}
func TestDelete_QuotaRelease_NilQuota(t *testing.T) {
@@ -0,0 +1,37 @@
package service
import "context"
// SoraObjectStorage 是 Sora 媒体文件的通用对象存储接口。
// S3 和 Google Drive 等存储后端均实现此接口。
type SoraObjectStorage interface {
// Enabled 返回存储是否已启用且配置有效。
Enabled(ctx context.Context) bool
// IsHealthy 返回存储健康状态(带短缓存)。
IsHealthy(ctx context.Context) bool
// TestConnection 测试存储连接。
TestConnection(ctx context.Context) error
// UploadFromURL 从上游 URL 下载并上传到存储。
// 返回 object key(S3 key 或 GDrive file ID)、文件大小、实际使用的存储类型。
UploadFromURL(ctx context.Context, userID int64, sourceURL string) (objectKey string, sizeBytes int64, storageType string, err error)
// DeleteObjects 删除一组存储对象。
DeleteObjects(ctx context.Context, objectKeys []string) error
// GetAccessURL 获取存储文件的访问 URL。
GetAccessURL(ctx context.Context, objectKey string) (string, error)
// RefreshClient 清除缓存客户端,配置变更时调用。
RefreshClient()
// StorageType 返回存储类型标识("s3" / "gdrive")。
StorageType() string
}
// IsObjectStorageType 判断是否为对象存储类型(S3 或 Google Drive)。
func IsObjectStorageType(t string) bool {
return t == SoraStorageTypeS3 || t == SoraStorageTypeGDrive
}
+14 -9
View File
@@ -212,29 +212,29 @@ func (s *SoraS3Storage) GenerateObjectKey(prefix string, userID int64, ext strin
}
// UploadFromURL 从上游 URL 下载并流式上传到 S3。
// 返回 S3 object key。
func (s *SoraS3Storage) UploadFromURL(ctx context.Context, userID int64, sourceURL string) (string, int64, error) {
// 返回 S3 object key、文件大小、存储类型。
func (s *SoraS3Storage) UploadFromURL(ctx context.Context, userID int64, sourceURL string) (string, int64, string, error) {
client, cfg, err := s.getClient(ctx)
if err != nil {
return "", 0, err
return "", 0, "", err
}
// 下载源文件
req, err := http.NewRequestWithContext(ctx, http.MethodGet, sourceURL, nil)
if err != nil {
return "", 0, fmt.Errorf("create download request: %w", err)
return "", 0, "", fmt.Errorf("create download request: %w", err)
}
httpClient := &http.Client{Timeout: 5 * time.Minute}
resp, err := httpClient.Do(req)
if err != nil {
return "", 0, fmt.Errorf("download from upstream: %w", err)
return "", 0, "", fmt.Errorf("download from upstream: %w", err)
}
defer func() {
_ = resp.Body.Close()
}()
if resp.StatusCode != http.StatusOK {
return "", 0, &UpstreamDownloadError{StatusCode: resp.StatusCode}
return "", 0, "", &UpstreamDownloadError{StatusCode: resp.StatusCode}
}
// 推断文件扩展名
@@ -275,14 +275,14 @@ func (s *SoraS3Storage) UploadFromURL(ctx context.Context, userID int64, sourceU
_ = writer.CloseWithError(copyErr)
uploadErr := <-uploadErrCh
if copyErr != nil {
return "", 0, fmt.Errorf("stream upload copy failed: %w", copyErr)
return "", 0, "", fmt.Errorf("stream upload copy failed: %w", copyErr)
}
if uploadErr != nil {
return "", 0, fmt.Errorf("s3 upload: %w", uploadErr)
return "", 0, "", fmt.Errorf("s3 upload: %w", uploadErr)
}
logger.LegacyPrintf("service.sora_s3", "[SoraS3] 上传完成 key=%s size=%d", objectKey, written)
return objectKey, written, nil
return objectKey, written, SoraStorageTypeS3, nil
}
func buildSoraS3Client(ctx context.Context, cfg *SoraS3Settings) (*s3.Client, string, error) {
@@ -380,6 +380,11 @@ func (s *SoraS3Storage) GeneratePresignedURL(ctx context.Context, objectKey stri
return result.URL, nil
}
// StorageType 返回存储类型标识。
func (s *SoraS3Storage) StorageType() string {
return SoraStorageTypeS3
}
// GetMediaType 从 object key 推断媒体类型(image/video)。
func GetMediaTypeFromKey(objectKey string) string {
ext := strings.ToLower(path.Ext(objectKey))
@@ -238,7 +238,7 @@ func TestTestConnection_GetClientError(t *testing.T) {
func TestUploadFromURL_GetClientError(t *testing.T) {
s := NewSoraS3Storage(nil)
_, _, err := s.UploadFromURL(context.Background(), 1, "https://example.com/file.mp4")
_, _, _, err := s.UploadFromURL(context.Background(), 1, "https://example.com/file.mp4")
require.Error(t, err)
}
+40 -18
View File
@@ -316,20 +316,36 @@ func (c *SoraSDKClient) GetCameoStatus(ctx context.Context, account *Account, ca
if err != nil {
return nil, err
}
sdkClient, err := c.getSDKClient(account)
if err != nil {
return nil, err
}
status, err := sdkClient.GetCameoStatus(ctx, token, cameoID)
// 直接调用 Sora 后端 API 而非 SDK,以获取 SDK 未暴露的字段
// (status_message、instruction_set_hint、instruction_set)。
path := "/project_y/cameos/in_progress/" + cameoID
raw, err := c.doSoraBackendJSON(ctx, account, http.MethodGet, path, token, "", nil)
if err != nil {
return nil, c.wrapSDKError(err, account)
}
return &SoraCameoStatus{
Status: status.Status,
DisplayNameHint: status.DisplayNameHint,
UsernameHint: status.UsernameHint,
ProfileAssetURL: status.ProfileAssetURL,
}, nil
return parseCameoStatusFromRaw(raw), nil
}
// parseCameoStatusFromRaw 从原始 JSON 解析 SoraCameoStatus,
// 包含 SDK 未暴露的 status_message / instruction_set_hint / instruction_set 字段。
func parseCameoStatusFromRaw(raw []byte) *SoraCameoStatus {
result := gjson.ParseBytes(raw)
cameoStatus := &SoraCameoStatus{
Status: strings.TrimSpace(result.Get("status").String()),
StatusMessage: strings.TrimSpace(result.Get("status_message").String()),
DisplayNameHint: strings.TrimSpace(result.Get("display_name_hint").String()),
UsernameHint: strings.TrimSpace(result.Get("username_hint").String()),
ProfileAssetURL: strings.TrimSpace(result.Get("profile_asset_url").String()),
}
if v := result.Get("instruction_set_hint"); v.Exists() {
cameoStatus.InstructionSetHint = v.Value()
}
if v := result.Get("instruction_set"); v.Exists() {
cameoStatus.InstructionSet = v.Value()
}
return cameoStatus
}
func (c *SoraSDKClient) DownloadCharacterImage(ctx context.Context, account *Account, imageURL string) ([]byte, error) {
@@ -925,26 +941,32 @@ func (c *SoraSDKClient) exchangeSessionToken(ctx context.Context, account *Accou
return accessToken, expiresAt, nil
}
// applyRecoveredToken 将恢复的 token 写入账号内存和数据库
// applyRecoveredToken 将恢复的 token 写入账号内存和数据库。
// 使用 copy-on-write 避免并发 map 写入 panic:创建新 map 后整体替换指针。
func (c *SoraSDKClient) applyRecoveredToken(ctx context.Context, account *Account, accessToken, refreshToken, expiresAt, sessionToken string) {
if account == nil {
return
}
if account.Credentials == nil {
account.Credentials = make(map[string]any)
// Copy-on-write: 复制旧 map 并写入新值,最后整体替换
oldCreds := account.Credentials
newCreds := make(map[string]any, len(oldCreds)+4)
for k, v := range oldCreds {
newCreds[k] = v
}
if strings.TrimSpace(accessToken) != "" {
account.Credentials["access_token"] = accessToken
newCreds["access_token"] = accessToken
}
if strings.TrimSpace(refreshToken) != "" {
account.Credentials["refresh_token"] = refreshToken
newCreds["refresh_token"] = refreshToken
}
if strings.TrimSpace(expiresAt) != "" {
account.Credentials["expires_at"] = expiresAt
newCreds["expires_at"] = expiresAt
}
if strings.TrimSpace(sessionToken) != "" {
account.Credentials["session_token"] = sessionToken
newCreds["session_token"] = sessionToken
}
account.Credentials = newCreds
if c.accountRepo != nil {
if err := c.accountRepo.Update(ctx, account); err != nil && c.debugEnabled() {
@@ -0,0 +1,129 @@
package service
import (
"context"
"fmt"
"github.com/Wei-Shaw/sub2api/internal/pkg/logger"
)
// SoraStorageRouter 根据激活 profile 的 provider 字段路由到对应存储实现。
// 实现 SoraObjectStorage 接口。
type SoraStorageRouter struct {
settingService *SettingService
s3Storage *SoraS3Storage
gdriveStorage SoraObjectStorage // 可为 nil(GDrive 未实现时)
}
// NewSoraStorageRouter 创建存储路由。
func NewSoraStorageRouter(
settingService *SettingService,
s3Storage *SoraS3Storage,
gdriveStorage SoraObjectStorage,
) *SoraStorageRouter {
return &SoraStorageRouter{
settingService: settingService,
s3Storage: s3Storage,
gdriveStorage: gdriveStorage,
}
}
// activeBackend 返回当前激活 profile 对应的存储后端。
func (r *SoraStorageRouter) activeBackend(ctx context.Context) SoraObjectStorage {
if r.settingService == nil {
return r.s3Storage // 默认 S3
}
profile, err := r.settingService.GetActiveStorageProfile(ctx)
if err != nil || profile == nil {
return r.s3Storage // 默认 S3
}
switch profile.GetProvider() {
case SoraStorageTypeGDrive:
if r.gdriveStorage != nil {
return r.gdriveStorage
}
logger.LegacyPrintf("service.storage_router", "[StorageRouter] GDrive 后端未初始化,降级到 S3")
return r.s3Storage
default:
return r.s3Storage
}
}
func (r *SoraStorageRouter) Enabled(ctx context.Context) bool {
backend := r.activeBackend(ctx)
if backend == nil {
return false
}
return backend.Enabled(ctx)
}
func (r *SoraStorageRouter) IsHealthy(ctx context.Context) bool {
backend := r.activeBackend(ctx)
if backend == nil {
return false
}
return backend.IsHealthy(ctx)
}
func (r *SoraStorageRouter) TestConnection(ctx context.Context) error {
backend := r.activeBackend(ctx)
if backend == nil {
return fmt.Errorf("no storage backend available")
}
return backend.TestConnection(ctx)
}
func (r *SoraStorageRouter) UploadFromURL(ctx context.Context, userID int64, sourceURL string) (string, int64, string, error) {
backend := r.activeBackend(ctx)
if backend == nil {
return "", 0, "", fmt.Errorf("no storage backend available")
}
return backend.UploadFromURL(ctx, userID, sourceURL)
}
func (r *SoraStorageRouter) DeleteObjects(ctx context.Context, objectKeys []string) error {
backend := r.activeBackend(ctx)
if backend == nil {
return fmt.Errorf("no storage backend available")
}
return backend.DeleteObjects(ctx, objectKeys)
}
func (r *SoraStorageRouter) GetAccessURL(ctx context.Context, objectKey string) (string, error) {
backend := r.activeBackend(ctx)
if backend == nil {
return "", fmt.Errorf("no storage backend available")
}
return backend.GetAccessURL(ctx, objectKey)
}
func (r *SoraStorageRouter) RefreshClient() {
if r.s3Storage != nil {
r.s3Storage.RefreshClient()
}
if r.gdriveStorage != nil {
r.gdriveStorage.RefreshClient()
}
}
// RefreshAll 刷新所有后端客户端(用作配置变更回调)。
func (r *SoraStorageRouter) RefreshAll() {
r.RefreshClient()
}
func (r *SoraStorageRouter) StorageType() string {
// 不带 context 的方法,返回默认值
// 真实的 StorageType 在 activeBackend 中动态确定
return SoraStorageTypeS3
}
// StorageTypeWithContext 返回当前激活后端的存储类型。
func (r *SoraStorageRouter) StorageTypeWithContext(ctx context.Context) string {
backend := r.activeBackend(ctx)
if backend == nil {
return SoraStorageTypeS3
}
return backend.StorageType()
}

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