mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-09-24 16:05:44 +08:00
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:
@@ -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
@@ -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/
|
||||
|
||||
|
||||
@@ -1 +1 @@
|
||||
0.1.88
|
||||
0.1.98.1
|
||||
|
||||
@@ -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{
|
||||
|
||||
@@ -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
@@ -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()
|
||||
}
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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 配置列表响应
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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, "aErr) {
|
||||
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"
|
||||
}
|
||||
}
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
`
|
||||
|
||||
@@ -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{
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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 验证)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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()
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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{
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
|
||||
@@ -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
Reference in New Issue
Block a user