diff --git a/Dockerfile b/Dockerfile index bae531ac6e..13a6b8700d 100644 --- a/Dockerfile +++ b/Dockerfile @@ -1,3 +1,4 @@ +# syntax=docker/dockerfile:1.7 # ============================================================================= # Sub2API Multi-Stage Dockerfile # ============================================================================= @@ -12,11 +13,13 @@ ARG ALPINE_IMAGE=alpine:3.21 ARG POSTGRES_IMAGE=postgres:18-alpine ARG GOPROXY=https://goproxy.cn,direct ARG GOSUMDB=sum.golang.google.cn +ARG NPM_CONFIG_REGISTRY= # ----------------------------------------------------------------------------- # Stage 1: Frontend Builder # ----------------------------------------------------------------------------- FROM ${NODE_IMAGE} AS frontend-builder +ARG NPM_CONFIG_REGISTRY WORKDIR /app/frontend @@ -25,7 +28,9 @@ RUN corepack enable && corepack prepare pnpm@9 --activate # Install dependencies first (better caching) COPY frontend/package.json frontend/pnpm-lock.yaml ./ -RUN pnpm install --frozen-lockfile +RUN --mount=type=cache,id=sub2api-pnpm-store,target=/root/.local/share/pnpm/store \ + if [ -n "${NPM_CONFIG_REGISTRY}" ]; then pnpm config set registry "${NPM_CONFIG_REGISTRY}"; fi && \ + pnpm install --frozen-lockfile --prefer-offline # Copy frontend source and build. # LegalDocumentView.vue (admin-compliance gate) build-time imports diff --git a/backend/cmd/server/wire_gen.go b/backend/cmd/server/wire_gen.go index d6f0bdfee1..aae4c405b9 100644 --- a/backend/cmd/server/wire_gen.go +++ b/backend/cmd/server/wire_gen.go @@ -138,10 +138,10 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) { channelService := service.NewChannelService(channelRepository, groupRepository, apiKeyAuthCacheInvalidator, pricingService) modelPricingResolver := service.NewModelPricingResolver(channelService, billingService) batchImageModelPricingResolver := service.ProvideBatchImageModelPricingResolver(modelPricingResolver) - batchImagePublicService := service.NewBatchImagePublicService(batchImageRepository, accountRepository, batchImageQueue, batchImageModelPricingResolver, configConfig) + batchImagePublicService := service.NewBatchImagePublicService(batchImageRepository, accountRepository, groupRepository, userGroupRateRepository, batchImageQueue, batchImageModelPricingResolver, usageBillingRepository, apiKeyAuthCacheInvalidator, configConfig) batchImageDownloadService := service.NewBatchImageDownloadService(batchImageRepository, accountRepository, batchImageDownloadLimiter, configConfig) batchImageCleanupService := service.ProvideBatchImageCleanupService(batchImageRepository, accountRepository, configConfig) - batchImageWorkerRuntime := service.ProvideBatchImageWorkerRuntime(batchImageRepository, accountRepository, batchImageQueue, usageBillingRepository, batchImageModelPricingResolver, configConfig) + batchImageWorkerRuntime := service.ProvideBatchImageWorkerRuntime(batchImageRepository, accountRepository, batchImageQueue, usageBillingRepository, usageLogRepository, batchImageModelPricingResolver, apiKeyAuthCacheInvalidator, configConfig) notificationEmailService := service.NewNotificationEmailService(settingRepository, emailService) balanceNotifyService := service.ProvideBalanceNotifyService(emailService, settingRepository, accountRepository, notificationEmailService) gatewayService := service.NewGatewayService(accountRepository, groupRepository, usageLogRepository, usageBillingRepository, userRepository, userSubscriptionRepository, userGroupRateRepository, gatewayCache, configConfig, schedulerSnapshotService, concurrencyService, billingService, rateLimitService, billingCacheService, identityService, httpUpstream, deferredService, claudeTokenProvider, sessionLimitCache, rpmCache, digestSessionStore, settingService, tlsFingerprintProfileService, channelService, modelPricingResolver, balanceNotifyService, serviceUserPlatformQuotaRepository) diff --git a/backend/ent/batchimagejob.go b/backend/ent/batchimagejob.go index b63ad6c6df..29f09bb289 100644 --- a/backend/ent/batchimagejob.go +++ b/backend/ent/batchimagejob.go @@ -29,6 +29,8 @@ type BatchImageJob struct { Provider string `json:"provider,omitempty"` // Model holds the value of the "model" field. Model string `json:"model,omitempty"` + // TaskName holds the value of the "task_name" field. + TaskName string `json:"task_name,omitempty"` // Status holds the value of the "status" field. Status string `json:"status,omitempty"` // ProviderJobName holds the value of the "provider_job_name" field. @@ -75,6 +77,10 @@ type BatchImageJob struct { InputDeletedAt *time.Time `json:"input_deleted_at,omitempty"` // OutputDeletedAt holds the value of the "output_deleted_at" field. OutputDeletedAt *time.Time `json:"output_deleted_at,omitempty"` + // DownloadedAt holds the value of the "downloaded_at" field. + DownloadedAt *time.Time `json:"downloaded_at,omitempty"` + // UserDeletedAt holds the value of the "user_deleted_at" field. + UserDeletedAt *time.Time `json:"user_deleted_at,omitempty"` // LastErrorCode holds the value of the "last_error_code" field. LastErrorCode *string `json:"last_error_code,omitempty"` // LastErrorMessage holds the value of the "last_error_message" field. @@ -103,9 +109,9 @@ func (*BatchImageJob) scanValues(columns []string) ([]any, error) { values[i] = new(sql.NullFloat64) case batchimagejob.FieldID, batchimagejob.FieldUserID, batchimagejob.FieldAPIKeyID, batchimagejob.FieldAccountID, batchimagejob.FieldItemCount, batchimagejob.FieldSuccessCount, batchimagejob.FieldFailCount, batchimagejob.FieldCancelledCount, batchimagejob.FieldRetryCount, batchimagejob.FieldVersion: values[i] = new(sql.NullInt64) - case batchimagejob.FieldBatchID, batchimagejob.FieldProvider, batchimagejob.FieldModel, batchimagejob.FieldStatus, batchimagejob.FieldProviderJobName, batchimagejob.FieldProviderInputRef, batchimagejob.FieldProviderOutputRef, batchimagejob.FieldGcsInputURI, batchimagejob.FieldGcsOutputURI, batchimagejob.FieldCurrency, batchimagejob.FieldHoldID, batchimagejob.FieldIdempotencyKey, batchimagejob.FieldRequestHash, batchimagejob.FieldManifestHash, batchimagejob.FieldLastErrorCode, batchimagejob.FieldLastErrorMessage: + case batchimagejob.FieldBatchID, batchimagejob.FieldProvider, batchimagejob.FieldModel, batchimagejob.FieldTaskName, batchimagejob.FieldStatus, batchimagejob.FieldProviderJobName, batchimagejob.FieldProviderInputRef, batchimagejob.FieldProviderOutputRef, batchimagejob.FieldGcsInputURI, batchimagejob.FieldGcsOutputURI, batchimagejob.FieldCurrency, batchimagejob.FieldHoldID, batchimagejob.FieldIdempotencyKey, batchimagejob.FieldRequestHash, batchimagejob.FieldManifestHash, batchimagejob.FieldLastErrorCode, batchimagejob.FieldLastErrorMessage: values[i] = new(sql.NullString) - case batchimagejob.FieldOutputExpiresAt, batchimagejob.FieldInputDeletedAt, batchimagejob.FieldOutputDeletedAt, batchimagejob.FieldCreatedAt, batchimagejob.FieldUpdatedAt, batchimagejob.FieldSubmittedAt, batchimagejob.FieldStartedAt, batchimagejob.FieldFinishedAt, batchimagejob.FieldSettledAt: + case batchimagejob.FieldOutputExpiresAt, batchimagejob.FieldInputDeletedAt, batchimagejob.FieldOutputDeletedAt, batchimagejob.FieldDownloadedAt, batchimagejob.FieldUserDeletedAt, batchimagejob.FieldCreatedAt, batchimagejob.FieldUpdatedAt, batchimagejob.FieldSubmittedAt, batchimagejob.FieldStartedAt, batchimagejob.FieldFinishedAt, batchimagejob.FieldSettledAt: values[i] = new(sql.NullTime) default: values[i] = new(sql.UnknownType) @@ -166,6 +172,12 @@ func (_m *BatchImageJob) assignValues(columns []string, values []any) error { } else if value.Valid { _m.Model = value.String } + case batchimagejob.FieldTaskName: + if value, ok := values[i].(*sql.NullString); !ok { + return fmt.Errorf("unexpected type %T for field task_name", values[i]) + } else if value.Valid { + _m.TaskName = value.String + } case batchimagejob.FieldStatus: if value, ok := values[i].(*sql.NullString); !ok { return fmt.Errorf("unexpected type %T for field status", values[i]) @@ -318,6 +330,20 @@ func (_m *BatchImageJob) assignValues(columns []string, values []any) error { _m.OutputDeletedAt = new(time.Time) *_m.OutputDeletedAt = value.Time } + case batchimagejob.FieldDownloadedAt: + if value, ok := values[i].(*sql.NullTime); !ok { + return fmt.Errorf("unexpected type %T for field downloaded_at", values[i]) + } else if value.Valid { + _m.DownloadedAt = new(time.Time) + *_m.DownloadedAt = value.Time + } + case batchimagejob.FieldUserDeletedAt: + if value, ok := values[i].(*sql.NullTime); !ok { + return fmt.Errorf("unexpected type %T for field user_deleted_at", values[i]) + } else if value.Valid { + _m.UserDeletedAt = new(time.Time) + *_m.UserDeletedAt = value.Time + } case batchimagejob.FieldLastErrorCode: if value, ok := values[i].(*sql.NullString); !ok { return fmt.Errorf("unexpected type %T for field last_error_code", values[i]) @@ -430,6 +456,9 @@ func (_m *BatchImageJob) String() string { builder.WriteString("model=") builder.WriteString(_m.Model) builder.WriteString(", ") + builder.WriteString("task_name=") + builder.WriteString(_m.TaskName) + builder.WriteString(", ") builder.WriteString("status=") builder.WriteString(_m.Status) builder.WriteString(", ") @@ -527,6 +556,16 @@ func (_m *BatchImageJob) String() string { builder.WriteString(v.Format(time.ANSIC)) } builder.WriteString(", ") + if v := _m.DownloadedAt; v != nil { + builder.WriteString("downloaded_at=") + builder.WriteString(v.Format(time.ANSIC)) + } + builder.WriteString(", ") + if v := _m.UserDeletedAt; v != nil { + builder.WriteString("user_deleted_at=") + builder.WriteString(v.Format(time.ANSIC)) + } + builder.WriteString(", ") if v := _m.LastErrorCode; v != nil { builder.WriteString("last_error_code=") builder.WriteString(*v) diff --git a/backend/ent/batchimagejob/batchimagejob.go b/backend/ent/batchimagejob/batchimagejob.go index 19c7d03131..be819183da 100644 --- a/backend/ent/batchimagejob/batchimagejob.go +++ b/backend/ent/batchimagejob/batchimagejob.go @@ -25,6 +25,8 @@ const ( FieldProvider = "provider" // FieldModel holds the string denoting the model field in the database. FieldModel = "model" + // FieldTaskName holds the string denoting the task_name field in the database. + FieldTaskName = "task_name" // FieldStatus holds the string denoting the status field in the database. FieldStatus = "status" // FieldProviderJobName holds the string denoting the provider_job_name field in the database. @@ -71,6 +73,10 @@ const ( FieldInputDeletedAt = "input_deleted_at" // FieldOutputDeletedAt holds the string denoting the output_deleted_at field in the database. FieldOutputDeletedAt = "output_deleted_at" + // FieldDownloadedAt holds the string denoting the downloaded_at field in the database. + FieldDownloadedAt = "downloaded_at" + // FieldUserDeletedAt holds the string denoting the user_deleted_at field in the database. + FieldUserDeletedAt = "user_deleted_at" // FieldLastErrorCode holds the string denoting the last_error_code field in the database. FieldLastErrorCode = "last_error_code" // FieldLastErrorMessage holds the string denoting the last_error_message field in the database. @@ -100,6 +106,7 @@ var Columns = []string{ FieldAccountID, FieldProvider, FieldModel, + FieldTaskName, FieldStatus, FieldProviderJobName, FieldProviderInputRef, @@ -123,6 +130,8 @@ var Columns = []string{ FieldOutputExpiresAt, FieldInputDeletedAt, FieldOutputDeletedAt, + FieldDownloadedAt, + FieldUserDeletedAt, FieldLastErrorCode, FieldLastErrorMessage, FieldCreatedAt, @@ -150,6 +159,10 @@ var ( ProviderValidator func(string) error // ModelValidator is a validator for the "model" field. It is called by the builders before save. ModelValidator func(string) error + // DefaultTaskName holds the default value on creation for the "task_name" field. + DefaultTaskName string + // TaskNameValidator is a validator for the "task_name" field. It is called by the builders before save. + TaskNameValidator func(string) error // DefaultStatus holds the default value on creation for the "status" field. DefaultStatus string // StatusValidator is a validator for the "status" field. It is called by the builders before save. @@ -236,6 +249,11 @@ func ByModel(opts ...sql.OrderTermOption) OrderOption { return sql.OrderByField(FieldModel, opts...).ToFunc() } +// ByTaskName orders the results by the task_name field. +func ByTaskName(opts ...sql.OrderTermOption) OrderOption { + return sql.OrderByField(FieldTaskName, opts...).ToFunc() +} + // ByStatus orders the results by the status field. func ByStatus(opts ...sql.OrderTermOption) OrderOption { return sql.OrderByField(FieldStatus, opts...).ToFunc() @@ -351,6 +369,16 @@ func ByOutputDeletedAt(opts ...sql.OrderTermOption) OrderOption { return sql.OrderByField(FieldOutputDeletedAt, opts...).ToFunc() } +// ByDownloadedAt orders the results by the downloaded_at field. +func ByDownloadedAt(opts ...sql.OrderTermOption) OrderOption { + return sql.OrderByField(FieldDownloadedAt, opts...).ToFunc() +} + +// ByUserDeletedAt orders the results by the user_deleted_at field. +func ByUserDeletedAt(opts ...sql.OrderTermOption) OrderOption { + return sql.OrderByField(FieldUserDeletedAt, opts...).ToFunc() +} + // ByLastErrorCode orders the results by the last_error_code field. func ByLastErrorCode(opts ...sql.OrderTermOption) OrderOption { return sql.OrderByField(FieldLastErrorCode, opts...).ToFunc() diff --git a/backend/ent/batchimagejob/where.go b/backend/ent/batchimagejob/where.go index a8d66994fb..b94722e41d 100644 --- a/backend/ent/batchimagejob/where.go +++ b/backend/ent/batchimagejob/where.go @@ -84,6 +84,11 @@ func Model(v string) predicate.BatchImageJob { return predicate.BatchImageJob(sql.FieldEQ(FieldModel, v)) } +// TaskName applies equality check predicate on the "task_name" field. It's identical to TaskNameEQ. +func TaskName(v string) predicate.BatchImageJob { + return predicate.BatchImageJob(sql.FieldEQ(FieldTaskName, v)) +} + // Status applies equality check predicate on the "status" field. It's identical to StatusEQ. func Status(v string) predicate.BatchImageJob { return predicate.BatchImageJob(sql.FieldEQ(FieldStatus, v)) @@ -199,6 +204,16 @@ func OutputDeletedAt(v time.Time) predicate.BatchImageJob { return predicate.BatchImageJob(sql.FieldEQ(FieldOutputDeletedAt, v)) } +// DownloadedAt applies equality check predicate on the "downloaded_at" field. It's identical to DownloadedAtEQ. +func DownloadedAt(v time.Time) predicate.BatchImageJob { + return predicate.BatchImageJob(sql.FieldEQ(FieldDownloadedAt, v)) +} + +// UserDeletedAt applies equality check predicate on the "user_deleted_at" field. It's identical to UserDeletedAtEQ. +func UserDeletedAt(v time.Time) predicate.BatchImageJob { + return predicate.BatchImageJob(sql.FieldEQ(FieldUserDeletedAt, v)) +} + // LastErrorCode applies equality check predicate on the "last_error_code" field. It's identical to LastErrorCodeEQ. func LastErrorCode(v string) predicate.BatchImageJob { return predicate.BatchImageJob(sql.FieldEQ(FieldLastErrorCode, v)) @@ -574,6 +589,71 @@ func ModelContainsFold(v string) predicate.BatchImageJob { return predicate.BatchImageJob(sql.FieldContainsFold(FieldModel, v)) } +// TaskNameEQ applies the EQ predicate on the "task_name" field. +func TaskNameEQ(v string) predicate.BatchImageJob { + return predicate.BatchImageJob(sql.FieldEQ(FieldTaskName, v)) +} + +// TaskNameNEQ applies the NEQ predicate on the "task_name" field. +func TaskNameNEQ(v string) predicate.BatchImageJob { + return predicate.BatchImageJob(sql.FieldNEQ(FieldTaskName, v)) +} + +// TaskNameIn applies the In predicate on the "task_name" field. +func TaskNameIn(vs ...string) predicate.BatchImageJob { + return predicate.BatchImageJob(sql.FieldIn(FieldTaskName, vs...)) +} + +// TaskNameNotIn applies the NotIn predicate on the "task_name" field. +func TaskNameNotIn(vs ...string) predicate.BatchImageJob { + return predicate.BatchImageJob(sql.FieldNotIn(FieldTaskName, vs...)) +} + +// TaskNameGT applies the GT predicate on the "task_name" field. +func TaskNameGT(v string) predicate.BatchImageJob { + return predicate.BatchImageJob(sql.FieldGT(FieldTaskName, v)) +} + +// TaskNameGTE applies the GTE predicate on the "task_name" field. +func TaskNameGTE(v string) predicate.BatchImageJob { + return predicate.BatchImageJob(sql.FieldGTE(FieldTaskName, v)) +} + +// TaskNameLT applies the LT predicate on the "task_name" field. +func TaskNameLT(v string) predicate.BatchImageJob { + return predicate.BatchImageJob(sql.FieldLT(FieldTaskName, v)) +} + +// TaskNameLTE applies the LTE predicate on the "task_name" field. +func TaskNameLTE(v string) predicate.BatchImageJob { + return predicate.BatchImageJob(sql.FieldLTE(FieldTaskName, v)) +} + +// TaskNameContains applies the Contains predicate on the "task_name" field. +func TaskNameContains(v string) predicate.BatchImageJob { + return predicate.BatchImageJob(sql.FieldContains(FieldTaskName, v)) +} + +// TaskNameHasPrefix applies the HasPrefix predicate on the "task_name" field. +func TaskNameHasPrefix(v string) predicate.BatchImageJob { + return predicate.BatchImageJob(sql.FieldHasPrefix(FieldTaskName, v)) +} + +// TaskNameHasSuffix applies the HasSuffix predicate on the "task_name" field. +func TaskNameHasSuffix(v string) predicate.BatchImageJob { + return predicate.BatchImageJob(sql.FieldHasSuffix(FieldTaskName, v)) +} + +// TaskNameEqualFold applies the EqualFold predicate on the "task_name" field. +func TaskNameEqualFold(v string) predicate.BatchImageJob { + return predicate.BatchImageJob(sql.FieldEqualFold(FieldTaskName, v)) +} + +// TaskNameContainsFold applies the ContainsFold predicate on the "task_name" field. +func TaskNameContainsFold(v string) predicate.BatchImageJob { + return predicate.BatchImageJob(sql.FieldContainsFold(FieldTaskName, v)) +} + // StatusEQ applies the EQ predicate on the "status" field. func StatusEQ(v string) predicate.BatchImageJob { return predicate.BatchImageJob(sql.FieldEQ(FieldStatus, v)) @@ -1909,6 +1989,106 @@ func OutputDeletedAtNotNil() predicate.BatchImageJob { return predicate.BatchImageJob(sql.FieldNotNull(FieldOutputDeletedAt)) } +// DownloadedAtEQ applies the EQ predicate on the "downloaded_at" field. +func DownloadedAtEQ(v time.Time) predicate.BatchImageJob { + return predicate.BatchImageJob(sql.FieldEQ(FieldDownloadedAt, v)) +} + +// DownloadedAtNEQ applies the NEQ predicate on the "downloaded_at" field. +func DownloadedAtNEQ(v time.Time) predicate.BatchImageJob { + return predicate.BatchImageJob(sql.FieldNEQ(FieldDownloadedAt, v)) +} + +// DownloadedAtIn applies the In predicate on the "downloaded_at" field. +func DownloadedAtIn(vs ...time.Time) predicate.BatchImageJob { + return predicate.BatchImageJob(sql.FieldIn(FieldDownloadedAt, vs...)) +} + +// DownloadedAtNotIn applies the NotIn predicate on the "downloaded_at" field. +func DownloadedAtNotIn(vs ...time.Time) predicate.BatchImageJob { + return predicate.BatchImageJob(sql.FieldNotIn(FieldDownloadedAt, vs...)) +} + +// DownloadedAtGT applies the GT predicate on the "downloaded_at" field. +func DownloadedAtGT(v time.Time) predicate.BatchImageJob { + return predicate.BatchImageJob(sql.FieldGT(FieldDownloadedAt, v)) +} + +// DownloadedAtGTE applies the GTE predicate on the "downloaded_at" field. +func DownloadedAtGTE(v time.Time) predicate.BatchImageJob { + return predicate.BatchImageJob(sql.FieldGTE(FieldDownloadedAt, v)) +} + +// DownloadedAtLT applies the LT predicate on the "downloaded_at" field. +func DownloadedAtLT(v time.Time) predicate.BatchImageJob { + return predicate.BatchImageJob(sql.FieldLT(FieldDownloadedAt, v)) +} + +// DownloadedAtLTE applies the LTE predicate on the "downloaded_at" field. +func DownloadedAtLTE(v time.Time) predicate.BatchImageJob { + return predicate.BatchImageJob(sql.FieldLTE(FieldDownloadedAt, v)) +} + +// DownloadedAtIsNil applies the IsNil predicate on the "downloaded_at" field. +func DownloadedAtIsNil() predicate.BatchImageJob { + return predicate.BatchImageJob(sql.FieldIsNull(FieldDownloadedAt)) +} + +// DownloadedAtNotNil applies the NotNil predicate on the "downloaded_at" field. +func DownloadedAtNotNil() predicate.BatchImageJob { + return predicate.BatchImageJob(sql.FieldNotNull(FieldDownloadedAt)) +} + +// UserDeletedAtEQ applies the EQ predicate on the "user_deleted_at" field. +func UserDeletedAtEQ(v time.Time) predicate.BatchImageJob { + return predicate.BatchImageJob(sql.FieldEQ(FieldUserDeletedAt, v)) +} + +// UserDeletedAtNEQ applies the NEQ predicate on the "user_deleted_at" field. +func UserDeletedAtNEQ(v time.Time) predicate.BatchImageJob { + return predicate.BatchImageJob(sql.FieldNEQ(FieldUserDeletedAt, v)) +} + +// UserDeletedAtIn applies the In predicate on the "user_deleted_at" field. +func UserDeletedAtIn(vs ...time.Time) predicate.BatchImageJob { + return predicate.BatchImageJob(sql.FieldIn(FieldUserDeletedAt, vs...)) +} + +// UserDeletedAtNotIn applies the NotIn predicate on the "user_deleted_at" field. +func UserDeletedAtNotIn(vs ...time.Time) predicate.BatchImageJob { + return predicate.BatchImageJob(sql.FieldNotIn(FieldUserDeletedAt, vs...)) +} + +// UserDeletedAtGT applies the GT predicate on the "user_deleted_at" field. +func UserDeletedAtGT(v time.Time) predicate.BatchImageJob { + return predicate.BatchImageJob(sql.FieldGT(FieldUserDeletedAt, v)) +} + +// UserDeletedAtGTE applies the GTE predicate on the "user_deleted_at" field. +func UserDeletedAtGTE(v time.Time) predicate.BatchImageJob { + return predicate.BatchImageJob(sql.FieldGTE(FieldUserDeletedAt, v)) +} + +// UserDeletedAtLT applies the LT predicate on the "user_deleted_at" field. +func UserDeletedAtLT(v time.Time) predicate.BatchImageJob { + return predicate.BatchImageJob(sql.FieldLT(FieldUserDeletedAt, v)) +} + +// UserDeletedAtLTE applies the LTE predicate on the "user_deleted_at" field. +func UserDeletedAtLTE(v time.Time) predicate.BatchImageJob { + return predicate.BatchImageJob(sql.FieldLTE(FieldUserDeletedAt, v)) +} + +// UserDeletedAtIsNil applies the IsNil predicate on the "user_deleted_at" field. +func UserDeletedAtIsNil() predicate.BatchImageJob { + return predicate.BatchImageJob(sql.FieldIsNull(FieldUserDeletedAt)) +} + +// UserDeletedAtNotNil applies the NotNil predicate on the "user_deleted_at" field. +func UserDeletedAtNotNil() predicate.BatchImageJob { + return predicate.BatchImageJob(sql.FieldNotNull(FieldUserDeletedAt)) +} + // LastErrorCodeEQ applies the EQ predicate on the "last_error_code" field. func LastErrorCodeEQ(v string) predicate.BatchImageJob { return predicate.BatchImageJob(sql.FieldEQ(FieldLastErrorCode, v)) diff --git a/backend/ent/batchimagejob_create.go b/backend/ent/batchimagejob_create.go index 26df896d1c..88c1197b15 100644 --- a/backend/ent/batchimagejob_create.go +++ b/backend/ent/batchimagejob_create.go @@ -74,6 +74,20 @@ func (_c *BatchImageJobCreate) SetModel(v string) *BatchImageJobCreate { return _c } +// SetTaskName sets the "task_name" field. +func (_c *BatchImageJobCreate) SetTaskName(v string) *BatchImageJobCreate { + _c.mutation.SetTaskName(v) + return _c +} + +// SetNillableTaskName sets the "task_name" field if the given value is not nil. +func (_c *BatchImageJobCreate) SetNillableTaskName(v *string) *BatchImageJobCreate { + if v != nil { + _c.SetTaskName(*v) + } + return _c +} + // SetStatus sets the "status" field. func (_c *BatchImageJobCreate) SetStatus(v string) *BatchImageJobCreate { _c.mutation.SetStatus(v) @@ -388,6 +402,34 @@ func (_c *BatchImageJobCreate) SetNillableOutputDeletedAt(v *time.Time) *BatchIm return _c } +// SetDownloadedAt sets the "downloaded_at" field. +func (_c *BatchImageJobCreate) SetDownloadedAt(v time.Time) *BatchImageJobCreate { + _c.mutation.SetDownloadedAt(v) + return _c +} + +// SetNillableDownloadedAt sets the "downloaded_at" field if the given value is not nil. +func (_c *BatchImageJobCreate) SetNillableDownloadedAt(v *time.Time) *BatchImageJobCreate { + if v != nil { + _c.SetDownloadedAt(*v) + } + return _c +} + +// SetUserDeletedAt sets the "user_deleted_at" field. +func (_c *BatchImageJobCreate) SetUserDeletedAt(v time.Time) *BatchImageJobCreate { + _c.mutation.SetUserDeletedAt(v) + return _c +} + +// SetNillableUserDeletedAt sets the "user_deleted_at" field if the given value is not nil. +func (_c *BatchImageJobCreate) SetNillableUserDeletedAt(v *time.Time) *BatchImageJobCreate { + if v != nil { + _c.SetUserDeletedAt(*v) + } + return _c +} + // SetLastErrorCode sets the "last_error_code" field. func (_c *BatchImageJobCreate) SetLastErrorCode(v string) *BatchImageJobCreate { _c.mutation.SetLastErrorCode(v) @@ -535,6 +577,10 @@ func (_c *BatchImageJobCreate) ExecX(ctx context.Context) { // defaults sets the default values of the builder before save. func (_c *BatchImageJobCreate) defaults() { + if _, ok := _c.mutation.TaskName(); !ok { + v := batchimagejob.DefaultTaskName + _c.mutation.SetTaskName(v) + } if _, ok := _c.mutation.Status(); !ok { v := batchimagejob.DefaultStatus _c.mutation.SetStatus(v) @@ -606,6 +652,14 @@ func (_c *BatchImageJobCreate) check() error { return &ValidationError{Name: "model", err: fmt.Errorf(`ent: validator failed for field "BatchImageJob.model": %w`, err)} } } + if _, ok := _c.mutation.TaskName(); !ok { + return &ValidationError{Name: "task_name", err: errors.New(`ent: missing required field "BatchImageJob.task_name"`)} + } + if v, ok := _c.mutation.TaskName(); ok { + if err := batchimagejob.TaskNameValidator(v); err != nil { + return &ValidationError{Name: "task_name", err: fmt.Errorf(`ent: validator failed for field "BatchImageJob.task_name": %w`, err)} + } + } if _, ok := _c.mutation.Status(); !ok { return &ValidationError{Name: "status", err: errors.New(`ent: missing required field "BatchImageJob.status"`)} } @@ -750,6 +804,10 @@ func (_c *BatchImageJobCreate) createSpec() (*BatchImageJob, *sqlgraph.CreateSpe _spec.SetField(batchimagejob.FieldModel, field.TypeString, value) _node.Model = value } + if value, ok := _c.mutation.TaskName(); ok { + _spec.SetField(batchimagejob.FieldTaskName, field.TypeString, value) + _node.TaskName = value + } if value, ok := _c.mutation.Status(); ok { _spec.SetField(batchimagejob.FieldStatus, field.TypeString, value) _node.Status = value @@ -842,6 +900,14 @@ func (_c *BatchImageJobCreate) createSpec() (*BatchImageJob, *sqlgraph.CreateSpe _spec.SetField(batchimagejob.FieldOutputDeletedAt, field.TypeTime, value) _node.OutputDeletedAt = &value } + if value, ok := _c.mutation.DownloadedAt(); ok { + _spec.SetField(batchimagejob.FieldDownloadedAt, field.TypeTime, value) + _node.DownloadedAt = &value + } + if value, ok := _c.mutation.UserDeletedAt(); ok { + _spec.SetField(batchimagejob.FieldUserDeletedAt, field.TypeTime, value) + _node.UserDeletedAt = &value + } if value, ok := _c.mutation.LastErrorCode(); ok { _spec.SetField(batchimagejob.FieldLastErrorCode, field.TypeString, value) _node.LastErrorCode = &value @@ -1016,6 +1082,18 @@ func (u *BatchImageJobUpsert) UpdateModel() *BatchImageJobUpsert { return u } +// SetTaskName sets the "task_name" field. +func (u *BatchImageJobUpsert) SetTaskName(v string) *BatchImageJobUpsert { + u.Set(batchimagejob.FieldTaskName, v) + return u +} + +// UpdateTaskName sets the "task_name" field to the value that was provided on create. +func (u *BatchImageJobUpsert) UpdateTaskName() *BatchImageJobUpsert { + u.SetExcluded(batchimagejob.FieldTaskName) + return u +} + // SetStatus sets the "status" field. func (u *BatchImageJobUpsert) SetStatus(v string) *BatchImageJobUpsert { u.Set(batchimagejob.FieldStatus, v) @@ -1430,6 +1508,42 @@ func (u *BatchImageJobUpsert) ClearOutputDeletedAt() *BatchImageJobUpsert { return u } +// SetDownloadedAt sets the "downloaded_at" field. +func (u *BatchImageJobUpsert) SetDownloadedAt(v time.Time) *BatchImageJobUpsert { + u.Set(batchimagejob.FieldDownloadedAt, v) + return u +} + +// UpdateDownloadedAt sets the "downloaded_at" field to the value that was provided on create. +func (u *BatchImageJobUpsert) UpdateDownloadedAt() *BatchImageJobUpsert { + u.SetExcluded(batchimagejob.FieldDownloadedAt) + return u +} + +// ClearDownloadedAt clears the value of the "downloaded_at" field. +func (u *BatchImageJobUpsert) ClearDownloadedAt() *BatchImageJobUpsert { + u.SetNull(batchimagejob.FieldDownloadedAt) + return u +} + +// SetUserDeletedAt sets the "user_deleted_at" field. +func (u *BatchImageJobUpsert) SetUserDeletedAt(v time.Time) *BatchImageJobUpsert { + u.Set(batchimagejob.FieldUserDeletedAt, v) + return u +} + +// UpdateUserDeletedAt sets the "user_deleted_at" field to the value that was provided on create. +func (u *BatchImageJobUpsert) UpdateUserDeletedAt() *BatchImageJobUpsert { + u.SetExcluded(batchimagejob.FieldUserDeletedAt) + return u +} + +// ClearUserDeletedAt clears the value of the "user_deleted_at" field. +func (u *BatchImageJobUpsert) ClearUserDeletedAt() *BatchImageJobUpsert { + u.SetNull(batchimagejob.FieldUserDeletedAt) + return u +} + // SetLastErrorCode sets the "last_error_code" field. func (u *BatchImageJobUpsert) SetLastErrorCode(v string) *BatchImageJobUpsert { u.Set(batchimagejob.FieldLastErrorCode, v) @@ -1703,6 +1817,20 @@ func (u *BatchImageJobUpsertOne) UpdateModel() *BatchImageJobUpsertOne { }) } +// SetTaskName sets the "task_name" field. +func (u *BatchImageJobUpsertOne) SetTaskName(v string) *BatchImageJobUpsertOne { + return u.Update(func(s *BatchImageJobUpsert) { + s.SetTaskName(v) + }) +} + +// UpdateTaskName sets the "task_name" field to the value that was provided on create. +func (u *BatchImageJobUpsertOne) UpdateTaskName() *BatchImageJobUpsertOne { + return u.Update(func(s *BatchImageJobUpsert) { + s.UpdateTaskName() + }) +} + // SetStatus sets the "status" field. func (u *BatchImageJobUpsertOne) SetStatus(v string) *BatchImageJobUpsertOne { return u.Update(func(s *BatchImageJobUpsert) { @@ -2186,6 +2314,48 @@ func (u *BatchImageJobUpsertOne) ClearOutputDeletedAt() *BatchImageJobUpsertOne }) } +// SetDownloadedAt sets the "downloaded_at" field. +func (u *BatchImageJobUpsertOne) SetDownloadedAt(v time.Time) *BatchImageJobUpsertOne { + return u.Update(func(s *BatchImageJobUpsert) { + s.SetDownloadedAt(v) + }) +} + +// UpdateDownloadedAt sets the "downloaded_at" field to the value that was provided on create. +func (u *BatchImageJobUpsertOne) UpdateDownloadedAt() *BatchImageJobUpsertOne { + return u.Update(func(s *BatchImageJobUpsert) { + s.UpdateDownloadedAt() + }) +} + +// ClearDownloadedAt clears the value of the "downloaded_at" field. +func (u *BatchImageJobUpsertOne) ClearDownloadedAt() *BatchImageJobUpsertOne { + return u.Update(func(s *BatchImageJobUpsert) { + s.ClearDownloadedAt() + }) +} + +// SetUserDeletedAt sets the "user_deleted_at" field. +func (u *BatchImageJobUpsertOne) SetUserDeletedAt(v time.Time) *BatchImageJobUpsertOne { + return u.Update(func(s *BatchImageJobUpsert) { + s.SetUserDeletedAt(v) + }) +} + +// UpdateUserDeletedAt sets the "user_deleted_at" field to the value that was provided on create. +func (u *BatchImageJobUpsertOne) UpdateUserDeletedAt() *BatchImageJobUpsertOne { + return u.Update(func(s *BatchImageJobUpsert) { + s.UpdateUserDeletedAt() + }) +} + +// ClearUserDeletedAt clears the value of the "user_deleted_at" field. +func (u *BatchImageJobUpsertOne) ClearUserDeletedAt() *BatchImageJobUpsertOne { + return u.Update(func(s *BatchImageJobUpsert) { + s.ClearUserDeletedAt() + }) +} + // SetLastErrorCode sets the "last_error_code" field. func (u *BatchImageJobUpsertOne) SetLastErrorCode(v string) *BatchImageJobUpsertOne { return u.Update(func(s *BatchImageJobUpsert) { @@ -2645,6 +2815,20 @@ func (u *BatchImageJobUpsertBulk) UpdateModel() *BatchImageJobUpsertBulk { }) } +// SetTaskName sets the "task_name" field. +func (u *BatchImageJobUpsertBulk) SetTaskName(v string) *BatchImageJobUpsertBulk { + return u.Update(func(s *BatchImageJobUpsert) { + s.SetTaskName(v) + }) +} + +// UpdateTaskName sets the "task_name" field to the value that was provided on create. +func (u *BatchImageJobUpsertBulk) UpdateTaskName() *BatchImageJobUpsertBulk { + return u.Update(func(s *BatchImageJobUpsert) { + s.UpdateTaskName() + }) +} + // SetStatus sets the "status" field. func (u *BatchImageJobUpsertBulk) SetStatus(v string) *BatchImageJobUpsertBulk { return u.Update(func(s *BatchImageJobUpsert) { @@ -3128,6 +3312,48 @@ func (u *BatchImageJobUpsertBulk) ClearOutputDeletedAt() *BatchImageJobUpsertBul }) } +// SetDownloadedAt sets the "downloaded_at" field. +func (u *BatchImageJobUpsertBulk) SetDownloadedAt(v time.Time) *BatchImageJobUpsertBulk { + return u.Update(func(s *BatchImageJobUpsert) { + s.SetDownloadedAt(v) + }) +} + +// UpdateDownloadedAt sets the "downloaded_at" field to the value that was provided on create. +func (u *BatchImageJobUpsertBulk) UpdateDownloadedAt() *BatchImageJobUpsertBulk { + return u.Update(func(s *BatchImageJobUpsert) { + s.UpdateDownloadedAt() + }) +} + +// ClearDownloadedAt clears the value of the "downloaded_at" field. +func (u *BatchImageJobUpsertBulk) ClearDownloadedAt() *BatchImageJobUpsertBulk { + return u.Update(func(s *BatchImageJobUpsert) { + s.ClearDownloadedAt() + }) +} + +// SetUserDeletedAt sets the "user_deleted_at" field. +func (u *BatchImageJobUpsertBulk) SetUserDeletedAt(v time.Time) *BatchImageJobUpsertBulk { + return u.Update(func(s *BatchImageJobUpsert) { + s.SetUserDeletedAt(v) + }) +} + +// UpdateUserDeletedAt sets the "user_deleted_at" field to the value that was provided on create. +func (u *BatchImageJobUpsertBulk) UpdateUserDeletedAt() *BatchImageJobUpsertBulk { + return u.Update(func(s *BatchImageJobUpsert) { + s.UpdateUserDeletedAt() + }) +} + +// ClearUserDeletedAt clears the value of the "user_deleted_at" field. +func (u *BatchImageJobUpsertBulk) ClearUserDeletedAt() *BatchImageJobUpsertBulk { + return u.Update(func(s *BatchImageJobUpsert) { + s.ClearUserDeletedAt() + }) +} + // SetLastErrorCode sets the "last_error_code" field. func (u *BatchImageJobUpsertBulk) SetLastErrorCode(v string) *BatchImageJobUpsertBulk { return u.Update(func(s *BatchImageJobUpsert) { diff --git a/backend/ent/batchimagejob_update.go b/backend/ent/batchimagejob_update.go index 96572b3b22..8df7302500 100644 --- a/backend/ent/batchimagejob_update.go +++ b/backend/ent/batchimagejob_update.go @@ -131,6 +131,20 @@ func (_u *BatchImageJobUpdate) SetNillableModel(v *string) *BatchImageJobUpdate return _u } +// SetTaskName sets the "task_name" field. +func (_u *BatchImageJobUpdate) SetTaskName(v string) *BatchImageJobUpdate { + _u.mutation.SetTaskName(v) + return _u +} + +// SetNillableTaskName sets the "task_name" field if the given value is not nil. +func (_u *BatchImageJobUpdate) SetNillableTaskName(v *string) *BatchImageJobUpdate { + if v != nil { + _u.SetTaskName(*v) + } + return _u +} + // SetStatus sets the "status" field. func (_u *BatchImageJobUpdate) SetStatus(v string) *BatchImageJobUpdate { _u.mutation.SetStatus(v) @@ -600,6 +614,46 @@ func (_u *BatchImageJobUpdate) ClearOutputDeletedAt() *BatchImageJobUpdate { return _u } +// SetDownloadedAt sets the "downloaded_at" field. +func (_u *BatchImageJobUpdate) SetDownloadedAt(v time.Time) *BatchImageJobUpdate { + _u.mutation.SetDownloadedAt(v) + return _u +} + +// SetNillableDownloadedAt sets the "downloaded_at" field if the given value is not nil. +func (_u *BatchImageJobUpdate) SetNillableDownloadedAt(v *time.Time) *BatchImageJobUpdate { + if v != nil { + _u.SetDownloadedAt(*v) + } + return _u +} + +// ClearDownloadedAt clears the value of the "downloaded_at" field. +func (_u *BatchImageJobUpdate) ClearDownloadedAt() *BatchImageJobUpdate { + _u.mutation.ClearDownloadedAt() + return _u +} + +// SetUserDeletedAt sets the "user_deleted_at" field. +func (_u *BatchImageJobUpdate) SetUserDeletedAt(v time.Time) *BatchImageJobUpdate { + _u.mutation.SetUserDeletedAt(v) + return _u +} + +// SetNillableUserDeletedAt sets the "user_deleted_at" field if the given value is not nil. +func (_u *BatchImageJobUpdate) SetNillableUserDeletedAt(v *time.Time) *BatchImageJobUpdate { + if v != nil { + _u.SetUserDeletedAt(*v) + } + return _u +} + +// ClearUserDeletedAt clears the value of the "user_deleted_at" field. +func (_u *BatchImageJobUpdate) ClearUserDeletedAt() *BatchImageJobUpdate { + _u.mutation.ClearUserDeletedAt() + return _u +} + // SetLastErrorCode sets the "last_error_code" field. func (_u *BatchImageJobUpdate) SetLastErrorCode(v string) *BatchImageJobUpdate { _u.mutation.SetLastErrorCode(v) @@ -779,6 +833,11 @@ func (_u *BatchImageJobUpdate) check() error { return &ValidationError{Name: "model", err: fmt.Errorf(`ent: validator failed for field "BatchImageJob.model": %w`, err)} } } + if v, ok := _u.mutation.TaskName(); ok { + if err := batchimagejob.TaskNameValidator(v); err != nil { + return &ValidationError{Name: "task_name", err: fmt.Errorf(`ent: validator failed for field "BatchImageJob.task_name": %w`, err)} + } + } if v, ok := _u.mutation.Status(); ok { if err := batchimagejob.StatusValidator(v); err != nil { return &ValidationError{Name: "status", err: fmt.Errorf(`ent: validator failed for field "BatchImageJob.status": %w`, err)} @@ -884,6 +943,9 @@ func (_u *BatchImageJobUpdate) sqlSave(ctx context.Context) (_node int, err erro if value, ok := _u.mutation.Model(); ok { _spec.SetField(batchimagejob.FieldModel, field.TypeString, value) } + if value, ok := _u.mutation.TaskName(); ok { + _spec.SetField(batchimagejob.FieldTaskName, field.TypeString, value) + } if value, ok := _u.mutation.Status(); ok { _spec.SetField(batchimagejob.FieldStatus, field.TypeString, value) } @@ -1022,6 +1084,18 @@ func (_u *BatchImageJobUpdate) sqlSave(ctx context.Context) (_node int, err erro if _u.mutation.OutputDeletedAtCleared() { _spec.ClearField(batchimagejob.FieldOutputDeletedAt, field.TypeTime) } + if value, ok := _u.mutation.DownloadedAt(); ok { + _spec.SetField(batchimagejob.FieldDownloadedAt, field.TypeTime, value) + } + if _u.mutation.DownloadedAtCleared() { + _spec.ClearField(batchimagejob.FieldDownloadedAt, field.TypeTime) + } + if value, ok := _u.mutation.UserDeletedAt(); ok { + _spec.SetField(batchimagejob.FieldUserDeletedAt, field.TypeTime, value) + } + if _u.mutation.UserDeletedAtCleared() { + _spec.ClearField(batchimagejob.FieldUserDeletedAt, field.TypeTime) + } if value, ok := _u.mutation.LastErrorCode(); ok { _spec.SetField(batchimagejob.FieldLastErrorCode, field.TypeString, value) } @@ -1184,6 +1258,20 @@ func (_u *BatchImageJobUpdateOne) SetNillableModel(v *string) *BatchImageJobUpda return _u } +// SetTaskName sets the "task_name" field. +func (_u *BatchImageJobUpdateOne) SetTaskName(v string) *BatchImageJobUpdateOne { + _u.mutation.SetTaskName(v) + return _u +} + +// SetNillableTaskName sets the "task_name" field if the given value is not nil. +func (_u *BatchImageJobUpdateOne) SetNillableTaskName(v *string) *BatchImageJobUpdateOne { + if v != nil { + _u.SetTaskName(*v) + } + return _u +} + // SetStatus sets the "status" field. func (_u *BatchImageJobUpdateOne) SetStatus(v string) *BatchImageJobUpdateOne { _u.mutation.SetStatus(v) @@ -1653,6 +1741,46 @@ func (_u *BatchImageJobUpdateOne) ClearOutputDeletedAt() *BatchImageJobUpdateOne return _u } +// SetDownloadedAt sets the "downloaded_at" field. +func (_u *BatchImageJobUpdateOne) SetDownloadedAt(v time.Time) *BatchImageJobUpdateOne { + _u.mutation.SetDownloadedAt(v) + return _u +} + +// SetNillableDownloadedAt sets the "downloaded_at" field if the given value is not nil. +func (_u *BatchImageJobUpdateOne) SetNillableDownloadedAt(v *time.Time) *BatchImageJobUpdateOne { + if v != nil { + _u.SetDownloadedAt(*v) + } + return _u +} + +// ClearDownloadedAt clears the value of the "downloaded_at" field. +func (_u *BatchImageJobUpdateOne) ClearDownloadedAt() *BatchImageJobUpdateOne { + _u.mutation.ClearDownloadedAt() + return _u +} + +// SetUserDeletedAt sets the "user_deleted_at" field. +func (_u *BatchImageJobUpdateOne) SetUserDeletedAt(v time.Time) *BatchImageJobUpdateOne { + _u.mutation.SetUserDeletedAt(v) + return _u +} + +// SetNillableUserDeletedAt sets the "user_deleted_at" field if the given value is not nil. +func (_u *BatchImageJobUpdateOne) SetNillableUserDeletedAt(v *time.Time) *BatchImageJobUpdateOne { + if v != nil { + _u.SetUserDeletedAt(*v) + } + return _u +} + +// ClearUserDeletedAt clears the value of the "user_deleted_at" field. +func (_u *BatchImageJobUpdateOne) ClearUserDeletedAt() *BatchImageJobUpdateOne { + _u.mutation.ClearUserDeletedAt() + return _u +} + // SetLastErrorCode sets the "last_error_code" field. func (_u *BatchImageJobUpdateOne) SetLastErrorCode(v string) *BatchImageJobUpdateOne { _u.mutation.SetLastErrorCode(v) @@ -1845,6 +1973,11 @@ func (_u *BatchImageJobUpdateOne) check() error { return &ValidationError{Name: "model", err: fmt.Errorf(`ent: validator failed for field "BatchImageJob.model": %w`, err)} } } + if v, ok := _u.mutation.TaskName(); ok { + if err := batchimagejob.TaskNameValidator(v); err != nil { + return &ValidationError{Name: "task_name", err: fmt.Errorf(`ent: validator failed for field "BatchImageJob.task_name": %w`, err)} + } + } if v, ok := _u.mutation.Status(); ok { if err := batchimagejob.StatusValidator(v); err != nil { return &ValidationError{Name: "status", err: fmt.Errorf(`ent: validator failed for field "BatchImageJob.status": %w`, err)} @@ -1967,6 +2100,9 @@ func (_u *BatchImageJobUpdateOne) sqlSave(ctx context.Context) (_node *BatchImag if value, ok := _u.mutation.Model(); ok { _spec.SetField(batchimagejob.FieldModel, field.TypeString, value) } + if value, ok := _u.mutation.TaskName(); ok { + _spec.SetField(batchimagejob.FieldTaskName, field.TypeString, value) + } if value, ok := _u.mutation.Status(); ok { _spec.SetField(batchimagejob.FieldStatus, field.TypeString, value) } @@ -2105,6 +2241,18 @@ func (_u *BatchImageJobUpdateOne) sqlSave(ctx context.Context) (_node *BatchImag if _u.mutation.OutputDeletedAtCleared() { _spec.ClearField(batchimagejob.FieldOutputDeletedAt, field.TypeTime) } + if value, ok := _u.mutation.DownloadedAt(); ok { + _spec.SetField(batchimagejob.FieldDownloadedAt, field.TypeTime, value) + } + if _u.mutation.DownloadedAtCleared() { + _spec.ClearField(batchimagejob.FieldDownloadedAt, field.TypeTime) + } + if value, ok := _u.mutation.UserDeletedAt(); ok { + _spec.SetField(batchimagejob.FieldUserDeletedAt, field.TypeTime, value) + } + if _u.mutation.UserDeletedAtCleared() { + _spec.ClearField(batchimagejob.FieldUserDeletedAt, field.TypeTime) + } if value, ok := _u.mutation.LastErrorCode(); ok { _spec.SetField(batchimagejob.FieldLastErrorCode, field.TypeString, value) } diff --git a/backend/ent/group.go b/backend/ent/group.go index 5624d47d83..2a0eb4d3ac 100644 --- a/backend/ent/group.go +++ b/backend/ent/group.go @@ -57,6 +57,8 @@ type Group struct { DefaultValidityDays int `json:"default_validity_days,omitempty"` // 是否允许该分组使用图片生成能力 AllowImageGeneration bool `json:"allow_image_generation,omitempty"` + // 是否允许该分组使用批量图片生成能力 + AllowBatchImageGeneration bool `json:"allow_batch_image_generation,omitempty"` // 图片生成是否使用独立倍率;false 表示共享分组有效倍率 ImageRateIndependent bool `json:"image_rate_independent,omitempty"` // 图片生成独立倍率,仅 image_rate_independent=true 时生效 @@ -67,6 +69,10 @@ type Group struct { ImagePrice2k *float64 `json:"image_price_2k,omitempty"` // ImagePrice4k holds the value of the "image_price_4k" field. ImagePrice4k *float64 `json:"image_price_4k,omitempty"` + // 批量图片生成折扣倍率,最终单价会乘以该值;0 表示免费 + BatchImageDiscountMultiplier float64 `json:"batch_image_discount_multiplier,omitempty"` + // 批量图片生成冻结价格比例,按普通生图原价乘以该比例冻结,结算后释放差额 + BatchImageHoldMultiplier float64 `json:"batch_image_hold_multiplier,omitempty"` // 是否仅允许 Claude Code 客户端 ClaudeCodeOnly bool `json:"claude_code_only,omitempty"` // 非 Claude Code 请求降级使用的分组 ID @@ -205,9 +211,9 @@ func (*Group) scanValues(columns []string) ([]any, error) { switch columns[i] { case group.FieldModelRouting, group.FieldSupportedModelScopes, group.FieldMessagesDispatchModelConfig, group.FieldModelsListConfig: values[i] = new([]byte) - case group.FieldPeakRateEnabled, group.FieldIsExclusive, group.FieldAllowImageGeneration, group.FieldImageRateIndependent, group.FieldClaudeCodeOnly, group.FieldModelRoutingEnabled, group.FieldMcpXMLInject, group.FieldAllowMessagesDispatch, group.FieldRequireOauthOnly, group.FieldRequirePrivacySet: + case group.FieldPeakRateEnabled, group.FieldIsExclusive, group.FieldAllowImageGeneration, group.FieldAllowBatchImageGeneration, group.FieldImageRateIndependent, group.FieldClaudeCodeOnly, group.FieldModelRoutingEnabled, group.FieldMcpXMLInject, group.FieldAllowMessagesDispatch, group.FieldRequireOauthOnly, group.FieldRequirePrivacySet: values[i] = new(sql.NullBool) - case group.FieldRateMultiplier, group.FieldPeakRateMultiplier, group.FieldDailyLimitUsd, group.FieldWeeklyLimitUsd, group.FieldMonthlyLimitUsd, group.FieldImageRateMultiplier, group.FieldImagePrice1k, group.FieldImagePrice2k, group.FieldImagePrice4k: + case group.FieldRateMultiplier, group.FieldPeakRateMultiplier, group.FieldDailyLimitUsd, group.FieldWeeklyLimitUsd, group.FieldMonthlyLimitUsd, group.FieldImageRateMultiplier, group.FieldImagePrice1k, group.FieldImagePrice2k, group.FieldImagePrice4k, group.FieldBatchImageDiscountMultiplier, group.FieldBatchImageHoldMultiplier: values[i] = new(sql.NullFloat64) case group.FieldID, group.FieldDefaultValidityDays, group.FieldFallbackGroupID, group.FieldFallbackGroupIDOnInvalidRequest, group.FieldSortOrder, group.FieldRpmLimit: values[i] = new(sql.NullInt64) @@ -355,6 +361,12 @@ func (_m *Group) assignValues(columns []string, values []any) error { } else if value.Valid { _m.AllowImageGeneration = value.Bool } + case group.FieldAllowBatchImageGeneration: + if value, ok := values[i].(*sql.NullBool); !ok { + return fmt.Errorf("unexpected type %T for field allow_batch_image_generation", values[i]) + } else if value.Valid { + _m.AllowBatchImageGeneration = value.Bool + } case group.FieldImageRateIndependent: if value, ok := values[i].(*sql.NullBool); !ok { return fmt.Errorf("unexpected type %T for field image_rate_independent", values[i]) @@ -388,6 +400,18 @@ func (_m *Group) assignValues(columns []string, values []any) error { _m.ImagePrice4k = new(float64) *_m.ImagePrice4k = value.Float64 } + case group.FieldBatchImageDiscountMultiplier: + if value, ok := values[i].(*sql.NullFloat64); !ok { + return fmt.Errorf("unexpected type %T for field batch_image_discount_multiplier", values[i]) + } else if value.Valid { + _m.BatchImageDiscountMultiplier = value.Float64 + } + case group.FieldBatchImageHoldMultiplier: + if value, ok := values[i].(*sql.NullFloat64); !ok { + return fmt.Errorf("unexpected type %T for field batch_image_hold_multiplier", values[i]) + } else if value.Valid { + _m.BatchImageHoldMultiplier = value.Float64 + } case group.FieldClaudeCodeOnly: if value, ok := values[i].(*sql.NullBool); !ok { return fmt.Errorf("unexpected type %T for field claude_code_only", values[i]) @@ -631,6 +655,9 @@ func (_m *Group) String() string { builder.WriteString("allow_image_generation=") builder.WriteString(fmt.Sprintf("%v", _m.AllowImageGeneration)) builder.WriteString(", ") + builder.WriteString("allow_batch_image_generation=") + builder.WriteString(fmt.Sprintf("%v", _m.AllowBatchImageGeneration)) + builder.WriteString(", ") builder.WriteString("image_rate_independent=") builder.WriteString(fmt.Sprintf("%v", _m.ImageRateIndependent)) builder.WriteString(", ") @@ -652,6 +679,12 @@ func (_m *Group) String() string { builder.WriteString(fmt.Sprintf("%v", *v)) } builder.WriteString(", ") + builder.WriteString("batch_image_discount_multiplier=") + builder.WriteString(fmt.Sprintf("%v", _m.BatchImageDiscountMultiplier)) + builder.WriteString(", ") + builder.WriteString("batch_image_hold_multiplier=") + builder.WriteString(fmt.Sprintf("%v", _m.BatchImageHoldMultiplier)) + builder.WriteString(", ") builder.WriteString("claude_code_only=") builder.WriteString(fmt.Sprintf("%v", _m.ClaudeCodeOnly)) builder.WriteString(", ") diff --git a/backend/ent/group/group.go b/backend/ent/group/group.go index bc95af71b6..540ce8f9f5 100644 --- a/backend/ent/group/group.go +++ b/backend/ent/group/group.go @@ -54,6 +54,8 @@ const ( FieldDefaultValidityDays = "default_validity_days" // FieldAllowImageGeneration holds the string denoting the allow_image_generation field in the database. FieldAllowImageGeneration = "allow_image_generation" + // FieldAllowBatchImageGeneration holds the string denoting the allow_batch_image_generation field in the database. + FieldAllowBatchImageGeneration = "allow_batch_image_generation" // FieldImageRateIndependent holds the string denoting the image_rate_independent field in the database. FieldImageRateIndependent = "image_rate_independent" // FieldImageRateMultiplier holds the string denoting the image_rate_multiplier field in the database. @@ -64,6 +66,10 @@ const ( FieldImagePrice2k = "image_price_2k" // FieldImagePrice4k holds the string denoting the image_price_4k field in the database. FieldImagePrice4k = "image_price_4k" + // FieldBatchImageDiscountMultiplier holds the string denoting the batch_image_discount_multiplier field in the database. + FieldBatchImageDiscountMultiplier = "batch_image_discount_multiplier" + // FieldBatchImageHoldMultiplier holds the string denoting the batch_image_hold_multiplier field in the database. + FieldBatchImageHoldMultiplier = "batch_image_hold_multiplier" // FieldClaudeCodeOnly holds the string denoting the claude_code_only field in the database. FieldClaudeCodeOnly = "claude_code_only" // FieldFallbackGroupID holds the string denoting the fallback_group_id field in the database. @@ -188,11 +194,14 @@ var Columns = []string{ FieldMonthlyLimitUsd, FieldDefaultValidityDays, FieldAllowImageGeneration, + FieldAllowBatchImageGeneration, FieldImageRateIndependent, FieldImageRateMultiplier, FieldImagePrice1k, FieldImagePrice2k, FieldImagePrice4k, + FieldBatchImageDiscountMultiplier, + FieldBatchImageHoldMultiplier, FieldClaudeCodeOnly, FieldFallbackGroupID, FieldFallbackGroupIDOnInvalidRequest, @@ -277,10 +286,16 @@ var ( DefaultDefaultValidityDays int // DefaultAllowImageGeneration holds the default value on creation for the "allow_image_generation" field. DefaultAllowImageGeneration bool + // DefaultAllowBatchImageGeneration holds the default value on creation for the "allow_batch_image_generation" field. + DefaultAllowBatchImageGeneration bool // DefaultImageRateIndependent holds the default value on creation for the "image_rate_independent" field. DefaultImageRateIndependent bool // DefaultImageRateMultiplier holds the default value on creation for the "image_rate_multiplier" field. DefaultImageRateMultiplier float64 + // DefaultBatchImageDiscountMultiplier holds the default value on creation for the "batch_image_discount_multiplier" field. + DefaultBatchImageDiscountMultiplier float64 + // DefaultBatchImageHoldMultiplier holds the default value on creation for the "batch_image_hold_multiplier" field. + DefaultBatchImageHoldMultiplier float64 // DefaultClaudeCodeOnly holds the default value on creation for the "claude_code_only" field. DefaultClaudeCodeOnly bool // DefaultModelRoutingEnabled holds the default value on creation for the "model_routing_enabled" field. @@ -412,6 +427,11 @@ func ByAllowImageGeneration(opts ...sql.OrderTermOption) OrderOption { return sql.OrderByField(FieldAllowImageGeneration, opts...).ToFunc() } +// ByAllowBatchImageGeneration orders the results by the allow_batch_image_generation field. +func ByAllowBatchImageGeneration(opts ...sql.OrderTermOption) OrderOption { + return sql.OrderByField(FieldAllowBatchImageGeneration, opts...).ToFunc() +} + // ByImageRateIndependent orders the results by the image_rate_independent field. func ByImageRateIndependent(opts ...sql.OrderTermOption) OrderOption { return sql.OrderByField(FieldImageRateIndependent, opts...).ToFunc() @@ -437,6 +457,16 @@ func ByImagePrice4k(opts ...sql.OrderTermOption) OrderOption { return sql.OrderByField(FieldImagePrice4k, opts...).ToFunc() } +// ByBatchImageDiscountMultiplier orders the results by the batch_image_discount_multiplier field. +func ByBatchImageDiscountMultiplier(opts ...sql.OrderTermOption) OrderOption { + return sql.OrderByField(FieldBatchImageDiscountMultiplier, opts...).ToFunc() +} + +// ByBatchImageHoldMultiplier orders the results by the batch_image_hold_multiplier field. +func ByBatchImageHoldMultiplier(opts ...sql.OrderTermOption) OrderOption { + return sql.OrderByField(FieldBatchImageHoldMultiplier, opts...).ToFunc() +} + // ByClaudeCodeOnly orders the results by the claude_code_only field. func ByClaudeCodeOnly(opts ...sql.OrderTermOption) OrderOption { return sql.OrderByField(FieldClaudeCodeOnly, opts...).ToFunc() diff --git a/backend/ent/group/where.go b/backend/ent/group/where.go index 4a7fc01991..a76d3a8783 100644 --- a/backend/ent/group/where.go +++ b/backend/ent/group/where.go @@ -150,6 +150,11 @@ func AllowImageGeneration(v bool) predicate.Group { return predicate.Group(sql.FieldEQ(FieldAllowImageGeneration, v)) } +// AllowBatchImageGeneration applies equality check predicate on the "allow_batch_image_generation" field. It's identical to AllowBatchImageGenerationEQ. +func AllowBatchImageGeneration(v bool) predicate.Group { + return predicate.Group(sql.FieldEQ(FieldAllowBatchImageGeneration, v)) +} + // ImageRateIndependent applies equality check predicate on the "image_rate_independent" field. It's identical to ImageRateIndependentEQ. func ImageRateIndependent(v bool) predicate.Group { return predicate.Group(sql.FieldEQ(FieldImageRateIndependent, v)) @@ -175,6 +180,16 @@ func ImagePrice4k(v float64) predicate.Group { return predicate.Group(sql.FieldEQ(FieldImagePrice4k, v)) } +// BatchImageDiscountMultiplier applies equality check predicate on the "batch_image_discount_multiplier" field. It's identical to BatchImageDiscountMultiplierEQ. +func BatchImageDiscountMultiplier(v float64) predicate.Group { + return predicate.Group(sql.FieldEQ(FieldBatchImageDiscountMultiplier, v)) +} + +// BatchImageHoldMultiplier applies equality check predicate on the "batch_image_hold_multiplier" field. It's identical to BatchImageHoldMultiplierEQ. +func BatchImageHoldMultiplier(v float64) predicate.Group { + return predicate.Group(sql.FieldEQ(FieldBatchImageHoldMultiplier, v)) +} + // ClaudeCodeOnly applies equality check predicate on the "claude_code_only" field. It's identical to ClaudeCodeOnlyEQ. func ClaudeCodeOnly(v bool) predicate.Group { return predicate.Group(sql.FieldEQ(FieldClaudeCodeOnly, v)) @@ -1125,6 +1140,16 @@ func AllowImageGenerationNEQ(v bool) predicate.Group { return predicate.Group(sql.FieldNEQ(FieldAllowImageGeneration, v)) } +// AllowBatchImageGenerationEQ applies the EQ predicate on the "allow_batch_image_generation" field. +func AllowBatchImageGenerationEQ(v bool) predicate.Group { + return predicate.Group(sql.FieldEQ(FieldAllowBatchImageGeneration, v)) +} + +// AllowBatchImageGenerationNEQ applies the NEQ predicate on the "allow_batch_image_generation" field. +func AllowBatchImageGenerationNEQ(v bool) predicate.Group { + return predicate.Group(sql.FieldNEQ(FieldAllowBatchImageGeneration, v)) +} + // ImageRateIndependentEQ applies the EQ predicate on the "image_rate_independent" field. func ImageRateIndependentEQ(v bool) predicate.Group { return predicate.Group(sql.FieldEQ(FieldImageRateIndependent, v)) @@ -1325,6 +1350,86 @@ func ImagePrice4kNotNil() predicate.Group { return predicate.Group(sql.FieldNotNull(FieldImagePrice4k)) } +// BatchImageDiscountMultiplierEQ applies the EQ predicate on the "batch_image_discount_multiplier" field. +func BatchImageDiscountMultiplierEQ(v float64) predicate.Group { + return predicate.Group(sql.FieldEQ(FieldBatchImageDiscountMultiplier, v)) +} + +// BatchImageDiscountMultiplierNEQ applies the NEQ predicate on the "batch_image_discount_multiplier" field. +func BatchImageDiscountMultiplierNEQ(v float64) predicate.Group { + return predicate.Group(sql.FieldNEQ(FieldBatchImageDiscountMultiplier, v)) +} + +// BatchImageDiscountMultiplierIn applies the In predicate on the "batch_image_discount_multiplier" field. +func BatchImageDiscountMultiplierIn(vs ...float64) predicate.Group { + return predicate.Group(sql.FieldIn(FieldBatchImageDiscountMultiplier, vs...)) +} + +// BatchImageDiscountMultiplierNotIn applies the NotIn predicate on the "batch_image_discount_multiplier" field. +func BatchImageDiscountMultiplierNotIn(vs ...float64) predicate.Group { + return predicate.Group(sql.FieldNotIn(FieldBatchImageDiscountMultiplier, vs...)) +} + +// BatchImageDiscountMultiplierGT applies the GT predicate on the "batch_image_discount_multiplier" field. +func BatchImageDiscountMultiplierGT(v float64) predicate.Group { + return predicate.Group(sql.FieldGT(FieldBatchImageDiscountMultiplier, v)) +} + +// BatchImageDiscountMultiplierGTE applies the GTE predicate on the "batch_image_discount_multiplier" field. +func BatchImageDiscountMultiplierGTE(v float64) predicate.Group { + return predicate.Group(sql.FieldGTE(FieldBatchImageDiscountMultiplier, v)) +} + +// BatchImageDiscountMultiplierLT applies the LT predicate on the "batch_image_discount_multiplier" field. +func BatchImageDiscountMultiplierLT(v float64) predicate.Group { + return predicate.Group(sql.FieldLT(FieldBatchImageDiscountMultiplier, v)) +} + +// BatchImageDiscountMultiplierLTE applies the LTE predicate on the "batch_image_discount_multiplier" field. +func BatchImageDiscountMultiplierLTE(v float64) predicate.Group { + return predicate.Group(sql.FieldLTE(FieldBatchImageDiscountMultiplier, v)) +} + +// BatchImageHoldMultiplierEQ applies the EQ predicate on the "batch_image_hold_multiplier" field. +func BatchImageHoldMultiplierEQ(v float64) predicate.Group { + return predicate.Group(sql.FieldEQ(FieldBatchImageHoldMultiplier, v)) +} + +// BatchImageHoldMultiplierNEQ applies the NEQ predicate on the "batch_image_hold_multiplier" field. +func BatchImageHoldMultiplierNEQ(v float64) predicate.Group { + return predicate.Group(sql.FieldNEQ(FieldBatchImageHoldMultiplier, v)) +} + +// BatchImageHoldMultiplierIn applies the In predicate on the "batch_image_hold_multiplier" field. +func BatchImageHoldMultiplierIn(vs ...float64) predicate.Group { + return predicate.Group(sql.FieldIn(FieldBatchImageHoldMultiplier, vs...)) +} + +// BatchImageHoldMultiplierNotIn applies the NotIn predicate on the "batch_image_hold_multiplier" field. +func BatchImageHoldMultiplierNotIn(vs ...float64) predicate.Group { + return predicate.Group(sql.FieldNotIn(FieldBatchImageHoldMultiplier, vs...)) +} + +// BatchImageHoldMultiplierGT applies the GT predicate on the "batch_image_hold_multiplier" field. +func BatchImageHoldMultiplierGT(v float64) predicate.Group { + return predicate.Group(sql.FieldGT(FieldBatchImageHoldMultiplier, v)) +} + +// BatchImageHoldMultiplierGTE applies the GTE predicate on the "batch_image_hold_multiplier" field. +func BatchImageHoldMultiplierGTE(v float64) predicate.Group { + return predicate.Group(sql.FieldGTE(FieldBatchImageHoldMultiplier, v)) +} + +// BatchImageHoldMultiplierLT applies the LT predicate on the "batch_image_hold_multiplier" field. +func BatchImageHoldMultiplierLT(v float64) predicate.Group { + return predicate.Group(sql.FieldLT(FieldBatchImageHoldMultiplier, v)) +} + +// BatchImageHoldMultiplierLTE applies the LTE predicate on the "batch_image_hold_multiplier" field. +func BatchImageHoldMultiplierLTE(v float64) predicate.Group { + return predicate.Group(sql.FieldLTE(FieldBatchImageHoldMultiplier, v)) +} + // ClaudeCodeOnlyEQ applies the EQ predicate on the "claude_code_only" field. func ClaudeCodeOnlyEQ(v bool) predicate.Group { return predicate.Group(sql.FieldEQ(FieldClaudeCodeOnly, v)) diff --git a/backend/ent/group_create.go b/backend/ent/group_create.go index 0f35f070ee..9c635847d0 100644 --- a/backend/ent/group_create.go +++ b/backend/ent/group_create.go @@ -287,6 +287,20 @@ func (_c *GroupCreate) SetNillableAllowImageGeneration(v *bool) *GroupCreate { return _c } +// SetAllowBatchImageGeneration sets the "allow_batch_image_generation" field. +func (_c *GroupCreate) SetAllowBatchImageGeneration(v bool) *GroupCreate { + _c.mutation.SetAllowBatchImageGeneration(v) + return _c +} + +// SetNillableAllowBatchImageGeneration sets the "allow_batch_image_generation" field if the given value is not nil. +func (_c *GroupCreate) SetNillableAllowBatchImageGeneration(v *bool) *GroupCreate { + if v != nil { + _c.SetAllowBatchImageGeneration(*v) + } + return _c +} + // SetImageRateIndependent sets the "image_rate_independent" field. func (_c *GroupCreate) SetImageRateIndependent(v bool) *GroupCreate { _c.mutation.SetImageRateIndependent(v) @@ -357,6 +371,34 @@ func (_c *GroupCreate) SetNillableImagePrice4k(v *float64) *GroupCreate { return _c } +// SetBatchImageDiscountMultiplier sets the "batch_image_discount_multiplier" field. +func (_c *GroupCreate) SetBatchImageDiscountMultiplier(v float64) *GroupCreate { + _c.mutation.SetBatchImageDiscountMultiplier(v) + return _c +} + +// SetNillableBatchImageDiscountMultiplier sets the "batch_image_discount_multiplier" field if the given value is not nil. +func (_c *GroupCreate) SetNillableBatchImageDiscountMultiplier(v *float64) *GroupCreate { + if v != nil { + _c.SetBatchImageDiscountMultiplier(*v) + } + return _c +} + +// SetBatchImageHoldMultiplier sets the "batch_image_hold_multiplier" field. +func (_c *GroupCreate) SetBatchImageHoldMultiplier(v float64) *GroupCreate { + _c.mutation.SetBatchImageHoldMultiplier(v) + return _c +} + +// SetNillableBatchImageHoldMultiplier sets the "batch_image_hold_multiplier" field if the given value is not nil. +func (_c *GroupCreate) SetNillableBatchImageHoldMultiplier(v *float64) *GroupCreate { + if v != nil { + _c.SetBatchImageHoldMultiplier(*v) + } + return _c +} + // SetClaudeCodeOnly sets the "claude_code_only" field. func (_c *GroupCreate) SetClaudeCodeOnly(v bool) *GroupCreate { _c.mutation.SetClaudeCodeOnly(v) @@ -736,6 +778,10 @@ func (_c *GroupCreate) defaults() error { v := group.DefaultAllowImageGeneration _c.mutation.SetAllowImageGeneration(v) } + if _, ok := _c.mutation.AllowBatchImageGeneration(); !ok { + v := group.DefaultAllowBatchImageGeneration + _c.mutation.SetAllowBatchImageGeneration(v) + } if _, ok := _c.mutation.ImageRateIndependent(); !ok { v := group.DefaultImageRateIndependent _c.mutation.SetImageRateIndependent(v) @@ -744,6 +790,14 @@ func (_c *GroupCreate) defaults() error { v := group.DefaultImageRateMultiplier _c.mutation.SetImageRateMultiplier(v) } + if _, ok := _c.mutation.BatchImageDiscountMultiplier(); !ok { + v := group.DefaultBatchImageDiscountMultiplier + _c.mutation.SetBatchImageDiscountMultiplier(v) + } + if _, ok := _c.mutation.BatchImageHoldMultiplier(); !ok { + v := group.DefaultBatchImageHoldMultiplier + _c.mutation.SetBatchImageHoldMultiplier(v) + } if _, ok := _c.mutation.ClaudeCodeOnly(); !ok { v := group.DefaultClaudeCodeOnly _c.mutation.SetClaudeCodeOnly(v) @@ -869,12 +923,21 @@ func (_c *GroupCreate) check() error { if _, ok := _c.mutation.AllowImageGeneration(); !ok { return &ValidationError{Name: "allow_image_generation", err: errors.New(`ent: missing required field "Group.allow_image_generation"`)} } + if _, ok := _c.mutation.AllowBatchImageGeneration(); !ok { + return &ValidationError{Name: "allow_batch_image_generation", err: errors.New(`ent: missing required field "Group.allow_batch_image_generation"`)} + } if _, ok := _c.mutation.ImageRateIndependent(); !ok { return &ValidationError{Name: "image_rate_independent", err: errors.New(`ent: missing required field "Group.image_rate_independent"`)} } if _, ok := _c.mutation.ImageRateMultiplier(); !ok { return &ValidationError{Name: "image_rate_multiplier", err: errors.New(`ent: missing required field "Group.image_rate_multiplier"`)} } + if _, ok := _c.mutation.BatchImageDiscountMultiplier(); !ok { + return &ValidationError{Name: "batch_image_discount_multiplier", err: errors.New(`ent: missing required field "Group.batch_image_discount_multiplier"`)} + } + if _, ok := _c.mutation.BatchImageHoldMultiplier(); !ok { + return &ValidationError{Name: "batch_image_hold_multiplier", err: errors.New(`ent: missing required field "Group.batch_image_hold_multiplier"`)} + } if _, ok := _c.mutation.ClaudeCodeOnly(); !ok { return &ValidationError{Name: "claude_code_only", err: errors.New(`ent: missing required field "Group.claude_code_only"`)} } @@ -1019,6 +1082,10 @@ func (_c *GroupCreate) createSpec() (*Group, *sqlgraph.CreateSpec) { _spec.SetField(group.FieldAllowImageGeneration, field.TypeBool, value) _node.AllowImageGeneration = value } + if value, ok := _c.mutation.AllowBatchImageGeneration(); ok { + _spec.SetField(group.FieldAllowBatchImageGeneration, field.TypeBool, value) + _node.AllowBatchImageGeneration = value + } if value, ok := _c.mutation.ImageRateIndependent(); ok { _spec.SetField(group.FieldImageRateIndependent, field.TypeBool, value) _node.ImageRateIndependent = value @@ -1039,6 +1106,14 @@ func (_c *GroupCreate) createSpec() (*Group, *sqlgraph.CreateSpec) { _spec.SetField(group.FieldImagePrice4k, field.TypeFloat64, value) _node.ImagePrice4k = &value } + if value, ok := _c.mutation.BatchImageDiscountMultiplier(); ok { + _spec.SetField(group.FieldBatchImageDiscountMultiplier, field.TypeFloat64, value) + _node.BatchImageDiscountMultiplier = value + } + if value, ok := _c.mutation.BatchImageHoldMultiplier(); ok { + _spec.SetField(group.FieldBatchImageHoldMultiplier, field.TypeFloat64, value) + _node.BatchImageHoldMultiplier = value + } if value, ok := _c.mutation.ClaudeCodeOnly(); ok { _spec.SetField(group.FieldClaudeCodeOnly, field.TypeBool, value) _node.ClaudeCodeOnly = value @@ -1537,6 +1612,18 @@ func (u *GroupUpsert) UpdateAllowImageGeneration() *GroupUpsert { return u } +// SetAllowBatchImageGeneration sets the "allow_batch_image_generation" field. +func (u *GroupUpsert) SetAllowBatchImageGeneration(v bool) *GroupUpsert { + u.Set(group.FieldAllowBatchImageGeneration, v) + return u +} + +// UpdateAllowBatchImageGeneration sets the "allow_batch_image_generation" field to the value that was provided on create. +func (u *GroupUpsert) UpdateAllowBatchImageGeneration() *GroupUpsert { + u.SetExcluded(group.FieldAllowBatchImageGeneration) + return u +} + // SetImageRateIndependent sets the "image_rate_independent" field. func (u *GroupUpsert) SetImageRateIndependent(v bool) *GroupUpsert { u.Set(group.FieldImageRateIndependent, v) @@ -1639,6 +1726,42 @@ func (u *GroupUpsert) ClearImagePrice4k() *GroupUpsert { return u } +// SetBatchImageDiscountMultiplier sets the "batch_image_discount_multiplier" field. +func (u *GroupUpsert) SetBatchImageDiscountMultiplier(v float64) *GroupUpsert { + u.Set(group.FieldBatchImageDiscountMultiplier, v) + return u +} + +// UpdateBatchImageDiscountMultiplier sets the "batch_image_discount_multiplier" field to the value that was provided on create. +func (u *GroupUpsert) UpdateBatchImageDiscountMultiplier() *GroupUpsert { + u.SetExcluded(group.FieldBatchImageDiscountMultiplier) + return u +} + +// AddBatchImageDiscountMultiplier adds v to the "batch_image_discount_multiplier" field. +func (u *GroupUpsert) AddBatchImageDiscountMultiplier(v float64) *GroupUpsert { + u.Add(group.FieldBatchImageDiscountMultiplier, v) + return u +} + +// SetBatchImageHoldMultiplier sets the "batch_image_hold_multiplier" field. +func (u *GroupUpsert) SetBatchImageHoldMultiplier(v float64) *GroupUpsert { + u.Set(group.FieldBatchImageHoldMultiplier, v) + return u +} + +// UpdateBatchImageHoldMultiplier sets the "batch_image_hold_multiplier" field to the value that was provided on create. +func (u *GroupUpsert) UpdateBatchImageHoldMultiplier() *GroupUpsert { + u.SetExcluded(group.FieldBatchImageHoldMultiplier) + return u +} + +// AddBatchImageHoldMultiplier adds v to the "batch_image_hold_multiplier" field. +func (u *GroupUpsert) AddBatchImageHoldMultiplier(v float64) *GroupUpsert { + u.Add(group.FieldBatchImageHoldMultiplier, v) + return u +} + // SetClaudeCodeOnly sets the "claude_code_only" field. func (u *GroupUpsert) SetClaudeCodeOnly(v bool) *GroupUpsert { u.Set(group.FieldClaudeCodeOnly, v) @@ -2235,6 +2358,20 @@ func (u *GroupUpsertOne) UpdateAllowImageGeneration() *GroupUpsertOne { }) } +// SetAllowBatchImageGeneration sets the "allow_batch_image_generation" field. +func (u *GroupUpsertOne) SetAllowBatchImageGeneration(v bool) *GroupUpsertOne { + return u.Update(func(s *GroupUpsert) { + s.SetAllowBatchImageGeneration(v) + }) +} + +// UpdateAllowBatchImageGeneration sets the "allow_batch_image_generation" field to the value that was provided on create. +func (u *GroupUpsertOne) UpdateAllowBatchImageGeneration() *GroupUpsertOne { + return u.Update(func(s *GroupUpsert) { + s.UpdateAllowBatchImageGeneration() + }) +} + // SetImageRateIndependent sets the "image_rate_independent" field. func (u *GroupUpsertOne) SetImageRateIndependent(v bool) *GroupUpsertOne { return u.Update(func(s *GroupUpsert) { @@ -2354,6 +2491,48 @@ func (u *GroupUpsertOne) ClearImagePrice4k() *GroupUpsertOne { }) } +// SetBatchImageDiscountMultiplier sets the "batch_image_discount_multiplier" field. +func (u *GroupUpsertOne) SetBatchImageDiscountMultiplier(v float64) *GroupUpsertOne { + return u.Update(func(s *GroupUpsert) { + s.SetBatchImageDiscountMultiplier(v) + }) +} + +// AddBatchImageDiscountMultiplier adds v to the "batch_image_discount_multiplier" field. +func (u *GroupUpsertOne) AddBatchImageDiscountMultiplier(v float64) *GroupUpsertOne { + return u.Update(func(s *GroupUpsert) { + s.AddBatchImageDiscountMultiplier(v) + }) +} + +// UpdateBatchImageDiscountMultiplier sets the "batch_image_discount_multiplier" field to the value that was provided on create. +func (u *GroupUpsertOne) UpdateBatchImageDiscountMultiplier() *GroupUpsertOne { + return u.Update(func(s *GroupUpsert) { + s.UpdateBatchImageDiscountMultiplier() + }) +} + +// SetBatchImageHoldMultiplier sets the "batch_image_hold_multiplier" field. +func (u *GroupUpsertOne) SetBatchImageHoldMultiplier(v float64) *GroupUpsertOne { + return u.Update(func(s *GroupUpsert) { + s.SetBatchImageHoldMultiplier(v) + }) +} + +// AddBatchImageHoldMultiplier adds v to the "batch_image_hold_multiplier" field. +func (u *GroupUpsertOne) AddBatchImageHoldMultiplier(v float64) *GroupUpsertOne { + return u.Update(func(s *GroupUpsert) { + s.AddBatchImageHoldMultiplier(v) + }) +} + +// UpdateBatchImageHoldMultiplier sets the "batch_image_hold_multiplier" field to the value that was provided on create. +func (u *GroupUpsertOne) UpdateBatchImageHoldMultiplier() *GroupUpsertOne { + return u.Update(func(s *GroupUpsert) { + s.UpdateBatchImageHoldMultiplier() + }) +} + // SetClaudeCodeOnly sets the "claude_code_only" field. func (u *GroupUpsertOne) SetClaudeCodeOnly(v bool) *GroupUpsertOne { return u.Update(func(s *GroupUpsert) { @@ -3153,6 +3332,20 @@ func (u *GroupUpsertBulk) UpdateAllowImageGeneration() *GroupUpsertBulk { }) } +// SetAllowBatchImageGeneration sets the "allow_batch_image_generation" field. +func (u *GroupUpsertBulk) SetAllowBatchImageGeneration(v bool) *GroupUpsertBulk { + return u.Update(func(s *GroupUpsert) { + s.SetAllowBatchImageGeneration(v) + }) +} + +// UpdateAllowBatchImageGeneration sets the "allow_batch_image_generation" field to the value that was provided on create. +func (u *GroupUpsertBulk) UpdateAllowBatchImageGeneration() *GroupUpsertBulk { + return u.Update(func(s *GroupUpsert) { + s.UpdateAllowBatchImageGeneration() + }) +} + // SetImageRateIndependent sets the "image_rate_independent" field. func (u *GroupUpsertBulk) SetImageRateIndependent(v bool) *GroupUpsertBulk { return u.Update(func(s *GroupUpsert) { @@ -3272,6 +3465,48 @@ func (u *GroupUpsertBulk) ClearImagePrice4k() *GroupUpsertBulk { }) } +// SetBatchImageDiscountMultiplier sets the "batch_image_discount_multiplier" field. +func (u *GroupUpsertBulk) SetBatchImageDiscountMultiplier(v float64) *GroupUpsertBulk { + return u.Update(func(s *GroupUpsert) { + s.SetBatchImageDiscountMultiplier(v) + }) +} + +// AddBatchImageDiscountMultiplier adds v to the "batch_image_discount_multiplier" field. +func (u *GroupUpsertBulk) AddBatchImageDiscountMultiplier(v float64) *GroupUpsertBulk { + return u.Update(func(s *GroupUpsert) { + s.AddBatchImageDiscountMultiplier(v) + }) +} + +// UpdateBatchImageDiscountMultiplier sets the "batch_image_discount_multiplier" field to the value that was provided on create. +func (u *GroupUpsertBulk) UpdateBatchImageDiscountMultiplier() *GroupUpsertBulk { + return u.Update(func(s *GroupUpsert) { + s.UpdateBatchImageDiscountMultiplier() + }) +} + +// SetBatchImageHoldMultiplier sets the "batch_image_hold_multiplier" field. +func (u *GroupUpsertBulk) SetBatchImageHoldMultiplier(v float64) *GroupUpsertBulk { + return u.Update(func(s *GroupUpsert) { + s.SetBatchImageHoldMultiplier(v) + }) +} + +// AddBatchImageHoldMultiplier adds v to the "batch_image_hold_multiplier" field. +func (u *GroupUpsertBulk) AddBatchImageHoldMultiplier(v float64) *GroupUpsertBulk { + return u.Update(func(s *GroupUpsert) { + s.AddBatchImageHoldMultiplier(v) + }) +} + +// UpdateBatchImageHoldMultiplier sets the "batch_image_hold_multiplier" field to the value that was provided on create. +func (u *GroupUpsertBulk) UpdateBatchImageHoldMultiplier() *GroupUpsertBulk { + return u.Update(func(s *GroupUpsert) { + s.UpdateBatchImageHoldMultiplier() + }) +} + // SetClaudeCodeOnly sets the "claude_code_only" field. func (u *GroupUpsertBulk) SetClaudeCodeOnly(v bool) *GroupUpsertBulk { return u.Update(func(s *GroupUpsert) { diff --git a/backend/ent/group_update.go b/backend/ent/group_update.go index 55555f323c..6f1831b1ea 100644 --- a/backend/ent/group_update.go +++ b/backend/ent/group_update.go @@ -352,6 +352,20 @@ func (_u *GroupUpdate) SetNillableAllowImageGeneration(v *bool) *GroupUpdate { return _u } +// SetAllowBatchImageGeneration sets the "allow_batch_image_generation" field. +func (_u *GroupUpdate) SetAllowBatchImageGeneration(v bool) *GroupUpdate { + _u.mutation.SetAllowBatchImageGeneration(v) + return _u +} + +// SetNillableAllowBatchImageGeneration sets the "allow_batch_image_generation" field if the given value is not nil. +func (_u *GroupUpdate) SetNillableAllowBatchImageGeneration(v *bool) *GroupUpdate { + if v != nil { + _u.SetAllowBatchImageGeneration(*v) + } + return _u +} + // SetImageRateIndependent sets the "image_rate_independent" field. func (_u *GroupUpdate) SetImageRateIndependent(v bool) *GroupUpdate { _u.mutation.SetImageRateIndependent(v) @@ -468,6 +482,48 @@ func (_u *GroupUpdate) ClearImagePrice4k() *GroupUpdate { return _u } +// SetBatchImageDiscountMultiplier sets the "batch_image_discount_multiplier" field. +func (_u *GroupUpdate) SetBatchImageDiscountMultiplier(v float64) *GroupUpdate { + _u.mutation.ResetBatchImageDiscountMultiplier() + _u.mutation.SetBatchImageDiscountMultiplier(v) + return _u +} + +// SetNillableBatchImageDiscountMultiplier sets the "batch_image_discount_multiplier" field if the given value is not nil. +func (_u *GroupUpdate) SetNillableBatchImageDiscountMultiplier(v *float64) *GroupUpdate { + if v != nil { + _u.SetBatchImageDiscountMultiplier(*v) + } + return _u +} + +// AddBatchImageDiscountMultiplier adds value to the "batch_image_discount_multiplier" field. +func (_u *GroupUpdate) AddBatchImageDiscountMultiplier(v float64) *GroupUpdate { + _u.mutation.AddBatchImageDiscountMultiplier(v) + return _u +} + +// SetBatchImageHoldMultiplier sets the "batch_image_hold_multiplier" field. +func (_u *GroupUpdate) SetBatchImageHoldMultiplier(v float64) *GroupUpdate { + _u.mutation.ResetBatchImageHoldMultiplier() + _u.mutation.SetBatchImageHoldMultiplier(v) + return _u +} + +// SetNillableBatchImageHoldMultiplier sets the "batch_image_hold_multiplier" field if the given value is not nil. +func (_u *GroupUpdate) SetNillableBatchImageHoldMultiplier(v *float64) *GroupUpdate { + if v != nil { + _u.SetBatchImageHoldMultiplier(*v) + } + return _u +} + +// AddBatchImageHoldMultiplier adds value to the "batch_image_hold_multiplier" field. +func (_u *GroupUpdate) AddBatchImageHoldMultiplier(v float64) *GroupUpdate { + _u.mutation.AddBatchImageHoldMultiplier(v) + return _u +} + // SetClaudeCodeOnly sets the "claude_code_only" field. func (_u *GroupUpdate) SetClaudeCodeOnly(v bool) *GroupUpdate { _u.mutation.SetClaudeCodeOnly(v) @@ -1116,6 +1172,9 @@ func (_u *GroupUpdate) sqlSave(ctx context.Context) (_node int, err error) { if value, ok := _u.mutation.AllowImageGeneration(); ok { _spec.SetField(group.FieldAllowImageGeneration, field.TypeBool, value) } + if value, ok := _u.mutation.AllowBatchImageGeneration(); ok { + _spec.SetField(group.FieldAllowBatchImageGeneration, field.TypeBool, value) + } if value, ok := _u.mutation.ImageRateIndependent(); ok { _spec.SetField(group.FieldImageRateIndependent, field.TypeBool, value) } @@ -1152,6 +1211,18 @@ func (_u *GroupUpdate) sqlSave(ctx context.Context) (_node int, err error) { if _u.mutation.ImagePrice4kCleared() { _spec.ClearField(group.FieldImagePrice4k, field.TypeFloat64) } + if value, ok := _u.mutation.BatchImageDiscountMultiplier(); ok { + _spec.SetField(group.FieldBatchImageDiscountMultiplier, field.TypeFloat64, value) + } + if value, ok := _u.mutation.AddedBatchImageDiscountMultiplier(); ok { + _spec.AddField(group.FieldBatchImageDiscountMultiplier, field.TypeFloat64, value) + } + if value, ok := _u.mutation.BatchImageHoldMultiplier(); ok { + _spec.SetField(group.FieldBatchImageHoldMultiplier, field.TypeFloat64, value) + } + if value, ok := _u.mutation.AddedBatchImageHoldMultiplier(); ok { + _spec.AddField(group.FieldBatchImageHoldMultiplier, field.TypeFloat64, value) + } if value, ok := _u.mutation.ClaudeCodeOnly(); ok { _spec.SetField(group.FieldClaudeCodeOnly, field.TypeBool, value) } @@ -1853,6 +1924,20 @@ func (_u *GroupUpdateOne) SetNillableAllowImageGeneration(v *bool) *GroupUpdateO return _u } +// SetAllowBatchImageGeneration sets the "allow_batch_image_generation" field. +func (_u *GroupUpdateOne) SetAllowBatchImageGeneration(v bool) *GroupUpdateOne { + _u.mutation.SetAllowBatchImageGeneration(v) + return _u +} + +// SetNillableAllowBatchImageGeneration sets the "allow_batch_image_generation" field if the given value is not nil. +func (_u *GroupUpdateOne) SetNillableAllowBatchImageGeneration(v *bool) *GroupUpdateOne { + if v != nil { + _u.SetAllowBatchImageGeneration(*v) + } + return _u +} + // SetImageRateIndependent sets the "image_rate_independent" field. func (_u *GroupUpdateOne) SetImageRateIndependent(v bool) *GroupUpdateOne { _u.mutation.SetImageRateIndependent(v) @@ -1969,6 +2054,48 @@ func (_u *GroupUpdateOne) ClearImagePrice4k() *GroupUpdateOne { return _u } +// SetBatchImageDiscountMultiplier sets the "batch_image_discount_multiplier" field. +func (_u *GroupUpdateOne) SetBatchImageDiscountMultiplier(v float64) *GroupUpdateOne { + _u.mutation.ResetBatchImageDiscountMultiplier() + _u.mutation.SetBatchImageDiscountMultiplier(v) + return _u +} + +// SetNillableBatchImageDiscountMultiplier sets the "batch_image_discount_multiplier" field if the given value is not nil. +func (_u *GroupUpdateOne) SetNillableBatchImageDiscountMultiplier(v *float64) *GroupUpdateOne { + if v != nil { + _u.SetBatchImageDiscountMultiplier(*v) + } + return _u +} + +// AddBatchImageDiscountMultiplier adds value to the "batch_image_discount_multiplier" field. +func (_u *GroupUpdateOne) AddBatchImageDiscountMultiplier(v float64) *GroupUpdateOne { + _u.mutation.AddBatchImageDiscountMultiplier(v) + return _u +} + +// SetBatchImageHoldMultiplier sets the "batch_image_hold_multiplier" field. +func (_u *GroupUpdateOne) SetBatchImageHoldMultiplier(v float64) *GroupUpdateOne { + _u.mutation.ResetBatchImageHoldMultiplier() + _u.mutation.SetBatchImageHoldMultiplier(v) + return _u +} + +// SetNillableBatchImageHoldMultiplier sets the "batch_image_hold_multiplier" field if the given value is not nil. +func (_u *GroupUpdateOne) SetNillableBatchImageHoldMultiplier(v *float64) *GroupUpdateOne { + if v != nil { + _u.SetBatchImageHoldMultiplier(*v) + } + return _u +} + +// AddBatchImageHoldMultiplier adds value to the "batch_image_hold_multiplier" field. +func (_u *GroupUpdateOne) AddBatchImageHoldMultiplier(v float64) *GroupUpdateOne { + _u.mutation.AddBatchImageHoldMultiplier(v) + return _u +} + // SetClaudeCodeOnly sets the "claude_code_only" field. func (_u *GroupUpdateOne) SetClaudeCodeOnly(v bool) *GroupUpdateOne { _u.mutation.SetClaudeCodeOnly(v) @@ -2647,6 +2774,9 @@ func (_u *GroupUpdateOne) sqlSave(ctx context.Context) (_node *Group, err error) if value, ok := _u.mutation.AllowImageGeneration(); ok { _spec.SetField(group.FieldAllowImageGeneration, field.TypeBool, value) } + if value, ok := _u.mutation.AllowBatchImageGeneration(); ok { + _spec.SetField(group.FieldAllowBatchImageGeneration, field.TypeBool, value) + } if value, ok := _u.mutation.ImageRateIndependent(); ok { _spec.SetField(group.FieldImageRateIndependent, field.TypeBool, value) } @@ -2683,6 +2813,18 @@ func (_u *GroupUpdateOne) sqlSave(ctx context.Context) (_node *Group, err error) if _u.mutation.ImagePrice4kCleared() { _spec.ClearField(group.FieldImagePrice4k, field.TypeFloat64) } + if value, ok := _u.mutation.BatchImageDiscountMultiplier(); ok { + _spec.SetField(group.FieldBatchImageDiscountMultiplier, field.TypeFloat64, value) + } + if value, ok := _u.mutation.AddedBatchImageDiscountMultiplier(); ok { + _spec.AddField(group.FieldBatchImageDiscountMultiplier, field.TypeFloat64, value) + } + if value, ok := _u.mutation.BatchImageHoldMultiplier(); ok { + _spec.SetField(group.FieldBatchImageHoldMultiplier, field.TypeFloat64, value) + } + if value, ok := _u.mutation.AddedBatchImageHoldMultiplier(); ok { + _spec.AddField(group.FieldBatchImageHoldMultiplier, field.TypeFloat64, value) + } if value, ok := _u.mutation.ClaudeCodeOnly(); ok { _spec.SetField(group.FieldClaudeCodeOnly, field.TypeBool, value) } diff --git a/backend/ent/migrate/schema.go b/backend/ent/migrate/schema.go index 15228dcced..a584cbe39d 100644 --- a/backend/ent/migrate/schema.go +++ b/backend/ent/migrate/schema.go @@ -523,6 +523,7 @@ var ( {Name: "account_id", Type: field.TypeInt64, Nullable: true}, {Name: "provider", Type: field.TypeString, Size: 32}, {Name: "model", Type: field.TypeString, Size: 128}, + {Name: "task_name", Type: field.TypeString, Size: 255, Default: ""}, {Name: "status", Type: field.TypeString, Size: 32, Default: "created"}, {Name: "provider_job_name", Type: field.TypeString, Nullable: true, Size: 512}, {Name: "provider_input_ref", Type: field.TypeString, Nullable: true, Size: 1024}, @@ -546,6 +547,8 @@ var ( {Name: "output_expires_at", Type: field.TypeTime, Nullable: true, SchemaType: map[string]string{"postgres": "timestamptz"}}, {Name: "input_deleted_at", Type: field.TypeTime, Nullable: true, SchemaType: map[string]string{"postgres": "timestamptz"}}, {Name: "output_deleted_at", Type: field.TypeTime, Nullable: true, SchemaType: map[string]string{"postgres": "timestamptz"}}, + {Name: "downloaded_at", Type: field.TypeTime, Nullable: true, SchemaType: map[string]string{"postgres": "timestamptz"}}, + {Name: "user_deleted_at", Type: field.TypeTime, Nullable: true, SchemaType: map[string]string{"postgres": "timestamptz"}}, {Name: "last_error_code", Type: field.TypeString, Nullable: true, Size: 128}, {Name: "last_error_message", Type: field.TypeString, Nullable: true, SchemaType: map[string]string{"postgres": "text"}}, {Name: "created_at", Type: field.TypeTime, SchemaType: map[string]string{"postgres": "timestamptz"}}, @@ -569,22 +572,22 @@ var ( { Name: "batchimagejob_user_id_created_at", Unique: false, - Columns: []*schema.Column{BatchImageJobsColumns[2], BatchImageJobsColumns[32]}, + Columns: []*schema.Column{BatchImageJobsColumns[2], BatchImageJobsColumns[35]}, }, { Name: "batchimagejob_status", Unique: false, - Columns: []*schema.Column{BatchImageJobsColumns[7]}, + Columns: []*schema.Column{BatchImageJobsColumns[8]}, }, { Name: "batchimagejob_provider_status", Unique: false, - Columns: []*schema.Column{BatchImageJobsColumns[5], BatchImageJobsColumns[7]}, + Columns: []*schema.Column{BatchImageJobsColumns[5], BatchImageJobsColumns[8]}, }, { Name: "batchimagejob_idempotency_key", Unique: false, - Columns: []*schema.Column{BatchImageJobsColumns[22]}, + Columns: []*schema.Column{BatchImageJobsColumns[23]}, Annotation: &entsql.IndexAnnotation{ Where: "idempotency_key IS NOT NULL AND idempotency_key <> ''", }, @@ -592,7 +595,7 @@ var ( { Name: "batchimagejob_manifest_hash", Unique: true, - Columns: []*schema.Column{BatchImageJobsColumns[24]}, + Columns: []*schema.Column{BatchImageJobsColumns[25]}, Annotation: &entsql.IndexAnnotation{ Where: "manifest_hash IS NOT NULL AND manifest_hash <> ''", }, @@ -600,7 +603,17 @@ var ( { Name: "batchimagejob_output_expires_at", Unique: false, - Columns: []*schema.Column{BatchImageJobsColumns[27]}, + Columns: []*schema.Column{BatchImageJobsColumns[28]}, + }, + { + Name: "batchimagejob_downloaded_at", + Unique: false, + Columns: []*schema.Column{BatchImageJobsColumns[31]}, + }, + { + Name: "batchimagejob_user_deleted_at", + Unique: false, + Columns: []*schema.Column{BatchImageJobsColumns[32]}, }, }, } @@ -839,11 +852,14 @@ var ( {Name: "monthly_limit_usd", Type: field.TypeFloat64, Nullable: true, SchemaType: map[string]string{"postgres": "decimal(20,8)"}}, {Name: "default_validity_days", Type: field.TypeInt, Default: 30}, {Name: "allow_image_generation", Type: field.TypeBool, Default: false}, + {Name: "allow_batch_image_generation", Type: field.TypeBool, Default: false}, {Name: "image_rate_independent", Type: field.TypeBool, Default: false}, {Name: "image_rate_multiplier", Type: field.TypeFloat64, Default: 1, SchemaType: map[string]string{"postgres": "decimal(10,4)"}}, {Name: "image_price_1k", Type: field.TypeFloat64, Nullable: true, SchemaType: map[string]string{"postgres": "decimal(20,8)"}}, {Name: "image_price_2k", Type: field.TypeFloat64, Nullable: true, SchemaType: map[string]string{"postgres": "decimal(20,8)"}}, {Name: "image_price_4k", Type: field.TypeFloat64, Nullable: true, SchemaType: map[string]string{"postgres": "decimal(20,8)"}}, + {Name: "batch_image_discount_multiplier", Type: field.TypeFloat64, Default: 0.5, SchemaType: map[string]string{"postgres": "decimal(10,4)"}}, + {Name: "batch_image_hold_multiplier", Type: field.TypeFloat64, Default: 0.6, SchemaType: map[string]string{"postgres": "decimal(10,4)"}}, {Name: "claude_code_only", Type: field.TypeBool, Default: false}, {Name: "fallback_group_id", Type: field.TypeInt64, Nullable: true}, {Name: "fallback_group_id_on_invalid_request", Type: field.TypeInt64, Nullable: true}, @@ -894,7 +910,7 @@ var ( { Name: "group_sort_order", Unique: false, - Columns: []*schema.Column{GroupsColumns[32]}, + Columns: []*schema.Column{GroupsColumns[35]}, }, }, } @@ -1669,6 +1685,7 @@ var ( {Name: "password_hash", Type: field.TypeString, Size: 255}, {Name: "role", Type: field.TypeString, Size: 20, Default: "user"}, {Name: "balance", Type: field.TypeFloat64, Default: 0, SchemaType: map[string]string{"postgres": "decimal(20,8)"}}, + {Name: "frozen_balance", Type: field.TypeFloat64, Default: 0, SchemaType: map[string]string{"postgres": "decimal(20,8)"}}, {Name: "concurrency", Type: field.TypeInt, Default: 5}, {Name: "status", Type: field.TypeString, Size: 20, Default: "active"}, {Name: "username", Type: field.TypeString, Size: 100, Default: ""}, @@ -1695,7 +1712,7 @@ var ( { Name: "user_status", Unique: false, - Columns: []*schema.Column{UsersColumns[9]}, + Columns: []*schema.Column{UsersColumns[10]}, }, { Name: "user_deleted_at", diff --git a/backend/ent/mutation.go b/backend/ent/mutation.go index 7cd434d274..987ec4146b 100644 --- a/backend/ent/mutation.go +++ b/backend/ent/mutation.go @@ -11317,6 +11317,7 @@ type BatchImageJobMutation struct { addaccount_id *int64 provider *string model *string + task_name *string status *string provider_job_name *string provider_input_ref *string @@ -11349,6 +11350,8 @@ type BatchImageJobMutation struct { output_expires_at *time.Time input_deleted_at *time.Time output_deleted_at *time.Time + downloaded_at *time.Time + user_deleted_at *time.Time last_error_code *string last_error_message *string created_at *time.Time @@ -11765,6 +11768,42 @@ func (m *BatchImageJobMutation) ResetModel() { m.model = nil } +// SetTaskName sets the "task_name" field. +func (m *BatchImageJobMutation) SetTaskName(s string) { + m.task_name = &s +} + +// TaskName returns the value of the "task_name" field in the mutation. +func (m *BatchImageJobMutation) TaskName() (r string, exists bool) { + v := m.task_name + if v == nil { + return + } + return *v, true +} + +// OldTaskName returns the old "task_name" field's value of the BatchImageJob entity. +// If the BatchImageJob 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 *BatchImageJobMutation) OldTaskName(ctx context.Context) (v string, err error) { + if !m.op.Is(OpUpdateOne) { + return v, errors.New("OldTaskName is only allowed on UpdateOne operations") + } + if m.id == nil || m.oldValue == nil { + return v, errors.New("OldTaskName requires an ID field in the mutation") + } + oldValue, err := m.oldValue(ctx) + if err != nil { + return v, fmt.Errorf("querying old value for OldTaskName: %w", err) + } + return oldValue.TaskName, nil +} + +// ResetTaskName resets all changes to the "task_name" field. +func (m *BatchImageJobMutation) ResetTaskName() { + m.task_name = nil +} + // SetStatus sets the "status" field. func (m *BatchImageJobMutation) SetStatus(s string) { m.status = &s @@ -12957,6 +12996,104 @@ func (m *BatchImageJobMutation) ResetOutputDeletedAt() { delete(m.clearedFields, batchimagejob.FieldOutputDeletedAt) } +// SetDownloadedAt sets the "downloaded_at" field. +func (m *BatchImageJobMutation) SetDownloadedAt(t time.Time) { + m.downloaded_at = &t +} + +// DownloadedAt returns the value of the "downloaded_at" field in the mutation. +func (m *BatchImageJobMutation) DownloadedAt() (r time.Time, exists bool) { + v := m.downloaded_at + if v == nil { + return + } + return *v, true +} + +// OldDownloadedAt returns the old "downloaded_at" field's value of the BatchImageJob entity. +// If the BatchImageJob 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 *BatchImageJobMutation) OldDownloadedAt(ctx context.Context) (v *time.Time, err error) { + if !m.op.Is(OpUpdateOne) { + return v, errors.New("OldDownloadedAt is only allowed on UpdateOne operations") + } + if m.id == nil || m.oldValue == nil { + return v, errors.New("OldDownloadedAt requires an ID field in the mutation") + } + oldValue, err := m.oldValue(ctx) + if err != nil { + return v, fmt.Errorf("querying old value for OldDownloadedAt: %w", err) + } + return oldValue.DownloadedAt, nil +} + +// ClearDownloadedAt clears the value of the "downloaded_at" field. +func (m *BatchImageJobMutation) ClearDownloadedAt() { + m.downloaded_at = nil + m.clearedFields[batchimagejob.FieldDownloadedAt] = struct{}{} +} + +// DownloadedAtCleared returns if the "downloaded_at" field was cleared in this mutation. +func (m *BatchImageJobMutation) DownloadedAtCleared() bool { + _, ok := m.clearedFields[batchimagejob.FieldDownloadedAt] + return ok +} + +// ResetDownloadedAt resets all changes to the "downloaded_at" field. +func (m *BatchImageJobMutation) ResetDownloadedAt() { + m.downloaded_at = nil + delete(m.clearedFields, batchimagejob.FieldDownloadedAt) +} + +// SetUserDeletedAt sets the "user_deleted_at" field. +func (m *BatchImageJobMutation) SetUserDeletedAt(t time.Time) { + m.user_deleted_at = &t +} + +// UserDeletedAt returns the value of the "user_deleted_at" field in the mutation. +func (m *BatchImageJobMutation) UserDeletedAt() (r time.Time, exists bool) { + v := m.user_deleted_at + if v == nil { + return + } + return *v, true +} + +// OldUserDeletedAt returns the old "user_deleted_at" field's value of the BatchImageJob entity. +// If the BatchImageJob 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 *BatchImageJobMutation) OldUserDeletedAt(ctx context.Context) (v *time.Time, err error) { + if !m.op.Is(OpUpdateOne) { + return v, errors.New("OldUserDeletedAt is only allowed on UpdateOne operations") + } + if m.id == nil || m.oldValue == nil { + return v, errors.New("OldUserDeletedAt requires an ID field in the mutation") + } + oldValue, err := m.oldValue(ctx) + if err != nil { + return v, fmt.Errorf("querying old value for OldUserDeletedAt: %w", err) + } + return oldValue.UserDeletedAt, nil +} + +// ClearUserDeletedAt clears the value of the "user_deleted_at" field. +func (m *BatchImageJobMutation) ClearUserDeletedAt() { + m.user_deleted_at = nil + m.clearedFields[batchimagejob.FieldUserDeletedAt] = struct{}{} +} + +// UserDeletedAtCleared returns if the "user_deleted_at" field was cleared in this mutation. +func (m *BatchImageJobMutation) UserDeletedAtCleared() bool { + _, ok := m.clearedFields[batchimagejob.FieldUserDeletedAt] + return ok +} + +// ResetUserDeletedAt resets all changes to the "user_deleted_at" field. +func (m *BatchImageJobMutation) ResetUserDeletedAt() { + m.user_deleted_at = nil + delete(m.clearedFields, batchimagejob.FieldUserDeletedAt) +} + // SetLastErrorCode sets the "last_error_code" field. func (m *BatchImageJobMutation) SetLastErrorCode(s string) { m.last_error_code = &s @@ -13357,7 +13494,7 @@ func (m *BatchImageJobMutation) Type() string { // order to get all numeric fields that were incremented/decremented, call // AddedFields(). func (m *BatchImageJobMutation) Fields() []string { - fields := make([]string, 0, 37) + fields := make([]string, 0, 40) if m.batch_id != nil { fields = append(fields, batchimagejob.FieldBatchID) } @@ -13376,6 +13513,9 @@ func (m *BatchImageJobMutation) Fields() []string { if m.model != nil { fields = append(fields, batchimagejob.FieldModel) } + if m.task_name != nil { + fields = append(fields, batchimagejob.FieldTaskName) + } if m.status != nil { fields = append(fields, batchimagejob.FieldStatus) } @@ -13445,6 +13585,12 @@ func (m *BatchImageJobMutation) Fields() []string { if m.output_deleted_at != nil { fields = append(fields, batchimagejob.FieldOutputDeletedAt) } + if m.downloaded_at != nil { + fields = append(fields, batchimagejob.FieldDownloadedAt) + } + if m.user_deleted_at != nil { + fields = append(fields, batchimagejob.FieldUserDeletedAt) + } if m.last_error_code != nil { fields = append(fields, batchimagejob.FieldLastErrorCode) } @@ -13489,6 +13635,8 @@ func (m *BatchImageJobMutation) Field(name string) (ent.Value, bool) { return m.Provider() case batchimagejob.FieldModel: return m.Model() + case batchimagejob.FieldTaskName: + return m.TaskName() case batchimagejob.FieldStatus: return m.Status() case batchimagejob.FieldProviderJobName: @@ -13535,6 +13683,10 @@ func (m *BatchImageJobMutation) Field(name string) (ent.Value, bool) { return m.InputDeletedAt() case batchimagejob.FieldOutputDeletedAt: return m.OutputDeletedAt() + case batchimagejob.FieldDownloadedAt: + return m.DownloadedAt() + case batchimagejob.FieldUserDeletedAt: + return m.UserDeletedAt() case batchimagejob.FieldLastErrorCode: return m.LastErrorCode() case batchimagejob.FieldLastErrorMessage: @@ -13572,6 +13724,8 @@ func (m *BatchImageJobMutation) OldField(ctx context.Context, name string) (ent. return m.OldProvider(ctx) case batchimagejob.FieldModel: return m.OldModel(ctx) + case batchimagejob.FieldTaskName: + return m.OldTaskName(ctx) case batchimagejob.FieldStatus: return m.OldStatus(ctx) case batchimagejob.FieldProviderJobName: @@ -13618,6 +13772,10 @@ func (m *BatchImageJobMutation) OldField(ctx context.Context, name string) (ent. return m.OldInputDeletedAt(ctx) case batchimagejob.FieldOutputDeletedAt: return m.OldOutputDeletedAt(ctx) + case batchimagejob.FieldDownloadedAt: + return m.OldDownloadedAt(ctx) + case batchimagejob.FieldUserDeletedAt: + return m.OldUserDeletedAt(ctx) case batchimagejob.FieldLastErrorCode: return m.OldLastErrorCode(ctx) case batchimagejob.FieldLastErrorMessage: @@ -13685,6 +13843,13 @@ func (m *BatchImageJobMutation) SetField(name string, value ent.Value) error { } m.SetModel(v) return nil + case batchimagejob.FieldTaskName: + v, ok := value.(string) + if !ok { + return fmt.Errorf("unexpected type %T for field %s", value, name) + } + m.SetTaskName(v) + return nil case batchimagejob.FieldStatus: v, ok := value.(string) if !ok { @@ -13846,6 +14011,20 @@ func (m *BatchImageJobMutation) SetField(name string, value ent.Value) error { } m.SetOutputDeletedAt(v) return nil + case batchimagejob.FieldDownloadedAt: + v, ok := value.(time.Time) + if !ok { + return fmt.Errorf("unexpected type %T for field %s", value, name) + } + m.SetDownloadedAt(v) + return nil + case batchimagejob.FieldUserDeletedAt: + v, ok := value.(time.Time) + if !ok { + return fmt.Errorf("unexpected type %T for field %s", value, name) + } + m.SetUserDeletedAt(v) + return nil case batchimagejob.FieldLastErrorCode: v, ok := value.(string) if !ok { @@ -14127,6 +14306,12 @@ func (m *BatchImageJobMutation) ClearedFields() []string { if m.FieldCleared(batchimagejob.FieldOutputDeletedAt) { fields = append(fields, batchimagejob.FieldOutputDeletedAt) } + if m.FieldCleared(batchimagejob.FieldDownloadedAt) { + fields = append(fields, batchimagejob.FieldDownloadedAt) + } + if m.FieldCleared(batchimagejob.FieldUserDeletedAt) { + fields = append(fields, batchimagejob.FieldUserDeletedAt) + } if m.FieldCleared(batchimagejob.FieldLastErrorCode) { fields = append(fields, batchimagejob.FieldLastErrorCode) } @@ -14207,6 +14392,12 @@ func (m *BatchImageJobMutation) ClearField(name string) error { case batchimagejob.FieldOutputDeletedAt: m.ClearOutputDeletedAt() return nil + case batchimagejob.FieldDownloadedAt: + m.ClearDownloadedAt() + return nil + case batchimagejob.FieldUserDeletedAt: + m.ClearUserDeletedAt() + return nil case batchimagejob.FieldLastErrorCode: m.ClearLastErrorCode() return nil @@ -14251,6 +14442,9 @@ func (m *BatchImageJobMutation) ResetField(name string) error { case batchimagejob.FieldModel: m.ResetModel() return nil + case batchimagejob.FieldTaskName: + m.ResetTaskName() + return nil case batchimagejob.FieldStatus: m.ResetStatus() return nil @@ -14320,6 +14514,12 @@ func (m *BatchImageJobMutation) ResetField(name string) error { case batchimagejob.FieldOutputDeletedAt: m.ResetOutputDeletedAt() return nil + case batchimagejob.FieldDownloadedAt: + m.ResetDownloadedAt() + return nil + case batchimagejob.FieldUserDeletedAt: + m.ResetUserDeletedAt() + return nil case batchimagejob.FieldLastErrorCode: m.ResetLastErrorCode() return nil @@ -20619,6 +20819,7 @@ type GroupMutation struct { default_validity_days *int adddefault_validity_days *int allow_image_generation *bool + allow_batch_image_generation *bool image_rate_independent *bool image_rate_multiplier *float64 addimage_rate_multiplier *float64 @@ -20628,6 +20829,10 @@ type GroupMutation struct { addimage_price_2k *float64 image_price_4k *float64 addimage_price_4k *float64 + batch_image_discount_multiplier *float64 + addbatch_image_discount_multiplier *float64 + batch_image_hold_multiplier *float64 + addbatch_image_hold_multiplier *float64 claude_code_only *bool fallback_group_id *int64 addfallback_group_id *int64 @@ -21642,6 +21847,42 @@ func (m *GroupMutation) ResetAllowImageGeneration() { m.allow_image_generation = nil } +// SetAllowBatchImageGeneration sets the "allow_batch_image_generation" field. +func (m *GroupMutation) SetAllowBatchImageGeneration(b bool) { + m.allow_batch_image_generation = &b +} + +// AllowBatchImageGeneration returns the value of the "allow_batch_image_generation" field in the mutation. +func (m *GroupMutation) AllowBatchImageGeneration() (r bool, exists bool) { + v := m.allow_batch_image_generation + if v == nil { + return + } + return *v, true +} + +// OldAllowBatchImageGeneration returns the old "allow_batch_image_generation" 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) OldAllowBatchImageGeneration(ctx context.Context) (v bool, err error) { + if !m.op.Is(OpUpdateOne) { + return v, errors.New("OldAllowBatchImageGeneration is only allowed on UpdateOne operations") + } + if m.id == nil || m.oldValue == nil { + return v, errors.New("OldAllowBatchImageGeneration requires an ID field in the mutation") + } + oldValue, err := m.oldValue(ctx) + if err != nil { + return v, fmt.Errorf("querying old value for OldAllowBatchImageGeneration: %w", err) + } + return oldValue.AllowBatchImageGeneration, nil +} + +// ResetAllowBatchImageGeneration resets all changes to the "allow_batch_image_generation" field. +func (m *GroupMutation) ResetAllowBatchImageGeneration() { + m.allow_batch_image_generation = nil +} + // SetImageRateIndependent sets the "image_rate_independent" field. func (m *GroupMutation) SetImageRateIndependent(b bool) { m.image_rate_independent = &b @@ -21944,6 +22185,118 @@ func (m *GroupMutation) ResetImagePrice4k() { delete(m.clearedFields, group.FieldImagePrice4k) } +// SetBatchImageDiscountMultiplier sets the "batch_image_discount_multiplier" field. +func (m *GroupMutation) SetBatchImageDiscountMultiplier(f float64) { + m.batch_image_discount_multiplier = &f + m.addbatch_image_discount_multiplier = nil +} + +// BatchImageDiscountMultiplier returns the value of the "batch_image_discount_multiplier" field in the mutation. +func (m *GroupMutation) BatchImageDiscountMultiplier() (r float64, exists bool) { + v := m.batch_image_discount_multiplier + if v == nil { + return + } + return *v, true +} + +// OldBatchImageDiscountMultiplier returns the old "batch_image_discount_multiplier" 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) OldBatchImageDiscountMultiplier(ctx context.Context) (v float64, err error) { + if !m.op.Is(OpUpdateOne) { + return v, errors.New("OldBatchImageDiscountMultiplier is only allowed on UpdateOne operations") + } + if m.id == nil || m.oldValue == nil { + return v, errors.New("OldBatchImageDiscountMultiplier requires an ID field in the mutation") + } + oldValue, err := m.oldValue(ctx) + if err != nil { + return v, fmt.Errorf("querying old value for OldBatchImageDiscountMultiplier: %w", err) + } + return oldValue.BatchImageDiscountMultiplier, nil +} + +// AddBatchImageDiscountMultiplier adds f to the "batch_image_discount_multiplier" field. +func (m *GroupMutation) AddBatchImageDiscountMultiplier(f float64) { + if m.addbatch_image_discount_multiplier != nil { + *m.addbatch_image_discount_multiplier += f + } else { + m.addbatch_image_discount_multiplier = &f + } +} + +// AddedBatchImageDiscountMultiplier returns the value that was added to the "batch_image_discount_multiplier" field in this mutation. +func (m *GroupMutation) AddedBatchImageDiscountMultiplier() (r float64, exists bool) { + v := m.addbatch_image_discount_multiplier + if v == nil { + return + } + return *v, true +} + +// ResetBatchImageDiscountMultiplier resets all changes to the "batch_image_discount_multiplier" field. +func (m *GroupMutation) ResetBatchImageDiscountMultiplier() { + m.batch_image_discount_multiplier = nil + m.addbatch_image_discount_multiplier = nil +} + +// SetBatchImageHoldMultiplier sets the "batch_image_hold_multiplier" field. +func (m *GroupMutation) SetBatchImageHoldMultiplier(f float64) { + m.batch_image_hold_multiplier = &f + m.addbatch_image_hold_multiplier = nil +} + +// BatchImageHoldMultiplier returns the value of the "batch_image_hold_multiplier" field in the mutation. +func (m *GroupMutation) BatchImageHoldMultiplier() (r float64, exists bool) { + v := m.batch_image_hold_multiplier + if v == nil { + return + } + return *v, true +} + +// OldBatchImageHoldMultiplier returns the old "batch_image_hold_multiplier" 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) OldBatchImageHoldMultiplier(ctx context.Context) (v float64, err error) { + if !m.op.Is(OpUpdateOne) { + return v, errors.New("OldBatchImageHoldMultiplier is only allowed on UpdateOne operations") + } + if m.id == nil || m.oldValue == nil { + return v, errors.New("OldBatchImageHoldMultiplier requires an ID field in the mutation") + } + oldValue, err := m.oldValue(ctx) + if err != nil { + return v, fmt.Errorf("querying old value for OldBatchImageHoldMultiplier: %w", err) + } + return oldValue.BatchImageHoldMultiplier, nil +} + +// AddBatchImageHoldMultiplier adds f to the "batch_image_hold_multiplier" field. +func (m *GroupMutation) AddBatchImageHoldMultiplier(f float64) { + if m.addbatch_image_hold_multiplier != nil { + *m.addbatch_image_hold_multiplier += f + } else { + m.addbatch_image_hold_multiplier = &f + } +} + +// AddedBatchImageHoldMultiplier returns the value that was added to the "batch_image_hold_multiplier" field in this mutation. +func (m *GroupMutation) AddedBatchImageHoldMultiplier() (r float64, exists bool) { + v := m.addbatch_image_hold_multiplier + if v == nil { + return + } + return *v, true +} + +// ResetBatchImageHoldMultiplier resets all changes to the "batch_image_hold_multiplier" field. +func (m *GroupMutation) ResetBatchImageHoldMultiplier() { + m.batch_image_hold_multiplier = nil + m.addbatch_image_hold_multiplier = nil +} + // SetClaudeCodeOnly sets the "claude_code_only" field. func (m *GroupMutation) SetClaudeCodeOnly(b bool) { m.claude_code_only = &b @@ -22978,7 +23331,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, 39) + fields := make([]string, 0, 42) if m.created_at != nil { fields = append(fields, group.FieldCreatedAt) } @@ -23036,6 +23389,9 @@ func (m *GroupMutation) Fields() []string { if m.allow_image_generation != nil { fields = append(fields, group.FieldAllowImageGeneration) } + if m.allow_batch_image_generation != nil { + fields = append(fields, group.FieldAllowBatchImageGeneration) + } if m.image_rate_independent != nil { fields = append(fields, group.FieldImageRateIndependent) } @@ -23051,6 +23407,12 @@ func (m *GroupMutation) Fields() []string { if m.image_price_4k != nil { fields = append(fields, group.FieldImagePrice4k) } + if m.batch_image_discount_multiplier != nil { + fields = append(fields, group.FieldBatchImageDiscountMultiplier) + } + if m.batch_image_hold_multiplier != nil { + fields = append(fields, group.FieldBatchImageHoldMultiplier) + } if m.claude_code_only != nil { fields = append(fields, group.FieldClaudeCodeOnly) } @@ -23142,6 +23504,8 @@ func (m *GroupMutation) Field(name string) (ent.Value, bool) { return m.DefaultValidityDays() case group.FieldAllowImageGeneration: return m.AllowImageGeneration() + case group.FieldAllowBatchImageGeneration: + return m.AllowBatchImageGeneration() case group.FieldImageRateIndependent: return m.ImageRateIndependent() case group.FieldImageRateMultiplier: @@ -23152,6 +23516,10 @@ func (m *GroupMutation) Field(name string) (ent.Value, bool) { return m.ImagePrice2k() case group.FieldImagePrice4k: return m.ImagePrice4k() + case group.FieldBatchImageDiscountMultiplier: + return m.BatchImageDiscountMultiplier() + case group.FieldBatchImageHoldMultiplier: + return m.BatchImageHoldMultiplier() case group.FieldClaudeCodeOnly: return m.ClaudeCodeOnly() case group.FieldFallbackGroupID: @@ -23229,6 +23597,8 @@ func (m *GroupMutation) OldField(ctx context.Context, name string) (ent.Value, e return m.OldDefaultValidityDays(ctx) case group.FieldAllowImageGeneration: return m.OldAllowImageGeneration(ctx) + case group.FieldAllowBatchImageGeneration: + return m.OldAllowBatchImageGeneration(ctx) case group.FieldImageRateIndependent: return m.OldImageRateIndependent(ctx) case group.FieldImageRateMultiplier: @@ -23239,6 +23609,10 @@ func (m *GroupMutation) OldField(ctx context.Context, name string) (ent.Value, e return m.OldImagePrice2k(ctx) case group.FieldImagePrice4k: return m.OldImagePrice4k(ctx) + case group.FieldBatchImageDiscountMultiplier: + return m.OldBatchImageDiscountMultiplier(ctx) + case group.FieldBatchImageHoldMultiplier: + return m.OldBatchImageHoldMultiplier(ctx) case group.FieldClaudeCodeOnly: return m.OldClaudeCodeOnly(ctx) case group.FieldFallbackGroupID: @@ -23411,6 +23785,13 @@ func (m *GroupMutation) SetField(name string, value ent.Value) error { } m.SetAllowImageGeneration(v) return nil + case group.FieldAllowBatchImageGeneration: + v, ok := value.(bool) + if !ok { + return fmt.Errorf("unexpected type %T for field %s", value, name) + } + m.SetAllowBatchImageGeneration(v) + return nil case group.FieldImageRateIndependent: v, ok := value.(bool) if !ok { @@ -23446,6 +23827,20 @@ func (m *GroupMutation) SetField(name string, value ent.Value) error { } m.SetImagePrice4k(v) return nil + case group.FieldBatchImageDiscountMultiplier: + v, ok := value.(float64) + if !ok { + return fmt.Errorf("unexpected type %T for field %s", value, name) + } + m.SetBatchImageDiscountMultiplier(v) + return nil + case group.FieldBatchImageHoldMultiplier: + v, ok := value.(float64) + if !ok { + return fmt.Errorf("unexpected type %T for field %s", value, name) + } + m.SetBatchImageHoldMultiplier(v) + return nil case group.FieldClaudeCodeOnly: v, ok := value.(bool) if !ok { @@ -23589,6 +23984,12 @@ func (m *GroupMutation) AddedFields() []string { if m.addimage_price_4k != nil { fields = append(fields, group.FieldImagePrice4k) } + if m.addbatch_image_discount_multiplier != nil { + fields = append(fields, group.FieldBatchImageDiscountMultiplier) + } + if m.addbatch_image_hold_multiplier != nil { + fields = append(fields, group.FieldBatchImageHoldMultiplier) + } if m.addfallback_group_id != nil { fields = append(fields, group.FieldFallbackGroupID) } @@ -23629,6 +24030,10 @@ func (m *GroupMutation) AddedField(name string) (ent.Value, bool) { return m.AddedImagePrice2k() case group.FieldImagePrice4k: return m.AddedImagePrice4k() + case group.FieldBatchImageDiscountMultiplier: + return m.AddedBatchImageDiscountMultiplier() + case group.FieldBatchImageHoldMultiplier: + return m.AddedBatchImageHoldMultiplier() case group.FieldFallbackGroupID: return m.AddedFallbackGroupID() case group.FieldFallbackGroupIDOnInvalidRequest: @@ -23716,6 +24121,20 @@ func (m *GroupMutation) AddField(name string, value ent.Value) error { } m.AddImagePrice4k(v) return nil + case group.FieldBatchImageDiscountMultiplier: + v, ok := value.(float64) + if !ok { + return fmt.Errorf("unexpected type %T for field %s", value, name) + } + m.AddBatchImageDiscountMultiplier(v) + return nil + case group.FieldBatchImageHoldMultiplier: + v, ok := value.(float64) + if !ok { + return fmt.Errorf("unexpected type %T for field %s", value, name) + } + m.AddBatchImageHoldMultiplier(v) + return nil case group.FieldFallbackGroupID: v, ok := value.(int64) if !ok { @@ -23897,6 +24316,9 @@ func (m *GroupMutation) ResetField(name string) error { case group.FieldAllowImageGeneration: m.ResetAllowImageGeneration() return nil + case group.FieldAllowBatchImageGeneration: + m.ResetAllowBatchImageGeneration() + return nil case group.FieldImageRateIndependent: m.ResetImageRateIndependent() return nil @@ -23912,6 +24334,12 @@ func (m *GroupMutation) ResetField(name string) error { case group.FieldImagePrice4k: m.ResetImagePrice4k() return nil + case group.FieldBatchImageDiscountMultiplier: + m.ResetBatchImageDiscountMultiplier() + return nil + case group.FieldBatchImageHoldMultiplier: + m.ResetBatchImageHoldMultiplier() + return nil case group.FieldClaudeCodeOnly: m.ResetClaudeCodeOnly() return nil @@ -44480,6 +44908,8 @@ type UserMutation struct { role *string balance *float64 addbalance *float64 + frozen_balance *float64 + addfrozen_balance *float64 concurrency *int addconcurrency *int status *string @@ -44928,6 +45358,62 @@ func (m *UserMutation) ResetBalance() { m.addbalance = nil } +// SetFrozenBalance sets the "frozen_balance" field. +func (m *UserMutation) SetFrozenBalance(f float64) { + m.frozen_balance = &f + m.addfrozen_balance = nil +} + +// FrozenBalance returns the value of the "frozen_balance" field in the mutation. +func (m *UserMutation) FrozenBalance() (r float64, exists bool) { + v := m.frozen_balance + if v == nil { + return + } + return *v, true +} + +// OldFrozenBalance returns the old "frozen_balance" field's value of the User entity. +// If the User 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 *UserMutation) OldFrozenBalance(ctx context.Context) (v float64, err error) { + if !m.op.Is(OpUpdateOne) { + return v, errors.New("OldFrozenBalance is only allowed on UpdateOne operations") + } + if m.id == nil || m.oldValue == nil { + return v, errors.New("OldFrozenBalance requires an ID field in the mutation") + } + oldValue, err := m.oldValue(ctx) + if err != nil { + return v, fmt.Errorf("querying old value for OldFrozenBalance: %w", err) + } + return oldValue.FrozenBalance, nil +} + +// AddFrozenBalance adds f to the "frozen_balance" field. +func (m *UserMutation) AddFrozenBalance(f float64) { + if m.addfrozen_balance != nil { + *m.addfrozen_balance += f + } else { + m.addfrozen_balance = &f + } +} + +// AddedFrozenBalance returns the value that was added to the "frozen_balance" field in this mutation. +func (m *UserMutation) AddedFrozenBalance() (r float64, exists bool) { + v := m.addfrozen_balance + if v == nil { + return + } + return *v, true +} + +// ResetFrozenBalance resets all changes to the "frozen_balance" field. +func (m *UserMutation) ResetFrozenBalance() { + m.frozen_balance = nil + m.addfrozen_balance = nil +} + // SetConcurrency sets the "concurrency" field. func (m *UserMutation) SetConcurrency(i int) { m.concurrency = &i @@ -46386,7 +46872,7 @@ func (m *UserMutation) Type() string { // order to get all numeric fields that were incremented/decremented, call // AddedFields(). func (m *UserMutation) Fields() []string { - fields := make([]string, 0, 23) + fields := make([]string, 0, 24) if m.created_at != nil { fields = append(fields, user.FieldCreatedAt) } @@ -46408,6 +46894,9 @@ func (m *UserMutation) Fields() []string { if m.balance != nil { fields = append(fields, user.FieldBalance) } + if m.frozen_balance != nil { + fields = append(fields, user.FieldFrozenBalance) + } if m.concurrency != nil { fields = append(fields, user.FieldConcurrency) } @@ -46478,6 +46967,8 @@ func (m *UserMutation) Field(name string) (ent.Value, bool) { return m.Role() case user.FieldBalance: return m.Balance() + case user.FieldFrozenBalance: + return m.FrozenBalance() case user.FieldConcurrency: return m.Concurrency() case user.FieldStatus: @@ -46533,6 +47024,8 @@ func (m *UserMutation) OldField(ctx context.Context, name string) (ent.Value, er return m.OldRole(ctx) case user.FieldBalance: return m.OldBalance(ctx) + case user.FieldFrozenBalance: + return m.OldFrozenBalance(ctx) case user.FieldConcurrency: return m.OldConcurrency(ctx) case user.FieldStatus: @@ -46623,6 +47116,13 @@ func (m *UserMutation) SetField(name string, value ent.Value) error { } m.SetBalance(v) return nil + case user.FieldFrozenBalance: + v, ok := value.(float64) + if !ok { + return fmt.Errorf("unexpected type %T for field %s", value, name) + } + m.SetFrozenBalance(v) + return nil case user.FieldConcurrency: v, ok := value.(int) if !ok { @@ -46746,6 +47246,9 @@ func (m *UserMutation) AddedFields() []string { if m.addbalance != nil { fields = append(fields, user.FieldBalance) } + if m.addfrozen_balance != nil { + fields = append(fields, user.FieldFrozenBalance) + } if m.addconcurrency != nil { fields = append(fields, user.FieldConcurrency) } @@ -46768,6 +47271,8 @@ func (m *UserMutation) AddedField(name string) (ent.Value, bool) { switch name { case user.FieldBalance: return m.AddedBalance() + case user.FieldFrozenBalance: + return m.AddedFrozenBalance() case user.FieldConcurrency: return m.AddedConcurrency() case user.FieldBalanceNotifyThreshold: @@ -46792,6 +47297,13 @@ func (m *UserMutation) AddField(name string, value ent.Value) error { } m.AddBalance(v) return nil + case user.FieldFrozenBalance: + v, ok := value.(float64) + if !ok { + return fmt.Errorf("unexpected type %T for field %s", value, name) + } + m.AddFrozenBalance(v) + return nil case user.FieldConcurrency: v, ok := value.(int) if !ok { @@ -46907,6 +47419,9 @@ func (m *UserMutation) ResetField(name string) error { case user.FieldBalance: m.ResetBalance() return nil + case user.FieldFrozenBalance: + m.ResetFrozenBalance() + return nil case user.FieldConcurrency: m.ResetConcurrency() return nil diff --git a/backend/ent/runtime/runtime.go b/backend/ent/runtime/runtime.go index a924e1fa4c..2c05c01c52 100644 --- a/backend/ent/runtime/runtime.go +++ b/backend/ent/runtime/runtime.go @@ -509,88 +509,94 @@ func init() { batchimagejobDescModel := batchimagejobFields[5].Descriptor() // batchimagejob.ModelValidator is a validator for the "model" field. It is called by the builders before save. batchimagejob.ModelValidator = batchimagejobDescModel.Validators[0].(func(string) error) + // batchimagejobDescTaskName is the schema descriptor for task_name field. + batchimagejobDescTaskName := batchimagejobFields[6].Descriptor() + // batchimagejob.DefaultTaskName holds the default value on creation for the task_name field. + batchimagejob.DefaultTaskName = batchimagejobDescTaskName.Default.(string) + // batchimagejob.TaskNameValidator is a validator for the "task_name" field. It is called by the builders before save. + batchimagejob.TaskNameValidator = batchimagejobDescTaskName.Validators[0].(func(string) error) // batchimagejobDescStatus is the schema descriptor for status field. - batchimagejobDescStatus := batchimagejobFields[6].Descriptor() + batchimagejobDescStatus := batchimagejobFields[7].Descriptor() // batchimagejob.DefaultStatus holds the default value on creation for the status field. batchimagejob.DefaultStatus = batchimagejobDescStatus.Default.(string) // batchimagejob.StatusValidator is a validator for the "status" field. It is called by the builders before save. batchimagejob.StatusValidator = batchimagejobDescStatus.Validators[0].(func(string) error) // batchimagejobDescProviderJobName is the schema descriptor for provider_job_name field. - batchimagejobDescProviderJobName := batchimagejobFields[7].Descriptor() + batchimagejobDescProviderJobName := batchimagejobFields[8].Descriptor() // batchimagejob.ProviderJobNameValidator is a validator for the "provider_job_name" field. It is called by the builders before save. batchimagejob.ProviderJobNameValidator = batchimagejobDescProviderJobName.Validators[0].(func(string) error) // batchimagejobDescProviderInputRef is the schema descriptor for provider_input_ref field. - batchimagejobDescProviderInputRef := batchimagejobFields[8].Descriptor() + batchimagejobDescProviderInputRef := batchimagejobFields[9].Descriptor() // batchimagejob.ProviderInputRefValidator is a validator for the "provider_input_ref" field. It is called by the builders before save. batchimagejob.ProviderInputRefValidator = batchimagejobDescProviderInputRef.Validators[0].(func(string) error) // batchimagejobDescProviderOutputRef is the schema descriptor for provider_output_ref field. - batchimagejobDescProviderOutputRef := batchimagejobFields[9].Descriptor() + batchimagejobDescProviderOutputRef := batchimagejobFields[10].Descriptor() // batchimagejob.ProviderOutputRefValidator is a validator for the "provider_output_ref" field. It is called by the builders before save. batchimagejob.ProviderOutputRefValidator = batchimagejobDescProviderOutputRef.Validators[0].(func(string) error) // batchimagejobDescGcsInputURI is the schema descriptor for gcs_input_uri field. - batchimagejobDescGcsInputURI := batchimagejobFields[10].Descriptor() + batchimagejobDescGcsInputURI := batchimagejobFields[11].Descriptor() // batchimagejob.GcsInputURIValidator is a validator for the "gcs_input_uri" field. It is called by the builders before save. batchimagejob.GcsInputURIValidator = batchimagejobDescGcsInputURI.Validators[0].(func(string) error) // batchimagejobDescGcsOutputURI is the schema descriptor for gcs_output_uri field. - batchimagejobDescGcsOutputURI := batchimagejobFields[11].Descriptor() + batchimagejobDescGcsOutputURI := batchimagejobFields[12].Descriptor() // batchimagejob.GcsOutputURIValidator is a validator for the "gcs_output_uri" field. It is called by the builders before save. batchimagejob.GcsOutputURIValidator = batchimagejobDescGcsOutputURI.Validators[0].(func(string) error) // batchimagejobDescSuccessCount is the schema descriptor for success_count field. - batchimagejobDescSuccessCount := batchimagejobFields[13].Descriptor() + batchimagejobDescSuccessCount := batchimagejobFields[14].Descriptor() // batchimagejob.DefaultSuccessCount holds the default value on creation for the success_count field. batchimagejob.DefaultSuccessCount = batchimagejobDescSuccessCount.Default.(int) // batchimagejobDescFailCount is the schema descriptor for fail_count field. - batchimagejobDescFailCount := batchimagejobFields[14].Descriptor() + batchimagejobDescFailCount := batchimagejobFields[15].Descriptor() // batchimagejob.DefaultFailCount holds the default value on creation for the fail_count field. batchimagejob.DefaultFailCount = batchimagejobDescFailCount.Default.(int) // batchimagejobDescCancelledCount is the schema descriptor for cancelled_count field. - batchimagejobDescCancelledCount := batchimagejobFields[15].Descriptor() + batchimagejobDescCancelledCount := batchimagejobFields[16].Descriptor() // batchimagejob.DefaultCancelledCount holds the default value on creation for the cancelled_count field. batchimagejob.DefaultCancelledCount = batchimagejobDescCancelledCount.Default.(int) // batchimagejobDescEstimatedCost is the schema descriptor for estimated_cost field. - batchimagejobDescEstimatedCost := batchimagejobFields[16].Descriptor() + batchimagejobDescEstimatedCost := batchimagejobFields[17].Descriptor() // batchimagejob.DefaultEstimatedCost holds the default value on creation for the estimated_cost field. batchimagejob.DefaultEstimatedCost = batchimagejobDescEstimatedCost.Default.(float64) // batchimagejobDescCurrency is the schema descriptor for currency field. - batchimagejobDescCurrency := batchimagejobFields[19].Descriptor() + batchimagejobDescCurrency := batchimagejobFields[20].Descriptor() // batchimagejob.DefaultCurrency holds the default value on creation for the currency field. batchimagejob.DefaultCurrency = batchimagejobDescCurrency.Default.(string) // batchimagejob.CurrencyValidator is a validator for the "currency" field. It is called by the builders before save. batchimagejob.CurrencyValidator = batchimagejobDescCurrency.Validators[0].(func(string) error) // batchimagejobDescHoldID is the schema descriptor for hold_id field. - batchimagejobDescHoldID := batchimagejobFields[20].Descriptor() + batchimagejobDescHoldID := batchimagejobFields[21].Descriptor() // batchimagejob.HoldIDValidator is a validator for the "hold_id" field. It is called by the builders before save. batchimagejob.HoldIDValidator = batchimagejobDescHoldID.Validators[0].(func(string) error) // batchimagejobDescIdempotencyKey is the schema descriptor for idempotency_key field. - batchimagejobDescIdempotencyKey := batchimagejobFields[21].Descriptor() + batchimagejobDescIdempotencyKey := batchimagejobFields[22].Descriptor() // batchimagejob.IdempotencyKeyValidator is a validator for the "idempotency_key" field. It is called by the builders before save. batchimagejob.IdempotencyKeyValidator = batchimagejobDescIdempotencyKey.Validators[0].(func(string) error) // batchimagejobDescRequestHash is the schema descriptor for request_hash field. - batchimagejobDescRequestHash := batchimagejobFields[22].Descriptor() + batchimagejobDescRequestHash := batchimagejobFields[23].Descriptor() // batchimagejob.RequestHashValidator is a validator for the "request_hash" field. It is called by the builders before save. batchimagejob.RequestHashValidator = batchimagejobDescRequestHash.Validators[0].(func(string) error) // batchimagejobDescManifestHash is the schema descriptor for manifest_hash field. - batchimagejobDescManifestHash := batchimagejobFields[23].Descriptor() + batchimagejobDescManifestHash := batchimagejobFields[24].Descriptor() // batchimagejob.ManifestHashValidator is a validator for the "manifest_hash" field. It is called by the builders before save. batchimagejob.ManifestHashValidator = batchimagejobDescManifestHash.Validators[0].(func(string) error) // batchimagejobDescRetryCount is the schema descriptor for retry_count field. - batchimagejobDescRetryCount := batchimagejobFields[24].Descriptor() + batchimagejobDescRetryCount := batchimagejobFields[25].Descriptor() // batchimagejob.DefaultRetryCount holds the default value on creation for the retry_count field. batchimagejob.DefaultRetryCount = batchimagejobDescRetryCount.Default.(int) // batchimagejobDescVersion is the schema descriptor for version field. - batchimagejobDescVersion := batchimagejobFields[25].Descriptor() + batchimagejobDescVersion := batchimagejobFields[26].Descriptor() // batchimagejob.DefaultVersion holds the default value on creation for the version field. batchimagejob.DefaultVersion = batchimagejobDescVersion.Default.(int) // batchimagejobDescLastErrorCode is the schema descriptor for last_error_code field. - batchimagejobDescLastErrorCode := batchimagejobFields[29].Descriptor() + batchimagejobDescLastErrorCode := batchimagejobFields[32].Descriptor() // batchimagejob.LastErrorCodeValidator is a validator for the "last_error_code" field. It is called by the builders before save. batchimagejob.LastErrorCodeValidator = batchimagejobDescLastErrorCode.Validators[0].(func(string) error) // batchimagejobDescCreatedAt is the schema descriptor for created_at field. - batchimagejobDescCreatedAt := batchimagejobFields[31].Descriptor() + batchimagejobDescCreatedAt := batchimagejobFields[34].Descriptor() // batchimagejob.DefaultCreatedAt holds the default value on creation for the created_at field. batchimagejob.DefaultCreatedAt = batchimagejobDescCreatedAt.Default.(func() time.Time) // batchimagejobDescUpdatedAt is the schema descriptor for updated_at field. - batchimagejobDescUpdatedAt := batchimagejobFields[32].Descriptor() + batchimagejobDescUpdatedAt := batchimagejobFields[35].Descriptor() // batchimagejob.DefaultUpdatedAt holds the default value on creation for the updated_at field. batchimagejob.DefaultUpdatedAt = batchimagejobDescUpdatedAt.Default.(func() time.Time) // batchimagejob.UpdateDefaultUpdatedAt holds the default value on update for the updated_at field. @@ -1009,62 +1015,74 @@ func init() { groupDescAllowImageGeneration := groupFields[15].Descriptor() // group.DefaultAllowImageGeneration holds the default value on creation for the allow_image_generation field. group.DefaultAllowImageGeneration = groupDescAllowImageGeneration.Default.(bool) + // groupDescAllowBatchImageGeneration is the schema descriptor for allow_batch_image_generation field. + groupDescAllowBatchImageGeneration := groupFields[16].Descriptor() + // group.DefaultAllowBatchImageGeneration holds the default value on creation for the allow_batch_image_generation field. + group.DefaultAllowBatchImageGeneration = groupDescAllowBatchImageGeneration.Default.(bool) // groupDescImageRateIndependent is the schema descriptor for image_rate_independent field. - groupDescImageRateIndependent := groupFields[16].Descriptor() + groupDescImageRateIndependent := groupFields[17].Descriptor() // group.DefaultImageRateIndependent holds the default value on creation for the image_rate_independent field. group.DefaultImageRateIndependent = groupDescImageRateIndependent.Default.(bool) // groupDescImageRateMultiplier is the schema descriptor for image_rate_multiplier field. - groupDescImageRateMultiplier := groupFields[17].Descriptor() + groupDescImageRateMultiplier := groupFields[18].Descriptor() // group.DefaultImageRateMultiplier holds the default value on creation for the image_rate_multiplier field. group.DefaultImageRateMultiplier = groupDescImageRateMultiplier.Default.(float64) + // groupDescBatchImageDiscountMultiplier is the schema descriptor for batch_image_discount_multiplier field. + groupDescBatchImageDiscountMultiplier := groupFields[22].Descriptor() + // group.DefaultBatchImageDiscountMultiplier holds the default value on creation for the batch_image_discount_multiplier field. + group.DefaultBatchImageDiscountMultiplier = groupDescBatchImageDiscountMultiplier.Default.(float64) + // groupDescBatchImageHoldMultiplier is the schema descriptor for batch_image_hold_multiplier field. + groupDescBatchImageHoldMultiplier := groupFields[23].Descriptor() + // group.DefaultBatchImageHoldMultiplier holds the default value on creation for the batch_image_hold_multiplier field. + group.DefaultBatchImageHoldMultiplier = groupDescBatchImageHoldMultiplier.Default.(float64) // groupDescClaudeCodeOnly is the schema descriptor for claude_code_only field. - groupDescClaudeCodeOnly := groupFields[21].Descriptor() + groupDescClaudeCodeOnly := groupFields[24].Descriptor() // group.DefaultClaudeCodeOnly holds the default value on creation for the claude_code_only field. group.DefaultClaudeCodeOnly = groupDescClaudeCodeOnly.Default.(bool) // groupDescModelRoutingEnabled is the schema descriptor for model_routing_enabled field. - groupDescModelRoutingEnabled := groupFields[25].Descriptor() + groupDescModelRoutingEnabled := groupFields[28].Descriptor() // group.DefaultModelRoutingEnabled holds the default value on creation for the model_routing_enabled field. group.DefaultModelRoutingEnabled = groupDescModelRoutingEnabled.Default.(bool) // groupDescMcpXMLInject is the schema descriptor for mcp_xml_inject field. - groupDescMcpXMLInject := groupFields[26].Descriptor() + groupDescMcpXMLInject := groupFields[29].Descriptor() // group.DefaultMcpXMLInject holds the default value on creation for the mcp_xml_inject field. group.DefaultMcpXMLInject = groupDescMcpXMLInject.Default.(bool) // groupDescSupportedModelScopes is the schema descriptor for supported_model_scopes field. - groupDescSupportedModelScopes := groupFields[27].Descriptor() + groupDescSupportedModelScopes := groupFields[30].Descriptor() // group.DefaultSupportedModelScopes holds the default value on creation for the supported_model_scopes field. group.DefaultSupportedModelScopes = groupDescSupportedModelScopes.Default.([]string) // groupDescSortOrder is the schema descriptor for sort_order field. - groupDescSortOrder := groupFields[28].Descriptor() + groupDescSortOrder := groupFields[31].Descriptor() // group.DefaultSortOrder holds the default value on creation for the sort_order field. group.DefaultSortOrder = groupDescSortOrder.Default.(int) // groupDescAllowMessagesDispatch is the schema descriptor for allow_messages_dispatch field. - groupDescAllowMessagesDispatch := groupFields[29].Descriptor() + groupDescAllowMessagesDispatch := groupFields[32].Descriptor() // group.DefaultAllowMessagesDispatch holds the default value on creation for the allow_messages_dispatch field. group.DefaultAllowMessagesDispatch = groupDescAllowMessagesDispatch.Default.(bool) // groupDescRequireOauthOnly is the schema descriptor for require_oauth_only field. - groupDescRequireOauthOnly := groupFields[30].Descriptor() + groupDescRequireOauthOnly := groupFields[33].Descriptor() // group.DefaultRequireOauthOnly holds the default value on creation for the require_oauth_only field. group.DefaultRequireOauthOnly = groupDescRequireOauthOnly.Default.(bool) // groupDescRequirePrivacySet is the schema descriptor for require_privacy_set field. - groupDescRequirePrivacySet := groupFields[31].Descriptor() + groupDescRequirePrivacySet := groupFields[34].Descriptor() // group.DefaultRequirePrivacySet holds the default value on creation for the require_privacy_set field. group.DefaultRequirePrivacySet = groupDescRequirePrivacySet.Default.(bool) // groupDescDefaultMappedModel is the schema descriptor for default_mapped_model field. - groupDescDefaultMappedModel := groupFields[32].Descriptor() + groupDescDefaultMappedModel := groupFields[35].Descriptor() // group.DefaultDefaultMappedModel holds the default value on creation for the default_mapped_model field. 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) // groupDescMessagesDispatchModelConfig is the schema descriptor for messages_dispatch_model_config field. - groupDescMessagesDispatchModelConfig := groupFields[33].Descriptor() + groupDescMessagesDispatchModelConfig := groupFields[36].Descriptor() // group.DefaultMessagesDispatchModelConfig holds the default value on creation for the messages_dispatch_model_config field. group.DefaultMessagesDispatchModelConfig = groupDescMessagesDispatchModelConfig.Default.(domain.OpenAIMessagesDispatchModelConfig) // groupDescModelsListConfig is the schema descriptor for models_list_config field. - groupDescModelsListConfig := groupFields[34].Descriptor() + groupDescModelsListConfig := groupFields[37].Descriptor() // group.DefaultModelsListConfig holds the default value on creation for the models_list_config field. group.DefaultModelsListConfig = groupDescModelsListConfig.Default.(domain.GroupModelsListConfig) // groupDescRpmLimit is the schema descriptor for rpm_limit field. - groupDescRpmLimit := groupFields[35].Descriptor() + groupDescRpmLimit := groupFields[38].Descriptor() // group.DefaultRpmLimit holds the default value on creation for the rpm_limit field. group.DefaultRpmLimit = groupDescRpmLimit.Default.(int) idempotencyrecordMixin := schema.IdempotencyRecord{}.Mixin() @@ -2023,54 +2041,58 @@ func init() { userDescBalance := userFields[3].Descriptor() // user.DefaultBalance holds the default value on creation for the balance field. user.DefaultBalance = userDescBalance.Default.(float64) + // userDescFrozenBalance is the schema descriptor for frozen_balance field. + userDescFrozenBalance := userFields[4].Descriptor() + // user.DefaultFrozenBalance holds the default value on creation for the frozen_balance field. + user.DefaultFrozenBalance = userDescFrozenBalance.Default.(float64) // userDescConcurrency is the schema descriptor for concurrency field. - userDescConcurrency := userFields[4].Descriptor() + userDescConcurrency := userFields[5].Descriptor() // user.DefaultConcurrency holds the default value on creation for the concurrency field. user.DefaultConcurrency = userDescConcurrency.Default.(int) // userDescStatus is the schema descriptor for status field. - userDescStatus := userFields[5].Descriptor() + userDescStatus := userFields[6].Descriptor() // user.DefaultStatus holds the default value on creation for the status field. user.DefaultStatus = userDescStatus.Default.(string) // user.StatusValidator is a validator for the "status" field. It is called by the builders before save. user.StatusValidator = userDescStatus.Validators[0].(func(string) error) // userDescUsername is the schema descriptor for username field. - userDescUsername := userFields[6].Descriptor() + userDescUsername := userFields[7].Descriptor() // user.DefaultUsername holds the default value on creation for the username field. user.DefaultUsername = userDescUsername.Default.(string) // user.UsernameValidator is a validator for the "username" field. It is called by the builders before save. user.UsernameValidator = userDescUsername.Validators[0].(func(string) error) // userDescNotes is the schema descriptor for notes field. - userDescNotes := userFields[7].Descriptor() + userDescNotes := userFields[8].Descriptor() // user.DefaultNotes holds the default value on creation for the notes field. user.DefaultNotes = userDescNotes.Default.(string) // userDescTotpEnabled is the schema descriptor for totp_enabled field. - userDescTotpEnabled := userFields[9].Descriptor() + userDescTotpEnabled := userFields[10].Descriptor() // user.DefaultTotpEnabled holds the default value on creation for the totp_enabled field. user.DefaultTotpEnabled = userDescTotpEnabled.Default.(bool) // userDescSignupSource is the schema descriptor for signup_source field. - userDescSignupSource := userFields[11].Descriptor() + userDescSignupSource := userFields[12].Descriptor() // user.DefaultSignupSource holds the default value on creation for the signup_source field. user.DefaultSignupSource = userDescSignupSource.Default.(string) // user.SignupSourceValidator is a validator for the "signup_source" field. It is called by the builders before save. user.SignupSourceValidator = userDescSignupSource.Validators[0].(func(string) error) // userDescBalanceNotifyEnabled is the schema descriptor for balance_notify_enabled field. - userDescBalanceNotifyEnabled := userFields[14].Descriptor() + userDescBalanceNotifyEnabled := userFields[15].Descriptor() // user.DefaultBalanceNotifyEnabled holds the default value on creation for the balance_notify_enabled field. user.DefaultBalanceNotifyEnabled = userDescBalanceNotifyEnabled.Default.(bool) // userDescBalanceNotifyThresholdType is the schema descriptor for balance_notify_threshold_type field. - userDescBalanceNotifyThresholdType := userFields[15].Descriptor() + userDescBalanceNotifyThresholdType := userFields[16].Descriptor() // user.DefaultBalanceNotifyThresholdType holds the default value on creation for the balance_notify_threshold_type field. user.DefaultBalanceNotifyThresholdType = userDescBalanceNotifyThresholdType.Default.(string) // userDescBalanceNotifyExtraEmails is the schema descriptor for balance_notify_extra_emails field. - userDescBalanceNotifyExtraEmails := userFields[17].Descriptor() + userDescBalanceNotifyExtraEmails := userFields[18].Descriptor() // user.DefaultBalanceNotifyExtraEmails holds the default value on creation for the balance_notify_extra_emails field. user.DefaultBalanceNotifyExtraEmails = userDescBalanceNotifyExtraEmails.Default.(string) // userDescTotalRecharged is the schema descriptor for total_recharged field. - userDescTotalRecharged := userFields[18].Descriptor() + userDescTotalRecharged := userFields[19].Descriptor() // user.DefaultTotalRecharged holds the default value on creation for the total_recharged field. user.DefaultTotalRecharged = userDescTotalRecharged.Default.(float64) // userDescRpmLimit is the schema descriptor for rpm_limit field. - userDescRpmLimit := userFields[19].Descriptor() + userDescRpmLimit := userFields[20].Descriptor() // user.DefaultRpmLimit holds the default value on creation for the rpm_limit field. user.DefaultRpmLimit = userDescRpmLimit.Default.(int) userallowedgroupFields := schema.UserAllowedGroup{}.Fields() diff --git a/backend/ent/schema/batch_image_job.go b/backend/ent/schema/batch_image_job.go index ba159f4cb8..a65156eaea 100644 --- a/backend/ent/schema/batch_image_job.go +++ b/backend/ent/schema/batch_image_job.go @@ -13,9 +13,9 @@ import ( // BatchImageJob holds the schema definition for asynchronous image batch jobs. // -// 删除策略:硬删除 -// 这张表是批量生图任务的账务和状态源,不使用软删除;输出清理通过 -// output_deleted 状态和删除时间字段表达。 +// 删除策略:账务源保留 +// 这张表是批量生图任务的账务和状态源;用户侧删除仅通过 user_deleted_at +// 从列表隐藏,输出清理通过 output_deleted 状态和删除时间字段表达。 type BatchImageJob struct { ent.Schema } @@ -34,6 +34,7 @@ func (BatchImageJob) Fields() []ent.Field { field.Int64("account_id").Optional().Nillable(), field.String("provider").MaxLen(32), field.String("model").MaxLen(128), + field.String("task_name").MaxLen(255).Default(""), field.String("status").MaxLen(32).Default("created"), field.String("provider_job_name").Optional().Nillable().MaxLen(512), field.String("provider_input_ref").Optional().Nillable().MaxLen(1024), @@ -57,6 +58,8 @@ func (BatchImageJob) Fields() []ent.Field { field.Time("output_expires_at").Optional().Nillable().SchemaType(map[string]string{dialect.Postgres: "timestamptz"}), field.Time("input_deleted_at").Optional().Nillable().SchemaType(map[string]string{dialect.Postgres: "timestamptz"}), field.Time("output_deleted_at").Optional().Nillable().SchemaType(map[string]string{dialect.Postgres: "timestamptz"}), + field.Time("downloaded_at").Optional().Nillable().SchemaType(map[string]string{dialect.Postgres: "timestamptz"}), + field.Time("user_deleted_at").Optional().Nillable().SchemaType(map[string]string{dialect.Postgres: "timestamptz"}), field.String("last_error_code").Optional().Nillable().MaxLen(128), field.String("last_error_message").Optional().Nillable().SchemaType(map[string]string{dialect.Postgres: "text"}), field.Time("created_at").Immutable().Default(time.Now).SchemaType(map[string]string{dialect.Postgres: "timestamptz"}), @@ -77,5 +80,7 @@ func (BatchImageJob) Indexes() []ent.Index { index.Fields("idempotency_key").Annotations(entsql.IndexWhere("idempotency_key IS NOT NULL AND idempotency_key <> ''")), index.Fields("manifest_hash").Unique().Annotations(entsql.IndexWhere("manifest_hash IS NOT NULL AND manifest_hash <> ''")), index.Fields("output_expires_at"), + index.Fields("downloaded_at"), + index.Fields("user_deleted_at"), } } diff --git a/backend/ent/schema/group.go b/backend/ent/schema/group.go index 2b8420db6d..d675ca52f1 100644 --- a/backend/ent/schema/group.go +++ b/backend/ent/schema/group.go @@ -93,6 +93,9 @@ func (Group) Fields() []ent.Field { field.Bool("allow_image_generation"). Default(false). Comment("是否允许该分组使用图片生成能力"), + field.Bool("allow_batch_image_generation"). + Default(false). + Comment("是否允许该分组使用批量图片生成能力"), field.Bool("image_rate_independent"). Default(false). Comment("图片生成是否使用独立倍率;false 表示共享分组有效倍率"), @@ -112,6 +115,14 @@ func (Group) Fields() []ent.Field { Optional(). Nillable(). SchemaType(map[string]string{dialect.Postgres: "decimal(20,8)"}), + field.Float("batch_image_discount_multiplier"). + SchemaType(map[string]string{dialect.Postgres: "decimal(10,4)"}). + Default(0.5). + Comment("批量图片生成折扣倍率,最终单价会乘以该值;0 表示免费"), + field.Float("batch_image_hold_multiplier"). + SchemaType(map[string]string{dialect.Postgres: "decimal(10,4)"}). + Default(0.6). + Comment("批量图片生成冻结价格比例,按普通生图原价乘以该比例冻结,结算后释放差额"), // Claude Code 客户端限制 (added by migration 029) field.Bool("claude_code_only"). diff --git a/backend/ent/schema/user.go b/backend/ent/schema/user.go index 127b5af9a7..baa7efbbd9 100644 --- a/backend/ent/schema/user.go +++ b/backend/ent/schema/user.go @@ -49,6 +49,9 @@ func (User) Fields() []ent.Field { field.Float("balance"). SchemaType(map[string]string{dialect.Postgres: "decimal(20,8)"}). Default(0), + field.Float("frozen_balance"). + SchemaType(map[string]string{dialect.Postgres: "decimal(20,8)"}). + Default(0), field.Int("concurrency"). Default(5), field.String("status"). diff --git a/backend/ent/user.go b/backend/ent/user.go index 486f2f64d9..299a8d627f 100644 --- a/backend/ent/user.go +++ b/backend/ent/user.go @@ -31,6 +31,8 @@ type User struct { Role string `json:"role,omitempty"` // Balance holds the value of the "balance" field. Balance float64 `json:"balance,omitempty"` + // FrozenBalance holds the value of the "frozen_balance" field. + FrozenBalance float64 `json:"frozen_balance,omitempty"` // Concurrency holds the value of the "concurrency" field. Concurrency int `json:"concurrency,omitempty"` // Status holds the value of the "status" field. @@ -237,7 +239,7 @@ func (*User) scanValues(columns []string) ([]any, error) { switch columns[i] { case user.FieldTotpEnabled, user.FieldBalanceNotifyEnabled: values[i] = new(sql.NullBool) - case user.FieldBalance, user.FieldBalanceNotifyThreshold, user.FieldTotalRecharged: + case user.FieldBalance, user.FieldFrozenBalance, user.FieldBalanceNotifyThreshold, user.FieldTotalRecharged: values[i] = new(sql.NullFloat64) case user.FieldID, user.FieldConcurrency, user.FieldRpmLimit: values[i] = new(sql.NullInt64) @@ -309,6 +311,12 @@ func (_m *User) assignValues(columns []string, values []any) error { } else if value.Valid { _m.Balance = value.Float64 } + case user.FieldFrozenBalance: + if value, ok := values[i].(*sql.NullFloat64); !ok { + return fmt.Errorf("unexpected type %T for field frozen_balance", values[i]) + } else if value.Valid { + _m.FrozenBalance = value.Float64 + } case user.FieldConcurrency: if value, ok := values[i].(*sql.NullInt64); !ok { return fmt.Errorf("unexpected type %T for field concurrency", values[i]) @@ -539,6 +547,9 @@ func (_m *User) String() string { builder.WriteString("balance=") builder.WriteString(fmt.Sprintf("%v", _m.Balance)) builder.WriteString(", ") + builder.WriteString("frozen_balance=") + builder.WriteString(fmt.Sprintf("%v", _m.FrozenBalance)) + builder.WriteString(", ") builder.WriteString("concurrency=") builder.WriteString(fmt.Sprintf("%v", _m.Concurrency)) builder.WriteString(", ") diff --git a/backend/ent/user/user.go b/backend/ent/user/user.go index ff40445bda..ae1a84494d 100644 --- a/backend/ent/user/user.go +++ b/backend/ent/user/user.go @@ -29,6 +29,8 @@ const ( FieldRole = "role" // FieldBalance holds the string denoting the balance field in the database. FieldBalance = "balance" + // FieldFrozenBalance holds the string denoting the frozen_balance field in the database. + FieldFrozenBalance = "frozen_balance" // FieldConcurrency holds the string denoting the concurrency field in the database. FieldConcurrency = "concurrency" // FieldStatus holds the string denoting the status field in the database. @@ -199,6 +201,7 @@ var Columns = []string{ FieldPasswordHash, FieldRole, FieldBalance, + FieldFrozenBalance, FieldConcurrency, FieldStatus, FieldUsername, @@ -257,6 +260,8 @@ var ( RoleValidator func(string) error // DefaultBalance holds the default value on creation for the "balance" field. DefaultBalance float64 + // DefaultFrozenBalance holds the default value on creation for the "frozen_balance" field. + DefaultFrozenBalance float64 // DefaultConcurrency holds the default value on creation for the "concurrency" field. DefaultConcurrency int // DefaultStatus holds the default value on creation for the "status" field. @@ -330,6 +335,11 @@ func ByBalance(opts ...sql.OrderTermOption) OrderOption { return sql.OrderByField(FieldBalance, opts...).ToFunc() } +// ByFrozenBalance orders the results by the frozen_balance field. +func ByFrozenBalance(opts ...sql.OrderTermOption) OrderOption { + return sql.OrderByField(FieldFrozenBalance, opts...).ToFunc() +} + // ByConcurrency orders the results by the concurrency field. func ByConcurrency(opts ...sql.OrderTermOption) OrderOption { return sql.OrderByField(FieldConcurrency, opts...).ToFunc() diff --git a/backend/ent/user/where.go b/backend/ent/user/where.go index a18cf49767..c2a71f6172 100644 --- a/backend/ent/user/where.go +++ b/backend/ent/user/where.go @@ -90,6 +90,11 @@ func Balance(v float64) predicate.User { return predicate.User(sql.FieldEQ(FieldBalance, v)) } +// FrozenBalance applies equality check predicate on the "frozen_balance" field. It's identical to FrozenBalanceEQ. +func FrozenBalance(v float64) predicate.User { + return predicate.User(sql.FieldEQ(FieldFrozenBalance, v)) +} + // Concurrency applies equality check predicate on the "concurrency" field. It's identical to ConcurrencyEQ. func Concurrency(v int) predicate.User { return predicate.User(sql.FieldEQ(FieldConcurrency, v)) @@ -535,6 +540,46 @@ func BalanceLTE(v float64) predicate.User { return predicate.User(sql.FieldLTE(FieldBalance, v)) } +// FrozenBalanceEQ applies the EQ predicate on the "frozen_balance" field. +func FrozenBalanceEQ(v float64) predicate.User { + return predicate.User(sql.FieldEQ(FieldFrozenBalance, v)) +} + +// FrozenBalanceNEQ applies the NEQ predicate on the "frozen_balance" field. +func FrozenBalanceNEQ(v float64) predicate.User { + return predicate.User(sql.FieldNEQ(FieldFrozenBalance, v)) +} + +// FrozenBalanceIn applies the In predicate on the "frozen_balance" field. +func FrozenBalanceIn(vs ...float64) predicate.User { + return predicate.User(sql.FieldIn(FieldFrozenBalance, vs...)) +} + +// FrozenBalanceNotIn applies the NotIn predicate on the "frozen_balance" field. +func FrozenBalanceNotIn(vs ...float64) predicate.User { + return predicate.User(sql.FieldNotIn(FieldFrozenBalance, vs...)) +} + +// FrozenBalanceGT applies the GT predicate on the "frozen_balance" field. +func FrozenBalanceGT(v float64) predicate.User { + return predicate.User(sql.FieldGT(FieldFrozenBalance, v)) +} + +// FrozenBalanceGTE applies the GTE predicate on the "frozen_balance" field. +func FrozenBalanceGTE(v float64) predicate.User { + return predicate.User(sql.FieldGTE(FieldFrozenBalance, v)) +} + +// FrozenBalanceLT applies the LT predicate on the "frozen_balance" field. +func FrozenBalanceLT(v float64) predicate.User { + return predicate.User(sql.FieldLT(FieldFrozenBalance, v)) +} + +// FrozenBalanceLTE applies the LTE predicate on the "frozen_balance" field. +func FrozenBalanceLTE(v float64) predicate.User { + return predicate.User(sql.FieldLTE(FieldFrozenBalance, v)) +} + // ConcurrencyEQ applies the EQ predicate on the "concurrency" field. func ConcurrencyEQ(v int) predicate.User { return predicate.User(sql.FieldEQ(FieldConcurrency, v)) diff --git a/backend/ent/user_create.go b/backend/ent/user_create.go index 92f1bd5e07..b5bdf986a0 100644 --- a/backend/ent/user_create.go +++ b/backend/ent/user_create.go @@ -116,6 +116,20 @@ func (_c *UserCreate) SetNillableBalance(v *float64) *UserCreate { return _c } +// SetFrozenBalance sets the "frozen_balance" field. +func (_c *UserCreate) SetFrozenBalance(v float64) *UserCreate { + _c.mutation.SetFrozenBalance(v) + return _c +} + +// SetNillableFrozenBalance sets the "frozen_balance" field if the given value is not nil. +func (_c *UserCreate) SetNillableFrozenBalance(v *float64) *UserCreate { + if v != nil { + _c.SetFrozenBalance(*v) + } + return _c +} + // SetConcurrency sets the "concurrency" field. func (_c *UserCreate) SetConcurrency(v int) *UserCreate { _c.mutation.SetConcurrency(v) @@ -594,6 +608,10 @@ func (_c *UserCreate) defaults() error { v := user.DefaultBalance _c.mutation.SetBalance(v) } + if _, ok := _c.mutation.FrozenBalance(); !ok { + v := user.DefaultFrozenBalance + _c.mutation.SetFrozenBalance(v) + } if _, ok := _c.mutation.Concurrency(); !ok { v := user.DefaultConcurrency _c.mutation.SetConcurrency(v) @@ -676,6 +694,9 @@ func (_c *UserCreate) check() error { if _, ok := _c.mutation.Balance(); !ok { return &ValidationError{Name: "balance", err: errors.New(`ent: missing required field "User.balance"`)} } + if _, ok := _c.mutation.FrozenBalance(); !ok { + return &ValidationError{Name: "frozen_balance", err: errors.New(`ent: missing required field "User.frozen_balance"`)} + } if _, ok := _c.mutation.Concurrency(); !ok { return &ValidationError{Name: "concurrency", err: errors.New(`ent: missing required field "User.concurrency"`)} } @@ -779,6 +800,10 @@ func (_c *UserCreate) createSpec() (*User, *sqlgraph.CreateSpec) { _spec.SetField(user.FieldBalance, field.TypeFloat64, value) _node.Balance = value } + if value, ok := _c.mutation.FrozenBalance(); ok { + _spec.SetField(user.FieldFrozenBalance, field.TypeFloat64, value) + _node.FrozenBalance = value + } if value, ok := _c.mutation.Concurrency(); ok { _spec.SetField(user.FieldConcurrency, field.TypeInt, value) _node.Concurrency = value @@ -1191,6 +1216,24 @@ func (u *UserUpsert) AddBalance(v float64) *UserUpsert { return u } +// SetFrozenBalance sets the "frozen_balance" field. +func (u *UserUpsert) SetFrozenBalance(v float64) *UserUpsert { + u.Set(user.FieldFrozenBalance, v) + return u +} + +// UpdateFrozenBalance sets the "frozen_balance" field to the value that was provided on create. +func (u *UserUpsert) UpdateFrozenBalance() *UserUpsert { + u.SetExcluded(user.FieldFrozenBalance) + return u +} + +// AddFrozenBalance adds v to the "frozen_balance" field. +func (u *UserUpsert) AddFrozenBalance(v float64) *UserUpsert { + u.Add(user.FieldFrozenBalance, v) + return u +} + // SetConcurrency sets the "concurrency" field. func (u *UserUpsert) SetConcurrency(v int) *UserUpsert { u.Set(user.FieldConcurrency, v) @@ -1580,6 +1623,27 @@ func (u *UserUpsertOne) UpdateBalance() *UserUpsertOne { }) } +// SetFrozenBalance sets the "frozen_balance" field. +func (u *UserUpsertOne) SetFrozenBalance(v float64) *UserUpsertOne { + return u.Update(func(s *UserUpsert) { + s.SetFrozenBalance(v) + }) +} + +// AddFrozenBalance adds v to the "frozen_balance" field. +func (u *UserUpsertOne) AddFrozenBalance(v float64) *UserUpsertOne { + return u.Update(func(s *UserUpsert) { + s.AddFrozenBalance(v) + }) +} + +// UpdateFrozenBalance sets the "frozen_balance" field to the value that was provided on create. +func (u *UserUpsertOne) UpdateFrozenBalance() *UserUpsertOne { + return u.Update(func(s *UserUpsert) { + s.UpdateFrozenBalance() + }) +} + // SetConcurrency sets the "concurrency" field. func (u *UserUpsertOne) SetConcurrency(v int) *UserUpsertOne { return u.Update(func(s *UserUpsert) { @@ -2176,6 +2240,27 @@ func (u *UserUpsertBulk) UpdateBalance() *UserUpsertBulk { }) } +// SetFrozenBalance sets the "frozen_balance" field. +func (u *UserUpsertBulk) SetFrozenBalance(v float64) *UserUpsertBulk { + return u.Update(func(s *UserUpsert) { + s.SetFrozenBalance(v) + }) +} + +// AddFrozenBalance adds v to the "frozen_balance" field. +func (u *UserUpsertBulk) AddFrozenBalance(v float64) *UserUpsertBulk { + return u.Update(func(s *UserUpsert) { + s.AddFrozenBalance(v) + }) +} + +// UpdateFrozenBalance sets the "frozen_balance" field to the value that was provided on create. +func (u *UserUpsertBulk) UpdateFrozenBalance() *UserUpsertBulk { + return u.Update(func(s *UserUpsert) { + s.UpdateFrozenBalance() + }) +} + // SetConcurrency sets the "concurrency" field. func (u *UserUpsertBulk) SetConcurrency(v int) *UserUpsertBulk { return u.Update(func(s *UserUpsert) { diff --git a/backend/ent/user_update.go b/backend/ent/user_update.go index 67d3f8e6bb..6df9b320da 100644 --- a/backend/ent/user_update.go +++ b/backend/ent/user_update.go @@ -129,6 +129,27 @@ func (_u *UserUpdate) AddBalance(v float64) *UserUpdate { return _u } +// SetFrozenBalance sets the "frozen_balance" field. +func (_u *UserUpdate) SetFrozenBalance(v float64) *UserUpdate { + _u.mutation.ResetFrozenBalance() + _u.mutation.SetFrozenBalance(v) + return _u +} + +// SetNillableFrozenBalance sets the "frozen_balance" field if the given value is not nil. +func (_u *UserUpdate) SetNillableFrozenBalance(v *float64) *UserUpdate { + if v != nil { + _u.SetFrozenBalance(*v) + } + return _u +} + +// AddFrozenBalance adds value to the "frozen_balance" field. +func (_u *UserUpdate) AddFrozenBalance(v float64) *UserUpdate { + _u.mutation.AddFrozenBalance(v) + return _u +} + // SetConcurrency sets the "concurrency" field. func (_u *UserUpdate) SetConcurrency(v int) *UserUpdate { _u.mutation.ResetConcurrency() @@ -997,6 +1018,12 @@ func (_u *UserUpdate) sqlSave(ctx context.Context) (_node int, err error) { if value, ok := _u.mutation.AddedBalance(); ok { _spec.AddField(user.FieldBalance, field.TypeFloat64, value) } + if value, ok := _u.mutation.FrozenBalance(); ok { + _spec.SetField(user.FieldFrozenBalance, field.TypeFloat64, value) + } + if value, ok := _u.mutation.AddedFrozenBalance(); ok { + _spec.AddField(user.FieldFrozenBalance, field.TypeFloat64, value) + } if value, ok := _u.mutation.Concurrency(); ok { _spec.SetField(user.FieldConcurrency, field.TypeInt, value) } @@ -1778,6 +1805,27 @@ func (_u *UserUpdateOne) AddBalance(v float64) *UserUpdateOne { return _u } +// SetFrozenBalance sets the "frozen_balance" field. +func (_u *UserUpdateOne) SetFrozenBalance(v float64) *UserUpdateOne { + _u.mutation.ResetFrozenBalance() + _u.mutation.SetFrozenBalance(v) + return _u +} + +// SetNillableFrozenBalance sets the "frozen_balance" field if the given value is not nil. +func (_u *UserUpdateOne) SetNillableFrozenBalance(v *float64) *UserUpdateOne { + if v != nil { + _u.SetFrozenBalance(*v) + } + return _u +} + +// AddFrozenBalance adds value to the "frozen_balance" field. +func (_u *UserUpdateOne) AddFrozenBalance(v float64) *UserUpdateOne { + _u.mutation.AddFrozenBalance(v) + return _u +} + // SetConcurrency sets the "concurrency" field. func (_u *UserUpdateOne) SetConcurrency(v int) *UserUpdateOne { _u.mutation.ResetConcurrency() @@ -2676,6 +2724,12 @@ func (_u *UserUpdateOne) sqlSave(ctx context.Context) (_node *User, err error) { if value, ok := _u.mutation.AddedBalance(); ok { _spec.AddField(user.FieldBalance, field.TypeFloat64, value) } + if value, ok := _u.mutation.FrozenBalance(); ok { + _spec.SetField(user.FieldFrozenBalance, field.TypeFloat64, value) + } + if value, ok := _u.mutation.AddedFrozenBalance(); ok { + _spec.AddField(user.FieldFrozenBalance, field.TypeFloat64, value) + } if value, ok := _u.mutation.Concurrency(); ok { _spec.SetField(user.FieldConcurrency, field.TypeInt, value) } diff --git a/backend/internal/config/config.go b/backend/internal/config/config.go index 1f6d710d41..0e94c82527 100644 --- a/backend/internal/config/config.go +++ b/backend/internal/config/config.go @@ -29,7 +29,7 @@ const ( // DefaultCSPPolicy is the default Content-Security-Policy with nonce support // __CSP_NONCE__ will be replaced with actual nonce at request time by the SecurityHeaders middleware -const DefaultCSPPolicy = "default-src 'self'; script-src 'self' __CSP_NONCE__ https://challenges.cloudflare.com https://static.cloudflareinsights.com https://*.stripe.com https://static.airwallex.com https://checkout.airwallex.com https://static-demo.airwallex.com https://checkout-demo.airwallex.com; style-src 'self' 'unsafe-inline' https://fonts.googleapis.com https://static.airwallex.com https://checkout.airwallex.com https://static-demo.airwallex.com https://checkout-demo.airwallex.com; img-src 'self' data: https:; font-src 'self' data: https://fonts.gstatic.com; connect-src 'self' https:; frame-src https://challenges.cloudflare.com https://*.stripe.com https://checkout.airwallex.com https://checkout-demo.airwallex.com; frame-ancestors 'none'; base-uri 'self'; form-action 'self'" +const DefaultCSPPolicy = "default-src 'self'; script-src 'self' __CSP_NONCE__ https://challenges.cloudflare.com https://static.cloudflareinsights.com https://*.stripe.com https://static.airwallex.com https://checkout.airwallex.com https://static-demo.airwallex.com https://checkout-demo.airwallex.com; style-src 'self' 'unsafe-inline' https://fonts.googleapis.com https://static.airwallex.com https://checkout.airwallex.com https://static-demo.airwallex.com https://checkout-demo.airwallex.com; img-src 'self' data: blob: https:; font-src 'self' data: https://fonts.gstatic.com; connect-src 'self' https:; frame-src https://challenges.cloudflare.com https://*.stripe.com https://checkout.airwallex.com https://checkout-demo.airwallex.com; frame-ancestors 'none'; base-uri 'self'; form-action 'self'" // UMQ(用户消息队列)模式常量 const ( diff --git a/backend/internal/handler/admin/group_handler.go b/backend/internal/handler/admin/group_handler.go index 0a98ad6784..4595adeb24 100644 --- a/backend/internal/handler/admin/group_handler.go +++ b/backend/internal/handler/admin/group_handler.go @@ -93,8 +93,11 @@ type CreateGroupRequest struct { MonthlyLimitUSD optionalLimitField `json:"monthly_limit_usd"` // 图片生成计费配置(antigravity 和 gemini 平台使用,负数表示清除配置) AllowImageGeneration bool `json:"allow_image_generation"` + AllowBatchImageGeneration bool `json:"allow_batch_image_generation"` ImageRateIndependent bool `json:"image_rate_independent"` ImageRateMultiplier *float64 `json:"image_rate_multiplier"` + BatchImageDiscountMultiplier *float64 `json:"batch_image_discount_multiplier"` + BatchImageHoldMultiplier *float64 `json:"batch_image_hold_multiplier"` PeakRateEnabled bool `json:"peak_rate_enabled"` PeakStart string `json:"peak_start"` PeakEnd string `json:"peak_end"` @@ -138,8 +141,11 @@ type UpdateGroupRequest struct { MonthlyLimitUSD optionalLimitField `json:"monthly_limit_usd"` // 图片生成计费配置(antigravity 和 gemini 平台使用,负数表示清除配置) AllowImageGeneration *bool `json:"allow_image_generation"` + AllowBatchImageGeneration *bool `json:"allow_batch_image_generation"` ImageRateIndependent *bool `json:"image_rate_independent"` ImageRateMultiplier *float64 `json:"image_rate_multiplier"` + BatchImageDiscountMultiplier *float64 `json:"batch_image_discount_multiplier"` + BatchImageHoldMultiplier *float64 `json:"batch_image_hold_multiplier"` PeakRateEnabled *bool `json:"peak_rate_enabled"` PeakStart *string `json:"peak_start"` PeakEnd *string `json:"peak_end"` @@ -301,8 +307,11 @@ func (h *GroupHandler) Create(c *gin.Context) { WeeklyLimitUSD: req.WeeklyLimitUSD.ToServiceInput(), MonthlyLimitUSD: req.MonthlyLimitUSD.ToServiceInput(), AllowImageGeneration: req.AllowImageGeneration, + AllowBatchImageGeneration: req.AllowBatchImageGeneration, ImageRateIndependent: req.ImageRateIndependent, ImageRateMultiplier: req.ImageRateMultiplier, + BatchImageDiscountMultiplier: req.BatchImageDiscountMultiplier, + BatchImageHoldMultiplier: req.BatchImageHoldMultiplier, PeakRateEnabled: req.PeakRateEnabled, PeakStart: req.PeakStart, PeakEnd: req.PeakEnd, @@ -361,8 +370,11 @@ func (h *GroupHandler) Update(c *gin.Context) { WeeklyLimitUSD: req.WeeklyLimitUSD.ToServiceInput(), MonthlyLimitUSD: req.MonthlyLimitUSD.ToServiceInput(), AllowImageGeneration: req.AllowImageGeneration, + AllowBatchImageGeneration: req.AllowBatchImageGeneration, ImageRateIndependent: req.ImageRateIndependent, ImageRateMultiplier: req.ImageRateMultiplier, + BatchImageDiscountMultiplier: req.BatchImageDiscountMultiplier, + BatchImageHoldMultiplier: req.BatchImageHoldMultiplier, PeakRateEnabled: req.PeakRateEnabled, PeakStart: req.PeakStart, PeakEnd: req.PeakEnd, diff --git a/backend/internal/handler/batch_image_handler.go b/backend/internal/handler/batch_image_handler.go index 9452e6b7f4..22c719bcb3 100644 --- a/backend/internal/handler/batch_image_handler.go +++ b/backend/internal/handler/batch_image_handler.go @@ -5,6 +5,7 @@ import ( "io" "net/http" "strconv" + "strings" infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" "github.com/Wei-Shaw/sub2api/internal/server/middleware" @@ -56,6 +57,43 @@ func (h *BatchImageHandler) Get(c *gin.Context) { c.JSON(http.StatusOK, got) } +func (h *BatchImageHandler) List(c *gin.Context) { + owner, ok := batchImageOwnerFromContext(c) + if !ok { + batchImageError(c, infraerrors.New(http.StatusUnauthorized, "API_KEY_REQUIRED", "API key is required")) + return + } + limit, _ := strconv.Atoi(c.Query("limit")) + got, err := h.service.List(c.Request.Context(), owner, service.BatchImageJobsQuery{ + Status: c.Query("status"), + TaskName: c.Query("task_name"), + Downloaded: c.Query("downloaded"), + From: c.Query("from"), + To: c.Query("to"), + Limit: limit, + Cursor: c.Query("cursor"), + }) + if err != nil { + batchImageError(c, err) + return + } + c.JSON(http.StatusOK, got) +} + +func (h *BatchImageHandler) Models(c *gin.Context) { + owner, ok := batchImageOwnerFromContext(c) + if !ok { + batchImageError(c, infraerrors.New(http.StatusUnauthorized, "API_KEY_REQUIRED", "API key is required")) + return + } + got, err := h.service.ListModels(c.Request.Context(), owner) + if err != nil { + batchImageError(c, err) + return + } + c.JSON(http.StatusOK, got) +} + func (h *BatchImageHandler) Items(c *gin.Context) { owner, ok := batchImageOwnerFromContext(c) if !ok { @@ -122,6 +160,7 @@ func (h *BatchImageHandler) ItemContent(c *gin.Context) { if _, err := io.Copy(c.Writer, stream.Reader); err != nil { return } + _ = h.service.MarkDownloaded(c.Request.Context(), owner, c.Param("id")) } func (h *BatchImageHandler) Download(c *gin.Context) { @@ -147,6 +186,20 @@ func (h *BatchImageHandler) Download(c *gin.Context) { } return } + _ = h.service.MarkDownloaded(c.Request.Context(), owner, c.Param("id")) +} + +func (h *BatchImageHandler) DeleteRecord(c *gin.Context) { + owner, ok := batchImageOwnerFromContext(c) + if !ok { + batchImageError(c, infraerrors.New(http.StatusUnauthorized, "API_KEY_REQUIRED", "API key is required")) + return + } + if err := h.service.DeleteRecord(c.Request.Context(), owner, c.Param("id")); err != nil { + batchImageError(c, err) + return + } + c.Status(http.StatusNoContent) } func (h *BatchImageHandler) DeleteOutputs(c *gin.Context) { @@ -184,7 +237,7 @@ func batchImageError(c *gin.Context, err error) { code = "INTERNAL_ERROR" message = "internal error" } - if status == 0 || status == http.StatusInternalServerError { + if status == 0 || (status == http.StatusInternalServerError && strings.TrimSpace(code) == "") { status = http.StatusInternalServerError code = "INTERNAL_ERROR" message = "internal error" diff --git a/backend/internal/handler/dto/mappers.go b/backend/internal/handler/dto/mappers.go index 5bbab4d45f..7949b278cd 100644 --- a/backend/internal/handler/dto/mappers.go +++ b/backend/internal/handler/dto/mappers.go @@ -18,6 +18,7 @@ func UserFromServiceShallow(u *service.User) *User { Username: u.Username, Role: u.Role, Balance: u.Balance, + FrozenBalance: u.FrozenBalance, Concurrency: u.Concurrency, Status: u.Status, AllowedGroups: u.AllowedGroups, @@ -179,8 +180,11 @@ func groupFromServiceBase(g *service.Group) Group { WeeklyLimitUSD: g.WeeklyLimitUSD, MonthlyLimitUSD: g.MonthlyLimitUSD, AllowImageGeneration: g.AllowImageGeneration, + AllowBatchImageGeneration: g.AllowBatchImageGeneration, ImageRateIndependent: g.ImageRateIndependent, ImageRateMultiplier: g.ImageRateMultiplier, + BatchImageDiscountMultiplier: g.BatchImageDiscountMultiplier, + BatchImageHoldMultiplier: g.BatchImageHoldMultiplier, PeakRateEnabled: g.PeakRateEnabled, PeakStart: g.PeakStart, PeakEnd: g.PeakEnd, diff --git a/backend/internal/handler/dto/types.go b/backend/internal/handler/dto/types.go index b08dea5680..3c705ed4b2 100644 --- a/backend/internal/handler/dto/types.go +++ b/backend/internal/handler/dto/types.go @@ -14,6 +14,7 @@ type User struct { Username string `json:"username"` Role string `json:"role"` Balance float64 `json:"balance"` + FrozenBalance float64 `json:"frozen_balance"` Concurrency int `json:"concurrency"` Status string `json:"status"` AllowedGroups []int64 `json:"allowed_groups"` @@ -97,9 +98,12 @@ type Group struct { MonthlyLimitUSD *float64 `json:"monthly_limit_usd"` // 图片生成计费配置(仅 antigravity 平台使用) - AllowImageGeneration bool `json:"allow_image_generation"` - ImageRateIndependent bool `json:"image_rate_independent"` - ImageRateMultiplier float64 `json:"image_rate_multiplier"` + AllowImageGeneration bool `json:"allow_image_generation"` + AllowBatchImageGeneration bool `json:"allow_batch_image_generation"` + ImageRateIndependent bool `json:"image_rate_independent"` + ImageRateMultiplier float64 `json:"image_rate_multiplier"` + BatchImageDiscountMultiplier float64 `json:"batch_image_discount_multiplier"` + BatchImageHoldMultiplier float64 `json:"batch_image_hold_multiplier"` // 高峰时段倍率配置 PeakRateEnabled bool `json:"peak_rate_enabled"` PeakStart string `json:"peak_start"` diff --git a/backend/internal/repository/api_key_repo.go b/backend/internal/repository/api_key_repo.go index 76cad809c3..877fc90353 100644 --- a/backend/internal/repository/api_key_repo.go +++ b/backend/internal/repository/api_key_repo.go @@ -177,6 +177,7 @@ func (r *apiKeyRepository) GetByKeyForAuth(ctx context.Context, key string) (*se group.FieldWeeklyLimitUsd, group.FieldMonthlyLimitUsd, group.FieldAllowImageGeneration, + group.FieldAllowBatchImageGeneration, group.FieldImageRateIndependent, group.FieldImageRateMultiplier, group.FieldImagePrice1k, @@ -755,6 +756,7 @@ func userEntityToService(u *dbent.User) *service.User { PasswordHash: u.PasswordHash, Role: u.Role, Balance: u.Balance, + FrozenBalance: u.FrozenBalance, Concurrency: u.Concurrency, Status: u.Status, SignupSource: u.SignupSource, @@ -797,11 +799,14 @@ func groupEntityToService(g *dbent.Group) *service.Group { WeeklyLimitUSD: g.WeeklyLimitUsd, MonthlyLimitUSD: g.MonthlyLimitUsd, AllowImageGeneration: g.AllowImageGeneration, + AllowBatchImageGeneration: g.AllowBatchImageGeneration, ImageRateIndependent: g.ImageRateIndependent, ImageRateMultiplier: g.ImageRateMultiplier, ImagePrice1K: g.ImagePrice1k, ImagePrice2K: g.ImagePrice2k, ImagePrice4K: g.ImagePrice4k, + BatchImageDiscountMultiplier: g.BatchImageDiscountMultiplier, + BatchImageHoldMultiplier: g.BatchImageHoldMultiplier, DefaultValidityDays: g.DefaultValidityDays, ClaudeCodeOnly: g.ClaudeCodeOnly, FallbackGroupID: g.FallbackGroupID, diff --git a/backend/internal/repository/batch_image_repo.go b/backend/internal/repository/batch_image_repo.go index 88e88637ef..932633eb7b 100644 --- a/backend/internal/repository/batch_image_repo.go +++ b/backend/internal/repository/batch_image_repo.go @@ -74,13 +74,61 @@ func (r *batchImageRepository) GetBatchImageJobByIdempotencyKey(ctx context.Cont func (r *batchImageRepository) GetBatchImageJobByBatchIDForOwner(ctx context.Context, userID, apiKeyID int64, batchID string) (*service.BatchImageJob, error) { job, err := scanBatchImageJob(r.sql.QueryRowContext(ctx, batchImageJobSelectSQL+` - WHERE batch_id = $1 AND user_id = $2 AND api_key_id = $3`, batchID, userID, apiKeyID)) + WHERE batch_id = $1 AND user_id = $2 AND api_key_id = $3 AND user_deleted_at IS NULL`, batchID, userID, apiKeyID)) if err != nil { return nil, translatePersistenceError(err, service.ErrBatchImageJobNotFound, nil) } return job, nil } +func (r *batchImageRepository) ListBatchImageJobsForOwner(ctx context.Context, userID, apiKeyID int64, filter service.BatchImageJobFilter) ([]*service.BatchImageJob, error) { + limit := filter.Limit + if limit <= 0 || limit > 100 { + limit = 20 + } + if filter.Offset < 0 { + filter.Offset = 0 + } + + query := batchImageJobSelectSQL + " WHERE user_id = $1 AND api_key_id = $2" + args := []any{userID, apiKeyID} + if filter.ExcludeDeleted { + query += " AND user_deleted_at IS NULL" + } + if filter.Status != "" { + query += " AND status = $" + strconv.Itoa(len(args)+1) + args = append(args, filter.Status) + } + if filter.TaskNameLike != "" { + query += " AND task_name ILIKE $" + strconv.Itoa(len(args)+1) + args = append(args, "%"+filter.TaskNameLike+"%") + } + if filter.Downloaded != nil { + if *filter.Downloaded { + query += " AND downloaded_at IS NOT NULL" + } else { + query += " AND downloaded_at IS NULL" + } + } + if filter.CreatedAfter != nil { + query += " AND created_at >= $" + strconv.Itoa(len(args)+1) + args = append(args, *filter.CreatedAfter) + } + if filter.CreatedBefore != nil { + query += " AND created_at < $" + strconv.Itoa(len(args)+1) + args = append(args, *filter.CreatedBefore) + } + query += " ORDER BY created_at DESC, id DESC LIMIT $" + strconv.Itoa(len(args)+1) + " OFFSET $" + strconv.Itoa(len(args)+2) + args = append(args, limit, filter.Offset) + + rows, err := r.sql.QueryContext(ctx, query, args...) + if err != nil { + return nil, err + } + defer rows.Close() + return scanBatchImageJobs(rows) +} + func (r *batchImageRepository) GetBatchImageJobByID(ctx context.Context, id int64) (*service.BatchImageJob, error) { job, err := scanBatchImageJob(r.sql.QueryRowContext(ctx, batchImageJobSelectSQL+" WHERE id = $1", id)) if err != nil { @@ -286,16 +334,16 @@ func (r *batchImageRepository) transitionBatchImageJobStatusWithSQL(ctx context. if _, err := sqlq.ExecContext(ctx, ` UPDATE batch_image_jobs SET - status = $2, + status = $2::varchar, version = version + 1, updated_at = $3, - last_error_code = CASE WHEN $2 = 'failed' THEN $4 ELSE last_error_code END, - last_error_message = CASE WHEN $2 = 'failed' THEN $5 ELSE last_error_message END, - submitted_at = CASE WHEN $2 = 'submitted' AND submitted_at IS NULL THEN $3 ELSE submitted_at END, - started_at = CASE WHEN $2 = 'running' AND started_at IS NULL THEN $3 ELSE started_at END, - finished_at = CASE WHEN $2 IN ('completed', 'failed', 'cancelled') AND finished_at IS NULL THEN $3 ELSE finished_at END, - settled_at = CASE WHEN $2 = 'completed' AND settled_at IS NULL THEN $3 ELSE settled_at END, - output_deleted_at = CASE WHEN $2 = 'output_deleted' AND output_deleted_at IS NULL THEN $3 ELSE output_deleted_at END + last_error_code = CASE WHEN $2::varchar = 'failed' THEN $4 ELSE last_error_code END, + last_error_message = CASE WHEN $2::varchar = 'failed' THEN $5 ELSE last_error_message END, + submitted_at = CASE WHEN $2::varchar = 'submitted' AND submitted_at IS NULL THEN $3 ELSE submitted_at END, + started_at = CASE WHEN $2::varchar = 'running' AND started_at IS NULL THEN $3 ELSE started_at END, + finished_at = CASE WHEN $2::varchar IN ('completed', 'failed', 'cancelled') AND finished_at IS NULL THEN $3 ELSE finished_at END, + settled_at = CASE WHEN $2::varchar = 'completed' AND settled_at IS NULL THEN $3 ELSE settled_at END, + output_deleted_at = CASE WHEN $2::varchar = 'output_deleted' AND output_deleted_at IS NULL THEN $3 ELSE output_deleted_at END WHERE batch_id = $1`, batchID, toStatus, now, opts.ErrorCode, opts.ErrorMessage); err != nil { return err } @@ -367,16 +415,25 @@ func (r *batchImageRepository) replaceBatchImageItemsForJobWithSQL(ctx context.C if err := sqlq.QueryRowContext(ctx, `SELECT id FROM batch_image_jobs WHERE batch_id = $1 FOR UPDATE`, batchID).Scan(&id); err != nil { return translatePersistenceError(err, service.ErrBatchImageJobNotFound, nil) } + promptPreviews, err := r.batchImageItemPromptPreviews(ctx, sqlq, batchID) + if err != nil { + return err + } if _, err := sqlq.ExecContext(ctx, `DELETE FROM batch_image_items WHERE job_id = $1`, batchID); err != nil { return err } for _, item := range items { item.JobID = batchID + if item.PromptPreview == nil { + if preview := promptPreviews[item.CustomID]; preview != "" { + item.PromptPreview = &preview + } + } if _, err := createBatchImageItemWithSQL(ctx, sqlq, item); err != nil { return translatePersistenceError(err, nil, service.ErrBatchImageItemExists) } } - _, err := sqlq.ExecContext(ctx, ` + _, err = sqlq.ExecContext(ctx, ` UPDATE batch_image_jobs SET success_count = $2, fail_count = $3, @@ -385,6 +442,26 @@ WHERE batch_id = $1`, batchID, counts.SuccessCount, counts.FailCount, time.Now() return err } +func (r *batchImageRepository) batchImageItemPromptPreviews(ctx context.Context, sqlq batchImageSQLExecutor, batchID string) (map[string]string, error) { + rows, err := sqlq.QueryContext(ctx, `SELECT custom_id, prompt_preview FROM batch_image_items WHERE job_id = $1 AND prompt_preview IS NOT NULL`, batchID) + if err != nil { + return nil, err + } + defer rows.Close() + out := make(map[string]string) + for rows.Next() { + var customID string + var preview sql.NullString + if err := rows.Scan(&customID, &preview); err != nil { + return nil, err + } + if preview.Valid && preview.String != "" { + out[customID] = preview.String + } + } + return out, rows.Err() +} + func (r *batchImageRepository) ListBatchImageItems(ctx context.Context, batchID string, filter service.BatchImageItemFilter) ([]*service.BatchImageItem, error) { limit := filter.Limit if limit <= 0 || limit > 500 { @@ -484,6 +561,24 @@ func (r *batchImageRepository) ListBatchImageJobsDueForOutputCleanup(ctx context return scanBatchImageJobs(rows) } +func (r *batchImageRepository) ListStaleUnsubmittedBatchImageJobs(ctx context.Context, cutoff time.Time, limit int) ([]*service.BatchImageJob, error) { + if limit <= 0 || limit > 1000 { + limit = 100 + } + rows, err := r.sql.QueryContext(ctx, batchImageJobSelectSQL+` + WHERE status IN ('created', 'uploading') + AND provider_job_name IS NULL + AND COALESCE(hold_amount, estimated_cost, 0) > 0 + AND updated_at <= $1 + ORDER BY updated_at ASC, id ASC + LIMIT $2`, cutoff, limit) + if err != nil { + return nil, err + } + defer rows.Close() + return scanBatchImageJobs(rows) +} + func (r *batchImageRepository) MarkBatchImageInputDeleted(ctx context.Context, batchID string, deletedAt time.Time) error { res, err := r.sql.ExecContext(ctx, ` UPDATE batch_image_jobs @@ -526,6 +621,48 @@ WHERE batch_id = $1`, batchID, deletedAt) }) } +func (r *batchImageRepository) MarkBatchImageDownloaded(ctx context.Context, batchID string, downloadedAt time.Time) error { + res, err := r.sql.ExecContext(ctx, ` +UPDATE batch_image_jobs +SET downloaded_at = CASE WHEN downloaded_at IS NULL THEN $2 ELSE downloaded_at END, + updated_at = $2 +WHERE batch_id = $1`, batchID, downloadedAt) + if err != nil { + return err + } + if affected, err := res.RowsAffected(); err == nil && affected == 0 { + return service.ErrBatchImageJobNotFound + } + return appendBatchImageEventWithSQL(ctx, r.sql, batchID, "download_completed", map[string]any{ + "batch_id": batchID, + "downloaded_at": downloadedAt.UTC().Format(time.RFC3339), + }) +} + +func (r *batchImageRepository) MarkBatchImageJobUserDeleted(ctx context.Context, userID, apiKeyID int64, batchID string, deletedAt time.Time) error { + res, err := r.sql.ExecContext(ctx, ` +UPDATE batch_image_jobs +SET user_deleted_at = CASE WHEN user_deleted_at IS NULL THEN $4 ELSE user_deleted_at END, + updated_at = $4 +WHERE batch_id = $1 + AND user_id = $2 + AND api_key_id = $3 + AND user_deleted_at IS NULL + AND status IN ('completed', 'failed', 'cancelled', 'output_deleted')`, batchID, userID, apiKeyID, deletedAt) + if err != nil { + return err + } + if affected, err := res.RowsAffected(); err == nil && affected == 0 { + return service.ErrBatchImageRecordDeleteNotReady + } + return appendBatchImageEventWithSQL(ctx, r.sql, batchID, "user_record_deleted", map[string]any{ + "batch_id": batchID, + "deleted_at": deletedAt.UTC().Format(time.RFC3339), + "user_id": userID, + "api_key_id": apiKeyID, + }) +} + func (r *batchImageRepository) SetBatchImageOutputExpiresAt(ctx context.Context, batchID string, expiresAt time.Time) error { res, err := r.sql.ExecContext(ctx, ` UPDATE batch_image_jobs @@ -562,23 +699,35 @@ func (r *batchImageRepository) AppendBatchImageEvent(ctx context.Context, batchI func createBatchImageJobWithSQL(ctx context.Context, sqlq batchImageSQLExecutor, params service.CreateBatchImageJobParams) (*service.BatchImageJob, error) { return scanBatchImageJob(sqlq.QueryRowContext(ctx, ` INSERT INTO batch_image_jobs ( - batch_id, user_id, api_key_id, account_id, provider, model, status, + batch_id, user_id, api_key_id, account_id, provider, model, task_name, parent_batch_id, status, provider_job_name, provider_input_ref, provider_output_ref, gcs_input_uri, gcs_output_uri, item_count, success_count, fail_count, cancelled_count, - estimated_cost, hold_amount, actual_cost, currency, hold_id, + estimated_cost, hold_amount, actual_cost, + base_unit_price, group_rate_multiplier, account_rate_multiplier, + batch_discount_multiplier, hold_multiplier, billable_unit_price, hold_unit_price, + pricing_snapshot_version, + currency, hold_id, idempotency_key, request_hash, manifest_hash, retry_count, output_expires_at ) VALUES ( - $1, $2, $3, $4, $5, $6, $7, - $8, $9, $10, $11, $12, - $13, $14, $15, $16, - $17, $18, $19, $20, $21, - $22, $23, $24, $25, $26 + $1, $2, $3, $4, $5, $6, $7, $8, $9, + $10, $11, $12, $13, $14, + $15, $16, $17, $18, + $19, $20, $21, + $22, $23, $24, + $25, $26, $27, $28, + $29, + $30, $31, + $32, $33, $34, $35, $36 ) RETURNING `+batchImageJobColumns, - params.BatchID, params.UserID, params.APIKeyID, params.AccountID, params.Provider, params.Model, params.Status, + params.BatchID, params.UserID, params.APIKeyID, params.AccountID, params.Provider, params.Model, params.TaskName, params.ParentBatchID, params.Status, params.ProviderJobName, params.ProviderInputRef, params.ProviderOutputRef, params.GCSInputURI, params.GCSOutputURI, params.ItemCount, params.SuccessCount, params.FailCount, params.CancelledCount, - params.EstimatedCost, params.HoldAmount, params.ActualCost, params.Currency, params.HoldID, + params.EstimatedCost, params.HoldAmount, params.ActualCost, + params.BaseUnitPrice, params.GroupRateMultiplier, params.AccountRateMultiplier, + params.BatchDiscountMultiplier, params.HoldMultiplier, params.BillableUnitPrice, params.HoldUnitPrice, + params.PricingSnapshotVersion, + params.Currency, params.HoldID, params.IdempotencyKey, params.RequestHash, params.ManifestHash, params.RetryCount, params.OutputExpiresAt, )) } @@ -624,12 +773,16 @@ type rowScanner interface { } const batchImageJobColumns = ` -id, batch_id, user_id, api_key_id, account_id, provider, model, status, +id, batch_id, user_id, api_key_id, account_id, provider, model, task_name, parent_batch_id, status, provider_job_name, provider_input_ref, provider_output_ref, gcs_input_uri, gcs_output_uri, item_count, success_count, fail_count, cancelled_count, -estimated_cost, hold_amount, actual_cost, currency, hold_id, +estimated_cost, hold_amount, actual_cost, +base_unit_price, group_rate_multiplier, account_rate_multiplier, +batch_discount_multiplier, hold_multiplier, billable_unit_price, hold_unit_price, +pricing_snapshot_version, +currency, hold_id, idempotency_key, request_hash, manifest_hash, -retry_count, version, output_expires_at, input_deleted_at, output_deleted_at, +retry_count, version, output_expires_at, input_deleted_at, output_deleted_at, downloaded_at, user_deleted_at, last_error_code, last_error_message, created_at, updated_at, submitted_at, started_at, finished_at, settled_at` @@ -639,19 +792,24 @@ func scanBatchImageJob(row rowScanner) (*service.BatchImageJob, error) { var job service.BatchImageJob var apiKeyID, accountID sql.NullInt64 var providerJobName, providerInputRef, providerOutputRef, gcsInputURI, gcsOutputURI sql.NullString + var parentBatchID sql.NullString var holdAmount, actualCost sql.NullFloat64 var holdID, idempotencyKey, requestHash, manifestHash sql.NullString - var outputExpiresAt, inputDeletedAt, outputDeletedAt sql.NullTime + var outputExpiresAt, inputDeletedAt, outputDeletedAt, downloadedAt, userDeletedAt sql.NullTime var lastErrorCode, lastErrorMessage sql.NullString var submittedAt, startedAt, finishedAt, settledAt sql.NullTime err := row.Scan( - &job.ID, &job.BatchID, &job.UserID, &apiKeyID, &accountID, &job.Provider, &job.Model, &job.Status, + &job.ID, &job.BatchID, &job.UserID, &apiKeyID, &accountID, &job.Provider, &job.Model, &job.TaskName, &parentBatchID, &job.Status, &providerJobName, &providerInputRef, &providerOutputRef, &gcsInputURI, &gcsOutputURI, &job.ItemCount, &job.SuccessCount, &job.FailCount, &job.CancelledCount, - &job.EstimatedCost, &holdAmount, &actualCost, &job.Currency, &holdID, + &job.EstimatedCost, &holdAmount, &actualCost, + &job.BaseUnitPrice, &job.GroupRateMultiplier, &job.AccountRateMultiplier, + &job.BatchDiscountMultiplier, &job.HoldMultiplier, &job.BillableUnitPrice, &job.HoldUnitPrice, + &job.PricingSnapshotVersion, + &job.Currency, &holdID, &idempotencyKey, &requestHash, &manifestHash, - &job.RetryCount, &job.Version, &outputExpiresAt, &inputDeletedAt, &outputDeletedAt, + &job.RetryCount, &job.Version, &outputExpiresAt, &inputDeletedAt, &outputDeletedAt, &downloadedAt, &userDeletedAt, &lastErrorCode, &lastErrorMessage, &job.CreatedAt, &job.UpdatedAt, &submittedAt, &startedAt, &finishedAt, &settledAt, ) @@ -664,6 +822,7 @@ func scanBatchImageJob(row rowScanner) (*service.BatchImageJob, error) { job.ProviderJobName = batchImageNullStringPtr(providerJobName) job.ProviderInputRef = batchImageNullStringPtr(providerInputRef) job.ProviderOutputRef = batchImageNullStringPtr(providerOutputRef) + job.ParentBatchID = batchImageNullStringPtr(parentBatchID) job.GCSInputURI = batchImageNullStringPtr(gcsInputURI) job.GCSOutputURI = batchImageNullStringPtr(gcsOutputURI) job.HoldAmount = batchImageNullFloat64Ptr(holdAmount) @@ -675,6 +834,8 @@ func scanBatchImageJob(row rowScanner) (*service.BatchImageJob, error) { job.OutputExpiresAt = batchImageNullTimePtr(outputExpiresAt) job.InputDeletedAt = batchImageNullTimePtr(inputDeletedAt) job.OutputDeletedAt = batchImageNullTimePtr(outputDeletedAt) + job.DownloadedAt = batchImageNullTimePtr(downloadedAt) + job.UserDeletedAt = batchImageNullTimePtr(userDeletedAt) job.LastErrorCode = batchImageNullStringPtr(lastErrorCode) job.LastErrorMessage = batchImageNullStringPtr(lastErrorMessage) job.SubmittedAt = batchImageNullTimePtr(submittedAt) diff --git a/backend/internal/repository/group_repo.go b/backend/internal/repository/group_repo.go index 4e839b6a12..cb4437cf56 100644 --- a/backend/internal/repository/group_repo.go +++ b/backend/internal/repository/group_repo.go @@ -50,11 +50,14 @@ func (r *groupRepository) Create(ctx context.Context, groupIn *service.Group) er SetNillableWeeklyLimitUsd(groupIn.WeeklyLimitUSD). SetNillableMonthlyLimitUsd(groupIn.MonthlyLimitUSD). SetAllowImageGeneration(groupIn.AllowImageGeneration). + SetAllowBatchImageGeneration(groupIn.AllowBatchImageGeneration). SetImageRateIndependent(groupIn.ImageRateIndependent). SetImageRateMultiplier(groupIn.ImageRateMultiplier). SetNillableImagePrice1k(groupIn.ImagePrice1K). SetNillableImagePrice2k(groupIn.ImagePrice2K). SetNillableImagePrice4k(groupIn.ImagePrice4K). + SetBatchImageDiscountMultiplier(groupIn.BatchImageDiscountMultiplier). + SetBatchImageHoldMultiplier(groupIn.BatchImageHoldMultiplier). SetDefaultValidityDays(groupIn.DefaultValidityDays). SetClaudeCodeOnly(groupIn.ClaudeCodeOnly). SetNillableFallbackGroupID(groupIn.FallbackGroupID). @@ -132,11 +135,14 @@ func (r *groupRepository) Update(ctx context.Context, groupIn *service.Group) er SetNillableWeeklyLimitUsd(groupIn.WeeklyLimitUSD). SetNillableMonthlyLimitUsd(groupIn.MonthlyLimitUSD). SetAllowImageGeneration(groupIn.AllowImageGeneration). + SetAllowBatchImageGeneration(groupIn.AllowBatchImageGeneration). SetImageRateIndependent(groupIn.ImageRateIndependent). SetImageRateMultiplier(groupIn.ImageRateMultiplier). SetNillableImagePrice1k(groupIn.ImagePrice1K). SetNillableImagePrice2k(groupIn.ImagePrice2K). SetNillableImagePrice4k(groupIn.ImagePrice4K). + SetBatchImageDiscountMultiplier(groupIn.BatchImageDiscountMultiplier). + SetBatchImageHoldMultiplier(groupIn.BatchImageHoldMultiplier). SetDefaultValidityDays(groupIn.DefaultValidityDays). SetClaudeCodeOnly(groupIn.ClaudeCodeOnly). SetModelRoutingEnabled(groupIn.ModelRoutingEnabled). diff --git a/backend/internal/repository/migrations_runner.go b/backend/internal/repository/migrations_runner.go index 285326537d..7c045fea74 100644 --- a/backend/internal/repository/migrations_runner.go +++ b/backend/internal/repository/migrations_runner.go @@ -77,6 +77,8 @@ var migrationChecksumCompatibilityRules = map[string]migrationChecksumCompatibil "119_enforce_payment_orders_out_trade_no_unique.sql": newMigrationChecksumCompatibilityRule("0bbe809ae48a9d811dabda1ba1c74955bd71c4a9cc610f9128816818dfa6c11e", "ebd2c67cce0116393fb4f1b5d5116a67c6aceb73820dfb5133d1ff6f36d72d34"), "120_enforce_payment_orders_out_trade_no_unique_notx.sql": newMigrationChecksumCompatibilityRule("34aadc0db59a4e390f92a12b73bd74642d9724f33124f73638ae00089ea5e074", "e77921f79d539bc24575cb9c16cbe566d2b23ce816190343d0a7568f6a3fcf61", "707431450603e70a43ce9fbd61e0c12fa67da4875158ccefabacea069587ab22", "04b082b5a239c525154fe9185d324ee2b05ff90da9297e10dba19f9be79aa59a"), "123_fix_legacy_auth_source_grant_on_signup_defaults.sql": newMigrationChecksumCompatibilityRule("2ce43c2cd89e9f9e1febd34a407ed9e84d177386c5544b6f02c1f58a21129f57", "6cd33422f215dcd1f486ab6f35c0ea5805d9ca69bb25906d94bc649156657145"), + "159_batch_image_foundation.sql": newMigrationChecksumCompatibilityRule("d902b70982025ec519749faf058aab7631e82c3f48167b9a4ae4db718eb72cce", "82da85b5d98e67a0507647b873a40373e84538e4adafdeed6767c0ac8b6570b2"), + "161_batch_image_pricing_snapshot.sql": newMigrationChecksumCompatibilityRule("4012af3e43636cb6af22e0176d59d1fcc70615c0f310194329461ae462c4fbd6", "96d915c9b7a6941ae99039e0ff3f1a61481eb9bddd933d11c6fadb2274554e87"), } // ApplyMigrations 将嵌入的 SQL 迁移文件应用到指定的数据库。 diff --git a/backend/internal/repository/usage_billing_repo.go b/backend/internal/repository/usage_billing_repo.go index 91ac536eee..f7e675439f 100644 --- a/backend/internal/repository/usage_billing_repo.go +++ b/backend/internal/repository/usage_billing_repo.go @@ -63,23 +63,27 @@ func (r *usageBillingRepository) Apply(ctx context.Context, cmd *service.UsageBi } func (r *usageBillingRepository) claimUsageBillingKey(ctx context.Context, tx *sql.Tx, cmd *service.UsageBillingCommand) (bool, error) { + return r.claimUsageBillingRequest(ctx, tx, cmd.RequestID, cmd.APIKeyID, cmd.RequestFingerprint) +} + +func (r *usageBillingRepository) claimUsageBillingRequest(ctx context.Context, tx *sql.Tx, requestID string, apiKeyID int64, requestFingerprint string) (bool, error) { var id int64 err := tx.QueryRowContext(ctx, ` INSERT INTO usage_billing_dedup (request_id, api_key_id, request_fingerprint) VALUES ($1, $2, $3) ON CONFLICT (request_id, api_key_id) DO NOTHING RETURNING id - `, cmd.RequestID, cmd.APIKeyID, cmd.RequestFingerprint).Scan(&id) + `, requestID, apiKeyID, requestFingerprint).Scan(&id) if errors.Is(err, sql.ErrNoRows) { var existingFingerprint string if err := tx.QueryRowContext(ctx, ` SELECT request_fingerprint FROM usage_billing_dedup WHERE request_id = $1 AND api_key_id = $2 - `, cmd.RequestID, cmd.APIKeyID).Scan(&existingFingerprint); err != nil { + `, requestID, apiKeyID).Scan(&existingFingerprint); err != nil { return false, err } - if strings.TrimSpace(existingFingerprint) != strings.TrimSpace(cmd.RequestFingerprint) { + if strings.TrimSpace(existingFingerprint) != strings.TrimSpace(requestFingerprint) { return false, service.ErrUsageBillingRequestConflict } return false, nil @@ -92,9 +96,9 @@ func (r *usageBillingRepository) claimUsageBillingKey(ctx context.Context, tx *s SELECT request_fingerprint FROM usage_billing_dedup_archive WHERE request_id = $1 AND api_key_id = $2 - `, cmd.RequestID, cmd.APIKeyID).Scan(&archivedFingerprint) + `, requestID, apiKeyID).Scan(&archivedFingerprint) if err == nil { - if strings.TrimSpace(archivedFingerprint) != strings.TrimSpace(cmd.RequestFingerprint) { + if strings.TrimSpace(archivedFingerprint) != strings.TrimSpace(requestFingerprint) { return false, service.ErrUsageBillingRequestConflict } return false, nil @@ -105,6 +109,68 @@ func (r *usageBillingRepository) claimUsageBillingKey(ctx context.Context, tx *s return true, nil } +func (r *usageBillingRepository) ReserveBatchImageBalance(ctx context.Context, cmd *service.BatchImageBalanceHoldCommand) (*service.BatchImageBalanceHoldResult, error) { + return r.applyBatchImageBalanceHold(ctx, cmd, reserveUsageBillingBatchImageBalance) +} + +func (r *usageBillingRepository) CaptureBatchImageBalance(ctx context.Context, cmd *service.BatchImageBalanceHoldCommand) (*service.BatchImageBalanceHoldResult, error) { + return r.applyBatchImageBalanceHold(ctx, cmd, captureUsageBillingBatchImageBalance) +} + +func (r *usageBillingRepository) ReleaseBatchImageBalance(ctx context.Context, cmd *service.BatchImageBalanceHoldCommand) (*service.BatchImageBalanceHoldResult, error) { + return r.applyBatchImageBalanceHold(ctx, cmd, releaseUsageBillingBatchImageBalance) +} + +func (r *usageBillingRepository) applyBatchImageBalanceHold( + ctx context.Context, + cmd *service.BatchImageBalanceHoldCommand, + apply func(context.Context, *sql.Tx, *service.BatchImageBalanceHoldCommand) (*service.BatchImageBalanceHoldResult, error), +) (_ *service.BatchImageBalanceHoldResult, err error) { + if cmd == nil { + return &service.BatchImageBalanceHoldResult{}, nil + } + if r == nil || r.db == nil { + return nil, errors.New("usage billing repository db is nil") + } + cmd.Normalize() + if cmd.RequestID == "" { + return nil, service.ErrUsageBillingRequestIDRequired + } + + tx, err := r.db.BeginTx(ctx, nil) + if err != nil { + return nil, err + } + defer func() { + if tx != nil { + _ = tx.Rollback() + } + }() + + applied, err := r.claimUsageBillingRequest(ctx, tx, cmd.RequestID, cmd.APIKeyID, cmd.RequestFingerprint) + if err != nil { + return nil, err + } + if !applied { + return &service.BatchImageBalanceHoldResult{Applied: false}, nil + } + + result, err := apply(ctx, tx, cmd) + if err != nil { + return nil, err + } + if result == nil { + result = &service.BatchImageBalanceHoldResult{} + } + result.Applied = true + + if err := tx.Commit(); err != nil { + return nil, err + } + tx = nil + return result, nil +} + func (r *usageBillingRepository) applyUsageBillingEffects(ctx context.Context, tx *sql.Tx, cmd *service.UsageBillingCommand, result *service.UsageBillingApplyResult) error { if cmd.SubscriptionCost > 0 && cmd.SubscriptionID != nil { if err := incrementUsageBillingSubscription(ctx, tx, *cmd.SubscriptionID, cmd.SubscriptionCost); err != nil { @@ -206,6 +272,108 @@ func deductUsageBillingBalance(ctx context.Context, tx *sql.Tx, userID int64, am return newBalance, false, nil } +func reserveUsageBillingBatchImageBalance(ctx context.Context, tx *sql.Tx, cmd *service.BatchImageBalanceHoldCommand) (*service.BatchImageBalanceHoldResult, error) { + if cmd.HoldAmount <= 0 { + return &service.BatchImageBalanceHoldResult{}, nil + } + var balance, frozen float64 + err := tx.QueryRowContext(ctx, ` + UPDATE users + SET balance = balance - $1, + frozen_balance = COALESCE(frozen_balance, 0) + $1, + updated_at = NOW() + WHERE id = $2 AND deleted_at IS NULL AND balance >= $1 + RETURNING balance, frozen_balance + `, cmd.HoldAmount, cmd.UserID).Scan(&balance, &frozen) + if err == nil { + return &service.BatchImageBalanceHoldResult{NewBalance: &balance, FrozenBalance: &frozen}, nil + } + if !errors.Is(err, sql.ErrNoRows) { + return nil, err + } + if exists, existsErr := userExistsForBilling(ctx, tx, cmd.UserID); existsErr != nil { + return nil, existsErr + } else if !exists { + return nil, service.ErrUserNotFound + } + return nil, service.ErrBatchImageInsufficientBalance +} + +func captureUsageBillingBatchImageBalance(ctx context.Context, tx *sql.Tx, cmd *service.BatchImageBalanceHoldCommand) (*service.BatchImageBalanceHoldResult, error) { + if cmd.HoldAmount <= 0 && cmd.ActualAmount <= 0 { + return &service.BatchImageBalanceHoldResult{}, nil + } + if cmd.ActualAmount-cmd.HoldAmount > 0.00000001 { + return nil, service.ErrBatchImageSettlementCostExceedsHold + } + var balance, frozen float64 + err := tx.QueryRowContext(ctx, ` + UPDATE users + SET balance = balance + + CASE WHEN $1 > $2 THEN $1 - $2 ELSE 0 END + - CASE WHEN $2 > $1 THEN $2 - $1 ELSE 0 END, + frozen_balance = COALESCE(frozen_balance, 0) - $1, + updated_at = NOW() + WHERE id = $3 AND deleted_at IS NULL AND COALESCE(frozen_balance, 0) >= $1 + RETURNING balance, frozen_balance + `, cmd.HoldAmount, cmd.ActualAmount, cmd.UserID).Scan(&balance, &frozen) + if err == nil { + return &service.BatchImageBalanceHoldResult{NewBalance: &balance, FrozenBalance: &frozen}, nil + } + if !errors.Is(err, sql.ErrNoRows) { + return nil, err + } + if exists, existsErr := userExistsForBilling(ctx, tx, cmd.UserID); existsErr != nil { + return nil, existsErr + } else if !exists { + return nil, service.ErrUserNotFound + } + return nil, errors.New("batch image frozen balance is insufficient") +} + +func releaseUsageBillingBatchImageBalance(ctx context.Context, tx *sql.Tx, cmd *service.BatchImageBalanceHoldCommand) (*service.BatchImageBalanceHoldResult, error) { + if cmd.HoldAmount <= 0 { + return &service.BatchImageBalanceHoldResult{}, nil + } + var balance, frozen float64 + err := tx.QueryRowContext(ctx, ` + UPDATE users + SET balance = balance + $1, + frozen_balance = COALESCE(frozen_balance, 0) - $1, + updated_at = NOW() + WHERE id = $2 AND deleted_at IS NULL AND COALESCE(frozen_balance, 0) >= $1 + RETURNING balance, frozen_balance + `, cmd.HoldAmount, cmd.UserID).Scan(&balance, &frozen) + if err == nil { + return &service.BatchImageBalanceHoldResult{NewBalance: &balance, FrozenBalance: &frozen}, nil + } + if !errors.Is(err, sql.ErrNoRows) { + return nil, err + } + if exists, existsErr := userExistsForBilling(ctx, tx, cmd.UserID); existsErr != nil { + return nil, existsErr + } else if !exists { + return nil, service.ErrUserNotFound + } + return nil, errors.New("batch image frozen balance is insufficient") +} + +func userExistsForBilling(ctx context.Context, tx *sql.Tx, userID int64) (bool, error) { + var exists int + err := tx.QueryRowContext(ctx, ` + SELECT 1 + FROM users + WHERE id = $1 AND deleted_at IS NULL + `, userID).Scan(&exists) + if errors.Is(err, sql.ErrNoRows) { + return false, nil + } + if err != nil { + return false, err + } + return true, nil +} + func incrementUsageBillingAPIKeyQuota(ctx context.Context, tx *sql.Tx, apiKeyID int64, amount float64) (bool, error) { var exhausted bool err := tx.QueryRowContext(ctx, ` diff --git a/backend/internal/repository/usage_billing_repo_unit_test.go b/backend/internal/repository/usage_billing_repo_unit_test.go index 8ed5530a8f..0c469db899 100644 --- a/backend/internal/repository/usage_billing_repo_unit_test.go +++ b/backend/internal/repository/usage_billing_repo_unit_test.go @@ -16,6 +16,10 @@ import ( const ( conditionalBalanceDeductSQL = `(?s)UPDATE users\s+SET balance = balance - \$1,\s+updated_at = NOW\(\)\s+WHERE id = \$2 AND deleted_at IS NULL AND balance >= \$1\s+RETURNING balance` overdraftBalanceDeductSQL = `(?s)UPDATE users\s+SET balance = balance - \$1,\s+updated_at = NOW\(\)\s+WHERE id = \$2 AND deleted_at IS NULL\s+RETURNING balance` + reserveBatchImageHoldSQL = `(?s)UPDATE users\s+SET balance = balance - \$1,\s+frozen_balance = COALESCE\(frozen_balance, 0\) \+ \$1,\s+updated_at = NOW\(\)\s+WHERE id = \$2 AND deleted_at IS NULL AND balance >= \$1\s+RETURNING balance, frozen_balance` + captureBatchImageHoldSQL = `(?s)UPDATE users\s+SET balance = balance\s+\+ CASE WHEN \$1 > \$2 THEN \$1 - \$2 ELSE 0 END\s+- CASE WHEN \$2 > \$1 THEN \$2 - \$1 ELSE 0 END,\s+frozen_balance = COALESCE\(frozen_balance, 0\) - \$1,\s+updated_at = NOW\(\)\s+WHERE id = \$3 AND deleted_at IS NULL AND COALESCE\(frozen_balance, 0\) >= \$1\s+RETURNING balance, frozen_balance` + releaseBatchImageHoldSQL = `(?s)UPDATE users\s+SET balance = balance \+ \$1,\s+frozen_balance = COALESCE\(frozen_balance, 0\) - \$1,\s+updated_at = NOW\(\)\s+WHERE id = \$2 AND deleted_at IS NULL AND COALESCE\(frozen_balance, 0\) >= \$1\s+RETURNING balance, frozen_balance` + userExistsForBillingSQL = `(?s)SELECT 1\s+FROM users\s+WHERE id = \$1 AND deleted_at IS NULL` ) func TestDeductUsageBillingBalance_UsesSufficientBalanceGuard(t *testing.T) { @@ -117,3 +121,111 @@ func TestDeductUsageBillingBalance_ReturnsUserNotFoundWhenNoUserUpdated(t *testi require.NoError(t, tx.Rollback()) require.NoError(t, mock.ExpectationsWereMet()) } + +func TestReserveUsageBillingBatchImageBalance_MovesAvailableToFrozen(t *testing.T) { + ctx := context.Background() + db, mock, err := sqlmock.New() + require.NoError(t, err) + defer func() { _ = db.Close() }() + + mock.ExpectBegin() + tx, err := db.BeginTx(ctx, nil) + require.NoError(t, err) + mock.ExpectQuery(reserveBatchImageHoldSQL). + WithArgs(2.5, int64(42)). + WillReturnRows(sqlmock.NewRows([]string{"balance", "frozen_balance"}).AddRow(7.5, 2.5)) + mock.ExpectCommit() + + result, err := reserveUsageBillingBatchImageBalance(ctx, tx, &service.BatchImageBalanceHoldCommand{UserID: 42, HoldAmount: 2.5}) + require.NoError(t, err) + require.NotNil(t, result.NewBalance) + require.NotNil(t, result.FrozenBalance) + require.InDelta(t, 7.5, *result.NewBalance, 0.000001) + require.InDelta(t, 2.5, *result.FrozenBalance, 0.000001) + require.NoError(t, tx.Commit()) + require.NoError(t, mock.ExpectationsWereMet()) +} + +func TestReserveUsageBillingBatchImageBalance_InsufficientBalance(t *testing.T) { + ctx := context.Background() + db, mock, err := sqlmock.New() + require.NoError(t, err) + defer func() { _ = db.Close() }() + + mock.ExpectBegin() + tx, err := db.BeginTx(ctx, nil) + require.NoError(t, err) + mock.ExpectQuery(reserveBatchImageHoldSQL). + WithArgs(10.0, int64(42)). + WillReturnError(sql.ErrNoRows) + mock.ExpectQuery(userExistsForBillingSQL). + WithArgs(int64(42)). + WillReturnRows(sqlmock.NewRows([]string{"?column?"}).AddRow(1)) + mock.ExpectRollback() + + _, err = reserveUsageBillingBatchImageBalance(ctx, tx, &service.BatchImageBalanceHoldCommand{UserID: 42, HoldAmount: 10}) + require.ErrorIs(t, err, service.ErrBatchImageInsufficientBalance) + require.NoError(t, tx.Rollback()) + require.NoError(t, mock.ExpectationsWereMet()) +} + +func TestCaptureUsageBillingBatchImageBalance_ReleasesRemainder(t *testing.T) { + ctx := context.Background() + db, mock, err := sqlmock.New() + require.NoError(t, err) + defer func() { _ = db.Close() }() + + mock.ExpectBegin() + tx, err := db.BeginTx(ctx, nil) + require.NoError(t, err) + mock.ExpectQuery(captureBatchImageHoldSQL). + WithArgs(1.0, 0.25, int64(42)). + WillReturnRows(sqlmock.NewRows([]string{"balance", "frozen_balance"}).AddRow(9.75, 0.0)) + mock.ExpectCommit() + + result, err := captureUsageBillingBatchImageBalance(ctx, tx, &service.BatchImageBalanceHoldCommand{UserID: 42, HoldAmount: 1, ActualAmount: 0.25}) + require.NoError(t, err) + require.InDelta(t, 9.75, *result.NewBalance, 0.000001) + require.InDelta(t, 0.0, *result.FrozenBalance, 0.000001) + require.NoError(t, tx.Commit()) + require.NoError(t, mock.ExpectationsWereMet()) +} + +func TestCaptureUsageBillingBatchImageBalance_RejectsActualCostOverHold(t *testing.T) { + ctx := context.Background() + db, mock, err := sqlmock.New() + require.NoError(t, err) + defer func() { _ = db.Close() }() + + mock.ExpectBegin() + tx, err := db.BeginTx(ctx, nil) + require.NoError(t, err) + mock.ExpectRollback() + + _, err = captureUsageBillingBatchImageBalance(ctx, tx, &service.BatchImageBalanceHoldCommand{UserID: 42, HoldAmount: 0.5, ActualAmount: 1}) + require.ErrorIs(t, err, service.ErrBatchImageSettlementCostExceedsHold) + require.NoError(t, tx.Rollback()) + require.NoError(t, mock.ExpectationsWereMet()) +} + +func TestReleaseUsageBillingBatchImageBalance_ReturnsFrozenToAvailable(t *testing.T) { + ctx := context.Background() + db, mock, err := sqlmock.New() + require.NoError(t, err) + defer func() { _ = db.Close() }() + + mock.ExpectBegin() + tx, err := db.BeginTx(ctx, nil) + require.NoError(t, err) + mock.ExpectQuery(releaseBatchImageHoldSQL). + WithArgs(1.0, int64(42)). + WillReturnRows(sqlmock.NewRows([]string{"balance", "frozen_balance"}).AddRow(10.0, 0.0)) + mock.ExpectCommit() + + result, err := releaseUsageBillingBatchImageBalance(ctx, tx, &service.BatchImageBalanceHoldCommand{UserID: 42, HoldAmount: 1}) + require.NoError(t, err) + require.InDelta(t, 10.0, *result.NewBalance, 0.000001) + require.InDelta(t, 0.0, *result.FrozenBalance, 0.000001) + require.NoError(t, tx.Commit()) + require.NoError(t, mock.ExpectationsWereMet()) +} diff --git a/backend/internal/server/api_contract_test.go b/backend/internal/server/api_contract_test.go index 26f976c87b..fcf240b19d 100644 --- a/backend/internal/server/api_contract_test.go +++ b/backend/internal/server/api_contract_test.go @@ -52,9 +52,10 @@ func TestAPIContracts(t *testing.T) { "email": "alice@example.com", "email_bound": true, "username": "alice", - "role": "user", - "balance": 12.5, - "concurrency": 5, + "role": "user", + "balance": 12.5, + "frozen_balance": 0, + "concurrency": 5, "rpm_limit": 0, "status": "active", "allowed_groups": null, @@ -359,6 +360,9 @@ func TestAPIContracts(t *testing.T) { "image_price_2k": null, "image_price_4k": null, "allow_image_generation": false, + "allow_batch_image_generation": false, + "batch_image_discount_multiplier": 0, + "batch_image_hold_multiplier": 0, "image_rate_independent": false, "image_rate_multiplier": 0, "claude_code_only": false, diff --git a/backend/internal/server/middleware/api_key_auth.go b/backend/internal/server/middleware/api_key_auth.go index 2f0a3f1cf7..9bcb56d7fa 100644 --- a/backend/internal/server/middleware/api_key_auth.go +++ b/backend/internal/server/middleware/api_key_auth.go @@ -213,7 +213,7 @@ func apiKeyAuthWithSubscription(apiKeyService *service.APIKeyService, subscripti } } else { // 非订阅模式 或 订阅模式但 subscriptionService 未注入:回退到余额检查 - if apiKey.User.Balance <= 0 { + if apiKeyBalanceBelowAuthThreshold(apiKey.User.Balance, cfg) { AbortWithError(c, 403, "INSUFFICIENT_BALANCE", "Insufficient account balance") return } @@ -289,6 +289,16 @@ func setGroupContext(c *gin.Context, group *service.Group) { c.Request = c.Request.WithContext(ctx) } +func apiKeyBalanceBelowAuthThreshold(balance float64, cfg *config.Config) bool { + if balance <= 0 { + return true + } + if cfg == nil || cfg.Billing.MinimumBalanceReserve <= 0 { + return false + } + return balance < cfg.Billing.MinimumBalanceReserve +} + func abortIfAPIKeyGroupUnavailable(c *gin.Context, apiKey *service.APIKey) bool { code, message, ok := validateAPIKeyGroupAvailable(apiKey) if ok { diff --git a/backend/internal/server/middleware/api_key_auth_google.go b/backend/internal/server/middleware/api_key_auth_google.go index 97f3936c0c..5c5ee147a4 100644 --- a/backend/internal/server/middleware/api_key_auth_google.go +++ b/backend/internal/server/middleware/api_key_auth_google.go @@ -109,7 +109,7 @@ func APIKeyAuthWithSubscriptionGoogle(apiKeyService *service.APIKeyService, subs subscriptionService.DoWindowMaintenance(&maintenanceCopy) } } else { - if apiKey.User.Balance <= 0 { + if apiKeyBalanceBelowAuthThreshold(apiKey.User.Balance, cfg) { abortWithGoogleError(c, 403, "Insufficient account balance") return } diff --git a/backend/internal/server/middleware/api_key_auth_google_test.go b/backend/internal/server/middleware/api_key_auth_google_test.go index bf3909fcd4..899cd8bbe9 100644 --- a/backend/internal/server/middleware/api_key_auth_google_test.go +++ b/backend/internal/server/middleware/api_key_auth_google_test.go @@ -539,6 +539,42 @@ func TestApiKeyAuthWithSubscriptionGoogle_InsufficientBalance(t *testing.T) { require.Equal(t, "PERMISSION_DENIED", resp.Error.Status) } +func TestApiKeyAuthWithSubscriptionGoogle_BalanceBelowMinimumReserve(t *testing.T) { + gin.SetMode(gin.TestMode) + + r := gin.New() + apiKeyService := newTestAPIKeyService(fakeAPIKeyRepo{ + getByKey: func(ctx context.Context, key string) (*service.APIKey, error) { + return &service.APIKey{ + ID: 1, + Key: key, + Status: service.StatusActive, + User: &service.User{ + ID: 123, + Status: service.StatusActive, + Balance: 0.005, + }, + }, nil + }, + }) + cfg := &config.Config{} + cfg.Billing.MinimumBalanceReserve = 0.01 + r.Use(APIKeyAuthWithSubscriptionGoogle(apiKeyService, nil, cfg)) + r.GET("/v1beta/test", func(c *gin.Context) { c.JSON(200, gin.H{"ok": true}) }) + + req := httptest.NewRequest(http.MethodGet, "/v1beta/test", nil) + req.Header.Set("Authorization", "Bearer ok") + rec := httptest.NewRecorder() + r.ServeHTTP(rec, req) + + require.Equal(t, http.StatusForbidden, rec.Code) + var resp googleErrorResponse + require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &resp)) + require.Equal(t, http.StatusForbidden, resp.Error.Code) + require.Equal(t, "Insufficient account balance", resp.Error.Message) + require.Equal(t, "PERMISSION_DENIED", resp.Error.Status) +} + func TestApiKeyAuthWithSubscriptionGoogle_TouchesLastUsedOnSuccess(t *testing.T) { gin.SetMode(gin.TestMode) diff --git a/backend/internal/server/middleware/api_key_auth_test.go b/backend/internal/server/middleware/api_key_auth_test.go index 25c7db0aac..04ab9410ac 100644 --- a/backend/internal/server/middleware/api_key_auth_test.go +++ b/backend/internal/server/middleware/api_key_auth_test.go @@ -1000,6 +1000,49 @@ func TestAPIKeyAuthTouchesLastUsedInStandardMode(t *testing.T) { require.Equal(t, 1, touchCalls) } +func TestAPIKeyAuthRejectsBalanceBelowMinimumReserve(t *testing.T) { + gin.SetMode(gin.TestMode) + + user := &service.User{ + ID: 10, + Role: service.RoleUser, + Status: service.StatusActive, + Balance: 0.005, + Concurrency: 3, + } + apiKey := &service.APIKey{ + ID: 103, + UserID: user.ID, + Key: "held-balance-low", + Status: service.StatusActive, + User: user, + } + apiKeyRepo := &stubApiKeyRepo{ + getByKey: func(ctx context.Context, key string) (*service.APIKey, error) { + if key != apiKey.Key { + return nil, service.ErrAPIKeyNotFound + } + clone := *apiKey + userClone := *user + clone.User = &userClone + return &clone, nil + }, + } + + cfg := &config.Config{RunMode: config.RunModeStandard} + cfg.Billing.MinimumBalanceReserve = 0.01 + apiKeyService := service.NewAPIKeyService(apiKeyRepo, nil, nil, nil, nil, nil, cfg) + router := newAuthTestRouter(apiKeyService, nil, cfg) + + w := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodGet, "/t", nil) + req.Header.Set("x-api-key", apiKey.Key) + router.ServeHTTP(w, req) + + require.Equal(t, http.StatusForbidden, w.Code) + requireAPIKeyAuthError(t, w, "INSUFFICIENT_BALANCE", "Insufficient account balance") +} + func newAuthTestRouter(apiKeyService *service.APIKeyService, subscriptionService *service.SubscriptionService, cfg *config.Config) *gin.Engine { router := gin.New() router.Use(gin.HandlerFunc(NewAPIKeyAuthMiddleware(apiKeyService, subscriptionService, cfg))) diff --git a/backend/internal/server/routes/gateway.go b/backend/internal/server/routes/gateway.go index febbdc2682..d22e339c75 100644 --- a/backend/internal/server/routes/gateway.go +++ b/backend/internal/server/routes/gateway.go @@ -165,11 +165,14 @@ func RegisterGatewayRoutes( gateway.POST("/images/generations", imagesHandler) gateway.POST("/images/edits", imagesHandler) gateway.POST("/images/batches", h.BatchImage.Submit) + gateway.GET("/images/batches", h.BatchImage.List) + gateway.GET("/images/batches/models", h.BatchImage.Models) gateway.GET("/images/batches/:id", h.BatchImage.Get) gateway.GET("/images/batches/:id/items", h.BatchImage.Items) gateway.GET("/images/batches/:id/items/:custom_id/content", h.BatchImage.ItemContent) gateway.GET("/images/batches/:id/download", h.BatchImage.Download) gateway.POST("/images/batches/:id/cancel", h.BatchImage.Cancel) + gateway.DELETE("/images/batches/:id", h.BatchImage.DeleteRecord) gateway.DELETE("/images/batches/:id/outputs", h.BatchImage.DeleteOutputs) gateway.POST("/videos/generations", videoGenerationHandler) gateway.GET("/videos/:request_id", videoStatusHandler) diff --git a/backend/internal/service/admin_service.go b/backend/internal/service/admin_service.go index bacd134db4..18bf7b60ef 100644 --- a/backend/internal/service/admin_service.go +++ b/backend/internal/service/admin_service.go @@ -209,9 +209,12 @@ type CreateGroupInput struct { WeeklyLimitUSD *float64 // 周限额 (USD) MonthlyLimitUSD *float64 // 月限额 (USD) // 图片生成计费配置(仅 antigravity 平台使用) - AllowImageGeneration bool - ImageRateIndependent bool - ImageRateMultiplier *float64 + AllowImageGeneration bool + AllowBatchImageGeneration bool + ImageRateIndependent bool + ImageRateMultiplier *float64 + BatchImageDiscountMultiplier *float64 + BatchImageHoldMultiplier *float64 // 高峰时段倍率配置(PeakRateMultiplier 为 nil 时按 1.0 处理) PeakRateEnabled bool PeakStart string @@ -255,9 +258,12 @@ type UpdateGroupInput struct { WeeklyLimitUSD *float64 // 周限额 (USD) MonthlyLimitUSD *float64 // 月限额 (USD) // 图片生成计费配置(仅 antigravity 平台使用) - AllowImageGeneration *bool - ImageRateIndependent *bool - ImageRateMultiplier *float64 + AllowImageGeneration *bool + AllowBatchImageGeneration *bool + ImageRateIndependent *bool + ImageRateMultiplier *float64 + BatchImageDiscountMultiplier *float64 + BatchImageHoldMultiplier *float64 // 高峰时段倍率配置(nil 表示不修改) PeakRateEnabled *bool PeakStart *string @@ -1851,6 +1857,20 @@ func (s *adminServiceImpl) CreateGroup(ctx context.Context, input *CreateGroupIn } imageRateMultiplier = *input.ImageRateMultiplier } + batchImageDiscountMultiplier := defaultBatchImageDiscountMultiplier + if input.BatchImageDiscountMultiplier != nil { + if *input.BatchImageDiscountMultiplier < 0 { + return nil, errors.New("batch_image_discount_multiplier must be >= 0") + } + batchImageDiscountMultiplier = *input.BatchImageDiscountMultiplier + } + batchImageHoldMultiplier := defaultBatchImageHoldMultiplier + if input.BatchImageHoldMultiplier != nil { + if *input.BatchImageHoldMultiplier < 0 { + return nil, errors.New("batch_image_hold_multiplier must be >= 0") + } + batchImageHoldMultiplier = *input.BatchImageHoldMultiplier + } peakRateMultiplier := 1.0 if input.PeakRateMultiplier != nil { @@ -1886,6 +1906,7 @@ func (s *adminServiceImpl) CreateGroup(ctx context.Context, input *CreateGroupIn } allowImageGeneration := input.AllowImageGeneration || defaultAllowImageGenerationForPlatform(platform) + allowBatchImageGeneration := input.AllowBatchImageGeneration && allowImageGeneration // 如果指定了复制账号的源分组,先获取账号 ID 列表 var accountIDsToCopy []int64 @@ -1931,8 +1952,11 @@ func (s *adminServiceImpl) CreateGroup(ctx context.Context, input *CreateGroupIn WeeklyLimitUSD: weeklyLimit, MonthlyLimitUSD: monthlyLimit, AllowImageGeneration: allowImageGeneration, + AllowBatchImageGeneration: allowBatchImageGeneration, ImageRateIndependent: input.ImageRateIndependent, ImageRateMultiplier: imageRateMultiplier, + BatchImageDiscountMultiplier: batchImageDiscountMultiplier, + BatchImageHoldMultiplier: batchImageHoldMultiplier, PeakRateEnabled: peakRateEnabled, PeakStart: peakStart, PeakEnd: peakEnd, @@ -2117,6 +2141,12 @@ func (s *adminServiceImpl) UpdateGroup(ctx context.Context, id int64, input *Upd if input.AllowImageGeneration != nil { group.AllowImageGeneration = *input.AllowImageGeneration } + if input.AllowBatchImageGeneration != nil { + group.AllowBatchImageGeneration = *input.AllowBatchImageGeneration + } + if !group.AllowImageGeneration { + group.AllowBatchImageGeneration = false + } if input.ImageRateIndependent != nil { group.ImageRateIndependent = *input.ImageRateIndependent } @@ -2126,6 +2156,18 @@ func (s *adminServiceImpl) UpdateGroup(ctx context.Context, id int64, input *Upd } group.ImageRateMultiplier = *input.ImageRateMultiplier } + if input.BatchImageDiscountMultiplier != nil { + if *input.BatchImageDiscountMultiplier < 0 { + return nil, errors.New("batch_image_discount_multiplier must be >= 0") + } + group.BatchImageDiscountMultiplier = *input.BatchImageDiscountMultiplier + } + if input.BatchImageHoldMultiplier != nil { + if *input.BatchImageHoldMultiplier < 0 { + return nil, errors.New("batch_image_hold_multiplier must be >= 0") + } + group.BatchImageHoldMultiplier = *input.BatchImageHoldMultiplier + } if input.PeakRateEnabled != nil { group.PeakRateEnabled = *input.PeakRateEnabled } diff --git a/backend/internal/service/admin_service_group_test.go b/backend/internal/service/admin_service_group_test.go index 0b360c61c7..52485debff 100644 --- a/backend/internal/service/admin_service_group_test.go +++ b/backend/internal/service/admin_service_group_test.go @@ -232,6 +232,26 @@ func TestAdminService_CreateGroup_PreservesNonGrokImageGenerationDisabled(t *tes require.False(t, group.AllowImageGeneration) } +func TestAdminService_CreateGroup_DisablesBatchImageWhenImageGenerationDisabled(t *testing.T) { + repo := &groupRepoStubForAdmin{} + svc := &adminServiceImpl{groupRepo: repo} + + group, err := svc.CreateGroup(context.Background(), &CreateGroupInput{ + Name: "gemini-no-image", + Description: "Gemini group without image generation", + Platform: PlatformGemini, + RateMultiplier: 1.0, + AllowImageGeneration: false, + AllowBatchImageGeneration: true, + }) + require.NoError(t, err) + require.NotNil(t, group) + require.NotNil(t, repo.created) + require.False(t, repo.created.AllowImageGeneration) + require.False(t, repo.created.AllowBatchImageGeneration) + require.False(t, group.AllowBatchImageGeneration) +} + // TestAdminService_UpdateGroup_WithImagePricing 测试更新分组时 ImagePrice 字段正确更新 func TestAdminService_UpdateGroup_WithImagePricing(t *testing.T) { existingGroup := &Group{ @@ -326,6 +346,30 @@ func TestAdminService_UpdateGroup_PreservesImageGenerationControlsWhenOmitted(t require.InDelta(t, 0.5, repo.updated.ImageRateMultiplier, 1e-12) } +func TestAdminService_UpdateGroup_DisablesBatchImageWhenImageGenerationDisabled(t *testing.T) { + existingGroup := &Group{ + ID: 1, + Name: "existing-gemini", + Platform: PlatformGemini, + Status: StatusActive, + AllowImageGeneration: true, + AllowBatchImageGeneration: true, + } + repo := &groupRepoStubForAdmin{getByID: existingGroup} + svc := &adminServiceImpl{groupRepo: repo} + disabled := false + + group, err := svc.UpdateGroup(context.Background(), 1, &UpdateGroupInput{ + AllowImageGeneration: &disabled, + }) + require.NoError(t, err) + require.NotNil(t, group) + require.NotNil(t, repo.updated) + require.False(t, repo.updated.AllowImageGeneration) + require.False(t, repo.updated.AllowBatchImageGeneration) + require.False(t, group.AllowBatchImageGeneration) +} + func TestAdminService_UpdateGroup_ClearsDescriptionWhenEmptyString(t *testing.T) { existingGroup := &Group{ ID: 1, @@ -384,6 +428,58 @@ func TestAdminService_UpdateGroup_RejectsNegativeImageRateMultiplier(t *testing. require.Nil(t, repo.updated) } +func TestAdminService_CreateGroup_BatchImagePricingSettings(t *testing.T) { + repo := &groupRepoStubForAdmin{} + svc := &adminServiceImpl{groupRepo: repo} + discount := 0.8 + hold := 0.6 + + group, err := svc.CreateGroup(context.Background(), &CreateGroupInput{ + Name: "batch-image-pricing", + Platform: PlatformGemini, + RateMultiplier: 1, + BatchImageDiscountMultiplier: &discount, + BatchImageHoldMultiplier: &hold, + }) + require.NoError(t, err) + require.NotNil(t, group) + require.NotNil(t, repo.created) + require.InDelta(t, 0.8, repo.created.BatchImageDiscountMultiplier, 1e-12) + require.InDelta(t, 0.6, repo.created.BatchImageHoldMultiplier, 1e-12) +} + +func TestAdminService_GroupBatchImagePricingValidation(t *testing.T) { + tests := []struct { + name string + input *CreateGroupInput + }{ + { + name: "negative_discount", + input: func() *CreateGroupInput { + v := -0.1 + return &CreateGroupInput{Name: "bad-discount", RateMultiplier: 1, BatchImageDiscountMultiplier: &v} + }(), + }, + { + name: "negative_hold", + input: func() *CreateGroupInput { + v := -0.1 + return &CreateGroupInput{Name: "bad-hold", RateMultiplier: 1, BatchImageHoldMultiplier: &v} + }(), + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + repo := &groupRepoStubForAdmin{} + svc := &adminServiceImpl{groupRepo: repo} + + _, err := svc.CreateGroup(context.Background(), tt.input) + require.Error(t, err) + require.Nil(t, repo.created) + }) + } +} + func TestAdminService_UpdateGroup_InvalidatesAuthCacheOnRPMLimitChange(t *testing.T) { existingGroup := &Group{ ID: 1, diff --git a/backend/internal/service/api_key_auth_cache.go b/backend/internal/service/api_key_auth_cache.go index 32c3910c9d..6f927ff3b8 100644 --- a/backend/internal/service/api_key_auth_cache.go +++ b/backend/internal/service/api_key_auth_cache.go @@ -67,6 +67,7 @@ type APIKeyAuthGroupSnapshot struct { WeeklyLimitUSD *float64 `json:"weekly_limit_usd,omitempty"` MonthlyLimitUSD *float64 `json:"monthly_limit_usd,omitempty"` AllowImageGeneration bool `json:"allow_image_generation"` + AllowBatchImageGeneration bool `json:"allow_batch_image_generation"` ImageRateIndependent bool `json:"image_rate_independent"` ImageRateMultiplier float64 `json:"image_rate_multiplier"` ImagePrice1K *float64 `json:"image_price_1k,omitempty"` diff --git a/backend/internal/service/api_key_auth_cache_impl.go b/backend/internal/service/api_key_auth_cache_impl.go index b5aedf271e..f3da3df493 100644 --- a/backend/internal/service/api_key_auth_cache_impl.go +++ b/backend/internal/service/api_key_auth_cache_impl.go @@ -259,6 +259,7 @@ func (s *APIKeyService) snapshotFromAPIKey(ctx context.Context, apiKey *APIKey) WeeklyLimitUSD: apiKey.Group.WeeklyLimitUSD, MonthlyLimitUSD: apiKey.Group.MonthlyLimitUSD, AllowImageGeneration: apiKey.Group.AllowImageGeneration, + AllowBatchImageGeneration: apiKey.Group.AllowBatchImageGeneration, ImageRateIndependent: apiKey.Group.ImageRateIndependent, ImageRateMultiplier: apiKey.Group.ImageRateMultiplier, ImagePrice1K: apiKey.Group.ImagePrice1K, @@ -336,6 +337,7 @@ func (s *APIKeyService) snapshotToAPIKey(key string, snapshot *APIKeyAuthSnapsho WeeklyLimitUSD: snapshot.Group.WeeklyLimitUSD, MonthlyLimitUSD: snapshot.Group.MonthlyLimitUSD, AllowImageGeneration: snapshot.Group.AllowImageGeneration, + AllowBatchImageGeneration: snapshot.Group.AllowBatchImageGeneration, ImageRateIndependent: snapshot.Group.ImageRateIndependent, ImageRateMultiplier: snapshot.Group.ImageRateMultiplier, ImagePrice1K: snapshot.Group.ImagePrice1K, diff --git a/backend/internal/service/batch_image.go b/backend/internal/service/batch_image.go index 63d1913a0c..992f567841 100644 --- a/backend/internal/service/batch_image.go +++ b/backend/internal/service/batch_image.go @@ -29,6 +29,7 @@ const ( ) const ( + BatchImageItemStatusPending = "pending" BatchImageItemStatusSuccess = "success" BatchImageItemStatusFailed = "failed" BatchImageItemStatusCancelled = "cancelled" @@ -58,8 +59,12 @@ var ( ErrBatchImageSettlementMissingAPIKeyID = infraerrors.New(http.StatusBadRequest, "BATCH_IMAGE_SETTLEMENT_MISSING_API_KEY_ID", "batch image settlement api key id is missing") ErrBatchImageSettlementMissingAccountID = infraerrors.New(http.StatusBadRequest, "BATCH_IMAGE_SETTLEMENT_MISSING_ACCOUNT_ID", "batch image settlement account id is missing") ErrBatchImageSettlementInvalidCounts = infraerrors.New(http.StatusBadRequest, "BATCH_IMAGE_SETTLEMENT_INVALID_COUNTS", "batch image settlement counts are invalid") + ErrBatchImageSettlementCostExceedsHold = infraerrors.New(http.StatusConflict, "BATCH_IMAGE_SETTLEMENT_COST_EXCEEDS_HOLD", "batch image settlement cost exceeds held balance") + ErrBatchImageBillingHoldFailed = infraerrors.New(http.StatusBadGateway, "BATCH_IMAGE_BILLING_HOLD_FAILED", "batch image balance hold failed") + ErrBatchImageInsufficientBalance = infraerrors.New(http.StatusPaymentRequired, "BATCH_IMAGE_INSUFFICIENT_BALANCE", "insufficient balance for batch image hold") ErrBatchImageDisabled = infraerrors.New(http.StatusNotFound, "BATCH_IMAGE_DISABLED", "batch image API is disabled") + ErrBatchImageGroupDisabled = infraerrors.New(http.StatusForbidden, "BATCH_IMAGE_GROUP_DISABLED", "batch image API is disabled for this group") ErrBatchImageInvalidModel = infraerrors.New(http.StatusBadRequest, "BATCH_IMAGE_INVALID_MODEL", "batch image model is required") ErrBatchImageNoAccountAvailable = infraerrors.New(http.StatusBadGateway, "BATCH_IMAGE_NO_ACCOUNT_AVAILABLE", "no compatible batch image account is available") ErrBatchImageInvalidItems = infraerrors.New(http.StatusBadRequest, "BATCH_IMAGE_INVALID_ITEMS", "batch image items are invalid") @@ -69,6 +74,7 @@ var ( ErrBatchImageQueueFailed = infraerrors.New(http.StatusBadGateway, "BATCH_IMAGE_QUEUE_FAILED", "batch image queue failed") ErrBatchImageIdempotencyConflict = infraerrors.New(http.StatusConflict, "BATCH_IMAGE_IDEMPOTENCY_CONFLICT", "idempotency key reused with different batch image request") ErrBatchImageCancelFailed = infraerrors.New(http.StatusBadGateway, "BATCH_IMAGE_CANCEL_FAILED", "batch image cancel failed") + ErrBatchImageVertexGCSBucketMissing = infraerrors.New(http.StatusBadGateway, "BATCH_IMAGE_VERTEX_GCS_BUCKET_MISSING", "Vertex managed GCS bucket is not configured") ErrBatchImageNotReady = infraerrors.New(http.StatusConflict, "BATCH_IMAGE_NOT_READY", "batch image job is not completed") ErrBatchImageOutputDeleted = infraerrors.New(http.StatusGone, "BATCH_IMAGE_OUTPUT_DELETED", "batch image output has been deleted") @@ -80,6 +86,7 @@ var ( ErrBatchImageItemImageIndexOutOfRange = infraerrors.New(http.StatusBadRequest, "BATCH_IMAGE_ITEM_IMAGE_INDEX_OUT_OF_RANGE", "batch image item image index is out of range") ErrBatchImageZipTooManyItems = infraerrors.New(http.StatusBadRequest, "BATCH_IMAGE_ZIP_TOO_MANY_ITEMS", "batch image ZIP contains too many items; use single item downloads") ErrBatchImageOutputDeleteNotReady = infraerrors.New(http.StatusConflict, "BATCH_IMAGE_OUTPUT_DELETE_NOT_READY", "batch image output can only be deleted after completion") + ErrBatchImageRecordDeleteNotReady = infraerrors.New(http.StatusConflict, "BATCH_IMAGE_RECORD_DELETE_NOT_READY", "batch image record can only be deleted after the job finishes") ErrBatchImageCleanupFailed = infraerrors.New(http.StatusBadGateway, "BATCH_IMAGE_CLEANUP_FAILED", "batch image cleanup failed") ErrBatchImageCleanupUnsafePath = infraerrors.New(http.StatusBadRequest, "BATCH_IMAGE_CLEANUP_UNSAFE_PATH", "batch image cleanup path is unsafe") ErrBatchImageProviderCleanupFailed = infraerrors.New(http.StatusBadGateway, "BATCH_IMAGE_PROVIDER_CLEANUP_FAILED", "batch image provider cleanup failed") @@ -93,6 +100,8 @@ type BatchImageJob struct { AccountID *int64 Provider string Model string + TaskName string + ParentBatchID *string Status string ProviderJobName *string ProviderInputRef *string @@ -105,11 +114,19 @@ type BatchImageJob struct { FailCount int CancelledCount int - EstimatedCost float64 - HoldAmount *float64 - ActualCost *float64 - Currency string - HoldID *string + EstimatedCost float64 + HoldAmount *float64 + ActualCost *float64 + BaseUnitPrice float64 + GroupRateMultiplier float64 + AccountRateMultiplier float64 + BatchDiscountMultiplier float64 + HoldMultiplier float64 + BillableUnitPrice float64 + HoldUnitPrice float64 + PricingSnapshotVersion int + Currency string + HoldID *string IdempotencyKey *string RequestHash *string @@ -121,6 +138,8 @@ type BatchImageJob struct { OutputExpiresAt *time.Time InputDeletedAt *time.Time OutputDeletedAt *time.Time + DownloadedAt *time.Time + UserDeletedAt *time.Time LastErrorCode *string LastErrorMessage *string @@ -140,6 +159,8 @@ type CreateBatchImageJobParams struct { AccountID *int64 Provider string Model string + TaskName string + ParentBatchID *string Status string ProviderJobName *string ProviderInputRef *string @@ -152,11 +173,19 @@ type CreateBatchImageJobParams struct { FailCount int CancelledCount int - EstimatedCost float64 - HoldAmount *float64 - ActualCost *float64 - Currency string - HoldID *string + EstimatedCost float64 + HoldAmount *float64 + ActualCost *float64 + BaseUnitPrice float64 + GroupRateMultiplier float64 + AccountRateMultiplier float64 + BatchDiscountMultiplier float64 + HoldMultiplier float64 + BillableUnitPrice float64 + HoldUnitPrice float64 + PricingSnapshotVersion int + Currency string + HoldID *string IdempotencyKey *string RequestHash *string @@ -213,6 +242,17 @@ type BatchImageItemFilter struct { Offset int } +type BatchImageJobFilter struct { + Status string + TaskNameLike string + Downloaded *bool + CreatedAfter *time.Time + CreatedBefore *time.Time + ExcludeDeleted bool + Limit int + Offset int +} + type BatchImageCounts struct { SuccessCount int FailCount int @@ -260,6 +300,7 @@ type BatchImageRepository interface { GetBatchImageJobByIdempotencyKey(ctx context.Context, userID, apiKeyID int64, key string) (*BatchImageJob, error) GetBatchImageJobByBatchIDForOwner(ctx context.Context, userID, apiKeyID int64, batchID string) (*BatchImageJob, error) GetBatchImageJobByID(ctx context.Context, id int64) (*BatchImageJob, error) + ListBatchImageJobsForOwner(ctx context.Context, userID, apiKeyID int64, filter BatchImageJobFilter) ([]*BatchImageJob, error) TransitionBatchImageJobStatus(ctx context.Context, batchID, toStatus string, opts BatchImageTransitionOptions) error UpdateBatchImageJobProviderOutputRef(ctx context.Context, batchID, providerOutputRef string) error UpdateBatchImageJobProviderSubmit(ctx context.Context, params UpdateBatchImageJobProviderSubmitParams) error @@ -276,8 +317,11 @@ type BatchImageRepository interface { ListBatchImageItemsForDownload(ctx context.Context, batchID string, status string, limit int) ([]*BatchImageItem, error) ListBatchImageJobsDueForInputCleanup(ctx context.Context, cutoff time.Time, limit int) ([]*BatchImageJob, error) ListBatchImageJobsDueForOutputCleanup(ctx context.Context, now time.Time, limit int) ([]*BatchImageJob, error) + ListStaleUnsubmittedBatchImageJobs(ctx context.Context, cutoff time.Time, limit int) ([]*BatchImageJob, error) MarkBatchImageInputDeleted(ctx context.Context, batchID string, deletedAt time.Time) error MarkBatchImageOutputDeleted(ctx context.Context, batchID string, deletedAt time.Time) error + MarkBatchImageDownloaded(ctx context.Context, batchID string, downloadedAt time.Time) error + MarkBatchImageJobUserDeleted(ctx context.Context, userID, apiKeyID int64, batchID string, deletedAt time.Time) error SetBatchImageOutputExpiresAt(ctx context.Context, batchID string, expiresAt time.Time) error RecordBatchImageCleanupFailure(ctx context.Context, batchID, code, message string) error AppendBatchImageEvent(ctx context.Context, batchID, eventType string, payload any) error diff --git a/backend/internal/service/batch_image_billing_hold.go b/backend/internal/service/batch_image_billing_hold.go new file mode 100644 index 0000000000..af61117db0 --- /dev/null +++ b/backend/internal/service/batch_image_billing_hold.go @@ -0,0 +1,104 @@ +package service + +import ( + "context" + "errors" + "strings" +) + +const ( + batchImageHoldRequestPrefix = "batch_image_hold:" + batchImageCaptureRequestPrefix = "batch_image_capture:" + batchImageReleaseRequestPrefix = "batch_image_release:" +) + +func BatchImageHoldRequestID(batchID string) string { + return batchImageHoldRequestPrefix + strings.TrimSpace(batchID) +} + +func BatchImageCaptureRequestID(batchID string) string { + return batchImageCaptureRequestPrefix + strings.TrimSpace(batchID) +} + +func BatchImageReleaseRequestID(batchID string) string { + return batchImageReleaseRequestPrefix + strings.TrimSpace(batchID) +} + +func buildBatchImageHoldCommand(job *BatchImageJob, requestID string, actualAmount float64, payloadHash string) (*BatchImageBalanceHoldCommand, error) { + if job == nil { + return nil, ErrBatchImageBillingHoldFailed + } + if job.APIKeyID == nil || *job.APIKeyID <= 0 { + return nil, ErrBatchImageSettlementMissingAPIKeyID + } + holdAmount := job.EstimatedCost + if job.HoldAmount != nil { + holdAmount = *job.HoldAmount + } + if holdAmount < 0 { + holdAmount = 0 + } + if actualAmount < 0 { + actualAmount = 0 + } + return &BatchImageBalanceHoldCommand{ + RequestID: requestID, + APIKeyID: *job.APIKeyID, + UserID: job.UserID, + BatchID: job.BatchID, + HoldAmount: holdAmount, + ActualAmount: actualAmount, + RequestPayloadHash: strings.TrimSpace(payloadHash), + }, nil +} + +func reserveBatchImageBalanceHold(ctx context.Context, repo UsageBillingRepository, job *BatchImageJob, payloadHash string) error { + if repo == nil { + return ErrBatchImageBillingHoldFailed.WithCause(errors.New("batch image billing repository is not configured")) + } + cmd, err := buildBatchImageHoldCommand(job, BatchImageHoldRequestID(job.BatchID), 0, payloadHash) + if err != nil { + return err + } + if cmd.HoldAmount <= 0 { + return nil + } + if _, err := repo.ReserveBatchImageBalance(ctx, cmd); err != nil { + if errors.Is(err, ErrBatchImageInsufficientBalance) { + return ErrBatchImageInsufficientBalance + } + return ErrBatchImageBillingHoldFailed.WithCause(err) + } + return nil +} + +func captureBatchImageBalanceHold(ctx context.Context, repo UsageBillingRepository, job *BatchImageJob, actualAmount float64, payloadHash string) error { + if repo == nil { + return ErrBatchImageSettlementBillingFailed.WithCause(errors.New("batch image billing repository is not configured")) + } + cmd, err := buildBatchImageHoldCommand(job, BatchImageCaptureRequestID(job.BatchID), actualAmount, payloadHash) + if err != nil { + return err + } + if _, err := repo.CaptureBatchImageBalance(ctx, cmd); err != nil { + return ErrBatchImageSettlementBillingFailed.WithCause(err) + } + return nil +} + +func releaseBatchImageBalanceHold(ctx context.Context, repo UsageBillingRepository, job *BatchImageJob, payloadHash string) error { + if repo == nil || job == nil { + return nil + } + cmd, err := buildBatchImageHoldCommand(job, BatchImageReleaseRequestID(job.BatchID), 0, payloadHash) + if err != nil { + return err + } + if cmd.HoldAmount <= 0 { + return nil + } + if _, err := repo.ReleaseBatchImageBalance(ctx, cmd); err != nil { + return ErrBatchImageBillingHoldFailed.WithCause(err) + } + return nil +} diff --git a/backend/internal/service/batch_image_billing_recovery.go b/backend/internal/service/batch_image_billing_recovery.go new file mode 100644 index 0000000000..d89f508f47 --- /dev/null +++ b/backend/internal/service/batch_image_billing_recovery.go @@ -0,0 +1,62 @@ +package service + +import ( + "context" + "errors" + "time" +) + +const ( + defaultBatchImageBillingRecoveryStaleAfter = 10 * time.Minute + defaultBatchImageBillingRecoveryLimit = 100 +) + +type BatchImageBillingRecoveryService struct { + Repo BatchImageRepository + Billing UsageBillingRepository + AuthCache APIKeyAuthCacheInvalidator + StaleAfter time.Duration + Limit int +} + +func (s *BatchImageBillingRecoveryService) ReleaseStaleUnsubmittedOnce(ctx context.Context) (int, error) { + if s == nil || s.Repo == nil || s.Billing == nil { + return 0, nil + } + staleAfter := s.StaleAfter + if staleAfter <= 0 { + staleAfter = defaultBatchImageBillingRecoveryStaleAfter + } + limit := s.Limit + if limit <= 0 { + limit = defaultBatchImageBillingRecoveryLimit + } + jobs, err := s.Repo.ListStaleUnsubmittedBatchImageJobs(ctx, time.Now().Add(-staleAfter), limit) + if err != nil { + return 0, err + } + released := 0 + for _, job := range jobs { + if job == nil { + continue + } + msg := "batch image submission did not reach provider before recovery cutoff" + if err := s.Repo.TransitionBatchImageJobStatus(ctx, job.BatchID, BatchImageJobStatusFailed, BatchImageTransitionOptions{ + EventType: "billing_hold_recovery_failed_unsubmitted", + EventPayload: map[string]any{"batch_id": job.BatchID}, + ErrorCode: batchImageStringPtr("SUBMIT_STALE_BEFORE_PROVIDER"), + ErrorMessage: batchImageStringPtr(msg), + }); err != nil && !errors.Is(err, ErrBatchImageInvalidTransition) { + return released, err + } + job.Status = BatchImageJobStatusFailed + if err := releaseBatchImageBalanceHold(ctx, s.Billing, job, batchImageDerefString(job.RequestHash)); err != nil { + return released, err + } + if s.AuthCache != nil && job.UserID > 0 { + s.AuthCache.InvalidateAuthCacheByUserID(ctx, job.UserID) + } + released++ + } + return released, nil +} diff --git a/backend/internal/service/batch_image_billing_recovery_test.go b/backend/internal/service/batch_image_billing_recovery_test.go new file mode 100644 index 0000000000..2ab2783f83 --- /dev/null +++ b/backend/internal/service/batch_image_billing_recovery_test.go @@ -0,0 +1,52 @@ +//go:build unit + +package service + +import ( + "context" + "testing" + "time" + + "github.com/stretchr/testify/require" +) + +func TestBatchImageBillingRecoveryService_ReleasesStaleUnsubmittedHold(t *testing.T) { + repo := newFakeBatchImageRepository() + apiKeyID := int64(22) + holdAmount := 0.5 + stale := &BatchImageJob{ + BatchID: "imgbatch_stale_created", + UserID: 11, + APIKeyID: &apiKeyID, + Status: BatchImageJobStatusCreated, + EstimatedCost: holdAmount, + HoldAmount: &holdAmount, + CreatedAt: time.Now().Add(-time.Hour), + UpdatedAt: time.Now().Add(-time.Hour), + } + activeProviderName := "providers/job" + active := &BatchImageJob{ + BatchID: "imgbatch_has_provider", + UserID: 11, + APIKeyID: &apiKeyID, + Status: BatchImageJobStatusSubmitted, + ProviderJobName: &activeProviderName, + EstimatedCost: holdAmount, + HoldAmount: &holdAmount, + CreatedAt: time.Now().Add(-time.Hour), + UpdatedAt: time.Now().Add(-time.Hour), + } + repo.jobs[stale.BatchID] = stale + repo.jobs[active.BatchID] = active + billing := &fakeBatchImageBillingRepo{} + svc := &BatchImageBillingRecoveryService{Repo: repo, Billing: billing, StaleAfter: time.Minute, Limit: 10} + + released, err := svc.ReleaseStaleUnsubmittedOnce(context.Background()) + require.NoError(t, err) + require.Equal(t, 1, released) + require.Equal(t, BatchImageJobStatusFailed, repo.jobs[stale.BatchID].Status) + require.Equal(t, "SUBMIT_STALE_BEFORE_PROVIDER", batchImageDerefString(repo.jobs[stale.BatchID].LastErrorCode)) + require.Len(t, billing.releases, 1) + require.Equal(t, BatchImageReleaseRequestID(stale.BatchID), billing.releases[0].RequestID) + require.Equal(t, BatchImageJobStatusSubmitted, repo.jobs[active.BatchID].Status) +} diff --git a/backend/internal/service/batch_image_cleanup.go b/backend/internal/service/batch_image_cleanup.go index a6b6995f10..7b2b527080 100644 --- a/backend/internal/service/batch_image_cleanup.go +++ b/backend/internal/service/batch_image_cleanup.go @@ -32,7 +32,7 @@ type BatchImageCleanupService struct { func NewBatchImageCleanupService(repo BatchImageRepository, accountRepo AccountRepository, cfg *config.Config) *BatchImageCleanupService { return &BatchImageCleanupService{ Repo: repo, - ProviderRegistry: NewDefaultBatchImageProviderRegistry(), + ProviderRegistry: NewBatchImageProviderRegistryFromConfig(cfg), AccountResolver: &BatchImageAccountRepositoryResolver{Repo: accountRepo}, Config: cfg, } diff --git a/backend/internal/service/batch_image_download.go b/backend/internal/service/batch_image_download.go index f8933a23bb..d12540d9c9 100644 --- a/backend/internal/service/batch_image_download.go +++ b/backend/internal/service/batch_image_download.go @@ -77,7 +77,7 @@ type BatchImageDownloadService struct { func NewBatchImageDownloadService(repo BatchImageRepository, accountRepo AccountRepository, limiter BatchImageDownloadLimiter, cfg *config.Config) *BatchImageDownloadService { return &BatchImageDownloadService{ Repo: repo, - ProviderRegistry: NewDefaultBatchImageProviderRegistry(), + ProviderRegistry: NewBatchImageProviderRegistryFromConfig(cfg), AccountResolver: &BatchImageAccountRepositoryResolver{Repo: accountRepo}, Limiter: limiter, Config: cfg, diff --git a/backend/internal/service/batch_image_mvp_smoke_test.go b/backend/internal/service/batch_image_mvp_smoke_test.go index f4cb372602..b9daa8f457 100644 --- a/backend/internal/service/batch_image_mvp_smoke_test.go +++ b/backend/internal/service/batch_image_mvp_smoke_test.go @@ -10,6 +10,7 @@ import ( "io" "strings" "testing" + "time" "github.com/Wei-Shaw/sub2api/internal/config" "github.com/stretchr/testify/require" @@ -50,6 +51,7 @@ func TestBatchImageMVPFlow(t *testing.T) { Queue: queue, ProviderRegistry: registry, Pricing: pricing, + BillingRepo: billing, Config: cfg, } processor := &BatchImagePipelineProcessor{ @@ -57,6 +59,7 @@ func TestBatchImageMVPFlow(t *testing.T) { Repo: repo, ProviderRegistry: registry, AccountResolver: &fakeBatchImageAccountResolver{account: &accountRepo.accounts[0]}, + BillingRepo: billing, }, SettlementService: &BatchImageSettlementService{ Repo: repo, @@ -87,6 +90,9 @@ func TestBatchImageMVPFlow(t *testing.T) { require.Equal(t, 2, submitted.ItemCount) require.Equal(t, []string{submitted.ID}, queue.enqueued) require.Len(t, provider.submits, 1) + require.Len(t, billing.reserves, 1) + require.Equal(t, BatchImageHoldRequestID(submitted.ID), billing.reserves[0].RequestID) + require.InDelta(t, 0.3, billing.reserves[0].HoldAmount, 1e-12) requireBatchImagePublicJSONHasNoInternals(t, mustMarshalBatchImageSmokeJSON(t, submitted)) firstProcess, err := processor.Process(ctx, submitted.ID) @@ -96,7 +102,8 @@ func TestBatchImageMVPFlow(t *testing.T) { indexProcess, err := processor.Process(ctx, submitted.ID) require.NoError(t, err) - require.True(t, indexProcess.Terminal) + require.False(t, indexProcess.Terminal) + require.Equal(t, time.Millisecond, indexProcess.RequeueAfter) require.Equal(t, BatchImageJobStatusSettling, repo.jobs[submitted.ID].Status) require.Equal(t, BatchImageCounts{SuccessCount: 1, FailCount: 1}, repo.counts[submitted.ID]) @@ -108,15 +115,15 @@ func TestBatchImageMVPFlow(t *testing.T) { require.NotNil(t, job.OutputExpiresAt) require.Equal(t, 1, job.SuccessCount) require.Equal(t, 1, job.FailCount) - require.Len(t, billing.commands, 1) - require.Equal(t, BatchImageSettlementRequestID(submitted.ID), billing.commands[0].RequestID) - require.Equal(t, 1, billing.commands[0].ImageCount) - require.Equal(t, 0.25, billing.commands[0].BalanceCost) + require.Len(t, billing.captures, 1) + require.Equal(t, BatchImageCaptureRequestID(submitted.ID), billing.captures[0].RequestID) + require.InDelta(t, 0.3, billing.captures[0].HoldAmount, 1e-12) + require.InDelta(t, 0.125, billing.captures[0].ActualAmount, 1e-12) secondSettlement, err := processor.SettlementService.Settle(ctx, submitted.ID) require.NoError(t, err) require.True(t, secondSettlement.AlreadySettled) - require.Len(t, billing.commands, 1) + require.Len(t, billing.captures, 1) status, err := publicSvc.Get(ctx, owner, submitted.ID) require.NoError(t, err) diff --git a/backend/internal/service/batch_image_processor.go b/backend/internal/service/batch_image_processor.go index 4fbac83dae..b82bff85ef 100644 --- a/backend/internal/service/batch_image_processor.go +++ b/backend/internal/service/batch_image_processor.go @@ -13,6 +13,8 @@ import ( "time" infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" + "github.com/Wei-Shaw/sub2api/internal/pkg/logger" + "go.uber.org/zap" ) const ( @@ -48,6 +50,8 @@ type BatchImageProviderProcessor struct { ProviderRegistry *BatchImageProviderRegistry AccountResolver BatchImageAccountResolver Indexer *BatchImageResultIndexer + BillingRepo UsageBillingRepository + AuthCache APIKeyAuthCacheInvalidator DefaultRequeue time.Duration } @@ -61,6 +65,9 @@ func (p *BatchImageProviderProcessor) Process(ctx context.Context, batchID strin return BatchImageProcessResult{}, err } if isBatchImageProcessorDoneStatus(job.Status) { + if err := p.releaseTerminalHold(ctx, job); err != nil { + return BatchImageProcessResult{}, err + } return BatchImageProcessResult{Terminal: true}, nil } @@ -88,6 +95,12 @@ func (p *BatchImageProviderProcessor) Process(ctx context.Context, batchID strin status, err := provider.Get(ctx, job, account) if err != nil { + logger.L().Warn("batch_image.provider_status_check_failed", + zap.String("batch_id", job.BatchID), + zap.String("provider", job.Provider), + zap.String("provider_job_name", batchImageDerefString(job.ProviderJobName)), + zap.Error(err), + ) return BatchImageProcessResult{RequeueAfter: batchImageProviderErrorRequeue}, nil } if status == nil { @@ -139,6 +152,10 @@ func (p *BatchImageProviderProcessor) Process(ctx context.Context, batchID strin }); err != nil { return BatchImageProcessResult{}, err } + job.Status = BatchImageJobStatusFailed + if err := p.releaseTerminalHold(ctx, job); err != nil { + return BatchImageProcessResult{}, err + } return BatchImageProcessResult{Terminal: true}, nil case BatchProviderStateCancelled: if err := p.Repo.TransitionBatchImageJobStatus(ctx, job.BatchID, BatchImageJobStatusCancelled, BatchImageTransitionOptions{ @@ -147,6 +164,10 @@ func (p *BatchImageProviderProcessor) Process(ctx context.Context, batchID strin }); err != nil { return BatchImageProcessResult{}, err } + job.Status = BatchImageJobStatusCancelled + if err := p.releaseTerminalHold(ctx, job); err != nil { + return BatchImageProcessResult{}, err + } return BatchImageProcessResult{Terminal: true}, nil default: return BatchImageProcessResult{RequeueAfter: p.requeueDelay(status.SuggestedRequeueAfter)}, nil @@ -181,6 +202,10 @@ func (p *BatchImageProviderProcessor) indexAndSettle(ctx context.Context, job *B if transitionErr != nil { return BatchImageProcessResult{}, transitionErr } + job.Status = BatchImageJobStatusFailed + if err := p.releaseTerminalHold(ctx, job); err != nil { + return BatchImageProcessResult{}, err + } return BatchImageProcessResult{Terminal: true}, nil } @@ -194,7 +219,23 @@ func (p *BatchImageProviderProcessor) indexAndSettle(ctx context.Context, job *B }); err != nil { return BatchImageProcessResult{}, err } - return BatchImageProcessResult{Terminal: true}, nil + return BatchImageProcessResult{RequeueAfter: time.Millisecond}, nil +} + +func (p *BatchImageProviderProcessor) releaseTerminalHold(ctx context.Context, job *BatchImageJob) error { + if p == nil || job == nil { + return nil + } + if job.Status != BatchImageJobStatusFailed && job.Status != BatchImageJobStatusCancelled { + return nil + } + if err := releaseBatchImageBalanceHold(ctx, p.BillingRepo, job, batchImageDerefString(job.RequestHash)); err != nil { + return err + } + if p.AuthCache != nil && job.UserID > 0 { + p.AuthCache.InvalidateAuthCacheByUserID(ctx, job.UserID) + } + return nil } func (p *BatchImageProviderProcessor) persistProviderOutputRef(ctx context.Context, job *BatchImageJob, ref string) error { diff --git a/backend/internal/service/batch_image_processor_test.go b/backend/internal/service/batch_image_processor_test.go index 07268ca912..f4c96f3e19 100644 --- a/backend/internal/service/batch_image_processor_test.go +++ b/backend/internal/service/batch_image_processor_test.go @@ -230,7 +230,8 @@ func TestBatchImageProviderProcessor_StatusFlow(t *testing.T) { } got, err := newTestBatchImageProcessor(repo, provider).Process(ctx, "imgbatch_flow") require.NoError(t, err) - require.True(t, got.Terminal) + require.False(t, got.Terminal) + require.Equal(t, time.Millisecond, got.RequeueAfter) require.Equal(t, BatchImageJobStatusSettling, repo.jobs["imgbatch_flow"].Status) require.Equal(t, "files/output", batchImageDerefString(repo.jobs["imgbatch_flow"].ProviderOutputRef)) require.Equal(t, []string{BatchImageJobStatusIndexing, BatchImageJobStatusSettling}, repo.transitions["imgbatch_flow"]) @@ -251,11 +252,22 @@ func TestBatchImageProviderProcessor_StatusFlow(t *testing.T) { t.Run("cancelled provider marks job cancelled", func(t *testing.T) { repo := newFakeBatchImageRepository() repo.jobs["imgbatch_flow"] = newJob(BatchImageJobStatusRunning) + apiKeyID := int64(22) + holdAmount := 0.5 + repo.jobs["imgbatch_flow"].UserID = 11 + repo.jobs["imgbatch_flow"].APIKeyID = &apiKeyID + repo.jobs["imgbatch_flow"].EstimatedCost = holdAmount + repo.jobs["imgbatch_flow"].HoldAmount = &holdAmount provider := &fakeProcessorProvider{status: &BatchProviderStatus{InternalState: BatchProviderStateCancelled, RawState: "CANCELLED"}} - got, err := newTestBatchImageProcessor(repo, provider).Process(ctx, "imgbatch_flow") + processor := newTestBatchImageProcessor(repo, provider) + billing := &fakeBatchImageBillingRepo{} + processor.BillingRepo = billing + got, err := processor.Process(ctx, "imgbatch_flow") require.NoError(t, err) require.True(t, got.Terminal) require.Equal(t, BatchImageJobStatusCancelled, repo.jobs["imgbatch_flow"].Status) + require.Len(t, billing.releases, 1) + require.Equal(t, BatchImageReleaseRequestID("imgbatch_flow"), billing.releases[0].RequestID) }) } @@ -342,19 +354,31 @@ func newFakeBatchImageRepository() *fakeBatchImageRepository { func (r *fakeBatchImageRepository) CreateBatchImageJob(_ context.Context, params CreateBatchImageJobParams) (*BatchImageJob, error) { job := &BatchImageJob{ - BatchID: params.BatchID, - UserID: params.UserID, - APIKeyID: params.APIKeyID, - AccountID: params.AccountID, - Status: params.Status, - Provider: params.Provider, - Model: params.Model, - ProviderJobName: params.ProviderJobName, - ItemCount: params.ItemCount, - EstimatedCost: params.EstimatedCost, - IdempotencyKey: params.IdempotencyKey, - RequestHash: params.RequestHash, - CreatedAt: time.Now(), + BatchID: params.BatchID, + UserID: params.UserID, + APIKeyID: params.APIKeyID, + AccountID: params.AccountID, + Status: params.Status, + Provider: params.Provider, + Model: params.Model, + TaskName: params.TaskName, + ProviderJobName: params.ProviderJobName, + ItemCount: params.ItemCount, + EstimatedCost: params.EstimatedCost, + HoldAmount: params.HoldAmount, + HoldID: params.HoldID, + BaseUnitPrice: params.BaseUnitPrice, + GroupRateMultiplier: params.GroupRateMultiplier, + AccountRateMultiplier: params.AccountRateMultiplier, + BatchDiscountMultiplier: params.BatchDiscountMultiplier, + HoldMultiplier: params.HoldMultiplier, + BillableUnitPrice: params.BillableUnitPrice, + HoldUnitPrice: params.HoldUnitPrice, + PricingSnapshotVersion: params.PricingSnapshotVersion, + Currency: params.Currency, + IdempotencyKey: params.IdempotencyKey, + RequestHash: params.RequestHash, + CreatedAt: time.Now(), } r.jobs[job.BatchID] = job return job, nil @@ -385,6 +409,53 @@ func (r *fakeBatchImageRepository) GetBatchImageJobByBatchIDForOwner(_ context.C return job, nil } +func (r *fakeBatchImageRepository) ListBatchImageJobsForOwner(_ context.Context, userID, apiKeyID int64, filter BatchImageJobFilter) ([]*BatchImageJob, error) { + limit := filter.Limit + if limit <= 0 || limit > 100 { + limit = 20 + } + offset := filter.Offset + if offset < 0 { + offset = 0 + } + var jobs []*BatchImageJob + for _, job := range r.jobs { + if job.UserID != userID || job.APIKeyID == nil || *job.APIKeyID != apiKeyID { + continue + } + if filter.Status != "" && job.Status != filter.Status { + continue + } + if filter.TaskNameLike != "" && !strings.Contains(strings.ToLower(job.TaskName), strings.ToLower(filter.TaskNameLike)) { + continue + } + if filter.ExcludeDeleted && job.UserDeletedAt != nil { + continue + } + if filter.Downloaded != nil { + downloaded := job.DownloadedAt != nil + if downloaded != *filter.Downloaded { + continue + } + } + if filter.CreatedAfter != nil && job.CreatedAt.Before(*filter.CreatedAfter) { + continue + } + if filter.CreatedBefore != nil && !job.CreatedAt.Before(*filter.CreatedBefore) { + continue + } + if offset > 0 { + offset-- + continue + } + jobs = append(jobs, job) + if len(jobs) >= limit { + break + } + } + return jobs, nil +} + func (r *fakeBatchImageRepository) GetBatchImageJobByID(_ context.Context, id int64) (*BatchImageJob, error) { for _, job := range r.jobs { if job.ID == id { @@ -657,6 +728,33 @@ func (r *fakeBatchImageRepository) ListBatchImageJobsDueForOutputCleanup(_ conte return jobs, nil } +func (r *fakeBatchImageRepository) ListStaleUnsubmittedBatchImageJobs(_ context.Context, cutoff time.Time, limit int) ([]*BatchImageJob, error) { + if limit <= 0 { + limit = 100 + } + jobs := make([]*BatchImageJob, 0, limit) + for _, job := range r.jobs { + if len(jobs) >= limit { + break + } + if job.Status != BatchImageJobStatusCreated && job.Status != BatchImageJobStatusUploading { + continue + } + if batchImageDerefString(job.ProviderJobName) != "" { + continue + } + holdAmount := job.EstimatedCost + if job.HoldAmount != nil { + holdAmount = *job.HoldAmount + } + if holdAmount <= 0 || job.UpdatedAt.After(cutoff) { + continue + } + jobs = append(jobs, job) + } + return jobs, nil +} + func (r *fakeBatchImageRepository) MarkBatchImageInputDeleted(_ context.Context, batchID string, deletedAt time.Time) error { job, ok := r.jobs[batchID] if !ok { @@ -684,6 +782,33 @@ func (r *fakeBatchImageRepository) MarkBatchImageOutputDeleted(_ context.Context return nil } +func (r *fakeBatchImageRepository) MarkBatchImageDownloaded(_ context.Context, batchID string, downloadedAt time.Time) error { + job, ok := r.jobs[batchID] + if !ok { + return ErrBatchImageJobNotFound + } + if job.DownloadedAt == nil { + job.DownloadedAt = &downloadedAt + } + r.events[batchID] = append(r.events[batchID], "download_completed") + return nil +} + +func (r *fakeBatchImageRepository) MarkBatchImageJobUserDeleted(_ context.Context, userID, apiKeyID int64, batchID string, deletedAt time.Time) error { + job, ok := r.jobs[batchID] + if !ok || job.UserID != userID || job.APIKeyID == nil || *job.APIKeyID != apiKeyID { + return ErrBatchImageJobNotFound + } + if !isBatchImageProcessorDoneStatus(job.Status) { + return ErrBatchImageRecordDeleteNotReady + } + if job.UserDeletedAt == nil { + job.UserDeletedAt = &deletedAt + } + r.events[batchID] = append(r.events[batchID], "user_record_deleted") + return nil +} + func (r *fakeBatchImageRepository) SetBatchImageOutputExpiresAt(_ context.Context, batchID string, expiresAt time.Time) error { job, ok := r.jobs[batchID] if !ok { diff --git a/backend/internal/service/batch_image_provider.go b/backend/internal/service/batch_image_provider.go index 11700f5f68..4a638aa328 100644 --- a/backend/internal/service/batch_image_provider.go +++ b/backend/internal/service/batch_image_provider.go @@ -8,6 +8,7 @@ import ( "strings" "time" + "github.com/Wei-Shaw/sub2api/internal/config" infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" ) @@ -43,6 +44,13 @@ func NewDefaultBatchImageProviderRegistry() *BatchImageProviderRegistry { ) } +func NewBatchImageProviderRegistryFromConfig(cfg *config.Config) *BatchImageProviderRegistry { + return NewBatchImageProviderRegistry( + NewGeminiAPIBatchImageProvider(nil), + NewVertexBatchImageProviderFromConfig(cfg, nil, nil, nil), + ) +} + func (r *BatchImageProviderRegistry) Get(provider string) (BatchImageProvider, bool) { if r == nil { return nil, false diff --git a/backend/internal/service/batch_image_provider_vertex.go b/backend/internal/service/batch_image_provider_vertex.go index 727cd6f457..b37a0c35e8 100644 --- a/backend/internal/service/batch_image_provider_vertex.go +++ b/backend/internal/service/batch_image_provider_vertex.go @@ -626,7 +626,7 @@ func mapVertexClientError(err error) error { return vertexProviderError("VERTEX_INVALID_RESPONSE", "Vertex API request failed", nil) } } - return vertexProviderError("VERTEX_INVALID_RESPONSE", "Vertex API request failed", nil) + return vertexProviderError("VERTEX_INVALID_RESPONSE", "Vertex API request failed", err) } type vertexCombinedJSONLReadCloser struct { diff --git a/backend/internal/service/batch_image_public.go b/backend/internal/service/batch_image_public.go index c10c8d0246..19c2590837 100644 --- a/backend/internal/service/batch_image_public.go +++ b/backend/internal/service/batch_image_public.go @@ -13,14 +13,17 @@ import ( "time" "github.com/Wei-Shaw/sub2api/internal/config" + infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" ) const ( - defaultBatchImageMaxItems = 500 - defaultBatchImageMaxPromptChars = 8000 - defaultBatchImageResponseMime = "image/png" - defaultBatchImageImageSize = "1K" - maxBatchImagePublicErrorChars = 500 + defaultBatchImageMaxItems = 500 + defaultBatchImageMaxPromptChars = 8000 + defaultBatchImageResponseMime = "image/png" + defaultBatchImageImageSize = "1K" + defaultBatchImageDiscountMultiplier = 0.5 + defaultBatchImageHoldMultiplier = 0.6 + maxBatchImagePublicErrorChars = 500 ) type BatchImageAccountSelectionRepository interface { @@ -29,8 +32,18 @@ type BatchImageAccountSelectionRepository interface { ListSchedulableByGroupIDAndPlatform(ctx context.Context, groupID int64, platform string) ([]Account, error) } +type BatchImageGroupPricingRepository interface { + GetByIDLite(ctx context.Context, id int64) (*Group, error) +} + +type BatchImageUserGroupRateRepository interface { + GetByUserAndGroup(ctx context.Context, userID, groupID int64) (*float64, error) +} + type BatchImageSubmitRequest struct { Model string `json:"model"` + TaskName string `json:"task_name"` + ParentBatchID string `json:"parent_batch_id"` Provider string `json:"provider"` Items []BatchImageSubmitItem `json:"items"` ResponseMimeType string `json:"response_mime_type"` @@ -51,17 +64,35 @@ type BatchImageOwner struct { } type BatchImagePublicService struct { - Repo BatchImageRepository - AccountRepo BatchImageAccountSelectionRepository - Queue BatchImageQueue - ProviderRegistry *BatchImageProviderRegistry - Pricing BatchImagePricingResolver - Config *config.Config + Repo BatchImageRepository + AccountRepo BatchImageAccountSelectionRepository + GroupRepo BatchImageGroupPricingRepository + UserGroupRateRepo BatchImageUserGroupRateRepository + Queue BatchImageQueue + ProviderRegistry *BatchImageProviderRegistry + Pricing BatchImagePricingResolver + BillingRepo UsageBillingRepository + AuthCache APIKeyAuthCacheInvalidator + Config *config.Config +} + +type BatchImagePricingSnapshot struct { + BaseUnitPrice float64 + GroupRateMultiplier float64 + AccountRateMultiplier float64 + BatchDiscountMultiplier float64 + HoldMultiplier float64 + BillableUnitPrice float64 + HoldUnitPrice float64 + EstimatedCost float64 + HoldAmount float64 } type BatchImagePublicBatch struct { ID string `json:"id"` Object string `json:"object"` + TaskName string `json:"task_name"` + ParentBatchID *string `json:"parent_batch_id,omitempty"` Status string `json:"status"` Model string `json:"model"` Provider string `json:"provider"` @@ -69,16 +100,19 @@ type BatchImagePublicBatch struct { SuccessCount int `json:"success_count"` FailCount int `json:"fail_count"` EstimatedCost float64 `json:"estimated_cost"` + HoldAmount float64 `json:"hold_amount"` ActualCost *float64 `json:"actual_cost"` CreatedAt int64 `json:"created_at"` SubmittedAt *int64 `json:"submitted_at"` SettledAt *int64 `json:"settled_at"` + DownloadedAt *int64 `json:"downloaded_at,omitempty"` OutputDeletedAt *int64 `json:"output_deleted_at,omitempty"` } type BatchImagePublicItem struct { CustomID string `json:"custom_id"` Status string `json:"status"` + PromptPreview *string `json:"prompt_preview,omitempty"` MimeType *string `json:"mime_type"` FileExtension *string `json:"file_extension"` ImageCount int `json:"image_count"` @@ -88,6 +122,7 @@ type BatchImagePublicItem struct { type BatchImagePublicError struct { Code string `json:"code"` Message string `json:"message"` + Source string `json:"source,omitempty"` } type BatchImagePublicItemsResponse struct { @@ -96,20 +131,51 @@ type BatchImagePublicItemsResponse struct { HasMore bool `json:"has_more"` } +type BatchImagePublicListResponse struct { + Object string `json:"object"` + Data []*BatchImagePublicBatch `json:"data"` + HasMore bool `json:"has_more"` +} + +type BatchImagePublicModel struct { + ID string `json:"id"` + Object string `json:"object"` + Provider string `json:"provider"` +} + +type BatchImagePublicModelsResponse struct { + Object string `json:"object"` + Data []BatchImagePublicModel `json:"data"` +} + +type BatchImageJobsQuery struct { + Status string + TaskName string + Downloaded string + From string + To string + Limit int + Cursor string +} + type BatchImageItemsQuery struct { Status string Limit int Cursor string } -func NewBatchImagePublicService(repo BatchImageRepository, accountRepo AccountRepository, queue BatchImageQueue, pricing *BatchImageModelPricingResolver, cfg *config.Config) *BatchImagePublicService { +func NewBatchImagePublicService(repo BatchImageRepository, accountRepo AccountRepository, groupRepo GroupRepository, userGroupRateRepo UserGroupRateRepository, queue BatchImageQueue, pricing *BatchImageModelPricingResolver, billingRepo UsageBillingRepository, authCache APIKeyAuthCacheInvalidator, cfg *config.Config) *BatchImagePublicService { return &BatchImagePublicService{ - Repo: repo, - AccountRepo: accountRepo, - Queue: queue, - ProviderRegistry: NewDefaultBatchImageProviderRegistry(), - Pricing: pricing, - Config: cfg, + Repo: repo, + AccountRepo: accountRepo, + GroupRepo: groupRepo, + UserGroupRateRepo: userGroupRateRepo, + Queue: queue, + ProviderRegistry: NewBatchImageProviderRegistryFromConfig(cfg), + Pricing: pricing, + BillingRepo: billingRepo, + AuthCache: authCache, + Config: cfg, } } @@ -146,30 +212,75 @@ func (s *BatchImagePublicService) Submit(ctx context.Context, owner BatchImageOw if err != nil { return nil, err } - estimatedCost := s.estimateCost(ctx, normalized, provider.Name()) + pricingSnapshot, err := s.resolvePricingSnapshot(ctx, owner, normalized, provider.Name(), account) + if err != nil { + return nil, err + } + parentBatchID := batchImageOptionalStringPtr(normalized.ParentBatchID) + if parentBatchID != nil { + parent, parentErr := s.Repo.GetBatchImageJobByBatchIDForOwner(ctx, owner.UserID, owner.APIKeyID, *parentBatchID) + if parentErr != nil { + return nil, parentErr + } + if parent.ParentBatchID != nil && strings.TrimSpace(*parent.ParentBatchID) != "" { + parentBatchID = batchImageOptionalStringPtr(*parent.ParentBatchID) + } + } batchID, err := NewBatchImageID() if err != nil { return nil, err } apiKeyID := owner.APIKeyID accountID := account.ID + holdID := BatchImageHoldRequestID(batchID) + holdAmount := pricingSnapshot.HoldAmount job, err := s.Repo.CreateBatchImageJob(ctx, CreateBatchImageJobParams{ - BatchID: batchID, - UserID: owner.UserID, - APIKeyID: &apiKeyID, - AccountID: &accountID, - Provider: provider.Name(), - Model: normalized.Model, - Status: BatchImageJobStatusCreated, - ItemCount: len(normalized.Items), - EstimatedCost: estimatedCost, - Currency: "USD", - IdempotencyKey: batchImageOptionalStringPtr(idempotencyKey), - RequestHash: batchImageStringPtr(requestHash), + BatchID: batchID, + UserID: owner.UserID, + APIKeyID: &apiKeyID, + AccountID: &accountID, + Provider: provider.Name(), + Model: normalized.Model, + TaskName: normalized.TaskName, + ParentBatchID: parentBatchID, + Status: BatchImageJobStatusCreated, + ItemCount: len(normalized.Items), + EstimatedCost: pricingSnapshot.EstimatedCost, + HoldAmount: &holdAmount, + BaseUnitPrice: pricingSnapshot.BaseUnitPrice, + GroupRateMultiplier: pricingSnapshot.GroupRateMultiplier, + AccountRateMultiplier: pricingSnapshot.AccountRateMultiplier, + BatchDiscountMultiplier: pricingSnapshot.BatchDiscountMultiplier, + HoldMultiplier: pricingSnapshot.HoldMultiplier, + BillableUnitPrice: pricingSnapshot.BillableUnitPrice, + HoldUnitPrice: pricingSnapshot.HoldUnitPrice, + PricingSnapshotVersion: 1, + Currency: "USD", + HoldID: &holdID, + IdempotencyKey: batchImageOptionalStringPtr(idempotencyKey), + RequestHash: batchImageStringPtr(requestHash), }) if err != nil { return nil, err } + if err := reserveBatchImageBalanceHold(ctx, s.BillingRepo, job, requestHash); err != nil { + code := "BILLING_HOLD_FAILED" + if errors.Is(err, ErrBatchImageInsufficientBalance) { + code = "INSUFFICIENT_BALANCE" + } + _ = s.Repo.RecordBatchImageJobSubmitFailure(ctx, job.BatchID, code, sanitizeBatchImagePublicMessage(err.Error()), true) + s.hidePreUpstreamSubmitFailure(ctx, owner, job) + return nil, err + } + s.invalidateAuthCache(ctx, owner.UserID) + if err := s.createPendingItems(ctx, job.BatchID, requestHash, normalized.Items); err != nil { + if releaseErr := s.releaseFailedSubmitHold(ctx, job, requestHash); releaseErr != nil { + return nil, releaseErr + } + _ = s.Repo.RecordBatchImageJobSubmitFailure(ctx, job.BatchID, "ITEM_CREATE_FAILED", sanitizeBatchImagePublicMessage(err.Error()), true) + s.hidePreUpstreamSubmitFailure(ctx, owner, job) + return nil, ErrBatchImageQueueFailed + } input := BatchImageInput{ BatchID: job.BatchID, @@ -187,11 +298,21 @@ func (s *BatchImagePublicService) Submit(ctx context.Context, owner BatchImageOw providerJob, err := provider.Submit(ctx, job, account, input) if err != nil { - _ = s.Repo.RecordBatchImageJobSubmitFailure(ctx, job.BatchID, "PROVIDER_SUBMIT_FAILED", sanitizeBatchImagePublicMessage(err.Error()), true) - return nil, ErrBatchImageProviderSubmitFailed + if releaseErr := s.releaseFailedSubmitHold(ctx, job, requestHash); releaseErr != nil { + return nil, releaseErr + } + publicErr := batchImageProviderSubmitPublicError(err) + reason := batchImageProviderSubmitRecordCode(publicErr) + _ = s.Repo.RecordBatchImageJobSubmitFailure(ctx, job.BatchID, reason, sanitizeBatchImagePublicMessage(err.Error()), true) + s.hidePreUpstreamSubmitFailure(ctx, owner, job) + return nil, publicErr } if providerJob == nil || strings.TrimSpace(providerJob.ProviderJobName) == "" { + if releaseErr := s.releaseFailedSubmitHold(ctx, job, requestHash); releaseErr != nil { + return nil, releaseErr + } _ = s.Repo.RecordBatchImageJobSubmitFailure(ctx, job.BatchID, "PROVIDER_SUBMIT_FAILED", "provider job name missing", true) + s.hidePreUpstreamSubmitFailure(ctx, owner, job) return nil, ErrBatchImageProviderSubmitFailed } @@ -221,6 +342,54 @@ func (s *BatchImagePublicService) Submit(ctx context.Context, owner BatchImageOw return BatchImageJobToPublic(created), nil } +func (s *BatchImagePublicService) releaseFailedSubmitHold(ctx context.Context, job *BatchImageJob, requestHash string) error { + if err := releaseBatchImageBalanceHold(ctx, s.BillingRepo, job, requestHash); err != nil { + _ = s.Repo.RecordBatchImageJobSubmitFailure(ctx, job.BatchID, "BILLING_RELEASE_FAILED", sanitizeBatchImagePublicMessage(err.Error()), true) + s.enqueueBillingRetry(ctx, job.BatchID) + return ErrBatchImageBillingHoldFailed + } + s.invalidateAuthCache(ctx, job.UserID) + return nil +} + +func (s *BatchImagePublicService) createPendingItems(ctx context.Context, batchID, requestHash string, items []BatchImageSubmitItem) error { + if s == nil || s.Repo == nil || len(items) == 0 { + return nil + } + params := make([]CreateBatchImageItemParams, 0, len(items)) + for _, item := range items { + preview := truncateBatchImageMessage(item.Prompt, s.maxPromptChars()) + params = append(params, CreateBatchImageItemParams{ + JobID: batchID, + CustomID: item.CustomID, + Status: BatchImageItemStatusPending, + RequestHash: batchImageStringPtr(requestHash), + PromptPreview: batchImageStringPtr(preview), + ImageCount: 0, + }) + } + return s.Repo.BulkCreateBatchImageItems(ctx, params) +} + +func (s *BatchImagePublicService) enqueueBillingRetry(ctx context.Context, batchID string) { + if s == nil || s.Queue == nil { + return + } + if err := s.Queue.Enqueue(ctx, batchID); err != nil && !errors.Is(err, ErrBatchImageAlreadyQueued) { + _ = s.Repo.AppendBatchImageEvent(ctx, batchID, "billing_retry_enqueue_failed", map[string]any{ + "batch_id": batchID, + "error": sanitizeBatchImagePublicMessage(err.Error()), + }) + } +} + +func (s *BatchImagePublicService) hidePreUpstreamSubmitFailure(ctx context.Context, owner BatchImageOwner, job *BatchImageJob) { + if s == nil || s.Repo == nil || job == nil || job.ProviderJobName != nil { + return + } + _ = s.Repo.MarkBatchImageJobUserDeleted(ctx, owner.UserID, owner.APIKeyID, job.BatchID, time.Now()) +} + func (s *BatchImagePublicService) Get(ctx context.Context, owner BatchImageOwner, batchID string) (*BatchImagePublicBatch, error) { job, err := s.Repo.GetBatchImageJobByBatchIDForOwner(ctx, owner.UserID, owner.APIKeyID, batchID) if err != nil { @@ -229,12 +398,147 @@ func (s *BatchImagePublicService) Get(ctx context.Context, owner BatchImageOwner return BatchImageJobToPublic(job), nil } +func (s *BatchImagePublicService) List(ctx context.Context, owner BatchImageOwner, query BatchImageJobsQuery) (*BatchImagePublicListResponse, error) { + filter := BatchImageJobFilter{Limit: query.Limit, Offset: parseBatchImageCursor(query.Cursor), ExcludeDeleted: true} + filter.TaskNameLike = strings.TrimSpace(query.TaskName) + switch strings.TrimSpace(query.Status) { + case "", "all": + case "queued": + filter.Status = BatchImageJobStatusSubmitted + case "processing_results": + filter.Status = BatchImageJobStatusIndexing + case "completed": + filter.Status = BatchImageJobStatusCompleted + case "failed": + filter.Status = BatchImageJobStatusFailed + case "cancelled": + filter.Status = BatchImageJobStatusCancelled + case "output_deleted": + filter.Status = BatchImageJobStatusOutputDeleted + default: + filter.Status = strings.TrimSpace(query.Status) + } + switch strings.TrimSpace(strings.ToLower(query.Downloaded)) { + case "", "all": + case "true", "1", "yes", "downloaded": + downloaded := true + filter.Downloaded = &downloaded + case "false", "0", "no", "not_downloaded": + downloaded := false + filter.Downloaded = &downloaded + default: + return nil, ErrBatchImageInvalidItems + } + if from := parseBatchImageListTime(query.From); from != nil { + filter.CreatedAfter = from + } + if to := parseBatchImageListTime(query.To); to != nil { + filter.CreatedBefore = to + } + if filter.Limit <= 0 || filter.Limit > 100 { + filter.Limit = 20 + } + jobs, err := s.Repo.ListBatchImageJobsForOwner(ctx, owner.UserID, owner.APIKeyID, filter) + if err != nil { + return nil, err + } + data := make([]*BatchImagePublicBatch, 0, len(jobs)) + for _, job := range jobs { + data = append(data, BatchImageJobToPublic(job)) + } + return &BatchImagePublicListResponse{ + Object: "list", + Data: data, + HasMore: len(data) == filter.Limit, + }, nil +} + +func (s *BatchImagePublicService) MarkDownloaded(ctx context.Context, owner BatchImageOwner, batchID string) error { + job, err := s.Repo.GetBatchImageJobByBatchIDForOwner(ctx, owner.UserID, owner.APIKeyID, batchID) + if err != nil { + return err + } + return s.Repo.MarkBatchImageDownloaded(ctx, job.BatchID, time.Now()) +} + +func (s *BatchImagePublicService) DeleteRecord(ctx context.Context, owner BatchImageOwner, batchID string) error { + job, err := s.Repo.GetBatchImageJobByBatchIDForOwner(ctx, owner.UserID, owner.APIKeyID, batchID) + if err != nil { + return err + } + if !isBatchImageProcessorDoneStatus(job.Status) { + return ErrBatchImageRecordDeleteNotReady + } + return s.Repo.MarkBatchImageJobUserDeleted(ctx, owner.UserID, owner.APIKeyID, job.BatchID, time.Now()) +} + +func (s *BatchImagePublicService) ListModels(ctx context.Context, owner BatchImageOwner) (*BatchImagePublicModelsResponse, error) { + if !s.enabled() { + return nil, ErrBatchImageDisabled + } + if s.Pricing == nil { + return nil, ErrBatchImageSettlementPricingMissing + } + if err := s.ensureGroupAllowsBatchImage(ctx, owner.GroupID); err != nil { + return nil, err + } + + modelsByProvider := make(map[string]map[string]struct{}) + for _, providerName := range batchImageProviderSelectionOrder("") { + provider, ok := s.ProviderRegistry.Get(providerName) + if !ok || provider == nil { + continue + } + accounts, err := s.listCandidateAccounts(ctx, owner.GroupID, batchImageProviderPlatform(providerName)) + if err != nil { + return nil, err + } + for i := range accounts { + account := accounts[i] + if !account.IsSchedulable() || !provider.SupportsAccount(&account) { + continue + } + for _, model := range batchImageModelsFromAccountMapping(&account) { + if _, err := s.Pricing.BatchImageUnitPrice(ctx, &BatchImageJob{Provider: providerName, Model: model}); err != nil { + continue + } + if !account.IsModelSupported(model) { + continue + } + if modelsByProvider[providerName] == nil { + modelsByProvider[providerName] = make(map[string]struct{}) + } + modelsByProvider[providerName][model] = struct{}{} + } + } + } + + out := make([]BatchImagePublicModel, 0) + for _, providerName := range batchImageProviderSelectionOrder("") { + models := make([]string, 0, len(modelsByProvider[providerName])) + for model := range modelsByProvider[providerName] { + models = append(models, model) + } + sort.Strings(models) + for _, model := range models { + out = append(out, BatchImagePublicModel{ + ID: model, + Object: "image.batch.model", + Provider: providerName, + }) + } + } + return &BatchImagePublicModelsResponse{Object: "list", Data: out}, nil +} + func (s *BatchImagePublicService) ListItems(ctx context.Context, owner BatchImageOwner, batchID string, query BatchImageItemsQuery) (*BatchImagePublicItemsResponse, error) { filter := BatchImageItemFilter{Limit: query.Limit, Offset: parseBatchImageCursor(query.Cursor)} switch strings.TrimSpace(query.Status) { case "", "all": case "succeeded", "success": filter.Status = BatchImageItemStatusSuccess + case "pending": + filter.Status = BatchImageItemStatusPending case "failed": filter.Status = BatchImageItemStatusFailed default: @@ -264,6 +568,13 @@ func (s *BatchImagePublicService) Cancel(ctx context.Context, owner BatchImageOw return nil, err } if isBatchImageProcessorDoneStatus(job.Status) { + if job.Status == BatchImageJobStatusFailed || job.Status == BatchImageJobStatusCancelled { + if err := releaseBatchImageBalanceHold(ctx, s.BillingRepo, job, batchImageDerefString(job.RequestHash)); err != nil { + s.enqueueBillingRetry(ctx, job.BatchID) + return nil, ErrBatchImageCancelFailed + } + s.invalidateAuthCache(ctx, owner.UserID) + } return BatchImageJobToPublic(job), nil } if job.ProviderJobName != nil && strings.TrimSpace(*job.ProviderJobName) != "" { @@ -281,6 +592,17 @@ func (s *BatchImagePublicService) Cancel(ctx context.Context, owner BatchImageOw if err := provider.Cancel(ctx, job, account); err != nil { return nil, ErrBatchImageCancelFailed } + _ = s.Repo.AppendBatchImageEvent(ctx, job.BatchID, "job_cancel_requested", map[string]any{"batch_id": job.BatchID}) + if s.Queue != nil { + if err := s.Queue.Enqueue(ctx, job.BatchID); err != nil && !errors.Is(err, ErrBatchImageAlreadyQueued) { + return nil, ErrBatchImageCancelFailed + } + } + updated, err := s.Repo.GetBatchImageJobByBatchIDForOwner(ctx, owner.UserID, owner.APIKeyID, batchID) + if err != nil { + return nil, err + } + return BatchImageJobToPublic(updated), nil } if err := s.Repo.TransitionBatchImageJobStatus(ctx, job.BatchID, BatchImageJobStatusCancelled, BatchImageTransitionOptions{ EventType: "job_cancelled", @@ -288,6 +610,11 @@ func (s *BatchImagePublicService) Cancel(ctx context.Context, owner BatchImageOw }); err != nil { return nil, err } + if err := releaseBatchImageBalanceHold(ctx, s.BillingRepo, job, batchImageDerefString(job.RequestHash)); err != nil { + s.enqueueBillingRetry(ctx, job.BatchID) + return nil, ErrBatchImageCancelFailed + } + s.invalidateAuthCache(ctx, owner.UserID) updated, err := s.Repo.GetBatchImageJobByBatchIDForOwner(ctx, owner.UserID, owner.APIKeyID, batchID) if err != nil { return nil, err @@ -297,6 +624,8 @@ func (s *BatchImagePublicService) Cancel(ctx context.Context, owner BatchImageOw func (s *BatchImagePublicService) validateSubmitRequest(req BatchImageSubmitRequest) (BatchImageSubmitRequest, error) { req.Model = strings.TrimSpace(req.Model) + req.TaskName = strings.TrimSpace(req.TaskName) + req.ParentBatchID = strings.TrimSpace(req.ParentBatchID) req.Provider = strings.TrimSpace(req.Provider) req.ResponseMimeType = strings.TrimSpace(req.ResponseMimeType) req.AspectRatio = strings.TrimSpace(req.AspectRatio) @@ -304,6 +633,12 @@ func (s *BatchImagePublicService) validateSubmitRequest(req BatchImageSubmitRequ if req.Model == "" { return req, ErrBatchImageInvalidModel } + if req.TaskName == "" { + req.TaskName = defaultBatchImageTaskName(time.Now()) + } + if len(req.TaskName) > 255 { + req.TaskName = truncateBatchImageMessage(req.TaskName, 255) + } if req.Provider != "" && !IsSupportedBatchImageProvider(req.Provider) { return req, ErrBatchImageUnsupportedProvider } @@ -320,9 +655,10 @@ func (s *BatchImagePublicService) validateSubmitRequest(req BatchImageSubmitRequ if req.ImageSize == "" { req.ImageSize = s.defaultImageSize() } - if req.Provider == BatchImageProviderVertex && (strings.EqualFold(req.ImageSize, "2K") || strings.EqualFold(req.ImageSize, "4K")) { + if !strings.EqualFold(req.ImageSize, defaultBatchImageImageSize) { return req, ErrBatchImageInvalidItems } + req.ImageSize = defaultBatchImageImageSize req.Metadata = sanitizeBatchImageMetadata(req.Metadata) seen := make(map[string]struct{}, len(req.Items)) @@ -347,10 +683,7 @@ func (s *BatchImagePublicService) validateSubmitRequest(req BatchImageSubmitRequ } func (s *BatchImagePublicService) selectProviderAndAccount(ctx context.Context, owner BatchImageOwner, requestedProvider, model string) (BatchImageProvider, *Account, error) { - providers := []string{requestedProvider} - if strings.TrimSpace(requestedProvider) == "" { - providers = []string{BatchImageProviderGeminiAPI, BatchImageProviderVertex} - } + providers := batchImageProviderSelectionOrder(requestedProvider) for _, providerName := range providers { provider, ok := s.ProviderRegistry.Get(providerName) if !ok || provider == nil { @@ -392,21 +725,114 @@ func (s *BatchImagePublicService) listCandidateAccounts(ctx context.Context, gro return s.AccountRepo.ListSchedulableByPlatform(ctx, platform) } -func (s *BatchImagePublicService) estimateCost(ctx context.Context, req BatchImageSubmitRequest, provider string) float64 { - if s.Pricing == nil { - return 0 +func (s *BatchImagePublicService) ensureGroupAllowsBatchImage(ctx context.Context, groupID *int64) error { + if groupID == nil || *groupID <= 0 { + return nil } - unit, err := s.Pricing.BatchImageUnitPrice(ctx, &BatchImageJob{Provider: provider, Model: req.Model}) - if err != nil || unit < 0 { - return 0 + if s.GroupRepo == nil { + return ErrBatchImageSettlementPricingMissing } - return unit * float64(len(req.Items)) + group, err := s.GroupRepo.GetByIDLite(ctx, *groupID) + if err != nil || group == nil { + return ErrBatchImageSettlementPricingMissing + } + if !group.AllowBatchImageGeneration { + return ErrBatchImageGroupDisabled + } + return nil +} + +func (s *BatchImagePublicService) resolvePricingSnapshot(ctx context.Context, owner BatchImageOwner, req BatchImageSubmitRequest, provider string, account *Account) (*BatchImagePricingSnapshot, error) { + unit := -1.0 + groupMultiplier := 1.0 + discountMultiplier := defaultBatchImageDiscountMultiplier + holdMultiplier := defaultBatchImageHoldMultiplier + if owner.GroupID != nil && *owner.GroupID > 0 { + if s.GroupRepo == nil { + return nil, ErrBatchImageSettlementPricingMissing + } + group, err := s.GroupRepo.GetByIDLite(ctx, *owner.GroupID) + if err != nil || group == nil { + return nil, ErrBatchImageSettlementPricingMissing + } + if !group.AllowBatchImageGeneration { + return nil, ErrBatchImageGroupDisabled + } + groupDefaultMultiplier := group.RateMultiplier + if groupDefaultMultiplier < 0 { + groupDefaultMultiplier = 0 + } + effectiveGroupMultiplier := groupDefaultMultiplier + if s.UserGroupRateRepo != nil { + userRate, rateErr := s.UserGroupRateRepo.GetByUserAndGroup(ctx, owner.UserID, group.ID) + if rateErr != nil { + return nil, ErrBatchImageSettlementPricingMissing + } + if userRate != nil { + effectiveGroupMultiplier = *userRate + } + } + groupMultiplier = effectiveGroupMultiplier + if group.ImageRateIndependent { + groupMultiplier = group.ImageRateMultiplier + } + if groupMultiplier < 0 { + groupMultiplier = 0 + } + discountMultiplier = group.BatchImageDiscountMultiplier + if discountMultiplier < 0 { + discountMultiplier = 0 + } + if group.BatchImageHoldMultiplier >= 0 { + holdMultiplier = group.BatchImageHoldMultiplier + } + if configuredUnit := group.GetImagePrice(req.ImageSize); configuredUnit != nil && *configuredUnit >= 0 { + unit = *configuredUnit + } + } + if unit < 0 { + if s.Pricing == nil { + return nil, ErrBatchImageSettlementPricingMissing + } + resolvedUnit, err := s.Pricing.BatchImageUnitPrice(ctx, &BatchImageJob{Provider: provider, Model: req.Model}) + if err != nil || resolvedUnit < 0 { + return nil, ErrBatchImageSettlementPricingMissing + } + unit = resolvedUnit + } + accountMultiplier := 1.0 + if account != nil { + accountMultiplier = account.BillingRateMultiplier() + } + if accountMultiplier < 0 { + accountMultiplier = 0 + } + standardUnitPrice := unit * groupMultiplier * accountMultiplier + billableUnitPrice := standardUnitPrice * discountMultiplier + holdUnitPrice := standardUnitPrice * holdMultiplier + return &BatchImagePricingSnapshot{ + BaseUnitPrice: unit, + GroupRateMultiplier: groupMultiplier, + AccountRateMultiplier: accountMultiplier, + BatchDiscountMultiplier: discountMultiplier, + HoldMultiplier: holdMultiplier, + BillableUnitPrice: billableUnitPrice, + HoldUnitPrice: holdUnitPrice, + EstimatedCost: billableUnitPrice * float64(len(req.Items)), + HoldAmount: holdUnitPrice * float64(len(req.Items)), + }, nil } func (s *BatchImagePublicService) enabled() bool { return s != nil && s.Repo != nil && s.AccountRepo != nil && s.Config != nil && s.Config.BatchImage.Enabled } +func (s *BatchImagePublicService) invalidateAuthCache(ctx context.Context, userID int64) { + if s != nil && s.AuthCache != nil && userID > 0 { + s.AuthCache.InvalidateAuthCacheByUserID(ctx, userID) + } +} + func (s *BatchImagePublicService) maxItems() int { if s != nil && s.Config != nil && s.Config.BatchImage.MaxItemsPerJobDefault > 0 { return s.Config.BatchImage.MaxItemsPerJobDefault @@ -439,9 +865,15 @@ func BatchImageJobToPublic(job *BatchImageJob) *BatchImagePublicBatch { if job == nil { return nil } + holdAmount := job.EstimatedCost + if job.HoldAmount != nil { + holdAmount = *job.HoldAmount + } return &BatchImagePublicBatch{ ID: job.BatchID, Object: "image.batch", + TaskName: batchImagePublicTaskName(job), + ParentBatchID: job.ParentBatchID, Status: PublicBatchImageStatus(job.Status), Model: job.Model, Provider: job.Provider, @@ -449,10 +881,12 @@ func BatchImageJobToPublic(job *BatchImageJob) *BatchImagePublicBatch { SuccessCount: job.SuccessCount, FailCount: job.FailCount, EstimatedCost: job.EstimatedCost, + HoldAmount: holdAmount, ActualCost: job.ActualCost, CreatedAt: job.CreatedAt.Unix(), SubmittedAt: batchImageUnixPtr(job.SubmittedAt), SettledAt: batchImageUnixPtr(job.SettledAt), + DownloadedAt: batchImageUnixPtr(job.DownloadedAt), OutputDeletedAt: batchImageUnixPtr(job.OutputDeletedAt), } } @@ -461,10 +895,15 @@ func BatchImageItemToPublic(item *BatchImageItem) BatchImagePublicItem { out := BatchImagePublicItem{ CustomID: item.CustomID, Status: "failed", + PromptPreview: item.PromptPreview, MimeType: item.MimeType, FileExtension: item.FileExtension, ImageCount: item.ImageCount, } + if item.Status == BatchImageItemStatusPending { + out.Status = "pending" + return out + } if item.Status == BatchImageItemStatusSuccess { out.Status = "succeeded" return out @@ -472,10 +911,29 @@ func BatchImageItemToPublic(item *BatchImageItem) BatchImagePublicItem { out.Error = &BatchImagePublicError{ Code: batchImageDerefString(item.ErrorCode), Message: sanitizeBatchImagePublicMessage(batchImageDerefString(item.ErrorMessage)), + Source: batchImageItemErrorSource(item), } return out } +func batchImageItemErrorSource(item *BatchImageItem) string { + if item == nil || item.ErrorCode == nil { + return "" + } + code := strings.TrimSpace(*item.ErrorCode) + if batchImageDerefString(item.ProviderSourceObject) != "" { + return "provider" + } + switch code { + case "EMPTY_IMAGE_OUTPUT", "PROVIDER_ITEM_FAILED": + return "provider" + case "INDEX_OUTPUT_MISSING", "INDEX_PARSE_FAILED", "DUPLICATE_CUSTOM_ID_IN_OUTPUT": + return "system" + default: + return "" + } +} + func PublicBatchImageStatus(status string) string { switch status { case BatchImageJobStatusCreated, BatchImageJobStatusUploading, BatchImageJobStatusSubmitted: @@ -515,6 +973,57 @@ func batchImageProviderPlatform(provider string) string { } } +func batchImageProviderSelectionOrder(requestedProvider string) []string { + if strings.TrimSpace(requestedProvider) != "" { + return []string{strings.TrimSpace(requestedProvider)} + } + return []string{BatchImageProviderGeminiAPI, BatchImageProviderVertex} +} + +func batchImageModelsFromAccountMapping(account *Account) []string { + if account == nil { + return nil + } + mapping := account.GetModelMapping() + if len(mapping) == 0 { + return nil + } + models := make(map[string]struct{}) + for model := range mapping { + model = strings.TrimSpace(model) + if model == "" { + continue + } + if strings.ContainsAny(model, "*?") { + for _, candidate := range defaultBatchImageModelCandidates() { + if matchWildcard(model, candidate) { + models[candidate] = struct{}{} + } + } + continue + } + models[model] = struct{}{} + } + out := make([]string, 0, len(models)) + for model := range models { + out = append(out, model) + } + sort.Strings(out) + return out +} + +func defaultBatchImageModelCandidates() []string { + return []string{ + "gemini-2.0-flash-exp-image-generation", + "gemini-2.5-flash-image", + "gemini-3-pro-image", + "gemini-3-pro-image-preview", + "gemini-3.1-flash-image", + "gemini-3.1-flash-image-preview", + "gemini-3.1-flash-lite-image", + } +} + func batchImageGCSRef(provider, ref string) string { if provider == BatchImageProviderVertex && strings.HasPrefix(strings.TrimSpace(ref), "gs://") { return strings.TrimSpace(ref) @@ -522,6 +1031,65 @@ func batchImageGCSRef(provider, ref string) string { return "" } +func batchImageProviderSubmitPublicError(err error) error { + reason := strings.TrimSpace(infraerrors.Reason(err)) + switch reason { + case "VERTEX_MANAGED_GCS_BUCKET_MISSING": + return ErrBatchImageVertexGCSBucketMissing + case "BATCH_IMAGE_PROVIDER_MISSING_API_KEY": + return ErrBatchImageProviderMissingAPIKey + case "BATCH_IMAGE_PROVIDER_MISSING_SERVICE_ACCOUNT": + return ErrBatchImageProviderMissingServiceAccount + case "BATCH_IMAGE_PROVIDER_UNSUPPORTED_ACCOUNT": + return ErrBatchImageProviderUnsupportedAccount + default: + return ErrBatchImageProviderSubmitFailed + } +} + +func batchImagePublicTaskName(job *BatchImageJob) string { + if job == nil { + return "" + } + if strings.TrimSpace(job.TaskName) != "" { + return strings.TrimSpace(job.TaskName) + } + return defaultBatchImageTaskName(job.CreatedAt) +} + +func defaultBatchImageTaskName(now time.Time) string { + if now.IsZero() { + now = time.Now() + } + return now.Format("2006-01-02 15:04:05") +} + +func batchImageProviderSubmitRecordCode(err error) string { + reason := strings.TrimSpace(infraerrors.Reason(err)) + if reason == "" || reason == "BATCH_IMAGE_PROVIDER_SUBMIT_FAILED" { + return "PROVIDER_SUBMIT_FAILED" + } + return reason +} + +func parseBatchImageListTime(raw string) *time.Time { + raw = strings.TrimSpace(raw) + if raw == "" { + return nil + } + if unix, err := strconv.ParseInt(raw, 10, 64); err == nil && unix > 0 { + t := time.Unix(unix, 0) + return &t + } + if t, err := time.Parse(time.RFC3339, raw); err == nil { + return &t + } + if t, err := time.Parse("2006-01-02", raw); err == nil { + return &t + } + return nil +} + func sanitizeBatchImageMetadata(in map[string]string) map[string]string { if len(in) == 0 { return nil diff --git a/backend/internal/service/batch_image_public_test.go b/backend/internal/service/batch_image_public_test.go index 12f7d904f8..2d5c5a2533 100644 --- a/backend/internal/service/batch_image_public_test.go +++ b/backend/internal/service/batch_image_public_test.go @@ -34,10 +34,17 @@ func TestBatchImagePublicService_Submit(t *testing.T) { require.Equal(t, "queued", got.Status) require.Equal(t, BatchImageProviderGeminiAPI, got.Provider) require.Equal(t, 2, got.ItemCount) - require.Equal(t, 0.5, got.EstimatedCost) + require.Equal(t, 0.25, got.EstimatedCost) require.Len(t, repo.jobs, 1) require.Len(t, gemini.submits, 1) require.Equal(t, []string{got.ID}, queue.enqueued) + billing := svc.BillingRepo.(*fakeBatchImageBillingRepo) + require.Len(t, billing.reserves, 1) + require.Equal(t, BatchImageHoldRequestID(got.ID), billing.reserves[0].RequestID) + require.InDelta(t, 0.3, billing.reserves[0].HoldAmount, 1e-12) + require.Empty(t, billing.releases) + authCache := svc.AuthCache.(*fakeBatchImageAuthCacheInvalidator) + require.Equal(t, []int64{11}, authCache.userIDs) job := repo.jobs[got.ID] require.Equal(t, BatchImageJobStatusSubmitted, job.Status) @@ -46,6 +53,116 @@ func TestBatchImagePublicService_Submit(t *testing.T) { require.Equal(t, "files/gemini_api/output", batchImageDerefString(job.ProviderOutputRef)) require.NotNil(t, job.AccountID) require.Equal(t, int64(202), *job.AccountID) + require.Equal(t, 1, job.PricingSnapshotVersion) + require.InDelta(t, 0.25, job.BaseUnitPrice, 1e-12) + require.InDelta(t, 1.0, job.GroupRateMultiplier, 1e-12) + require.InDelta(t, 1.0, job.AccountRateMultiplier, 1e-12) + require.InDelta(t, 0.5, job.BatchDiscountMultiplier, 1e-12) + require.InDelta(t, 0.6, job.HoldMultiplier, 1e-12) + require.InDelta(t, 0.125, job.BillableUnitPrice, 1e-12) + require.InDelta(t, 0.15, job.HoldUnitPrice, 1e-12) + }) + + t.Run("combines user group image rate account rate discount and hold margin", func(t *testing.T) { + svc, repo, _, _, _ := newTestBatchImagePublicService(true) + groupID := int64(7) + accountMultiplier := 1.25 + accountRepo := svc.AccountRepo.(*publicBatchImageAccountRepo) + accountRepo.accounts[1].RateMultiplier = &accountMultiplier + svc.GroupRepo = &publicBatchImageGroupRepo{groups: map[int64]*Group{ + groupID: { + ID: groupID, + RateMultiplier: 2.0, + AllowBatchImageGeneration: true, + ImageRateIndependent: false, + BatchImageDiscountMultiplier: 0.8, + BatchImageHoldMultiplier: 0.6, + }, + }} + userRate := 0.5 + svc.UserGroupRateRepo = &publicBatchImageUserGroupRateRepo{rates: map[int64]*float64{groupID: &userRate}} + + got, err := svc.Submit(ctx, BatchImageOwner{UserID: 11, APIKeyID: 22, GroupID: &groupID}, validBatchImageSubmitRequest(), "") + require.NoError(t, err) + require.InDelta(t, 0.25, got.EstimatedCost, 1e-12) + + job := repo.jobs[got.ID] + require.InDelta(t, 0.25, job.BaseUnitPrice, 1e-12) + require.InDelta(t, 0.5, job.GroupRateMultiplier, 1e-12) + require.InDelta(t, 1.25, job.AccountRateMultiplier, 1e-12) + require.InDelta(t, 0.8, job.BatchDiscountMultiplier, 1e-12) + require.InDelta(t, 0.6, job.HoldMultiplier, 1e-12) + require.InDelta(t, 0.125, job.BillableUnitPrice, 1e-12) + require.InDelta(t, 0.09375, job.HoldUnitPrice, 1e-12) + require.InDelta(t, 0.1875, *job.HoldAmount, 1e-12) + }) + + t.Run("uses configured group 1k image price for batch image base price", func(t *testing.T) { + svc, repo, _, _, _ := newTestBatchImagePublicService(true) + groupID := int64(7) + imagePrice := 0.134 + svc.GroupRepo = &publicBatchImageGroupRepo{groups: map[int64]*Group{ + groupID: { + ID: groupID, + RateMultiplier: 1.0, + AllowBatchImageGeneration: true, + ImagePrice1K: &imagePrice, + BatchImageDiscountMultiplier: 0.5, + BatchImageHoldMultiplier: 0.6, + }, + }} + + got, err := svc.Submit(ctx, BatchImageOwner{UserID: 11, APIKeyID: 22, GroupID: &groupID}, validBatchImageSubmitRequest(), "") + require.NoError(t, err) + require.InDelta(t, 0.134, got.EstimatedCost, 1e-12) + + job := repo.jobs[got.ID] + require.InDelta(t, 0.134, job.BaseUnitPrice, 1e-12) + require.InDelta(t, 0.067, job.BillableUnitPrice, 1e-12) + require.InDelta(t, 0.0804, job.HoldUnitPrice, 1e-12) + require.InDelta(t, 0.1608, *job.HoldAmount, 1e-12) + }) + + t.Run("pricing missing rejects before provider submit", func(t *testing.T) { + svc, repo, queue, gemini, _ := newTestBatchImagePublicService(true) + svc.Pricing = &fakeBatchImagePricingResolver{err: ErrBatchImageSettlementPricingMissing} + + _, err := svc.Submit(ctx, testBatchImageOwner(), validBatchImageSubmitRequest(), "") + require.ErrorIs(t, err, ErrBatchImageSettlementPricingMissing) + require.Empty(t, repo.jobs) + require.Empty(t, queue.enqueued) + require.Empty(t, gemini.submits) + }) + + t.Run("group batch image disabled rejects before provider submit", func(t *testing.T) { + svc, repo, queue, gemini, _ := newTestBatchImagePublicService(true) + groupID := int64(7) + svc.GroupRepo = &publicBatchImageGroupRepo{groups: map[int64]*Group{ + groupID: { + ID: groupID, + RateMultiplier: 1, + AllowBatchImageGeneration: false, + BatchImageDiscountMultiplier: 0.5, + BatchImageHoldMultiplier: 0.6, + }, + }} + + _, err := svc.Submit(ctx, BatchImageOwner{UserID: 11, APIKeyID: 22, GroupID: &groupID}, validBatchImageSubmitRequest(), "") + require.ErrorIs(t, err, ErrBatchImageGroupDisabled) + require.Empty(t, repo.jobs) + require.Empty(t, queue.enqueued) + require.Empty(t, gemini.submits) + }) + + t.Run("group pricing load failure rejects before provider submit", func(t *testing.T) { + svc, repo, queue, gemini, _ := newTestBatchImagePublicService(true) + groupID := int64(404) + + _, err := svc.Submit(ctx, BatchImageOwner{UserID: 11, APIKeyID: 22, GroupID: &groupID}, validBatchImageSubmitRequest(), "") + require.ErrorIs(t, err, ErrBatchImageSettlementPricingMissing) + require.Empty(t, repo.jobs) + require.Empty(t, queue.enqueued) + require.Empty(t, gemini.submits) }) t.Run("generates custom ids deterministically", func(t *testing.T) { @@ -108,27 +225,72 @@ func TestBatchImagePublicService_Submit(t *testing.T) { require.Len(t, vertex.submits, 1) }) + t.Run("insufficient balance rejects before provider submit", func(t *testing.T) { + svc, repo, queue, gemini, _ := newTestBatchImagePublicService(true) + billing := &fakeBatchImageBillingRepo{err: ErrBatchImageInsufficientBalance} + svc.BillingRepo = billing + + _, err := svc.Submit(ctx, testBatchImageOwner(), validBatchImageSubmitRequest(), "") + require.ErrorIs(t, err, ErrBatchImageInsufficientBalance) + require.Empty(t, queue.enqueued) + require.Empty(t, gemini.submits) + require.Len(t, billing.reserves, 1) + require.Empty(t, billing.releases) + require.Len(t, repo.jobs, 1) + for _, job := range repo.jobs { + require.Equal(t, BatchImageJobStatusFailed, job.Status) + require.Equal(t, "INSUFFICIENT_BALANCE", batchImageDerefString(job.LastErrorCode)) + require.NotNil(t, job.UserDeletedAt) + } + }) + t.Run("provider failure marks failed and does not enqueue", func(t *testing.T) { svc, repo, queue, gemini, _ := newTestBatchImagePublicService(true) gemini.submitErr = errors.New("projects/secret-provider-job failed") + billing := svc.BillingRepo.(*fakeBatchImageBillingRepo) _, err := svc.Submit(ctx, testBatchImageOwner(), validBatchImageSubmitRequest(), "") require.ErrorIs(t, err, ErrBatchImageProviderSubmitFailed) require.Empty(t, queue.enqueued) + require.Len(t, billing.reserves, 1) + require.Len(t, billing.releases, 1) + require.Equal(t, BatchImageReleaseRequestID(billing.reserves[0].BatchID), billing.releases[0].RequestID) require.Len(t, repo.jobs, 1) for _, job := range repo.jobs { require.Equal(t, BatchImageJobStatusFailed, job.Status) require.Equal(t, "PROVIDER_SUBMIT_FAILED", batchImageDerefString(job.LastErrorCode)) require.Equal(t, "upstream provider operation failed", batchImageDerefString(job.LastErrorMessage)) + require.NotNil(t, job.UserDeletedAt) + } + }) + + t.Run("provider failure with release failure enqueues billing retry", func(t *testing.T) { + svc, repo, queue, gemini, _ := newTestBatchImagePublicService(true) + gemini.submitErr = errors.New("projects/secret-provider-job failed") + billing := svc.BillingRepo.(*fakeBatchImageBillingRepo) + billing.releaseErr = errors.New("billing database timeout") + + _, err := svc.Submit(ctx, testBatchImageOwner(), validBatchImageSubmitRequest(), "") + require.ErrorIs(t, err, ErrBatchImageBillingHoldFailed) + require.Len(t, billing.reserves, 1) + require.Len(t, billing.releases, 1) + require.Len(t, repo.jobs, 1) + for _, job := range repo.jobs { + require.Equal(t, BatchImageJobStatusFailed, job.Status) + require.Equal(t, "BILLING_RELEASE_FAILED", batchImageDerefString(job.LastErrorCode)) + require.Equal(t, []string{job.BatchID}, queue.enqueued) } }) t.Run("queue failure is recorded after provider submit", func(t *testing.T) { svc, repo, queue, _, _ := newTestBatchImagePublicService(true) queue.err = errors.New("redis unavailable") + billing := svc.BillingRepo.(*fakeBatchImageBillingRepo) _, err := svc.Submit(ctx, testBatchImageOwner(), validBatchImageSubmitRequest(), "") require.ErrorIs(t, err, ErrBatchImageQueueFailed) + require.Len(t, billing.reserves, 1) + require.Empty(t, billing.releases) require.Len(t, repo.jobs, 1) for _, job := range repo.jobs { require.Equal(t, BatchImageJobStatusSubmitted, job.Status) @@ -175,6 +337,136 @@ func TestBatchImagePublicService_Submit(t *testing.T) { }) } +func TestBatchImagePublicService_List(t *testing.T) { + ctx := context.Background() + svc, repo, _, _, _ := newTestBatchImagePublicService(true) + visibleKeyID := int64(22) + otherKeyID := int64(23) + + repo.jobs["visible-1"] = &BatchImageJob{ + BatchID: "visible-1", + UserID: 11, + APIKeyID: &visibleKeyID, + Status: BatchImageJobStatusCompleted, + Provider: BatchImageProviderVertex, + Model: "gemini-3.1-flash-lite-image", + ItemCount: 1, + CreatedAt: time.Now(), + } + repo.jobs["hidden-other-key"] = &BatchImageJob{ + BatchID: "hidden-other-key", + UserID: 11, + APIKeyID: &otherKeyID, + Status: BatchImageJobStatusCompleted, + Provider: BatchImageProviderVertex, + Model: "gemini-3.1-flash-lite-image", + ItemCount: 1, + CreatedAt: time.Now(), + } + + got, err := svc.List(ctx, BatchImageOwner{UserID: 11, APIKeyID: visibleKeyID}, BatchImageJobsQuery{Limit: 20}) + require.NoError(t, err) + require.Equal(t, "list", got.Object) + require.Len(t, got.Data, 1) + require.Equal(t, "visible-1", got.Data[0].ID) + require.False(t, got.HasMore) +} + +func TestBatchImagePublicService_ListModels(t *testing.T) { + ctx := context.Background() + + t.Run("requires explicit account model mapping", func(t *testing.T) { + svc, _, _, _, _ := newTestBatchImagePublicService(true) + + got, err := svc.ListModels(ctx, testBatchImageOwner()) + require.NoError(t, err) + require.Equal(t, "list", got.Object) + require.Empty(t, got.Data) + }) + + t.Run("returns priced models from selected account group", func(t *testing.T) { + svc, _, _, _, _ := newTestBatchImagePublicService(true) + groupID := int64(7) + svc.GroupRepo = &publicBatchImageGroupRepo{groups: map[int64]*Group{ + groupID: { + ID: groupID, + RateMultiplier: 1, + AllowBatchImageGeneration: true, + BatchImageDiscountMultiplier: 0.5, + BatchImageHoldMultiplier: 0.6, + }, + }} + accountRepo := svc.AccountRepo.(*publicBatchImageAccountRepo) + accountRepo.accounts = []Account{testBatchImageMappedAccount(303, AccountTypeAPIKey, map[string]any{ + "gemini-2.5-flash-image": "gemini-2.5-flash-image", + })} + + got, err := svc.ListModels(ctx, BatchImageOwner{UserID: 11, APIKeyID: 22, GroupID: &groupID}) + require.NoError(t, err) + require.Equal(t, []BatchImagePublicModel{{ + ID: "gemini-2.5-flash-image", + Object: "image.batch.model", + Provider: BatchImageProviderGeminiAPI, + }, { + ID: "gemini-2.5-flash-image", + Object: "image.batch.model", + Provider: BatchImageProviderVertex, + }}, got.Data) + }) + + t.Run("expands wildcard mappings against batch image candidates", func(t *testing.T) { + svc, _, _, _, _ := newTestBatchImagePublicService(true) + accountRepo := svc.AccountRepo.(*publicBatchImageAccountRepo) + accountRepo.accounts = []Account{testBatchImageMappedAccount(303, AccountTypeAPIKey, map[string]any{ + "gemini-3.1-*": "gemini-3.1-flash-lite-image", + })} + + got, err := svc.ListModels(ctx, testBatchImageOwner()) + require.NoError(t, err) + require.NotEmpty(t, got.Data) + ids := make([]string, 0, len(got.Data)) + for _, model := range got.Data { + ids = append(ids, model.ID) + } + require.Contains(t, ids, "gemini-3.1-flash-image") + require.Contains(t, ids, "gemini-3.1-flash-lite-image") + require.NotContains(t, ids, "gemini-2.5-flash-image") + }) + + t.Run("filters models without batch image pricing", func(t *testing.T) { + svc, _, _, _, _ := newTestBatchImagePublicService(true) + svc.Pricing = &fakeBatchImagePricingResolver{ + unitPrice: 0.25, + missingModels: map[string]bool{"gemini-3.1-flash-lite-image": true}, + } + accountRepo := svc.AccountRepo.(*publicBatchImageAccountRepo) + accountRepo.accounts = []Account{testBatchImageMappedAccount(303, AccountTypeAPIKey, map[string]any{ + "gemini-2.5-flash-image": "gemini-2.5-flash-image", + "gemini-3.1-flash-lite-image": "gemini-3.1-flash-lite-image", + })} + + got, err := svc.ListModels(ctx, testBatchImageOwner()) + require.NoError(t, err) + ids := make([]string, 0, len(got.Data)) + for _, model := range got.Data { + ids = append(ids, model.ID) + } + require.Contains(t, ids, "gemini-2.5-flash-image") + require.NotContains(t, ids, "gemini-3.1-flash-lite-image") + }) + + t.Run("rejects when group disables batch image", func(t *testing.T) { + svc, _, _, _, _ := newTestBatchImagePublicService(true) + groupID := int64(7) + svc.GroupRepo = &publicBatchImageGroupRepo{groups: map[int64]*Group{ + groupID: {ID: groupID, AllowBatchImageGeneration: false}, + }} + + _, err := svc.ListModels(ctx, BatchImageOwner{UserID: 11, APIKeyID: 22, GroupID: &groupID}) + require.ErrorIs(t, err, ErrBatchImageGroupDisabled) + }) +} + func TestBatchImagePublicService_StatusItemsAndCancel(t *testing.T) { ctx := context.Background() @@ -251,10 +543,12 @@ func TestBatchImagePublicService_StatusItemsAndCancel(t *testing.T) { require.ErrorIs(t, err, ErrBatchImageJobNotFound) }) - t.Run("cancel active job calls provider and marks cancelled", func(t *testing.T) { - svc, repo, _, gemini, _ := newTestBatchImagePublicService(true) + t.Run("cancel active job calls provider and waits for confirmed terminal state", func(t *testing.T) { + svc, repo, queue, gemini, _ := newTestBatchImagePublicService(true) apiKeyID := int64(22) accountID := int64(101) + holdAmount := 0.5 + holdID := BatchImageHoldRequestID("imgbatch_cancel") repo.jobs["imgbatch_cancel"] = &BatchImageJob{ BatchID: "imgbatch_cancel", UserID: 11, @@ -264,15 +558,21 @@ func TestBatchImagePublicService_StatusItemsAndCancel(t *testing.T) { Model: "gemini-2.5-flash-image", Status: BatchImageJobStatusSubmitted, ProviderJobName: batchImageStringPtr("providers/internal/job"), + EstimatedCost: holdAmount, + HoldAmount: &holdAmount, + HoldID: &holdID, CreatedAt: time.Now(), } got, err := svc.Cancel(ctx, testBatchImageOwner(), "imgbatch_cancel") require.NoError(t, err) - require.Equal(t, "cancelled", got.Status) + require.Equal(t, "queued", got.Status) require.Equal(t, 1, gemini.cancelCount) - require.Equal(t, BatchImageJobStatusCancelled, repo.jobs["imgbatch_cancel"].Status) - require.Contains(t, repo.events["imgbatch_cancel"], "job_cancelled") + billing := svc.BillingRepo.(*fakeBatchImageBillingRepo) + require.Empty(t, billing.releases) + require.Equal(t, []string{"imgbatch_cancel"}, queue.enqueued) + require.Equal(t, BatchImageJobStatusSubmitted, repo.jobs["imgbatch_cancel"].Status) + require.Contains(t, repo.events["imgbatch_cancel"], "job_cancel_requested") }) t.Run("cancel terminal job is idempotent", func(t *testing.T) { @@ -331,7 +631,9 @@ func newTestBatchImagePublicService(enabled bool) (*BatchImagePublicService, *fa gemini, vertex, ), - Pricing: &fakeBatchImagePricingResolver{unitPrice: 0.25}, + Pricing: &fakeBatchImagePricingResolver{unitPrice: 0.25}, + BillingRepo: &fakeBatchImageBillingRepo{}, + AuthCache: &fakeBatchImageAuthCacheInvalidator{}, Config: &config.Config{BatchImage: config.BatchImageConfig{ Enabled: enabled, MaxItemsPerJobDefault: 2, @@ -347,6 +649,24 @@ func testBatchImageOwner() BatchImageOwner { return BatchImageOwner{UserID: 11, APIKeyID: 22} } +type fakeBatchImageAuthCacheInvalidator struct { + keys []string + userIDs []int64 + groupIDs []int64 +} + +func (f *fakeBatchImageAuthCacheInvalidator) InvalidateAuthCacheByKey(_ context.Context, key string) { + f.keys = append(f.keys, key) +} + +func (f *fakeBatchImageAuthCacheInvalidator) InvalidateAuthCacheByUserID(_ context.Context, userID int64) { + f.userIDs = append(f.userIDs, userID) +} + +func (f *fakeBatchImageAuthCacheInvalidator) InvalidateAuthCacheByGroupID(_ context.Context, groupID int64) { + f.groupIDs = append(f.groupIDs, groupID) +} + func validBatchImageSubmitRequest() BatchImageSubmitRequest { return BatchImageSubmitRequest{ Model: "gemini-2.5-flash-image", @@ -376,6 +696,12 @@ func testBatchImageAccount(id int64, accountType string) Account { } } +func testBatchImageMappedAccount(id int64, accountType string, mapping map[string]any) Account { + account := testBatchImageAccount(id, accountType) + account.Credentials["model_mapping"] = mapping + return account +} + func requireBatchImagePublicJSONHasNoInternals(t *testing.T, body string) { t.Helper() for _, forbidden := range []string{ @@ -517,3 +843,30 @@ func (p *publicBatchImageProvider) Cleanup(_ context.Context, _ *BatchImageJob, var _ BatchImageAccountSelectionRepository = (*publicBatchImageAccountRepo)(nil) var _ BatchImageQueue = (*publicBatchImageQueue)(nil) var _ BatchImageProvider = (*publicBatchImageProvider)(nil) + +type publicBatchImageGroupRepo struct { + groups map[int64]*Group +} + +func (r *publicBatchImageGroupRepo) GetByIDLite(_ context.Context, id int64) (*Group, error) { + if r != nil && r.groups != nil { + if group, ok := r.groups[id]; ok { + return group, nil + } + } + return nil, ErrGroupNotFound +} + +type publicBatchImageUserGroupRateRepo struct { + rates map[int64]*float64 +} + +func (r *publicBatchImageUserGroupRateRepo) GetByUserAndGroup(_ context.Context, _ int64, groupID int64) (*float64, error) { + if r != nil && r.rates != nil { + return r.rates[groupID], nil + } + return nil, nil +} + +var _ BatchImageGroupPricingRepository = (*publicBatchImageGroupRepo)(nil) +var _ BatchImageUserGroupRateRepository = (*publicBatchImageUserGroupRateRepo)(nil) diff --git a/backend/internal/service/batch_image_settlement.go b/backend/internal/service/batch_image_settlement.go index 5c2e477f7b..870883b360 100644 --- a/backend/internal/service/batch_image_settlement.go +++ b/backend/internal/service/batch_image_settlement.go @@ -16,6 +16,7 @@ import ( const ( batchImageSettlementRequestPrefix = "batch_image_settlement:" batchImageSettlementRetryDelay = time.Minute + batchImageCostEpsilon = 0.00000001 ) type BatchImagePricingResolver interface { @@ -51,10 +52,12 @@ func (r *BatchImageModelPricingResolver) BatchImageUnitPrice(ctx context.Context } type BatchImageSettlementService struct { - Repo BatchImageRepository - BillingRepo UsageBillingRepository - Pricing BatchImagePricingResolver - Config *config.Config + Repo BatchImageRepository + BillingRepo UsageBillingRepository + UsageLogRepo UsageLogRepository + Pricing BatchImagePricingResolver + AuthCache APIKeyAuthCacheInvalidator + Config *config.Config } type BatchImageSettlementResult struct { @@ -82,7 +85,7 @@ func (s *BatchImageSettlementService) Settle(ctx context.Context, batchID string SuccessCount: job.SuccessCount, FailCount: job.FailCount, ManifestHash: manifestHash, - RequestID: BatchImageSettlementRequestID(job.BatchID), + RequestID: BatchImageCaptureRequestID(job.BatchID), } if job.ActualCost != nil { result.ActualCost = *job.ActualCost @@ -94,7 +97,7 @@ func (s *BatchImageSettlementService) Settle(ctx context.Context, batchID string if job.Status != BatchImageJobStatusSettling { return nil, ErrBatchImageSettlementInvalidStatus } - if job.SuccessCount < 0 || job.FailCount < 0 || job.ItemCount < 0 { + if job.SuccessCount < 0 || job.FailCount < 0 || job.ItemCount < 0 || job.SuccessCount+job.FailCount > job.ItemCount { return nil, ErrBatchImageSettlementInvalidCounts } if strings.TrimSpace(batchImageDerefString(job.ManifestHash)) != "" && batchImageDerefString(job.ManifestHash) != manifestHash { @@ -107,7 +110,7 @@ func (s *BatchImageSettlementService) Settle(ctx context.Context, batchID string return nil, ErrBatchImageSettlementMissingAccountID } - unitPrice, err := s.Pricing.BatchImageUnitPrice(ctx, job) + unitPrice, err := s.settlementUnitPrice(ctx, job) if err != nil { return nil, err } @@ -116,24 +119,22 @@ func (s *BatchImageSettlementService) Settle(ctx context.Context, batchID string } actualCost := float64(job.SuccessCount) * unitPrice result.ActualCost = actualCost - - cmd := &UsageBillingCommand{ - RequestID: result.RequestID, - APIKeyID: *job.APIKeyID, - RequestPayloadHash: manifestHash, - UserID: job.UserID, - AccountID: *job.AccountID, - Model: job.Model, - BillingType: BillingTypeBalance, - ImageCount: job.SuccessCount, - MediaType: "image", - BalanceCost: actualCost, + holdAmount := job.EstimatedCost + if job.HoldAmount != nil { + holdAmount = *job.HoldAmount } - if _, err := s.BillingRepo.Apply(ctx, cmd); err != nil { + if actualCost-holdAmount > batchImageCostEpsilon { + msg := fmt.Sprintf("actual cost %.10f exceeds held amount %.10f", actualCost, holdAmount) + _ = s.Repo.SetBatchImageJobSettlementFailed(ctx, job.BatchID, "SETTLEMENT_COST_EXCEEDS_HOLD", msg) + return nil, ErrBatchImageSettlementCostExceedsHold + } + + if err := captureBatchImageBalanceHold(ctx, s.BillingRepo, job, actualCost, manifestHash); err != nil { msg := truncateBatchImageMessage(err.Error(), batchImageMaxErrorMessageLength) _ = s.Repo.SetBatchImageJobSettlementFailed(ctx, job.BatchID, "SETTLEMENT_BILLING_FAILED", msg) - return nil, ErrBatchImageSettlementBillingFailed.WithCause(err) + return nil, err } + s.invalidateAuthCache(ctx, job.UserID) now := time.Now() outputExpiresAt := now.Add(s.outputRetentionAfterTerminal()) @@ -154,10 +155,64 @@ func (s *BatchImageSettlementService) Settle(ctx context.Context, batchID string }); err != nil { return nil, err } + s.recordUsageLog(ctx, job, actualCost, result.RequestID, now) return result, nil } +func (s *BatchImageSettlementService) recordUsageLog(ctx context.Context, job *BatchImageJob, actualCost float64, requestID string, createdAt time.Time) { + if s == nil || s.UsageLogRepo == nil || job == nil || job.APIKeyID == nil || job.AccountID == nil { + return + } + billingMode := string(BillingModeImage) + accountRateMultiplier := job.AccountRateMultiplier + inboundEndpoint := "/v1/images/batches" + upstreamEndpoint := "vertex:batchPredictionJobs" + imageSize := "1K" + usageLog := &UsageLog{ + UserID: job.UserID, + APIKeyID: *job.APIKeyID, + AccountID: *job.AccountID, + RequestID: strings.TrimSpace(requestID), + Model: job.Model, + RequestedModel: job.Model, + InboundEndpoint: &inboundEndpoint, + UpstreamEndpoint: &upstreamEndpoint, + ImageCount: job.SuccessCount, + ImageOutputCost: actualCost, + TotalCost: actualCost, + ActualCost: actualCost, + RateMultiplier: job.GroupRateMultiplier * job.BatchDiscountMultiplier, + AccountRateMultiplier: &accountRateMultiplier, + BillingType: BillingTypeBalance, + RequestType: RequestTypeSync, + BillingMode: &billingMode, + ImageSize: &imageSize, + CreatedAt: createdAt, + } + writeUsageLogBestEffort(ctx, s.UsageLogRepo, usageLog, "service.batch_image_settlement") +} + +func (s *BatchImageSettlementService) invalidateAuthCache(ctx context.Context, userID int64) { + if s != nil && s.AuthCache != nil && userID > 0 { + s.AuthCache.InvalidateAuthCacheByUserID(ctx, userID) + } +} + +func (s *BatchImageSettlementService) settlementUnitPrice(ctx context.Context, job *BatchImageJob) (float64, error) { + if job != nil && job.PricingSnapshotVersion >= 1 { + if job.BillableUnitPrice < 0 { + return 0, ErrBatchImageSettlementPricingMissing + } + return job.BillableUnitPrice, nil + } + unitPrice, err := s.Pricing.BatchImageUnitPrice(ctx, job) + if err != nil { + return 0, err + } + return unitPrice, nil +} + func (s *BatchImageSettlementService) outputRetentionAfterTerminal() time.Duration { if s != nil && s.Config != nil && s.Config.BatchImage.OutputRetentionAfterTerminalHours > 0 { return time.Duration(s.Config.BatchImage.OutputRetentionAfterTerminalHours) * time.Hour diff --git a/backend/internal/service/batch_image_settlement_test.go b/backend/internal/service/batch_image_settlement_test.go index a3a60c0c0a..c09f0a6cf7 100644 --- a/backend/internal/service/batch_image_settlement_test.go +++ b/backend/internal/service/batch_image_settlement_test.go @@ -25,24 +25,22 @@ func TestBatchImageSettlementService_SettlesAndChargesSuccessfulImagesOnly(t *te result, err := svc.Settle(context.Background(), job.BatchID) require.NoError(t, err) require.Equal(t, 0.75, result.ActualCost) - require.Equal(t, "batch_image_settlement:"+job.BatchID, result.RequestID) + require.Equal(t, BatchImageCaptureRequestID(job.BatchID), result.RequestID) require.False(t, result.AlreadySettled) require.Equal(t, BatchImageJobStatusCompleted, repo.jobs[job.BatchID].Status) require.NotNil(t, repo.jobs[job.BatchID].ActualCost) require.Equal(t, 0.75, *repo.jobs[job.BatchID].ActualCost) require.NotEmpty(t, batchImageDerefString(repo.jobs[job.BatchID].ManifestHash)) require.NotNil(t, repo.jobs[job.BatchID].SettledAt) - require.Len(t, billing.commands, 1) - require.Equal(t, int64(321), billing.commands[0].APIKeyID) - require.Equal(t, job.UserID, billing.commands[0].UserID) - require.Equal(t, int64(654), billing.commands[0].AccountID) - require.Equal(t, job.Model, billing.commands[0].Model) - require.Equal(t, 3, billing.commands[0].ImageCount) - require.Equal(t, 0.75, billing.commands[0].BalanceCost) - require.Equal(t, "image", billing.commands[0].MediaType) - require.NotContains(t, fmt.Sprintf("%+v", billing.commands[0]), batchImageTestData) - require.NotContains(t, fmt.Sprintf("%+v", billing.commands[0]), "gs://") - require.NotContains(t, fmt.Sprintf("%+v", billing.commands[0]), "prompt") + require.Len(t, billing.captures, 1) + require.Equal(t, int64(321), billing.captures[0].APIKeyID) + require.Equal(t, job.UserID, billing.captures[0].UserID) + require.Equal(t, job.BatchID, billing.captures[0].BatchID) + require.Equal(t, 0.75, billing.captures[0].ActualAmount) + require.Equal(t, 1.25, billing.captures[0].HoldAmount) + require.NotContains(t, fmt.Sprintf("%+v", billing.captures[0]), batchImageTestData) + require.NotContains(t, fmt.Sprintf("%+v", billing.captures[0]), "gs://") + require.NotContains(t, fmt.Sprintf("%+v", billing.captures[0]), "prompt") } func TestBatchImageSettlementService_ZeroSuccessCanComplete(t *testing.T) { @@ -59,8 +57,8 @@ func TestBatchImageSettlementService_ZeroSuccessCanComplete(t *testing.T) { require.NoError(t, err) require.Equal(t, 0.0, result.ActualCost) require.Equal(t, BatchImageJobStatusCompleted, repo.jobs[job.BatchID].Status) - require.Len(t, billing.commands, 1) - require.Equal(t, 0.0, billing.commands[0].BalanceCost) + require.Len(t, billing.captures, 1) + require.Equal(t, 0.0, billing.captures[0].ActualAmount) } func TestBatchImageSettlementService_CompletedJobReturnsAlreadySettledWithoutBilling(t *testing.T) { @@ -77,21 +75,21 @@ func TestBatchImageSettlementService_CompletedJobReturnsAlreadySettledWithoutBil require.NoError(t, err) require.True(t, result.AlreadySettled) require.Equal(t, 0.5, result.ActualCost) - require.Empty(t, billing.commands) + require.Empty(t, billing.captures) } func TestBatchImageSettlementService_IdempotentAfterBillingCrash(t *testing.T) { repo := newFakeBatchImageRepository() job := testSettlingBatchImageJob("imgbatch_crash") repo.jobs[job.BatchID] = job - billing := &fakeBatchImageBillingRepo{alreadyApplied: map[string]bool{BatchImageSettlementRequestID(job.BatchID): true}} + billing := &fakeBatchImageBillingRepo{alreadyApplied: map[string]bool{BatchImageCaptureRequestID(job.BatchID): true}} svc := &BatchImageSettlementService{Repo: repo, BillingRepo: billing, Pricing: &fakeBatchImagePricingResolver{unitPrice: 0.25}} result, err := svc.Settle(context.Background(), job.BatchID) require.NoError(t, err) require.Equal(t, 0.5, result.ActualCost) require.Equal(t, BatchImageJobStatusCompleted, repo.jobs[job.BatchID].Status) - require.Len(t, billing.commands, 1) + require.Len(t, billing.captures, 1) } func TestBatchImageSettlementService_ValidationErrors(t *testing.T) { @@ -104,6 +102,7 @@ func TestBatchImageSettlementService_ValidationErrors(t *testing.T) { {name: "invalid_status", mutate: func(j *BatchImageJob) { j.Status = BatchImageJobStatusRunning }, want: ErrBatchImageSettlementInvalidStatus}, {name: "negative_success_count", mutate: func(j *BatchImageJob) { j.SuccessCount = -1 }, want: ErrBatchImageSettlementInvalidCounts}, {name: "negative_fail_count", mutate: func(j *BatchImageJob) { j.FailCount = -1 }, want: ErrBatchImageSettlementInvalidCounts}, + {name: "counts_exceed_item_count", mutate: func(j *BatchImageJob) { j.SuccessCount = 2; j.FailCount = 2; j.ItemCount = 3 }, want: ErrBatchImageSettlementInvalidCounts}, {name: "missing_api_key", mutate: func(j *BatchImageJob) { j.APIKeyID = nil }, want: ErrBatchImageSettlementMissingAPIKeyID}, {name: "missing_account", mutate: func(j *BatchImageJob) { j.AccountID = nil }, want: ErrBatchImageSettlementMissingAccountID}, {name: "pricing_missing", pricing: &fakeBatchImagePricingResolver{err: ErrBatchImageSettlementPricingMissing}, want: ErrBatchImageSettlementPricingMissing}, @@ -127,12 +126,61 @@ func TestBatchImageSettlementService_ValidationErrors(t *testing.T) { _, err := svc.Settle(context.Background(), job.BatchID) require.ErrorIs(t, err, tt.want) - require.Empty(t, billing.commands) + require.Empty(t, billing.captures) require.NotEqual(t, BatchImageJobStatusCompleted, repo.jobs[job.BatchID].Status) }) } } +func TestBatchImageSettlementService_CostExceedingHoldDoesNotCharge(t *testing.T) { + repo := newFakeBatchImageRepository() + job := testSettlingBatchImageJob("imgbatch_cost_over_hold") + job.SuccessCount = 2 + job.FailCount = 0 + job.ItemCount = 2 + holdAmount := 0.5 + job.HoldAmount = &holdAmount + job.EstimatedCost = holdAmount + repo.jobs[job.BatchID] = job + billing := &fakeBatchImageBillingRepo{} + svc := &BatchImageSettlementService{Repo: repo, BillingRepo: billing, Pricing: &fakeBatchImagePricingResolver{unitPrice: 0.50}} + + _, err := svc.Settle(context.Background(), job.BatchID) + require.ErrorIs(t, err, ErrBatchImageSettlementCostExceedsHold) + require.Empty(t, billing.captures) + require.Equal(t, BatchImageJobStatusSettling, repo.jobs[job.BatchID].Status) + require.Equal(t, "SETTLEMENT_COST_EXCEEDS_HOLD", batchImageDerefString(repo.jobs[job.BatchID].LastErrorCode)) +} + +func TestBatchImageSettlementService_UsesSubmittedPricingSnapshot(t *testing.T) { + repo := newFakeBatchImageRepository() + job := testSettlingBatchImageJob("imgbatch_snapshot") + job.SuccessCount = 2 + job.FailCount = 0 + job.ItemCount = 2 + job.PricingSnapshotVersion = 1 + job.BaseUnitPrice = 0.25 + job.GroupRateMultiplier = 1 + job.AccountRateMultiplier = 1 + job.BatchDiscountMultiplier = 1 + job.HoldMultiplier = 1.1 + job.BillableUnitPrice = 0.25 + job.HoldUnitPrice = 0.275 + holdAmount := 0.55 + job.HoldAmount = &holdAmount + job.EstimatedCost = 0.5 + repo.jobs[job.BatchID] = job + billing := &fakeBatchImageBillingRepo{} + svc := &BatchImageSettlementService{Repo: repo, BillingRepo: billing, Pricing: &fakeBatchImagePricingResolver{unitPrice: 0.50}} + + result, err := svc.Settle(context.Background(), job.BatchID) + require.NoError(t, err) + require.InDelta(t, 0.5, result.ActualCost, 1e-12) + require.Len(t, billing.captures, 1) + require.InDelta(t, 0.5, billing.captures[0].ActualAmount, 1e-12) + require.InDelta(t, 0.55, billing.captures[0].HoldAmount, 1e-12) +} + func TestBatchImageSettlementService_BillingFailureLeavesSettlingAndRecordsError(t *testing.T) { repo := newFakeBatchImageRepository() job := testSettlingBatchImageJob("imgbatch_billing_fail") @@ -145,7 +193,7 @@ func TestBatchImageSettlementService_BillingFailureLeavesSettlingAndRecordsError require.Equal(t, BatchImageJobStatusSettling, repo.jobs[job.BatchID].Status) require.Equal(t, "SETTLEMENT_BILLING_FAILED", batchImageDerefString(repo.jobs[job.BatchID].LastErrorCode)) require.Contains(t, batchImageDerefString(repo.jobs[job.BatchID].LastErrorMessage), "temporary billing timeout") - require.NotNil(t, billing.commands[0]) + require.NotNil(t, billing.captures[0]) } func TestBatchImagePipelineProcessor_SettlesQueuedSettlingJob(t *testing.T) { @@ -163,7 +211,7 @@ func TestBatchImagePipelineProcessor_SettlesQueuedSettlingJob(t *testing.T) { require.NoError(t, err) require.True(t, result.Terminal) require.Equal(t, BatchImageJobStatusCompleted, repo.jobs[job.BatchID].Status) - require.Len(t, billing.commands, 1) + require.Len(t, billing.captures, 1) } func TestBatchImagePipelineProcessor_RequeuesTransientSettlementFailure(t *testing.T) { @@ -214,10 +262,10 @@ func TestBatchImageSettlementBillingRequestIDs(t *testing.T) { _, err = svc.Settle(context.Background(), second.BatchID) require.NoError(t, err) - require.Len(t, billing.commands, 2) - require.Equal(t, "batch_image_settlement:"+first.BatchID, billing.commands[0].RequestID) - require.Equal(t, "batch_image_settlement:"+second.BatchID, billing.commands[1].RequestID) - require.NotEqual(t, billing.commands[0].RequestID, billing.commands[1].RequestID) + require.Len(t, billing.captures, 2) + require.Equal(t, BatchImageCaptureRequestID(first.BatchID), billing.captures[0].RequestID) + require.Equal(t, BatchImageCaptureRequestID(second.BatchID), billing.captures[1].RequestID) + require.NotEqual(t, billing.captures[0].RequestID, billing.captures[1].RequestID) require.Len(t, billing.seen, 2) } @@ -226,6 +274,8 @@ func testSettlingBatchImageJob(batchID string) *BatchImageJob { accountID := int64(654) providerJobName := "providers/job" outputRef := "files/output" + holdAmount := 1.25 + holdID := BatchImageHoldRequestID(batchID) return &BatchImageJob{ BatchID: batchID, UserID: 123, @@ -239,26 +289,39 @@ func testSettlingBatchImageJob(batchID string) *BatchImageJob { ItemCount: 3, SuccessCount: 2, FailCount: 1, + EstimatedCost: holdAmount, + HoldAmount: &holdAmount, + HoldID: &holdID, } } type fakeBatchImagePricingResolver struct { - unitPrice float64 - err error + unitPrice float64 + missingModels map[string]bool + err error } -func (r *fakeBatchImagePricingResolver) BatchImageUnitPrice(context.Context, *BatchImageJob) (float64, error) { +func (r *fakeBatchImagePricingResolver) BatchImageUnitPrice(_ context.Context, job *BatchImageJob) (float64, error) { if r.err != nil { return 0, r.err } + if job != nil && r.missingModels[job.Model] { + return 0, ErrBatchImageSettlementPricingMissing + } return r.unitPrice, nil } type fakeBatchImageBillingRepo struct { commands []*UsageBillingCommand + reserves []*BatchImageBalanceHoldCommand + captures []*BatchImageBalanceHoldCommand + releases []*BatchImageBalanceHoldCommand seen map[string]struct{} alreadyApplied map[string]bool err error + reserveErr error + captureErr error + releaseErr error } func (r *fakeBatchImageBillingRepo) Apply(_ context.Context, cmd *UsageBillingCommand) (*UsageBillingApplyResult, error) { @@ -281,6 +344,50 @@ func (r *fakeBatchImageBillingRepo) Apply(_ context.Context, cmd *UsageBillingCo return &UsageBillingApplyResult{Applied: true}, nil } +func (r *fakeBatchImageBillingRepo) ReserveBatchImageBalance(_ context.Context, cmd *BatchImageBalanceHoldCommand) (*BatchImageBalanceHoldResult, error) { + if r.reserveErr != nil { + r.reserves = append(r.reserves, cmd) + return nil, r.reserveErr + } + return r.applyHold(cmd, &r.reserves) +} + +func (r *fakeBatchImageBillingRepo) CaptureBatchImageBalance(_ context.Context, cmd *BatchImageBalanceHoldCommand) (*BatchImageBalanceHoldResult, error) { + if r.captureErr != nil { + r.captures = append(r.captures, cmd) + return nil, r.captureErr + } + return r.applyHold(cmd, &r.captures) +} + +func (r *fakeBatchImageBillingRepo) ReleaseBatchImageBalance(_ context.Context, cmd *BatchImageBalanceHoldCommand) (*BatchImageBalanceHoldResult, error) { + if r.releaseErr != nil { + r.releases = append(r.releases, cmd) + return nil, r.releaseErr + } + return r.applyHold(cmd, &r.releases) +} + +func (r *fakeBatchImageBillingRepo) applyHold(cmd *BatchImageBalanceHoldCommand, calls *[]*BatchImageBalanceHoldCommand) (*BatchImageBalanceHoldResult, error) { + if r.seen == nil { + r.seen = make(map[string]struct{}) + } + if r.err != nil { + *calls = append(*calls, cmd) + return nil, r.err + } + if cmd != nil { + cmd.Normalize() + if _, ok := r.seen[cmd.RequestID]; ok || r.alreadyApplied[cmd.RequestID] { + *calls = append(*calls, cmd) + return &BatchImageBalanceHoldResult{Applied: false}, nil + } + r.seen[cmd.RequestID] = struct{}{} + } + *calls = append(*calls, cmd) + return &BatchImageBalanceHoldResult{Applied: true}, nil +} + var _ UsageBillingRepository = (*fakeBatchImageBillingRepo)(nil) var _ BatchImagePricingResolver = (*fakeBatchImagePricingResolver)(nil) var _ = strings.TrimSpace diff --git a/backend/internal/service/batch_image_worker.go b/backend/internal/service/batch_image_worker.go index fca9681b5f..5027350689 100644 --- a/backend/internal/service/batch_image_worker.go +++ b/backend/internal/service/batch_image_worker.go @@ -6,6 +6,8 @@ import ( "time" "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/Wei-Shaw/sub2api/internal/pkg/logger" + "go.uber.org/zap" ) const ( @@ -156,6 +158,10 @@ func (w *BatchImageWorker) RunOnce(ctx context.Context) error { result, err := w.processor.Process(ctx, reserved.BatchID) if err != nil { + logger.L().Warn("batch_image.worker_process_failed", + zap.String("batch_id", reserved.BatchID), + zap.Error(err), + ) return w.queue.RequeueAfter(ctx, reserved.BatchID, w.opts.ErrorRetryDelay) } if result.Terminal { diff --git a/backend/internal/service/batch_image_worker_runtime.go b/backend/internal/service/batch_image_worker_runtime.go index e47a47c1a7..de3b4cb42a 100644 --- a/backend/internal/service/batch_image_worker_runtime.go +++ b/backend/internal/service/batch_image_worker_runtime.go @@ -8,8 +8,9 @@ import ( ) type BatchImageWorkerRuntime struct { - worker *BatchImageWorker - cfg *config.Config + worker *BatchImageWorker + billingRecovery *BatchImageBillingRecoveryService + cfg *config.Config mu sync.Mutex cancel context.CancelFunc @@ -25,23 +26,36 @@ func ProvideBatchImageWorkerRuntime( accountRepo AccountRepository, queue BatchImageQueue, billingRepo UsageBillingRepository, + usageLogRepo UsageLogRepository, pricing *BatchImageModelPricingResolver, + authCache APIKeyAuthCacheInvalidator, cfg *config.Config, ) *BatchImageWorkerRuntime { processor := &BatchImagePipelineProcessor{ ProviderProcessor: &BatchImageProviderProcessor{ Repo: repo, - ProviderRegistry: NewDefaultBatchImageProviderRegistry(), + ProviderRegistry: NewBatchImageProviderRegistryFromConfig(cfg), AccountResolver: &BatchImageAccountRepositoryResolver{Repo: accountRepo}, + BillingRepo: billingRepo, + AuthCache: authCache, }, SettlementService: &BatchImageSettlementService{ - Repo: repo, - BillingRepo: billingRepo, - Pricing: pricing, - Config: cfg, + Repo: repo, + BillingRepo: billingRepo, + UsageLogRepo: usageLogRepo, + Pricing: pricing, + AuthCache: authCache, + Config: cfg, }, } runtime := NewBatchImageWorkerRuntime(NewBatchImageWorker(queue, processor, NewBatchImageWorkerOptionsFromConfig(cfg)), cfg) + runtime.billingRecovery = &BatchImageBillingRecoveryService{ + Repo: repo, + Billing: billingRepo, + AuthCache: authCache, + StaleAfter: NewBatchImageWorkerOptionsFromConfig(cfg).StaleActiveAfter, + Limit: NewBatchImageWorkerOptionsFromConfig(cfg).RecoverLimit, + } runtime.Start() return runtime } @@ -62,7 +76,7 @@ func (r *BatchImageWorkerRuntime) Start() { r.done = done var wg sync.WaitGroup - wg.Add(3) + wg.Add(4) go func() { defer wg.Done() r.worker.Run(ctx) @@ -75,12 +89,30 @@ func (r *BatchImageWorkerRuntime) Start() { defer wg.Done() r.worker.RunStaleActiveRecovery(ctx) }() + go func() { + defer wg.Done() + r.runBillingRecovery(ctx) + }() go func() { wg.Wait() close(done) }() } +func (r *BatchImageWorkerRuntime) runBillingRecovery(ctx context.Context) { + if r == nil || r.worker == nil || r.billingRecovery == nil { + return + } + interval := r.worker.opts.RecoveryInterval + for { + if err := ctx.Err(); err != nil { + return + } + _, _ = r.billingRecovery.ReleaseStaleUnsubmittedOnce(ctx) + sleepOrDone(ctx, interval) + } +} + func (r *BatchImageWorkerRuntime) Stop() { if r == nil { return diff --git a/backend/internal/service/group.go b/backend/internal/service/group.go index 6d0a11f766..e3a1697b57 100644 --- a/backend/internal/service/group.go +++ b/backend/internal/service/group.go @@ -36,12 +36,15 @@ type Group struct { DefaultValidityDays int // 图片生成计费配置(antigravity 和 gemini 平台使用) - AllowImageGeneration bool - ImageRateIndependent bool - ImageRateMultiplier float64 - ImagePrice1K *float64 - ImagePrice2K *float64 - ImagePrice4K *float64 + AllowImageGeneration bool + AllowBatchImageGeneration bool + ImageRateIndependent bool + ImageRateMultiplier float64 + ImagePrice1K *float64 + ImagePrice2K *float64 + ImagePrice4K *float64 + BatchImageDiscountMultiplier float64 + BatchImageHoldMultiplier float64 // Claude Code 客户端限制 ClaudeCodeOnly bool diff --git a/backend/internal/service/pricing_service.go b/backend/internal/service/pricing_service.go index bd0c30df45..cc62248d08 100644 --- a/backend/internal/service/pricing_service.go +++ b/backend/internal/service/pricing_service.go @@ -319,6 +319,7 @@ func (s *PricingService) downloadPricingData() error { if err != nil { return fmt.Errorf("parse pricing data: %w", err) } + data = s.mergeFallbackPricingData(data) // 保存到本地文件 pricingFile := s.getPricingFilePath() @@ -373,7 +374,7 @@ func (s *PricingService) parsePricingData(body []byte) (map[string]*LiteLLMModel } // 只保留有有效价格的条目 - if entry.InputCostPerToken == nil && entry.OutputCostPerToken == nil { + if entry.InputCostPerToken == nil && entry.OutputCostPerToken == nil && entry.OutputCostPerImage == nil && entry.OutputCostPerImageToken == nil { continue } @@ -441,6 +442,7 @@ func (s *PricingService) loadPricingData(filePath string) error { if err != nil { return fmt.Errorf("parse pricing data: %w", err) } + pricingData = s.mergeFallbackPricingData(pricingData) // 计算哈希 hash := sha256.Sum256(data) @@ -462,6 +464,37 @@ func (s *PricingService) loadPricingData(filePath string) error { return nil } +func (s *PricingService) mergeFallbackPricingData(data map[string]*LiteLLMModelPricing) map[string]*LiteLLMModelPricing { + if data == nil { + data = make(map[string]*LiteLLMModelPricing) + } + if s == nil || s.cfg == nil || strings.TrimSpace(s.cfg.Pricing.FallbackFile) == "" { + return data + } + fallbackBody, err := os.ReadFile(s.cfg.Pricing.FallbackFile) + if err != nil { + logger.LegacyPrintf("service.pricing", "[Pricing] Fallback merge skipped: %v", err) + return data + } + fallbackData, err := s.parsePricingData(fallbackBody) + if err != nil { + logger.LegacyPrintf("service.pricing", "[Pricing] Fallback merge parse skipped: %v", err) + return data + } + merged := 0 + for modelName, pricing := range fallbackData { + if _, ok := data[modelName]; ok { + continue + } + data[modelName] = pricing + merged++ + } + if merged > 0 { + logger.LegacyPrintf("service.pricing", "[Pricing] Merged %d fallback-only models", merged) + } + return data +} + // useFallbackPricing 使用回退价格文件 func (s *PricingService) useFallbackPricing() error { fallbackFile := s.cfg.Pricing.FallbackFile diff --git a/backend/internal/service/pricing_service_test.go b/backend/internal/service/pricing_service_test.go index f4252f9540..11c1b58da9 100644 --- a/backend/internal/service/pricing_service_test.go +++ b/backend/internal/service/pricing_service_test.go @@ -6,6 +6,7 @@ import ( "path/filepath" "testing" + "github.com/Wei-Shaw/sub2api/internal/config" "github.com/stretchr/testify/require" ) @@ -37,6 +38,57 @@ func TestParsePricingData_ParsesPriorityAndServiceTierFields(t *testing.T) { require.True(t, pricing.SupportsServiceTier) } +func TestParsePricingData_KeepsImageOnlyPricing(t *testing.T) { + svc := &PricingService{} + body := []byte(`{ + "image-only-model": { + "output_cost_per_image": 0.034, + "litellm_provider": "vertex_ai-language-models", + "mode": "image_generation" + } + }`) + + data, err := svc.parsePricingData(body) + require.NoError(t, err) + pricing := data["image-only-model"] + require.NotNil(t, pricing) + require.InDelta(t, 0.034, pricing.OutputCostPerImage, 1e-12) + require.Equal(t, "image_generation", pricing.Mode) +} + +func TestPricingService_MergesFallbackOnlyModels(t *testing.T) { + dir := t.TempDir() + fallbackFile := filepath.Join(dir, "fallback.json") + require.NoError(t, os.WriteFile(fallbackFile, []byte(`{ + "remote-model": { + "input_cost_per_token": 0.000001, + "litellm_provider": "test", + "mode": "chat" + }, + "gemini-3.1-flash-lite-image": { + "output_cost_per_image": 0.034, + "litellm_provider": "vertex_ai-language-models", + "mode": "image_generation" + } + }`), 0644)) + + svc := &PricingService{cfg: &config.Config{}} + svc.cfg.Pricing.FallbackFile = fallbackFile + remoteData, err := svc.parsePricingData([]byte(`{ + "remote-model": { + "input_cost_per_token": 0.000002, + "litellm_provider": "test", + "mode": "chat" + } + }`)) + require.NoError(t, err) + + merged := svc.mergeFallbackPricingData(remoteData) + require.InDelta(t, 0.000002, merged["remote-model"].InputCostPerToken, 1e-12) + require.NotNil(t, merged["gemini-3.1-flash-lite-image"]) + require.InDelta(t, 0.034, merged["gemini-3.1-flash-lite-image"].OutputCostPerImage, 1e-12) +} + func TestGetModelPricing_Gpt53CodexSparkUsesGpt51CodexPricing(t *testing.T) { sparkPricing := &LiteLLMModelPricing{InputCostPerToken: 1} gpt53Pricing := &LiteLLMModelPricing{InputCostPerToken: 9} diff --git a/backend/internal/service/usage_billing.go b/backend/internal/service/usage_billing.go index accc7cb2cb..8d52c92d26 100644 --- a/backend/internal/service/usage_billing.go +++ b/backend/internal/service/usage_billing.go @@ -119,6 +119,57 @@ type UsageBillingApplyResult struct { QuotaState *AccountQuotaState // post-increment quota state (nil = no quota increment) } +// BatchImageBalanceHoldCommand describes an idempotent balance hold operation. +type BatchImageBalanceHoldCommand struct { + RequestID string + APIKeyID int64 + RequestFingerprint string + RequestPayloadHash string + UserID int64 + BatchID string + HoldAmount float64 + ActualAmount float64 +} + +func (c *BatchImageBalanceHoldCommand) Normalize() { + if c == nil { + return + } + c.RequestID = strings.TrimSpace(c.RequestID) + c.BatchID = strings.TrimSpace(c.BatchID) + if strings.TrimSpace(c.RequestFingerprint) == "" { + c.RequestFingerprint = buildBatchImageBalanceHoldFingerprint(c) + } +} + +func buildBatchImageBalanceHoldFingerprint(c *BatchImageBalanceHoldCommand) string { + if c == nil { + return "" + } + raw := fmt.Sprintf( + "%d|%d|%s|%0.10f|%0.10f", + c.UserID, + c.APIKeyID, + strings.TrimSpace(c.BatchID), + c.HoldAmount, + c.ActualAmount, + ) + if payloadHash := strings.TrimSpace(c.RequestPayloadHash); payloadHash != "" { + raw += "|" + payloadHash + } + sum := sha256.Sum256([]byte(raw)) + return hex.EncodeToString(sum[:]) +} + +type BatchImageBalanceHoldResult struct { + Applied bool + NewBalance *float64 + FrozenBalance *float64 +} + type UsageBillingRepository interface { Apply(ctx context.Context, cmd *UsageBillingCommand) (*UsageBillingApplyResult, error) + ReserveBatchImageBalance(ctx context.Context, cmd *BatchImageBalanceHoldCommand) (*BatchImageBalanceHoldResult, error) + CaptureBatchImageBalance(ctx context.Context, cmd *BatchImageBalanceHoldCommand) (*BatchImageBalanceHoldResult, error) + ReleaseBatchImageBalance(ctx context.Context, cmd *BatchImageBalanceHoldCommand) (*BatchImageBalanceHoldResult, error) } diff --git a/backend/internal/service/user.go b/backend/internal/service/user.go index edb944ee05..22a21a4634 100644 --- a/backend/internal/service/user.go +++ b/backend/internal/service/user.go @@ -19,6 +19,7 @@ type User struct { PasswordHash string Role string Balance float64 + FrozenBalance float64 Concurrency int Status string AllowedGroups []int64 diff --git a/backend/migrations/001_init.sql b/backend/migrations/001_init.sql index 64078c42df..9681fe9a56 100644 --- a/backend/migrations/001_init.sql +++ b/backend/migrations/001_init.sql @@ -43,7 +43,8 @@ CREATE TABLE IF NOT EXISTS users ( email VARCHAR(255) NOT NULL UNIQUE, password_hash VARCHAR(255) NOT NULL, role VARCHAR(20) NOT NULL DEFAULT 'user', -- admin/user - balance DECIMAL(20, 8) NOT NULL DEFAULT 0, -- 余额(可为负数) + balance DECIMAL(20, 8) NOT NULL DEFAULT 0, -- 可用余额(可为负数) + frozen_balance DECIMAL(20, 8) NOT NULL DEFAULT 0, -- 冻结余额 concurrency INT NOT NULL DEFAULT 5, -- 并发数限制 status VARCHAR(20) NOT NULL DEFAULT 'active', -- active/disabled allowed_groups BIGINT[] DEFAULT NULL, -- 允许绑定的分组ID列表 diff --git a/backend/migrations/134_image_generation_group_controls.sql b/backend/migrations/134_image_generation_group_controls.sql index 37941c001e..4d83702c59 100644 --- a/backend/migrations/134_image_generation_group_controls.sql +++ b/backend/migrations/134_image_generation_group_controls.sql @@ -7,6 +7,9 @@ ALTER TABLE groups ADD COLUMN IF NOT EXISTS allow_image_generation BOOLEAN NOT NULL DEFAULT false; +ALTER TABLE groups + ADD COLUMN IF NOT EXISTS allow_batch_image_generation BOOLEAN NOT NULL DEFAULT false; + ALTER TABLE groups ADD COLUMN IF NOT EXISTS image_rate_independent BOOLEAN NOT NULL DEFAULT false; @@ -22,5 +25,6 @@ SET image_rate_independent = false, image_rate_multiplier = 1.0; COMMENT ON COLUMN groups.allow_image_generation IS '是否允许该分组使用图片生成能力'; +COMMENT ON COLUMN groups.allow_batch_image_generation IS '是否允许该分组使用批量图片生成能力'; COMMENT ON COLUMN groups.image_rate_independent IS '图片生成是否使用独立倍率;false 表示共享分组有效倍率'; COMMENT ON COLUMN groups.image_rate_multiplier IS '图片生成独立倍率,仅 image_rate_independent=true 时生效'; diff --git a/backend/migrations/160_add_user_frozen_balance.sql b/backend/migrations/160_add_user_frozen_balance.sql new file mode 100644 index 0000000000..d113efc9f7 --- /dev/null +++ b/backend/migrations/160_add_user_frozen_balance.sql @@ -0,0 +1,2 @@ +ALTER TABLE users + ADD COLUMN IF NOT EXISTS frozen_balance DECIMAL(20,8) NOT NULL DEFAULT 0; diff --git a/backend/migrations/161_batch_image_pricing_snapshot.sql b/backend/migrations/161_batch_image_pricing_snapshot.sql new file mode 100644 index 0000000000..3ae6d1fbb6 --- /dev/null +++ b/backend/migrations/161_batch_image_pricing_snapshot.sql @@ -0,0 +1,25 @@ +ALTER TABLE groups + ADD COLUMN IF NOT EXISTS batch_image_discount_multiplier DECIMAL(10,4) NOT NULL DEFAULT 0.5, + ADD COLUMN IF NOT EXISTS batch_image_hold_multiplier DECIMAL(10,4) NOT NULL DEFAULT 0.6; + +COMMENT ON COLUMN groups.batch_image_discount_multiplier IS '批量图片生成折扣倍率,最终单价会乘以该值;0 表示免费'; +COMMENT ON COLUMN groups.batch_image_hold_multiplier IS '批量图片生成冻结价格比例,按普通生图原价乘以该比例冻结,结算后释放差额'; + +ALTER TABLE batch_image_jobs + ADD COLUMN IF NOT EXISTS base_unit_price DECIMAL(20,10) NOT NULL DEFAULT 0, + ADD COLUMN IF NOT EXISTS group_rate_multiplier DECIMAL(10,4) NOT NULL DEFAULT 1.0, + ADD COLUMN IF NOT EXISTS account_rate_multiplier DECIMAL(10,4) NOT NULL DEFAULT 1.0, + ADD COLUMN IF NOT EXISTS batch_discount_multiplier DECIMAL(10,4) NOT NULL DEFAULT 0.5, + ADD COLUMN IF NOT EXISTS hold_multiplier DECIMAL(10,4) NOT NULL DEFAULT 0.6, + ADD COLUMN IF NOT EXISTS billable_unit_price DECIMAL(20,10) NOT NULL DEFAULT 0, + ADD COLUMN IF NOT EXISTS hold_unit_price DECIMAL(20,10) NOT NULL DEFAULT 0, + ADD COLUMN IF NOT EXISTS pricing_snapshot_version INTEGER NOT NULL DEFAULT 0; + +COMMENT ON COLUMN batch_image_jobs.base_unit_price IS '提交时快照的基础批量图片单价'; +COMMENT ON COLUMN batch_image_jobs.group_rate_multiplier IS '提交时快照的分组/用户专属图片倍率'; +COMMENT ON COLUMN batch_image_jobs.account_rate_multiplier IS '提交时快照的账号倍率'; +COMMENT ON COLUMN batch_image_jobs.batch_discount_multiplier IS '提交时快照的批量折扣倍率'; +COMMENT ON COLUMN batch_image_jobs.hold_multiplier IS '提交时快照的冻结价格比例,按普通生图原价乘以该比例冻结'; +COMMENT ON COLUMN batch_image_jobs.billable_unit_price IS '提交时快照的实际结算单价'; +COMMENT ON COLUMN batch_image_jobs.hold_unit_price IS '提交时快照的冻结单价'; +COMMENT ON COLUMN batch_image_jobs.pricing_snapshot_version IS '批量图片任务价格快照版本;0 表示旧任务无快照'; diff --git a/backend/migrations/162_add_group_batch_image_generation_gate.sql b/backend/migrations/162_add_group_batch_image_generation_gate.sql new file mode 100644 index 0000000000..e96541b931 --- /dev/null +++ b/backend/migrations/162_add_group_batch_image_generation_gate.sql @@ -0,0 +1,4 @@ +ALTER TABLE groups + ADD COLUMN IF NOT EXISTS allow_batch_image_generation BOOLEAN NOT NULL DEFAULT false; + +COMMENT ON COLUMN groups.allow_batch_image_generation IS '是否允许该分组使用批量图片生成能力'; diff --git a/backend/migrations/163_batch_image_default_discount_and_hold_ratio.sql b/backend/migrations/163_batch_image_default_discount_and_hold_ratio.sql new file mode 100644 index 0000000000..65ac699ba9 --- /dev/null +++ b/backend/migrations/163_batch_image_default_discount_and_hold_ratio.sql @@ -0,0 +1,19 @@ +ALTER TABLE groups + ALTER COLUMN batch_image_discount_multiplier SET DEFAULT 0.5, + ALTER COLUMN batch_image_hold_multiplier SET DEFAULT 0.6; + +UPDATE groups +SET batch_image_discount_multiplier = 0.5 +WHERE batch_image_discount_multiplier = 1.0; + +UPDATE groups +SET batch_image_hold_multiplier = 0.6 +WHERE batch_image_hold_multiplier = 1.05; + +COMMENT ON COLUMN groups.batch_image_hold_multiplier IS '批量图片生成冻结价格比例,按普通生图原价乘以该比例冻结,结算后释放差额'; + +ALTER TABLE batch_image_jobs + ALTER COLUMN batch_discount_multiplier SET DEFAULT 0.5, + ALTER COLUMN hold_multiplier SET DEFAULT 0.6; + +COMMENT ON COLUMN batch_image_jobs.hold_multiplier IS '提交时快照的冻结价格比例,按普通生图原价乘以该比例冻结'; diff --git a/backend/migrations/164_batch_image_download_and_user_delete.sql b/backend/migrations/164_batch_image_download_and_user_delete.sql new file mode 100644 index 0000000000..56848b7c08 --- /dev/null +++ b/backend/migrations/164_batch_image_download_and_user_delete.sql @@ -0,0 +1,9 @@ +ALTER TABLE batch_image_jobs + ADD COLUMN IF NOT EXISTS downloaded_at TIMESTAMPTZ, + ADD COLUMN IF NOT EXISTS user_deleted_at TIMESTAMPTZ; + +CREATE INDEX IF NOT EXISTS batch_image_jobs_downloaded_at_idx ON batch_image_jobs (downloaded_at); +CREATE INDEX IF NOT EXISTS batch_image_jobs_user_deleted_at_idx ON batch_image_jobs (user_deleted_at); + +COMMENT ON COLUMN batch_image_jobs.downloaded_at IS '用户首次成功下载批量图片 ZIP 的时间'; +COMMENT ON COLUMN batch_image_jobs.user_deleted_at IS '用户侧删除/隐藏任务记录的时间;账务记录仍保留'; diff --git a/backend/migrations/165_hide_pre_upstream_batch_image_failures.sql b/backend/migrations/165_hide_pre_upstream_batch_image_failures.sql new file mode 100644 index 0000000000..3cd9293e74 --- /dev/null +++ b/backend/migrations/165_hide_pre_upstream_batch_image_failures.sql @@ -0,0 +1,16 @@ +UPDATE batch_image_jobs +SET user_deleted_at = COALESCE(user_deleted_at, updated_at, created_at, NOW()), + updated_at = NOW() +WHERE user_deleted_at IS NULL + AND provider_job_name IS NULL + AND status = 'failed' + AND last_error_code IN ( + 'INSUFFICIENT_BALANCE', + 'PROVIDER_SUBMIT_FAILED', + 'BATCH_IMAGE_PROVIDER_SUBMIT_FAILED', + 'BATCH_IMAGE_VERTEX_GCS_BUCKET_MISSING', + 'VERTEX_MANAGED_GCS_BUCKET_MISSING', + 'BATCH_IMAGE_PROVIDER_MISSING_API_KEY', + 'BATCH_IMAGE_PROVIDER_MISSING_SERVICE_ACCOUNT', + 'BATCH_IMAGE_PROVIDER_UNSUPPORTED_ACCOUNT' + ); diff --git a/backend/migrations/166_batch_image_task_name.sql b/backend/migrations/166_batch_image_task_name.sql new file mode 100644 index 0000000000..ef942d8dad --- /dev/null +++ b/backend/migrations/166_batch_image_task_name.sql @@ -0,0 +1,10 @@ +ALTER TABLE batch_image_jobs + ADD COLUMN IF NOT EXISTS task_name VARCHAR(255) NOT NULL DEFAULT ''; + +UPDATE batch_image_jobs +SET task_name = TO_CHAR(created_at AT TIME ZONE 'Asia/Shanghai', 'YYYY-MM-DD HH24:MI:SS') +WHERE task_name = ''; + +CREATE INDEX IF NOT EXISTS batch_image_jobs_task_name_idx ON batch_image_jobs (task_name); + +COMMENT ON COLUMN batch_image_jobs.task_name IS '用户可读的批量生图任务名称'; diff --git a/backend/migrations/167_clear_auto_batch_image_task_names.sql b/backend/migrations/167_clear_auto_batch_image_task_names.sql new file mode 100644 index 0000000000..d12eefb48c --- /dev/null +++ b/backend/migrations/167_clear_auto_batch_image_task_names.sql @@ -0,0 +1,5 @@ +UPDATE batch_image_jobs +SET task_name = '' +WHERE task_name = TO_CHAR(created_at AT TIME ZONE 'Asia/Shanghai', 'YYYY-MM-DD HH24:MI:SS'); + +COMMENT ON COLUMN batch_image_jobs.task_name IS '用户填写的批量生图任务名称;为空时用户侧显示未填写'; diff --git a/backend/migrations/168_restore_empty_batch_image_task_names.sql b/backend/migrations/168_restore_empty_batch_image_task_names.sql new file mode 100644 index 0000000000..7b2e34bb61 --- /dev/null +++ b/backend/migrations/168_restore_empty_batch_image_task_names.sql @@ -0,0 +1,5 @@ +UPDATE batch_image_jobs +SET task_name = TO_CHAR(created_at AT TIME ZONE 'Asia/Shanghai', 'YYYY-MM-DD HH24:MI:SS') +WHERE task_name = ''; + +COMMENT ON COLUMN batch_image_jobs.task_name IS '用户填写的批量生图任务名称;提交时为空则默认写入当前时间'; diff --git a/backend/migrations/169_batch_image_parent_batch.sql b/backend/migrations/169_batch_image_parent_batch.sql new file mode 100644 index 0000000000..e089c5e49a --- /dev/null +++ b/backend/migrations/169_batch_image_parent_batch.sql @@ -0,0 +1,8 @@ +ALTER TABLE batch_image_jobs + ADD COLUMN IF NOT EXISTS parent_batch_id VARCHAR(64); + +CREATE INDEX IF NOT EXISTS batch_image_jobs_parent_batch_id_idx + ON batch_image_jobs (parent_batch_id) + WHERE parent_batch_id IS NOT NULL AND parent_batch_id <> ''; + +COMMENT ON COLUMN batch_image_jobs.parent_batch_id IS '父批量生图任务 ID;失败项重试等子任务挂在主任务下展示'; diff --git a/backend/resources/model-pricing/model_prices_and_context_window.json b/backend/resources/model-pricing/model_prices_and_context_window.json index e88ed2da22..f35a91220e 100644 --- a/backend/resources/model-pricing/model_prices_and_context_window.json +++ b/backend/resources/model-pricing/model_prices_and_context_window.json @@ -873,7 +873,7 @@ "max_output_tokens": 32768, "max_tokens": 32768, "mode": "image_generation", - "output_cost_per_image": 0.039, + "output_cost_per_image": 0.034, "output_cost_per_token": 0.0, "source": "https://ai.google.dev/pricing", "supported_modalities": [ @@ -1625,6 +1625,47 @@ "supports_web_search": true, "web_search_billing_unit": "per_query" }, + "gemini-3-pro-image": { + "input_cost_per_image": 0.0011, + "input_cost_per_token": 2e-06, + "input_cost_per_token_batches": 1e-06, + "litellm_provider": "vertex_ai-language-models", + "max_input_tokens": 65536, + "max_output_tokens": 32768, + "max_tokens": 32768, + "mode": "image_generation", + "output_cost_per_image": 0.134, + "output_cost_per_image_token": 0.00012, + "output_cost_per_token": 1.2e-05, + "output_cost_per_token_batches": 6e-06, + "search_context_cost_per_query": { + "search_context_size_high": 0.014, + "search_context_size_low": 0.014, + "search_context_size_medium": 0.014 + }, + "source": "https://ai.google.dev/gemini-api/docs/pricing", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/completions", + "/v1/batch" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text", + "image" + ], + "supports_function_calling": false, + "supports_prompt_caching": true, + "supports_response_schema": true, + "supports_service_tier": true, + "supports_system_messages": true, + "supports_vision": true, + "supports_web_search": true, + "web_search_billing_unit": "per_query" + }, "gemini-3-pro-preview": { "cache_creation_input_token_cost_above_200k_tokens": 2.5e-07, "cache_read_input_token_cost": 2e-07, @@ -1726,6 +1767,39 @@ "supports_web_search": true, "web_search_billing_unit": "per_query" }, + "gemini-3.1-flash-lite-image": { + "input_cost_per_image": 0.0003, + "input_cost_per_token": 3e-07, + "litellm_provider": "vertex_ai-language-models", + "max_input_tokens": 32768, + "max_output_tokens": 32768, + "max_tokens": 32768, + "mode": "image_generation", + "output_cost_per_image": 0.034, + "output_cost_per_image_token": 3e-05, + "output_cost_per_token": 2.5e-06, + "source": "https://ai.google.dev/gemini-api/docs/pricing", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/completions", + "/v1/batch" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text", + "image" + ], + "supports_function_calling": false, + "supports_prompt_caching": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_vision": true, + "supports_web_search": true, + "web_search_billing_unit": "per_query" + }, "gemini-3.1-flash-image-preview": { "input_cost_per_image": 0.00056, "input_cost_per_token": 5e-07, diff --git a/deploy/docker-compose.dev.yml b/deploy/docker-compose.dev.yml index 7755fdbeff..e7b8b64c79 100644 --- a/deploy/docker-compose.dev.yml +++ b/deploy/docker-compose.dev.yml @@ -13,6 +13,8 @@ services: build: context: .. dockerfile: Dockerfile + args: + NPM_CONFIG_REGISTRY: ${NPM_CONFIG_REGISTRY:-https://registry.npmmirror.com} container_name: sub2api-dev restart: unless-stopped ports: @@ -40,6 +42,12 @@ services: - JWT_SECRET=${JWT_SECRET:-} - TOTP_ENCRYPTION_KEY=${TOTP_ENCRYPTION_KEY:-} - TZ=${TZ:-Asia/Shanghai} + # Local mainland-China development proxy. Containers cannot use + # 127.0.0.1 for the host proxy, so default to Docker Desktop's host name. + - HTTP_PROXY=${SUB2API_DEV_HTTP_PROXY:-http://host.docker.internal:7897} + - HTTPS_PROXY=${SUB2API_DEV_HTTPS_PROXY:-http://host.docker.internal:7897} + - ALL_PROXY=${SUB2API_DEV_ALL_PROXY:-socks5://host.docker.internal:7897} + - NO_PROXY=${SUB2API_DEV_NO_PROXY:-127.0.0.1,localhost,::1,postgres,redis,sub2api,192.168.0.0/16,10.0.0.0/8,172.16.0.0/12,.local} # OpenAI HTTP upstream protocol/timeout - GATEWAY_OPENAI_RESPONSE_HEADER_TIMEOUT=${GATEWAY_OPENAI_RESPONSE_HEADER_TIMEOUT:-0} - GATEWAY_OPENAI_HTTP2_ENABLED=${GATEWAY_OPENAI_HTTP2_ENABLED:-true} @@ -54,6 +62,15 @@ services: - GATEWAY_IMAGE_CONCURRENCY_OVERFLOW_MODE=${GATEWAY_IMAGE_CONCURRENCY_OVERFLOW_MODE:-reject} - GATEWAY_IMAGE_CONCURRENCY_WAIT_TIMEOUT_SECONDS=${GATEWAY_IMAGE_CONCURRENCY_WAIT_TIMEOUT_SECONDS:-30} - GATEWAY_IMAGE_CONCURRENCY_MAX_WAITING_REQUESTS=${GATEWAY_IMAGE_CONCURRENCY_MAX_WAITING_REQUESTS:-100} + - BATCH_IMAGE_ENABLED=${BATCH_IMAGE_ENABLED:-true} + - BATCH_IMAGE_QUEUE_ENABLED=${BATCH_IMAGE_QUEUE_ENABLED:-true} + - BATCH_IMAGE_VERTEX_ENABLED=${BATCH_IMAGE_VERTEX_ENABLED:-true} + - BATCH_IMAGE_VERTEX_PROJECT_ID=${BATCH_IMAGE_VERTEX_PROJECT_ID:-project-28424c50-8df2-46e2-a27} + - BATCH_IMAGE_VERTEX_LOCATION=${BATCH_IMAGE_VERTEX_LOCATION:-global} + - BATCH_IMAGE_VERTEX_MANAGED_GCS_BUCKET=${BATCH_IMAGE_VERTEX_MANAGED_GCS_BUCKET:-sub2-batch-image-prod-project-28424c50-8df2-46e2-a27} + - BATCH_IMAGE_VERTEX_MANAGED_GCS_PREFIX=${BATCH_IMAGE_VERTEX_MANAGED_GCS_PREFIX:-batch-image/prod/{batch_id}} + - BATCH_IMAGE_VERTEX_INPUT_RETENTION_HOURS=${BATCH_IMAGE_VERTEX_INPUT_RETENTION_HOURS:-24} + - BATCH_IMAGE_VERTEX_OUTPUT_RETENTION_HOURS=${BATCH_IMAGE_VERTEX_OUTPUT_RETENTION_HOURS:-72} depends_on: postgres: condition: service_healthy diff --git a/frontend/src/api/batchImage.ts b/frontend/src/api/batchImage.ts new file mode 100644 index 0000000000..e08743c74d --- /dev/null +++ b/frontend/src/api/batchImage.ts @@ -0,0 +1,235 @@ +import { buildGatewayUrl } from './client' + +export type BatchImageStatus = + | 'queued' + | 'running' + | 'indexing' + | 'processing_results' + | 'settling' + | 'completed' + | 'failed' + | 'cancelled' + | 'output_deleted' + | string + +export interface BatchImageSubmitItem { + custom_id: string + prompt: string +} + +export interface BatchImageSubmitRequest { + model: string + task_name?: string + parent_batch_id?: string + provider?: '' | 'gemini_api' | 'vertex' | string + image_size?: '1K' | '2K' | '4K' | string + response_mime_type?: string + aspect_ratio?: string + items: BatchImageSubmitItem[] + metadata?: Record +} + +export interface BatchImageJob { + id: string + object: string + task_name: string + parent_batch_id?: string | null + status: BatchImageStatus + model: string + provider: string + item_count: number + success_count: number + fail_count: number + estimated_cost: number + hold_amount: number + actual_cost: number | null + created_at: number + submitted_at: number | null + settled_at: number | null + downloaded_at?: number | null + output_deleted_at?: number | null +} + +export interface BatchImageItem { + batch_id?: string + source_task_name?: string + custom_id: string + status: string + prompt_preview?: string | null + mime_type: string | null + file_extension: string | null + image_count: number + error?: { + code: string + message: string + source?: 'provider' | 'system' | string + } | null +} + +export interface BatchImageItemsResponse { + object: string + data: BatchImageItem[] + has_more: boolean +} + +export interface BatchImageJobsResponse { + object: string + data: BatchImageJob[] + has_more: boolean +} + +export interface BatchImageModel { + id: string + object: string + provider: string +} + +export interface BatchImageModelsResponse { + object: string + data: BatchImageModel[] +} + +export interface BatchImageJobsListOptions { + limit?: number + cursor?: string + status?: string + taskName?: string + downloaded?: '' | 'true' | 'false' | string + from?: string + to?: string +} + +async function parseBatchImageError(response: Response): Promise { + try { + const body = await response.json() + const message = body?.error?.message || body?.message || response.statusText + const error = new Error(message) + ;(error as any).code = body?.error?.code || response.status + ;(error as any).status = response.status + ;(error as any).requestId = response.headers.get('X-Request-Id') || '' + return error + } catch { + const error = new Error(response.statusText || `HTTP ${response.status}`) + ;(error as any).code = response.status + ;(error as any).status = response.status + ;(error as any).requestId = response.headers.get('X-Request-Id') || '' + return error + } +} + +function authHeaders(apiKey: string, extra?: HeadersInit): HeadersInit { + return { + Authorization: `Bearer ${apiKey}`, + ...extra, + } +} + +export async function submitBatchImageJob( + apiKey: string, + payload: BatchImageSubmitRequest, + idempotencyKey: string, +): Promise { + const response = await fetch(buildGatewayUrl('/v1/images/batches'), { + method: 'POST', + headers: authHeaders(apiKey, { + 'Content-Type': 'application/json', + 'Idempotency-Key': idempotencyKey, + }), + body: JSON.stringify(payload), + }) + if (!response.ok) throw await parseBatchImageError(response) + return response.json() +} + +export async function getBatchImageJob(apiKey: string, batchId: string): Promise { + const response = await fetch(buildGatewayUrl(`/v1/images/batches/${encodeURIComponent(batchId)}`), { + headers: authHeaders(apiKey), + }) + if (!response.ok) throw await parseBatchImageError(response) + return response.json() +} + +export async function listBatchImageJobs(apiKey: string, options: number | BatchImageJobsListOptions = 20): Promise { + const params = new URLSearchParams() + if (typeof options === 'number') { + params.set('limit', String(options)) + } else { + params.set('limit', String(options.limit || 20)) + if (options.cursor) params.set('cursor', options.cursor) + if (options.status) params.set('status', options.status) + if (options.taskName) params.set('task_name', options.taskName) + if (options.downloaded) params.set('downloaded', options.downloaded) + if (options.from) params.set('from', options.from) + if (options.to) params.set('to', options.to) + } + const response = await fetch(buildGatewayUrl(`/v1/images/batches?${params.toString()}`), { + headers: authHeaders(apiKey), + }) + if (!response.ok) throw await parseBatchImageError(response) + return response.json() +} + +export async function listBatchImageModels(apiKey: string): Promise { + const response = await fetch(buildGatewayUrl('/v1/images/batches/models'), { + headers: authHeaders(apiKey), + }) + if (!response.ok) throw await parseBatchImageError(response) + return response.json() +} + +export async function listBatchImageItems( + apiKey: string, + batchId: string, + status = '', +): Promise { + const query = status ? `?status=${encodeURIComponent(status)}` : '' + const response = await fetch(buildGatewayUrl(`/v1/images/batches/${encodeURIComponent(batchId)}/items${query}`), { + headers: authHeaders(apiKey), + }) + if (!response.ok) throw await parseBatchImageError(response) + return response.json() +} + +export async function cancelBatchImageJob(apiKey: string, batchId: string): Promise { + const response = await fetch(buildGatewayUrl(`/v1/images/batches/${encodeURIComponent(batchId)}/cancel`), { + method: 'POST', + headers: authHeaders(apiKey), + }) + if (!response.ok) throw await parseBatchImageError(response) + return response.json() +} + +export async function downloadBatchImageZip(apiKey: string, batchId: string): Promise { + const response = await fetch(buildGatewayUrl(`/v1/images/batches/${encodeURIComponent(batchId)}/download`), { + headers: authHeaders(apiKey), + }) + if (!response.ok) throw await parseBatchImageError(response) + return response.blob() +} + +export async function getBatchImageItemContent(apiKey: string, batchId: string, customId: string, imageIndex = 0): Promise { + const response = await fetch(buildGatewayUrl(`/v1/images/batches/${encodeURIComponent(batchId)}/items/${encodeURIComponent(customId)}/content?image_index=${encodeURIComponent(String(imageIndex))}`), { + headers: authHeaders(apiKey), + }) + if (!response.ok) throw await parseBatchImageError(response) + return response.blob() +} + +export async function deleteBatchImageJobRecord(apiKey: string, batchId: string): Promise { + const response = await fetch(buildGatewayUrl(`/v1/images/batches/${encodeURIComponent(batchId)}`), { + method: 'DELETE', + headers: authHeaders(apiKey), + }) + if (!response.ok) throw await parseBatchImageError(response) +} + +export function saveBlob(blob: Blob, filename: string) { + const url = URL.createObjectURL(blob) + const link = document.createElement('a') + link.href = url + link.download = filename + document.body.appendChild(link) + link.click() + document.body.removeChild(link) + URL.revokeObjectURL(url) +} diff --git a/frontend/src/api/index.ts b/frontend/src/api/index.ts index 6702468d8e..71fa27e2a7 100644 --- a/frontend/src/api/index.ts +++ b/frontend/src/api/index.ts @@ -17,6 +17,7 @@ export { redeemAPI, type RedeemHistoryItem } from './redeem' export { paymentAPI } from './payment' export { userGroupsAPI } from './groups' export { userChannelsAPI } from './channels' +export * as batchImageAPI from './batchImage' export { totpAPI } from './totp' export { default as announcementsAPI } from './announcements' export { channelMonitorUserAPI } from './channelMonitor' diff --git a/frontend/src/components/common/BaseDialog.vue b/frontend/src/components/common/BaseDialog.vue index 6d9a08caa2..2a0f25870b 100644 --- a/frontend/src/components/common/BaseDialog.vue +++ b/frontend/src/components/common/BaseDialog.vue @@ -20,7 +20,7 @@ + + + + + +
diff --git a/frontend/src/views/admin/GroupsView.vue b/frontend/src/views/admin/GroupsView.vue index 56d21c86c1..4b82185270 100644 --- a/frontend/src/views/admin/GroupsView.vue +++ b/frontend/src/views/admin/GroupsView.vue @@ -889,6 +889,60 @@
+
+ +

+ {{ t("admin.groups.imagePricing.batchDisabledHint") }} +

+

+ {{ t("admin.groups.imagePricing.batchSectionHint") }} +

+
+
+ + +
+
+ + +
+
+
@@ -2228,6 +2282,60 @@ +
+ +

+ {{ t("admin.groups.imagePricing.batchDisabledHint") }} +

+

+ {{ t("admin.groups.imagePricing.batchSectionHint") }} +

+
+
+ + +
+
+ + +
+
+
@@ -3549,8 +3657,11 @@ const createForm = reactive({ monthly_limit_usd: null as number | null, // 图片生成计费配置 allow_image_generation: false, + allow_batch_image_generation: false, image_rate_independent: false, image_rate_multiplier: 1, + batch_image_discount_multiplier: 0.5, + batch_image_hold_multiplier: 0.6, image_price_1k: null as number | null, image_price_2k: null as number | null, image_price_4k: null as number | null, @@ -3885,8 +3996,11 @@ const editForm = reactive({ monthly_limit_usd: null as number | null, // 图片生成计费配置 allow_image_generation: false, + allow_batch_image_generation: false, image_rate_independent: false, image_rate_multiplier: 1, + batch_image_discount_multiplier: 0.5, + batch_image_hold_multiplier: 0.6, image_price_1k: null as number | null, image_price_2k: null as number | null, image_price_4k: null as number | null, @@ -3922,9 +4036,13 @@ const editForm = reactive({ }); type ImagePricingFormState = { + allow_image_generation: boolean; + allow_batch_image_generation: boolean; rate_multiplier: number; image_rate_independent: boolean; image_rate_multiplier: number; + batch_image_discount_multiplier: number; + batch_image_hold_multiplier: number; image_price_1k: number | string | null; image_price_2k: number | string | null; image_price_4k: number | string | null; @@ -3960,9 +4078,10 @@ const formatImagePricePreview = (value: number | string | null | undefined) => { }; const buildImageFinalPricePreview = (form: ImagePricingFormState) => { - const multiplier = form.image_rate_independent + const imageMultiplier = form.image_rate_independent ? normalizePreviewNumber(form.image_rate_multiplier, 1) : normalizePreviewNumber(form.rate_multiplier, 1); + const multiplier = imageMultiplier; return imagePricingTiers.map((tier) => { const basePrice = normalizePreviewNumber(form[tier.key]); return { @@ -3981,6 +4100,21 @@ const editImageFinalPricePreview = computed(() => buildImageFinalPricePreview(editForm), ); +const resetDisabledBatchImagePricing = ( + form: Pick< + ImagePricingFormState, + "allow_image_generation" | "allow_batch_image_generation" | "batch_image_discount_multiplier" | "batch_image_hold_multiplier" + >, +) => { + if (!form.allow_image_generation) { + form.allow_batch_image_generation = false; + } + if (!form.allow_batch_image_generation) { + form.batch_image_discount_multiplier = 0.5; + form.batch_image_hold_multiplier = 0.6; + } +}; + // 根据分组类型返回不同的删除确认消息 const deleteConfirmMessage = computed(() => { if (!deletingGroup.value) { @@ -4158,8 +4292,11 @@ const closeCreateModal = () => { createForm.weekly_limit_usd = null; createForm.monthly_limit_usd = null; createForm.allow_image_generation = false; + createForm.allow_batch_image_generation = false; createForm.image_rate_independent = false; createForm.image_rate_multiplier = 1; + createForm.batch_image_discount_multiplier = 0.5; + createForm.batch_image_hold_multiplier = 0.6; createForm.image_price_1k = null; createForm.image_price_2k = null; createForm.image_price_4k = null; @@ -4256,6 +4393,13 @@ const handleCreateGroup = async () => { requestData.image_rate_multiplier = normalizeRateMultiplier( requestData.image_rate_multiplier, ); + resetDisabledBatchImagePricing(requestData); + requestData.batch_image_discount_multiplier = normalizeRateMultiplier( + requestData.batch_image_discount_multiplier, + ); + requestData.batch_image_hold_multiplier = normalizeRateMultiplier( + requestData.batch_image_hold_multiplier, + ); requestData.peak_rate_enabled = createForm.peak_rate_enabled; requestData.peak_start = createForm.peak_start; requestData.peak_end = createForm.peak_end; @@ -4294,8 +4438,13 @@ const handleEdit = async (group: AdminGroup) => { editForm.weekly_limit_usd = group.weekly_limit_usd; editForm.monthly_limit_usd = group.monthly_limit_usd; editForm.allow_image_generation = group.allow_image_generation ?? false; + editForm.allow_batch_image_generation = + group.allow_batch_image_generation ?? false; editForm.image_rate_independent = group.image_rate_independent ?? false; editForm.image_rate_multiplier = group.image_rate_multiplier ?? 1; + editForm.batch_image_discount_multiplier = + group.batch_image_discount_multiplier ?? 0.5; + editForm.batch_image_hold_multiplier = group.batch_image_hold_multiplier ?? 0.6; editForm.image_price_1k = group.image_price_1k; editForm.image_price_2k = group.image_price_2k; editForm.image_price_4k = group.image_price_4k; @@ -4409,6 +4558,13 @@ const handleUpdateGroup = async () => { payload.image_rate_multiplier = normalizeRateMultiplier( payload.image_rate_multiplier, ); + resetDisabledBatchImagePricing(payload); + payload.batch_image_discount_multiplier = normalizeRateMultiplier( + payload.batch_image_discount_multiplier, + ); + payload.batch_image_hold_multiplier = normalizeRateMultiplier( + payload.batch_image_hold_multiplier, + ); payload.peak_rate_enabled = editForm.peak_rate_enabled; payload.peak_start = editForm.peak_start; payload.peak_end = editForm.peak_end; @@ -4532,6 +4688,20 @@ watch( }, ); +watch( + () => createForm.allow_image_generation, + () => { + resetDisabledBatchImagePricing(createForm); + }, +); + +watch( + () => createForm.allow_batch_image_generation, + () => { + resetDisabledBatchImagePricing(createForm); + }, +); + watch( () => editForm.platform, (newVal) => { @@ -4552,6 +4722,20 @@ watch( }, ); +watch( + () => editForm.allow_image_generation, + () => { + resetDisabledBatchImagePricing(editForm); + }, +); + +watch( + () => editForm.allow_batch_image_generation, + () => { + resetDisabledBatchImagePricing(editForm); + }, +); + watch( () => editForm.platform, (newVal) => { diff --git a/frontend/src/views/user/BatchImageGuideView.vue b/frontend/src/views/user/BatchImageGuideView.vue new file mode 100644 index 0000000000..8f1eebdbe3 --- /dev/null +++ b/frontend/src/views/user/BatchImageGuideView.vue @@ -0,0 +1,2563 @@ +