Merge pull request #3768 from Turtle-Li/feature/batch-image-foundation

feat: add batch image generation MVP
This commit is contained in:
Wesley Liddick
2026-07-07 16:32:45 +08:00
committed by GitHub
144 changed files with 42999 additions and 115 deletions
+6 -1
View File
@@ -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
+14
View File
@@ -85,6 +85,8 @@ func provideCleanup(
subscriptionExpiry *service.SubscriptionExpiryService,
usageCleanup *service.UsageCleanupService,
idempotencyCleanup *service.IdempotencyCleanupService,
batchImageCleanup *service.BatchImageCleanupService,
batchImageWorker *service.BatchImageWorkerRuntime,
pricing *service.PricingService,
emailQueue *service.EmailQueueService,
billingCache *service.BillingCacheService,
@@ -167,6 +169,18 @@ func provideCleanup(
}
return nil
}},
{"BatchImageCleanupService", func() error {
if batchImageCleanup != nil {
batchImageCleanup.Stop()
}
return nil
}},
{"BatchImageWorkerRuntime", func() error {
if batchImageWorker != nil {
batchImageWorker.Stop()
}
return nil
}},
{"TokenRefreshService", func() error {
tokenRefresh.Stop()
return nil
+25 -2
View File
@@ -96,6 +96,9 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) {
usageLogRepository := repository.NewUsageLogRepository(client, db)
usageService := service.NewUsageService(usageLogRepository, userRepository, client, apiKeyAuthCacheInvalidator)
opsRepository := repository.NewOpsRepository(db)
batchImageRepository := repository.NewBatchImageRepository(db)
batchImageQueue := repository.NewBatchImageQueue(redisClient, configConfig)
batchImageDownloadLimiter := repository.NewBatchImageDownloadLimiter(redisClient, configConfig)
usageBillingRepository := repository.NewUsageBillingRepository(client, db)
gatewayCache := repository.NewGatewayCache(redisClient)
schedulerOutboxRepository := repository.NewSchedulerOutboxRepository(db)
@@ -134,6 +137,11 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) {
channelRepository := repository.NewChannelRepository(db)
channelService := service.NewChannelService(channelRepository, groupRepository, apiKeyAuthCacheInvalidator, pricingService)
modelPricingResolver := service.NewModelPricingResolver(channelService, billingService)
batchImageModelPricingResolver := service.ProvideBatchImageModelPricingResolver(modelPricingResolver)
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, 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)
@@ -259,9 +267,10 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) {
handlerPaymentHandler := handler.NewPaymentHandler(paymentService, paymentConfigService, channelService)
paymentWebhookHandler := handler.NewPaymentWebhookHandler(paymentService, registry)
availableChannelHandler := handler.NewAvailableChannelHandler(channelService, apiKeyService, settingService)
batchImageHandler := handler.NewBatchImageHandler(batchImagePublicService, batchImageDownloadService, batchImageCleanupService)
idempotencyCoordinator := service.ProvideIdempotencyCoordinator(idempotencyRepository, configConfig)
idempotencyCleanupService := service.ProvideIdempotencyCleanupService(idempotencyRepository, configConfig)
handlers := handler.ProvideHandlers(authHandler, userHandler, apiKeyHandler, usageHandler, redeemHandler, subscriptionHandler, announcementHandler, channelMonitorUserHandler, adminHandlers, gatewayHandler, openAIGatewayHandler, handlerSettingHandler, totpHandler, handlerPaymentHandler, paymentWebhookHandler, availableChannelHandler, idempotencyCoordinator, idempotencyCleanupService)
handlers := handler.ProvideHandlers(authHandler, userHandler, apiKeyHandler, usageHandler, redeemHandler, subscriptionHandler, announcementHandler, channelMonitorUserHandler, adminHandlers, gatewayHandler, openAIGatewayHandler, handlerSettingHandler, totpHandler, handlerPaymentHandler, paymentWebhookHandler, availableChannelHandler, batchImageHandler, idempotencyCoordinator, idempotencyCleanupService)
jwtAuthMiddleware := middleware.NewJWTAuthMiddleware(authService, userService)
adminAuthMiddleware := middleware.NewAdminAuthMiddleware(authService, userService, settingService)
apiKeyAuthMiddleware := middleware.NewAPIKeyAuthMiddleware(apiKeyService, subscriptionService, configConfig)
@@ -280,7 +289,7 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) {
paymentOrderExpiryService := service.ProvidePaymentOrderExpiryService(paymentService, leaderLockCache, db)
channelMonitorRunner := service.ProvideChannelMonitorRunner(channelMonitorService, settingService)
userPlatformQuotaUsageFlusher := service.ProvideUserPlatformQuotaUsageFlusher(configConfig, billingCache, serviceUserPlatformQuotaRepository, timingWheelService)
v := provideCleanup(client, redisClient, opsMetricsCollector, opsAggregationService, opsAlertEvaluatorService, opsCleanupService, opsScheduledReportService, opsSystemLogSink, schedulerSnapshotService, tokenRefreshService, accountExpiryService, proxyExpiryService, subscriptionExpiryService, usageCleanupService, idempotencyCleanupService, pricingService, emailQueueService, billingCacheService, usageRecordWorkerPool, subscriptionService, oAuthService, openAIOAuthService, geminiOAuthService, antigravityOAuthService, grokOAuthService, openAIGatewayService, scheduledTestRunnerService, backupService, paymentOrderExpiryService, channelMonitorRunner, userPlatformQuotaUsageFlusher)
v := provideCleanup(client, redisClient, opsMetricsCollector, opsAggregationService, opsAlertEvaluatorService, opsCleanupService, opsScheduledReportService, opsSystemLogSink, schedulerSnapshotService, tokenRefreshService, accountExpiryService, proxyExpiryService, subscriptionExpiryService, usageCleanupService, idempotencyCleanupService, batchImageCleanupService, batchImageWorkerRuntime, pricingService, emailQueueService, billingCacheService, usageRecordWorkerPool, subscriptionService, oAuthService, openAIOAuthService, geminiOAuthService, antigravityOAuthService, grokOAuthService, openAIGatewayService, scheduledTestRunnerService, backupService, paymentOrderExpiryService, channelMonitorRunner, userPlatformQuotaUsageFlusher)
application := &Application{
Server: httpServer,
Cleanup: v,
@@ -322,6 +331,8 @@ func provideCleanup(
subscriptionExpiry *service.SubscriptionExpiryService,
usageCleanup *service.UsageCleanupService,
idempotencyCleanup *service.IdempotencyCleanupService,
batchImageCleanup *service.BatchImageCleanupService,
batchImageWorker *service.BatchImageWorkerRuntime,
pricing *service.PricingService,
emailQueue *service.EmailQueueService,
billingCache *service.BillingCacheService,
@@ -403,6 +414,18 @@ func provideCleanup(
}
return nil
}},
{"BatchImageCleanupService", func() error {
if batchImageCleanup != nil {
batchImageCleanup.Stop()
}
return nil
}},
{"BatchImageWorkerRuntime", func() error {
if batchImageWorker != nil {
batchImageWorker.Stop()
}
return nil
}},
{"TokenRefreshService", func() error {
tokenRefresh.Stop()
return nil
+2
View File
@@ -65,6 +65,8 @@ func TestProvideCleanup_WithMinimalDependencies_NoPanic(t *testing.T) {
subscriptionExpirySvc,
&service.UsageCleanupService{},
idempotencyCleanupSvc,
&service.BatchImageCleanupService{},
nil, // batchImageWorker
pricingSvc,
emailQueueSvc,
billingCacheSvc,
+158
View File
@@ -0,0 +1,158 @@
// Code generated by ent, DO NOT EDIT.
package ent
import (
"encoding/json"
"fmt"
"strings"
"time"
"entgo.io/ent"
"entgo.io/ent/dialect/sql"
"github.com/Wei-Shaw/sub2api/ent/batchimageevent"
)
// BatchImageEvent is the model entity for the BatchImageEvent schema.
type BatchImageEvent struct {
config `json:"-"`
// ID of the ent.
ID int64 `json:"id,omitempty"`
// JobID holds the value of the "job_id" field.
JobID string `json:"job_id,omitempty"`
// EventType holds the value of the "event_type" field.
EventType string `json:"event_type,omitempty"`
// Payload holds the value of the "payload" field.
Payload map[string]interface{} `json:"payload,omitempty"`
// EventHash holds the value of the "event_hash" field.
EventHash *string `json:"event_hash,omitempty"`
// CreatedAt holds the value of the "created_at" field.
CreatedAt time.Time `json:"created_at,omitempty"`
selectValues sql.SelectValues
}
// scanValues returns the types for scanning values from sql.Rows.
func (*BatchImageEvent) scanValues(columns []string) ([]any, error) {
values := make([]any, len(columns))
for i := range columns {
switch columns[i] {
case batchimageevent.FieldPayload:
values[i] = new([]byte)
case batchimageevent.FieldID:
values[i] = new(sql.NullInt64)
case batchimageevent.FieldJobID, batchimageevent.FieldEventType, batchimageevent.FieldEventHash:
values[i] = new(sql.NullString)
case batchimageevent.FieldCreatedAt:
values[i] = new(sql.NullTime)
default:
values[i] = new(sql.UnknownType)
}
}
return values, nil
}
// assignValues assigns the values that were returned from sql.Rows (after scanning)
// to the BatchImageEvent fields.
func (_m *BatchImageEvent) assignValues(columns []string, values []any) error {
if m, n := len(values), len(columns); m < n {
return fmt.Errorf("mismatch number of scan values: %d != %d", m, n)
}
for i := range columns {
switch columns[i] {
case batchimageevent.FieldID:
value, ok := values[i].(*sql.NullInt64)
if !ok {
return fmt.Errorf("unexpected type %T for field id", value)
}
_m.ID = int64(value.Int64)
case batchimageevent.FieldJobID:
if value, ok := values[i].(*sql.NullString); !ok {
return fmt.Errorf("unexpected type %T for field job_id", values[i])
} else if value.Valid {
_m.JobID = value.String
}
case batchimageevent.FieldEventType:
if value, ok := values[i].(*sql.NullString); !ok {
return fmt.Errorf("unexpected type %T for field event_type", values[i])
} else if value.Valid {
_m.EventType = value.String
}
case batchimageevent.FieldPayload:
if value, ok := values[i].(*[]byte); !ok {
return fmt.Errorf("unexpected type %T for field payload", values[i])
} else if value != nil && len(*value) > 0 {
if err := json.Unmarshal(*value, &_m.Payload); err != nil {
return fmt.Errorf("unmarshal field payload: %w", err)
}
}
case batchimageevent.FieldEventHash:
if value, ok := values[i].(*sql.NullString); !ok {
return fmt.Errorf("unexpected type %T for field event_hash", values[i])
} else if value.Valid {
_m.EventHash = new(string)
*_m.EventHash = value.String
}
case batchimageevent.FieldCreatedAt:
if value, ok := values[i].(*sql.NullTime); !ok {
return fmt.Errorf("unexpected type %T for field created_at", values[i])
} else if value.Valid {
_m.CreatedAt = value.Time
}
default:
_m.selectValues.Set(columns[i], values[i])
}
}
return nil
}
// Value returns the ent.Value that was dynamically selected and assigned to the BatchImageEvent.
// This includes values selected through modifiers, order, etc.
func (_m *BatchImageEvent) Value(name string) (ent.Value, error) {
return _m.selectValues.Get(name)
}
// Update returns a builder for updating this BatchImageEvent.
// Note that you need to call BatchImageEvent.Unwrap() before calling this method if this BatchImageEvent
// was returned from a transaction, and the transaction was committed or rolled back.
func (_m *BatchImageEvent) Update() *BatchImageEventUpdateOne {
return NewBatchImageEventClient(_m.config).UpdateOne(_m)
}
// Unwrap unwraps the BatchImageEvent entity that was returned from a transaction after it was closed,
// so that all future queries will be executed through the driver which created the transaction.
func (_m *BatchImageEvent) Unwrap() *BatchImageEvent {
_tx, ok := _m.config.driver.(*txDriver)
if !ok {
panic("ent: BatchImageEvent is not a transactional entity")
}
_m.config.driver = _tx.drv
return _m
}
// String implements the fmt.Stringer.
func (_m *BatchImageEvent) String() string {
var builder strings.Builder
builder.WriteString("BatchImageEvent(")
builder.WriteString(fmt.Sprintf("id=%v, ", _m.ID))
builder.WriteString("job_id=")
builder.WriteString(_m.JobID)
builder.WriteString(", ")
builder.WriteString("event_type=")
builder.WriteString(_m.EventType)
builder.WriteString(", ")
builder.WriteString("payload=")
builder.WriteString(fmt.Sprintf("%v", _m.Payload))
builder.WriteString(", ")
if v := _m.EventHash; v != nil {
builder.WriteString("event_hash=")
builder.WriteString(*v)
}
builder.WriteString(", ")
builder.WriteString("created_at=")
builder.WriteString(_m.CreatedAt.Format(time.ANSIC))
builder.WriteByte(')')
return builder.String()
}
// BatchImageEvents is a parsable slice of BatchImageEvent.
type BatchImageEvents []*BatchImageEvent
@@ -0,0 +1,87 @@
// Code generated by ent, DO NOT EDIT.
package batchimageevent
import (
"time"
"entgo.io/ent/dialect/sql"
)
const (
// Label holds the string label denoting the batchimageevent type in the database.
Label = "batch_image_event"
// FieldID holds the string denoting the id field in the database.
FieldID = "id"
// FieldJobID holds the string denoting the job_id field in the database.
FieldJobID = "job_id"
// FieldEventType holds the string denoting the event_type field in the database.
FieldEventType = "event_type"
// FieldPayload holds the string denoting the payload field in the database.
FieldPayload = "payload"
// FieldEventHash holds the string denoting the event_hash field in the database.
FieldEventHash = "event_hash"
// FieldCreatedAt holds the string denoting the created_at field in the database.
FieldCreatedAt = "created_at"
// Table holds the table name of the batchimageevent in the database.
Table = "batch_image_events"
)
// Columns holds all SQL columns for batchimageevent fields.
var Columns = []string{
FieldID,
FieldJobID,
FieldEventType,
FieldPayload,
FieldEventHash,
FieldCreatedAt,
}
// ValidColumn reports if the column name is valid (part of the table columns).
func ValidColumn(column string) bool {
for i := range Columns {
if column == Columns[i] {
return true
}
}
return false
}
var (
// JobIDValidator is a validator for the "job_id" field. It is called by the builders before save.
JobIDValidator func(string) error
// EventTypeValidator is a validator for the "event_type" field. It is called by the builders before save.
EventTypeValidator func(string) error
// EventHashValidator is a validator for the "event_hash" field. It is called by the builders before save.
EventHashValidator func(string) error
// DefaultCreatedAt holds the default value on creation for the "created_at" field.
DefaultCreatedAt func() time.Time
)
// OrderOption defines the ordering options for the BatchImageEvent queries.
type OrderOption func(*sql.Selector)
// ByID orders the results by the id field.
func ByID(opts ...sql.OrderTermOption) OrderOption {
return sql.OrderByField(FieldID, opts...).ToFunc()
}
// ByJobID orders the results by the job_id field.
func ByJobID(opts ...sql.OrderTermOption) OrderOption {
return sql.OrderByField(FieldJobID, opts...).ToFunc()
}
// ByEventType orders the results by the event_type field.
func ByEventType(opts ...sql.OrderTermOption) OrderOption {
return sql.OrderByField(FieldEventType, opts...).ToFunc()
}
// ByEventHash orders the results by the event_hash field.
func ByEventHash(opts ...sql.OrderTermOption) OrderOption {
return sql.OrderByField(FieldEventHash, opts...).ToFunc()
}
// ByCreatedAt orders the results by the created_at field.
func ByCreatedAt(opts ...sql.OrderTermOption) OrderOption {
return sql.OrderByField(FieldCreatedAt, opts...).ToFunc()
}
+345
View File
@@ -0,0 +1,345 @@
// Code generated by ent, DO NOT EDIT.
package batchimageevent
import (
"time"
"entgo.io/ent/dialect/sql"
"github.com/Wei-Shaw/sub2api/ent/predicate"
)
// ID filters vertices based on their ID field.
func ID(id int64) predicate.BatchImageEvent {
return predicate.BatchImageEvent(sql.FieldEQ(FieldID, id))
}
// IDEQ applies the EQ predicate on the ID field.
func IDEQ(id int64) predicate.BatchImageEvent {
return predicate.BatchImageEvent(sql.FieldEQ(FieldID, id))
}
// IDNEQ applies the NEQ predicate on the ID field.
func IDNEQ(id int64) predicate.BatchImageEvent {
return predicate.BatchImageEvent(sql.FieldNEQ(FieldID, id))
}
// IDIn applies the In predicate on the ID field.
func IDIn(ids ...int64) predicate.BatchImageEvent {
return predicate.BatchImageEvent(sql.FieldIn(FieldID, ids...))
}
// IDNotIn applies the NotIn predicate on the ID field.
func IDNotIn(ids ...int64) predicate.BatchImageEvent {
return predicate.BatchImageEvent(sql.FieldNotIn(FieldID, ids...))
}
// IDGT applies the GT predicate on the ID field.
func IDGT(id int64) predicate.BatchImageEvent {
return predicate.BatchImageEvent(sql.FieldGT(FieldID, id))
}
// IDGTE applies the GTE predicate on the ID field.
func IDGTE(id int64) predicate.BatchImageEvent {
return predicate.BatchImageEvent(sql.FieldGTE(FieldID, id))
}
// IDLT applies the LT predicate on the ID field.
func IDLT(id int64) predicate.BatchImageEvent {
return predicate.BatchImageEvent(sql.FieldLT(FieldID, id))
}
// IDLTE applies the LTE predicate on the ID field.
func IDLTE(id int64) predicate.BatchImageEvent {
return predicate.BatchImageEvent(sql.FieldLTE(FieldID, id))
}
// JobID applies equality check predicate on the "job_id" field. It's identical to JobIDEQ.
func JobID(v string) predicate.BatchImageEvent {
return predicate.BatchImageEvent(sql.FieldEQ(FieldJobID, v))
}
// EventType applies equality check predicate on the "event_type" field. It's identical to EventTypeEQ.
func EventType(v string) predicate.BatchImageEvent {
return predicate.BatchImageEvent(sql.FieldEQ(FieldEventType, v))
}
// EventHash applies equality check predicate on the "event_hash" field. It's identical to EventHashEQ.
func EventHash(v string) predicate.BatchImageEvent {
return predicate.BatchImageEvent(sql.FieldEQ(FieldEventHash, v))
}
// CreatedAt applies equality check predicate on the "created_at" field. It's identical to CreatedAtEQ.
func CreatedAt(v time.Time) predicate.BatchImageEvent {
return predicate.BatchImageEvent(sql.FieldEQ(FieldCreatedAt, v))
}
// JobIDEQ applies the EQ predicate on the "job_id" field.
func JobIDEQ(v string) predicate.BatchImageEvent {
return predicate.BatchImageEvent(sql.FieldEQ(FieldJobID, v))
}
// JobIDNEQ applies the NEQ predicate on the "job_id" field.
func JobIDNEQ(v string) predicate.BatchImageEvent {
return predicate.BatchImageEvent(sql.FieldNEQ(FieldJobID, v))
}
// JobIDIn applies the In predicate on the "job_id" field.
func JobIDIn(vs ...string) predicate.BatchImageEvent {
return predicate.BatchImageEvent(sql.FieldIn(FieldJobID, vs...))
}
// JobIDNotIn applies the NotIn predicate on the "job_id" field.
func JobIDNotIn(vs ...string) predicate.BatchImageEvent {
return predicate.BatchImageEvent(sql.FieldNotIn(FieldJobID, vs...))
}
// JobIDGT applies the GT predicate on the "job_id" field.
func JobIDGT(v string) predicate.BatchImageEvent {
return predicate.BatchImageEvent(sql.FieldGT(FieldJobID, v))
}
// JobIDGTE applies the GTE predicate on the "job_id" field.
func JobIDGTE(v string) predicate.BatchImageEvent {
return predicate.BatchImageEvent(sql.FieldGTE(FieldJobID, v))
}
// JobIDLT applies the LT predicate on the "job_id" field.
func JobIDLT(v string) predicate.BatchImageEvent {
return predicate.BatchImageEvent(sql.FieldLT(FieldJobID, v))
}
// JobIDLTE applies the LTE predicate on the "job_id" field.
func JobIDLTE(v string) predicate.BatchImageEvent {
return predicate.BatchImageEvent(sql.FieldLTE(FieldJobID, v))
}
// JobIDContains applies the Contains predicate on the "job_id" field.
func JobIDContains(v string) predicate.BatchImageEvent {
return predicate.BatchImageEvent(sql.FieldContains(FieldJobID, v))
}
// JobIDHasPrefix applies the HasPrefix predicate on the "job_id" field.
func JobIDHasPrefix(v string) predicate.BatchImageEvent {
return predicate.BatchImageEvent(sql.FieldHasPrefix(FieldJobID, v))
}
// JobIDHasSuffix applies the HasSuffix predicate on the "job_id" field.
func JobIDHasSuffix(v string) predicate.BatchImageEvent {
return predicate.BatchImageEvent(sql.FieldHasSuffix(FieldJobID, v))
}
// JobIDEqualFold applies the EqualFold predicate on the "job_id" field.
func JobIDEqualFold(v string) predicate.BatchImageEvent {
return predicate.BatchImageEvent(sql.FieldEqualFold(FieldJobID, v))
}
// JobIDContainsFold applies the ContainsFold predicate on the "job_id" field.
func JobIDContainsFold(v string) predicate.BatchImageEvent {
return predicate.BatchImageEvent(sql.FieldContainsFold(FieldJobID, v))
}
// EventTypeEQ applies the EQ predicate on the "event_type" field.
func EventTypeEQ(v string) predicate.BatchImageEvent {
return predicate.BatchImageEvent(sql.FieldEQ(FieldEventType, v))
}
// EventTypeNEQ applies the NEQ predicate on the "event_type" field.
func EventTypeNEQ(v string) predicate.BatchImageEvent {
return predicate.BatchImageEvent(sql.FieldNEQ(FieldEventType, v))
}
// EventTypeIn applies the In predicate on the "event_type" field.
func EventTypeIn(vs ...string) predicate.BatchImageEvent {
return predicate.BatchImageEvent(sql.FieldIn(FieldEventType, vs...))
}
// EventTypeNotIn applies the NotIn predicate on the "event_type" field.
func EventTypeNotIn(vs ...string) predicate.BatchImageEvent {
return predicate.BatchImageEvent(sql.FieldNotIn(FieldEventType, vs...))
}
// EventTypeGT applies the GT predicate on the "event_type" field.
func EventTypeGT(v string) predicate.BatchImageEvent {
return predicate.BatchImageEvent(sql.FieldGT(FieldEventType, v))
}
// EventTypeGTE applies the GTE predicate on the "event_type" field.
func EventTypeGTE(v string) predicate.BatchImageEvent {
return predicate.BatchImageEvent(sql.FieldGTE(FieldEventType, v))
}
// EventTypeLT applies the LT predicate on the "event_type" field.
func EventTypeLT(v string) predicate.BatchImageEvent {
return predicate.BatchImageEvent(sql.FieldLT(FieldEventType, v))
}
// EventTypeLTE applies the LTE predicate on the "event_type" field.
func EventTypeLTE(v string) predicate.BatchImageEvent {
return predicate.BatchImageEvent(sql.FieldLTE(FieldEventType, v))
}
// EventTypeContains applies the Contains predicate on the "event_type" field.
func EventTypeContains(v string) predicate.BatchImageEvent {
return predicate.BatchImageEvent(sql.FieldContains(FieldEventType, v))
}
// EventTypeHasPrefix applies the HasPrefix predicate on the "event_type" field.
func EventTypeHasPrefix(v string) predicate.BatchImageEvent {
return predicate.BatchImageEvent(sql.FieldHasPrefix(FieldEventType, v))
}
// EventTypeHasSuffix applies the HasSuffix predicate on the "event_type" field.
func EventTypeHasSuffix(v string) predicate.BatchImageEvent {
return predicate.BatchImageEvent(sql.FieldHasSuffix(FieldEventType, v))
}
// EventTypeEqualFold applies the EqualFold predicate on the "event_type" field.
func EventTypeEqualFold(v string) predicate.BatchImageEvent {
return predicate.BatchImageEvent(sql.FieldEqualFold(FieldEventType, v))
}
// EventTypeContainsFold applies the ContainsFold predicate on the "event_type" field.
func EventTypeContainsFold(v string) predicate.BatchImageEvent {
return predicate.BatchImageEvent(sql.FieldContainsFold(FieldEventType, v))
}
// PayloadIsNil applies the IsNil predicate on the "payload" field.
func PayloadIsNil() predicate.BatchImageEvent {
return predicate.BatchImageEvent(sql.FieldIsNull(FieldPayload))
}
// PayloadNotNil applies the NotNil predicate on the "payload" field.
func PayloadNotNil() predicate.BatchImageEvent {
return predicate.BatchImageEvent(sql.FieldNotNull(FieldPayload))
}
// EventHashEQ applies the EQ predicate on the "event_hash" field.
func EventHashEQ(v string) predicate.BatchImageEvent {
return predicate.BatchImageEvent(sql.FieldEQ(FieldEventHash, v))
}
// EventHashNEQ applies the NEQ predicate on the "event_hash" field.
func EventHashNEQ(v string) predicate.BatchImageEvent {
return predicate.BatchImageEvent(sql.FieldNEQ(FieldEventHash, v))
}
// EventHashIn applies the In predicate on the "event_hash" field.
func EventHashIn(vs ...string) predicate.BatchImageEvent {
return predicate.BatchImageEvent(sql.FieldIn(FieldEventHash, vs...))
}
// EventHashNotIn applies the NotIn predicate on the "event_hash" field.
func EventHashNotIn(vs ...string) predicate.BatchImageEvent {
return predicate.BatchImageEvent(sql.FieldNotIn(FieldEventHash, vs...))
}
// EventHashGT applies the GT predicate on the "event_hash" field.
func EventHashGT(v string) predicate.BatchImageEvent {
return predicate.BatchImageEvent(sql.FieldGT(FieldEventHash, v))
}
// EventHashGTE applies the GTE predicate on the "event_hash" field.
func EventHashGTE(v string) predicate.BatchImageEvent {
return predicate.BatchImageEvent(sql.FieldGTE(FieldEventHash, v))
}
// EventHashLT applies the LT predicate on the "event_hash" field.
func EventHashLT(v string) predicate.BatchImageEvent {
return predicate.BatchImageEvent(sql.FieldLT(FieldEventHash, v))
}
// EventHashLTE applies the LTE predicate on the "event_hash" field.
func EventHashLTE(v string) predicate.BatchImageEvent {
return predicate.BatchImageEvent(sql.FieldLTE(FieldEventHash, v))
}
// EventHashContains applies the Contains predicate on the "event_hash" field.
func EventHashContains(v string) predicate.BatchImageEvent {
return predicate.BatchImageEvent(sql.FieldContains(FieldEventHash, v))
}
// EventHashHasPrefix applies the HasPrefix predicate on the "event_hash" field.
func EventHashHasPrefix(v string) predicate.BatchImageEvent {
return predicate.BatchImageEvent(sql.FieldHasPrefix(FieldEventHash, v))
}
// EventHashHasSuffix applies the HasSuffix predicate on the "event_hash" field.
func EventHashHasSuffix(v string) predicate.BatchImageEvent {
return predicate.BatchImageEvent(sql.FieldHasSuffix(FieldEventHash, v))
}
// EventHashIsNil applies the IsNil predicate on the "event_hash" field.
func EventHashIsNil() predicate.BatchImageEvent {
return predicate.BatchImageEvent(sql.FieldIsNull(FieldEventHash))
}
// EventHashNotNil applies the NotNil predicate on the "event_hash" field.
func EventHashNotNil() predicate.BatchImageEvent {
return predicate.BatchImageEvent(sql.FieldNotNull(FieldEventHash))
}
// EventHashEqualFold applies the EqualFold predicate on the "event_hash" field.
func EventHashEqualFold(v string) predicate.BatchImageEvent {
return predicate.BatchImageEvent(sql.FieldEqualFold(FieldEventHash, v))
}
// EventHashContainsFold applies the ContainsFold predicate on the "event_hash" field.
func EventHashContainsFold(v string) predicate.BatchImageEvent {
return predicate.BatchImageEvent(sql.FieldContainsFold(FieldEventHash, v))
}
// CreatedAtEQ applies the EQ predicate on the "created_at" field.
func CreatedAtEQ(v time.Time) predicate.BatchImageEvent {
return predicate.BatchImageEvent(sql.FieldEQ(FieldCreatedAt, v))
}
// CreatedAtNEQ applies the NEQ predicate on the "created_at" field.
func CreatedAtNEQ(v time.Time) predicate.BatchImageEvent {
return predicate.BatchImageEvent(sql.FieldNEQ(FieldCreatedAt, v))
}
// CreatedAtIn applies the In predicate on the "created_at" field.
func CreatedAtIn(vs ...time.Time) predicate.BatchImageEvent {
return predicate.BatchImageEvent(sql.FieldIn(FieldCreatedAt, vs...))
}
// CreatedAtNotIn applies the NotIn predicate on the "created_at" field.
func CreatedAtNotIn(vs ...time.Time) predicate.BatchImageEvent {
return predicate.BatchImageEvent(sql.FieldNotIn(FieldCreatedAt, vs...))
}
// CreatedAtGT applies the GT predicate on the "created_at" field.
func CreatedAtGT(v time.Time) predicate.BatchImageEvent {
return predicate.BatchImageEvent(sql.FieldGT(FieldCreatedAt, v))
}
// CreatedAtGTE applies the GTE predicate on the "created_at" field.
func CreatedAtGTE(v time.Time) predicate.BatchImageEvent {
return predicate.BatchImageEvent(sql.FieldGTE(FieldCreatedAt, v))
}
// CreatedAtLT applies the LT predicate on the "created_at" field.
func CreatedAtLT(v time.Time) predicate.BatchImageEvent {
return predicate.BatchImageEvent(sql.FieldLT(FieldCreatedAt, v))
}
// CreatedAtLTE applies the LTE predicate on the "created_at" field.
func CreatedAtLTE(v time.Time) predicate.BatchImageEvent {
return predicate.BatchImageEvent(sql.FieldLTE(FieldCreatedAt, v))
}
// And groups predicates with the AND operator between them.
func And(predicates ...predicate.BatchImageEvent) predicate.BatchImageEvent {
return predicate.BatchImageEvent(sql.AndPredicates(predicates...))
}
// Or groups predicates with the OR operator between them.
func Or(predicates ...predicate.BatchImageEvent) predicate.BatchImageEvent {
return predicate.BatchImageEvent(sql.OrPredicates(predicates...))
}
// Not applies the not operator on the given predicate.
func Not(p predicate.BatchImageEvent) predicate.BatchImageEvent {
return predicate.BatchImageEvent(sql.NotPredicates(p))
}
+714
View File
@@ -0,0 +1,714 @@
// Code generated by ent, DO NOT EDIT.
package ent
import (
"context"
"errors"
"fmt"
"time"
"entgo.io/ent/dialect/sql"
"entgo.io/ent/dialect/sql/sqlgraph"
"entgo.io/ent/schema/field"
"github.com/Wei-Shaw/sub2api/ent/batchimageevent"
)
// BatchImageEventCreate is the builder for creating a BatchImageEvent entity.
type BatchImageEventCreate struct {
config
mutation *BatchImageEventMutation
hooks []Hook
conflict []sql.ConflictOption
}
// SetJobID sets the "job_id" field.
func (_c *BatchImageEventCreate) SetJobID(v string) *BatchImageEventCreate {
_c.mutation.SetJobID(v)
return _c
}
// SetEventType sets the "event_type" field.
func (_c *BatchImageEventCreate) SetEventType(v string) *BatchImageEventCreate {
_c.mutation.SetEventType(v)
return _c
}
// SetPayload sets the "payload" field.
func (_c *BatchImageEventCreate) SetPayload(v map[string]interface{}) *BatchImageEventCreate {
_c.mutation.SetPayload(v)
return _c
}
// SetEventHash sets the "event_hash" field.
func (_c *BatchImageEventCreate) SetEventHash(v string) *BatchImageEventCreate {
_c.mutation.SetEventHash(v)
return _c
}
// SetNillableEventHash sets the "event_hash" field if the given value is not nil.
func (_c *BatchImageEventCreate) SetNillableEventHash(v *string) *BatchImageEventCreate {
if v != nil {
_c.SetEventHash(*v)
}
return _c
}
// SetCreatedAt sets the "created_at" field.
func (_c *BatchImageEventCreate) SetCreatedAt(v time.Time) *BatchImageEventCreate {
_c.mutation.SetCreatedAt(v)
return _c
}
// SetNillableCreatedAt sets the "created_at" field if the given value is not nil.
func (_c *BatchImageEventCreate) SetNillableCreatedAt(v *time.Time) *BatchImageEventCreate {
if v != nil {
_c.SetCreatedAt(*v)
}
return _c
}
// Mutation returns the BatchImageEventMutation object of the builder.
func (_c *BatchImageEventCreate) Mutation() *BatchImageEventMutation {
return _c.mutation
}
// Save creates the BatchImageEvent in the database.
func (_c *BatchImageEventCreate) Save(ctx context.Context) (*BatchImageEvent, error) {
_c.defaults()
return withHooks(ctx, _c.sqlSave, _c.mutation, _c.hooks)
}
// SaveX calls Save and panics if Save returns an error.
func (_c *BatchImageEventCreate) SaveX(ctx context.Context) *BatchImageEvent {
v, err := _c.Save(ctx)
if err != nil {
panic(err)
}
return v
}
// Exec executes the query.
func (_c *BatchImageEventCreate) Exec(ctx context.Context) error {
_, err := _c.Save(ctx)
return err
}
// ExecX is like Exec, but panics if an error occurs.
func (_c *BatchImageEventCreate) ExecX(ctx context.Context) {
if err := _c.Exec(ctx); err != nil {
panic(err)
}
}
// defaults sets the default values of the builder before save.
func (_c *BatchImageEventCreate) defaults() {
if _, ok := _c.mutation.CreatedAt(); !ok {
v := batchimageevent.DefaultCreatedAt()
_c.mutation.SetCreatedAt(v)
}
}
// check runs all checks and user-defined validators on the builder.
func (_c *BatchImageEventCreate) check() error {
if _, ok := _c.mutation.JobID(); !ok {
return &ValidationError{Name: "job_id", err: errors.New(`ent: missing required field "BatchImageEvent.job_id"`)}
}
if v, ok := _c.mutation.JobID(); ok {
if err := batchimageevent.JobIDValidator(v); err != nil {
return &ValidationError{Name: "job_id", err: fmt.Errorf(`ent: validator failed for field "BatchImageEvent.job_id": %w`, err)}
}
}
if _, ok := _c.mutation.EventType(); !ok {
return &ValidationError{Name: "event_type", err: errors.New(`ent: missing required field "BatchImageEvent.event_type"`)}
}
if v, ok := _c.mutation.EventType(); ok {
if err := batchimageevent.EventTypeValidator(v); err != nil {
return &ValidationError{Name: "event_type", err: fmt.Errorf(`ent: validator failed for field "BatchImageEvent.event_type": %w`, err)}
}
}
if v, ok := _c.mutation.EventHash(); ok {
if err := batchimageevent.EventHashValidator(v); err != nil {
return &ValidationError{Name: "event_hash", err: fmt.Errorf(`ent: validator failed for field "BatchImageEvent.event_hash": %w`, err)}
}
}
if _, ok := _c.mutation.CreatedAt(); !ok {
return &ValidationError{Name: "created_at", err: errors.New(`ent: missing required field "BatchImageEvent.created_at"`)}
}
return nil
}
func (_c *BatchImageEventCreate) sqlSave(ctx context.Context) (*BatchImageEvent, error) {
if err := _c.check(); err != nil {
return nil, err
}
_node, _spec := _c.createSpec()
if err := sqlgraph.CreateNode(ctx, _c.driver, _spec); err != nil {
if sqlgraph.IsConstraintError(err) {
err = &ConstraintError{msg: err.Error(), wrap: err}
}
return nil, err
}
id := _spec.ID.Value.(int64)
_node.ID = int64(id)
_c.mutation.id = &_node.ID
_c.mutation.done = true
return _node, nil
}
func (_c *BatchImageEventCreate) createSpec() (*BatchImageEvent, *sqlgraph.CreateSpec) {
var (
_node = &BatchImageEvent{config: _c.config}
_spec = sqlgraph.NewCreateSpec(batchimageevent.Table, sqlgraph.NewFieldSpec(batchimageevent.FieldID, field.TypeInt64))
)
_spec.OnConflict = _c.conflict
if value, ok := _c.mutation.JobID(); ok {
_spec.SetField(batchimageevent.FieldJobID, field.TypeString, value)
_node.JobID = value
}
if value, ok := _c.mutation.EventType(); ok {
_spec.SetField(batchimageevent.FieldEventType, field.TypeString, value)
_node.EventType = value
}
if value, ok := _c.mutation.Payload(); ok {
_spec.SetField(batchimageevent.FieldPayload, field.TypeJSON, value)
_node.Payload = value
}
if value, ok := _c.mutation.EventHash(); ok {
_spec.SetField(batchimageevent.FieldEventHash, field.TypeString, value)
_node.EventHash = &value
}
if value, ok := _c.mutation.CreatedAt(); ok {
_spec.SetField(batchimageevent.FieldCreatedAt, field.TypeTime, value)
_node.CreatedAt = value
}
return _node, _spec
}
// OnConflict allows configuring the `ON CONFLICT` / `ON DUPLICATE KEY` clause
// of the `INSERT` statement. For example:
//
// client.BatchImageEvent.Create().
// SetJobID(v).
// OnConflict(
// // Update the row with the new values
// // the was proposed for insertion.
// sql.ResolveWithNewValues(),
// ).
// // Override some of the fields with custom
// // update values.
// Update(func(u *ent.BatchImageEventUpsert) {
// SetJobID(v+v).
// }).
// Exec(ctx)
func (_c *BatchImageEventCreate) OnConflict(opts ...sql.ConflictOption) *BatchImageEventUpsertOne {
_c.conflict = opts
return &BatchImageEventUpsertOne{
create: _c,
}
}
// OnConflictColumns calls `OnConflict` and configures the columns
// as conflict target. Using this option is equivalent to using:
//
// client.BatchImageEvent.Create().
// OnConflict(sql.ConflictColumns(columns...)).
// Exec(ctx)
func (_c *BatchImageEventCreate) OnConflictColumns(columns ...string) *BatchImageEventUpsertOne {
_c.conflict = append(_c.conflict, sql.ConflictColumns(columns...))
return &BatchImageEventUpsertOne{
create: _c,
}
}
type (
// BatchImageEventUpsertOne is the builder for "upsert"-ing
// one BatchImageEvent node.
BatchImageEventUpsertOne struct {
create *BatchImageEventCreate
}
// BatchImageEventUpsert is the "OnConflict" setter.
BatchImageEventUpsert struct {
*sql.UpdateSet
}
)
// SetJobID sets the "job_id" field.
func (u *BatchImageEventUpsert) SetJobID(v string) *BatchImageEventUpsert {
u.Set(batchimageevent.FieldJobID, v)
return u
}
// UpdateJobID sets the "job_id" field to the value that was provided on create.
func (u *BatchImageEventUpsert) UpdateJobID() *BatchImageEventUpsert {
u.SetExcluded(batchimageevent.FieldJobID)
return u
}
// SetEventType sets the "event_type" field.
func (u *BatchImageEventUpsert) SetEventType(v string) *BatchImageEventUpsert {
u.Set(batchimageevent.FieldEventType, v)
return u
}
// UpdateEventType sets the "event_type" field to the value that was provided on create.
func (u *BatchImageEventUpsert) UpdateEventType() *BatchImageEventUpsert {
u.SetExcluded(batchimageevent.FieldEventType)
return u
}
// SetPayload sets the "payload" field.
func (u *BatchImageEventUpsert) SetPayload(v map[string]interface{}) *BatchImageEventUpsert {
u.Set(batchimageevent.FieldPayload, v)
return u
}
// UpdatePayload sets the "payload" field to the value that was provided on create.
func (u *BatchImageEventUpsert) UpdatePayload() *BatchImageEventUpsert {
u.SetExcluded(batchimageevent.FieldPayload)
return u
}
// ClearPayload clears the value of the "payload" field.
func (u *BatchImageEventUpsert) ClearPayload() *BatchImageEventUpsert {
u.SetNull(batchimageevent.FieldPayload)
return u
}
// SetEventHash sets the "event_hash" field.
func (u *BatchImageEventUpsert) SetEventHash(v string) *BatchImageEventUpsert {
u.Set(batchimageevent.FieldEventHash, v)
return u
}
// UpdateEventHash sets the "event_hash" field to the value that was provided on create.
func (u *BatchImageEventUpsert) UpdateEventHash() *BatchImageEventUpsert {
u.SetExcluded(batchimageevent.FieldEventHash)
return u
}
// ClearEventHash clears the value of the "event_hash" field.
func (u *BatchImageEventUpsert) ClearEventHash() *BatchImageEventUpsert {
u.SetNull(batchimageevent.FieldEventHash)
return u
}
// UpdateNewValues updates the mutable fields using the new values that were set on create.
// Using this option is equivalent to using:
//
// client.BatchImageEvent.Create().
// OnConflict(
// sql.ResolveWithNewValues(),
// ).
// Exec(ctx)
func (u *BatchImageEventUpsertOne) UpdateNewValues() *BatchImageEventUpsertOne {
u.create.conflict = append(u.create.conflict, sql.ResolveWithNewValues())
u.create.conflict = append(u.create.conflict, sql.ResolveWith(func(s *sql.UpdateSet) {
if _, exists := u.create.mutation.CreatedAt(); exists {
s.SetIgnore(batchimageevent.FieldCreatedAt)
}
}))
return u
}
// Ignore sets each column to itself in case of conflict.
// Using this option is equivalent to using:
//
// client.BatchImageEvent.Create().
// OnConflict(sql.ResolveWithIgnore()).
// Exec(ctx)
func (u *BatchImageEventUpsertOne) Ignore() *BatchImageEventUpsertOne {
u.create.conflict = append(u.create.conflict, sql.ResolveWithIgnore())
return u
}
// DoNothing configures the conflict_action to `DO NOTHING`.
// Supported only by SQLite and PostgreSQL.
func (u *BatchImageEventUpsertOne) DoNothing() *BatchImageEventUpsertOne {
u.create.conflict = append(u.create.conflict, sql.DoNothing())
return u
}
// Update allows overriding fields `UPDATE` values. See the BatchImageEventCreate.OnConflict
// documentation for more info.
func (u *BatchImageEventUpsertOne) Update(set func(*BatchImageEventUpsert)) *BatchImageEventUpsertOne {
u.create.conflict = append(u.create.conflict, sql.ResolveWith(func(update *sql.UpdateSet) {
set(&BatchImageEventUpsert{UpdateSet: update})
}))
return u
}
// SetJobID sets the "job_id" field.
func (u *BatchImageEventUpsertOne) SetJobID(v string) *BatchImageEventUpsertOne {
return u.Update(func(s *BatchImageEventUpsert) {
s.SetJobID(v)
})
}
// UpdateJobID sets the "job_id" field to the value that was provided on create.
func (u *BatchImageEventUpsertOne) UpdateJobID() *BatchImageEventUpsertOne {
return u.Update(func(s *BatchImageEventUpsert) {
s.UpdateJobID()
})
}
// SetEventType sets the "event_type" field.
func (u *BatchImageEventUpsertOne) SetEventType(v string) *BatchImageEventUpsertOne {
return u.Update(func(s *BatchImageEventUpsert) {
s.SetEventType(v)
})
}
// UpdateEventType sets the "event_type" field to the value that was provided on create.
func (u *BatchImageEventUpsertOne) UpdateEventType() *BatchImageEventUpsertOne {
return u.Update(func(s *BatchImageEventUpsert) {
s.UpdateEventType()
})
}
// SetPayload sets the "payload" field.
func (u *BatchImageEventUpsertOne) SetPayload(v map[string]interface{}) *BatchImageEventUpsertOne {
return u.Update(func(s *BatchImageEventUpsert) {
s.SetPayload(v)
})
}
// UpdatePayload sets the "payload" field to the value that was provided on create.
func (u *BatchImageEventUpsertOne) UpdatePayload() *BatchImageEventUpsertOne {
return u.Update(func(s *BatchImageEventUpsert) {
s.UpdatePayload()
})
}
// ClearPayload clears the value of the "payload" field.
func (u *BatchImageEventUpsertOne) ClearPayload() *BatchImageEventUpsertOne {
return u.Update(func(s *BatchImageEventUpsert) {
s.ClearPayload()
})
}
// SetEventHash sets the "event_hash" field.
func (u *BatchImageEventUpsertOne) SetEventHash(v string) *BatchImageEventUpsertOne {
return u.Update(func(s *BatchImageEventUpsert) {
s.SetEventHash(v)
})
}
// UpdateEventHash sets the "event_hash" field to the value that was provided on create.
func (u *BatchImageEventUpsertOne) UpdateEventHash() *BatchImageEventUpsertOne {
return u.Update(func(s *BatchImageEventUpsert) {
s.UpdateEventHash()
})
}
// ClearEventHash clears the value of the "event_hash" field.
func (u *BatchImageEventUpsertOne) ClearEventHash() *BatchImageEventUpsertOne {
return u.Update(func(s *BatchImageEventUpsert) {
s.ClearEventHash()
})
}
// Exec executes the query.
func (u *BatchImageEventUpsertOne) Exec(ctx context.Context) error {
if len(u.create.conflict) == 0 {
return errors.New("ent: missing options for BatchImageEventCreate.OnConflict")
}
return u.create.Exec(ctx)
}
// ExecX is like Exec, but panics if an error occurs.
func (u *BatchImageEventUpsertOne) ExecX(ctx context.Context) {
if err := u.create.Exec(ctx); err != nil {
panic(err)
}
}
// Exec executes the UPSERT query and returns the inserted/updated ID.
func (u *BatchImageEventUpsertOne) ID(ctx context.Context) (id int64, err error) {
node, err := u.create.Save(ctx)
if err != nil {
return id, err
}
return node.ID, nil
}
// IDX is like ID, but panics if an error occurs.
func (u *BatchImageEventUpsertOne) IDX(ctx context.Context) int64 {
id, err := u.ID(ctx)
if err != nil {
panic(err)
}
return id
}
// BatchImageEventCreateBulk is the builder for creating many BatchImageEvent entities in bulk.
type BatchImageEventCreateBulk struct {
config
err error
builders []*BatchImageEventCreate
conflict []sql.ConflictOption
}
// Save creates the BatchImageEvent entities in the database.
func (_c *BatchImageEventCreateBulk) Save(ctx context.Context) ([]*BatchImageEvent, error) {
if _c.err != nil {
return nil, _c.err
}
specs := make([]*sqlgraph.CreateSpec, len(_c.builders))
nodes := make([]*BatchImageEvent, len(_c.builders))
mutators := make([]Mutator, len(_c.builders))
for i := range _c.builders {
func(i int, root context.Context) {
builder := _c.builders[i]
builder.defaults()
var mut Mutator = MutateFunc(func(ctx context.Context, m Mutation) (Value, error) {
mutation, ok := m.(*BatchImageEventMutation)
if !ok {
return nil, fmt.Errorf("unexpected mutation type %T", m)
}
if err := builder.check(); err != nil {
return nil, err
}
builder.mutation = mutation
var err error
nodes[i], specs[i] = builder.createSpec()
if i < len(mutators)-1 {
_, err = mutators[i+1].Mutate(root, _c.builders[i+1].mutation)
} else {
spec := &sqlgraph.BatchCreateSpec{Nodes: specs}
spec.OnConflict = _c.conflict
// Invoke the actual operation on the latest mutation in the chain.
if err = sqlgraph.BatchCreate(ctx, _c.driver, spec); err != nil {
if sqlgraph.IsConstraintError(err) {
err = &ConstraintError{msg: err.Error(), wrap: err}
}
}
}
if err != nil {
return nil, err
}
mutation.id = &nodes[i].ID
if specs[i].ID.Value != nil {
id := specs[i].ID.Value.(int64)
nodes[i].ID = int64(id)
}
mutation.done = true
return nodes[i], nil
})
for i := len(builder.hooks) - 1; i >= 0; i-- {
mut = builder.hooks[i](mut)
}
mutators[i] = mut
}(i, ctx)
}
if len(mutators) > 0 {
if _, err := mutators[0].Mutate(ctx, _c.builders[0].mutation); err != nil {
return nil, err
}
}
return nodes, nil
}
// SaveX is like Save, but panics if an error occurs.
func (_c *BatchImageEventCreateBulk) SaveX(ctx context.Context) []*BatchImageEvent {
v, err := _c.Save(ctx)
if err != nil {
panic(err)
}
return v
}
// Exec executes the query.
func (_c *BatchImageEventCreateBulk) Exec(ctx context.Context) error {
_, err := _c.Save(ctx)
return err
}
// ExecX is like Exec, but panics if an error occurs.
func (_c *BatchImageEventCreateBulk) ExecX(ctx context.Context) {
if err := _c.Exec(ctx); err != nil {
panic(err)
}
}
// OnConflict allows configuring the `ON CONFLICT` / `ON DUPLICATE KEY` clause
// of the `INSERT` statement. For example:
//
// client.BatchImageEvent.CreateBulk(builders...).
// OnConflict(
// // Update the row with the new values
// // the was proposed for insertion.
// sql.ResolveWithNewValues(),
// ).
// // Override some of the fields with custom
// // update values.
// Update(func(u *ent.BatchImageEventUpsert) {
// SetJobID(v+v).
// }).
// Exec(ctx)
func (_c *BatchImageEventCreateBulk) OnConflict(opts ...sql.ConflictOption) *BatchImageEventUpsertBulk {
_c.conflict = opts
return &BatchImageEventUpsertBulk{
create: _c,
}
}
// OnConflictColumns calls `OnConflict` and configures the columns
// as conflict target. Using this option is equivalent to using:
//
// client.BatchImageEvent.Create().
// OnConflict(sql.ConflictColumns(columns...)).
// Exec(ctx)
func (_c *BatchImageEventCreateBulk) OnConflictColumns(columns ...string) *BatchImageEventUpsertBulk {
_c.conflict = append(_c.conflict, sql.ConflictColumns(columns...))
return &BatchImageEventUpsertBulk{
create: _c,
}
}
// BatchImageEventUpsertBulk is the builder for "upsert"-ing
// a bulk of BatchImageEvent nodes.
type BatchImageEventUpsertBulk struct {
create *BatchImageEventCreateBulk
}
// UpdateNewValues updates the mutable fields using the new values that
// were set on create. Using this option is equivalent to using:
//
// client.BatchImageEvent.Create().
// OnConflict(
// sql.ResolveWithNewValues(),
// ).
// Exec(ctx)
func (u *BatchImageEventUpsertBulk) UpdateNewValues() *BatchImageEventUpsertBulk {
u.create.conflict = append(u.create.conflict, sql.ResolveWithNewValues())
u.create.conflict = append(u.create.conflict, sql.ResolveWith(func(s *sql.UpdateSet) {
for _, b := range u.create.builders {
if _, exists := b.mutation.CreatedAt(); exists {
s.SetIgnore(batchimageevent.FieldCreatedAt)
}
}
}))
return u
}
// Ignore sets each column to itself in case of conflict.
// Using this option is equivalent to using:
//
// client.BatchImageEvent.Create().
// OnConflict(sql.ResolveWithIgnore()).
// Exec(ctx)
func (u *BatchImageEventUpsertBulk) Ignore() *BatchImageEventUpsertBulk {
u.create.conflict = append(u.create.conflict, sql.ResolveWithIgnore())
return u
}
// DoNothing configures the conflict_action to `DO NOTHING`.
// Supported only by SQLite and PostgreSQL.
func (u *BatchImageEventUpsertBulk) DoNothing() *BatchImageEventUpsertBulk {
u.create.conflict = append(u.create.conflict, sql.DoNothing())
return u
}
// Update allows overriding fields `UPDATE` values. See the BatchImageEventCreateBulk.OnConflict
// documentation for more info.
func (u *BatchImageEventUpsertBulk) Update(set func(*BatchImageEventUpsert)) *BatchImageEventUpsertBulk {
u.create.conflict = append(u.create.conflict, sql.ResolveWith(func(update *sql.UpdateSet) {
set(&BatchImageEventUpsert{UpdateSet: update})
}))
return u
}
// SetJobID sets the "job_id" field.
func (u *BatchImageEventUpsertBulk) SetJobID(v string) *BatchImageEventUpsertBulk {
return u.Update(func(s *BatchImageEventUpsert) {
s.SetJobID(v)
})
}
// UpdateJobID sets the "job_id" field to the value that was provided on create.
func (u *BatchImageEventUpsertBulk) UpdateJobID() *BatchImageEventUpsertBulk {
return u.Update(func(s *BatchImageEventUpsert) {
s.UpdateJobID()
})
}
// SetEventType sets the "event_type" field.
func (u *BatchImageEventUpsertBulk) SetEventType(v string) *BatchImageEventUpsertBulk {
return u.Update(func(s *BatchImageEventUpsert) {
s.SetEventType(v)
})
}
// UpdateEventType sets the "event_type" field to the value that was provided on create.
func (u *BatchImageEventUpsertBulk) UpdateEventType() *BatchImageEventUpsertBulk {
return u.Update(func(s *BatchImageEventUpsert) {
s.UpdateEventType()
})
}
// SetPayload sets the "payload" field.
func (u *BatchImageEventUpsertBulk) SetPayload(v map[string]interface{}) *BatchImageEventUpsertBulk {
return u.Update(func(s *BatchImageEventUpsert) {
s.SetPayload(v)
})
}
// UpdatePayload sets the "payload" field to the value that was provided on create.
func (u *BatchImageEventUpsertBulk) UpdatePayload() *BatchImageEventUpsertBulk {
return u.Update(func(s *BatchImageEventUpsert) {
s.UpdatePayload()
})
}
// ClearPayload clears the value of the "payload" field.
func (u *BatchImageEventUpsertBulk) ClearPayload() *BatchImageEventUpsertBulk {
return u.Update(func(s *BatchImageEventUpsert) {
s.ClearPayload()
})
}
// SetEventHash sets the "event_hash" field.
func (u *BatchImageEventUpsertBulk) SetEventHash(v string) *BatchImageEventUpsertBulk {
return u.Update(func(s *BatchImageEventUpsert) {
s.SetEventHash(v)
})
}
// UpdateEventHash sets the "event_hash" field to the value that was provided on create.
func (u *BatchImageEventUpsertBulk) UpdateEventHash() *BatchImageEventUpsertBulk {
return u.Update(func(s *BatchImageEventUpsert) {
s.UpdateEventHash()
})
}
// ClearEventHash clears the value of the "event_hash" field.
func (u *BatchImageEventUpsertBulk) ClearEventHash() *BatchImageEventUpsertBulk {
return u.Update(func(s *BatchImageEventUpsert) {
s.ClearEventHash()
})
}
// Exec executes the query.
func (u *BatchImageEventUpsertBulk) Exec(ctx context.Context) error {
if u.create.err != nil {
return u.create.err
}
for i, b := range u.create.builders {
if len(b.conflict) != 0 {
return fmt.Errorf("ent: OnConflict was set for builder %d. Set it on the BatchImageEventCreateBulk instead", i)
}
}
if len(u.create.conflict) == 0 {
return errors.New("ent: missing options for BatchImageEventCreateBulk.OnConflict")
}
return u.create.Exec(ctx)
}
// ExecX is like Exec, but panics if an error occurs.
func (u *BatchImageEventUpsertBulk) ExecX(ctx context.Context) {
if err := u.create.Exec(ctx); err != nil {
panic(err)
}
}
+88
View File
@@ -0,0 +1,88 @@
// Code generated by ent, DO NOT EDIT.
package ent
import (
"context"
"entgo.io/ent/dialect/sql"
"entgo.io/ent/dialect/sql/sqlgraph"
"entgo.io/ent/schema/field"
"github.com/Wei-Shaw/sub2api/ent/batchimageevent"
"github.com/Wei-Shaw/sub2api/ent/predicate"
)
// BatchImageEventDelete is the builder for deleting a BatchImageEvent entity.
type BatchImageEventDelete struct {
config
hooks []Hook
mutation *BatchImageEventMutation
}
// Where appends a list predicates to the BatchImageEventDelete builder.
func (_d *BatchImageEventDelete) Where(ps ...predicate.BatchImageEvent) *BatchImageEventDelete {
_d.mutation.Where(ps...)
return _d
}
// Exec executes the deletion query and returns how many vertices were deleted.
func (_d *BatchImageEventDelete) Exec(ctx context.Context) (int, error) {
return withHooks(ctx, _d.sqlExec, _d.mutation, _d.hooks)
}
// ExecX is like Exec, but panics if an error occurs.
func (_d *BatchImageEventDelete) ExecX(ctx context.Context) int {
n, err := _d.Exec(ctx)
if err != nil {
panic(err)
}
return n
}
func (_d *BatchImageEventDelete) sqlExec(ctx context.Context) (int, error) {
_spec := sqlgraph.NewDeleteSpec(batchimageevent.Table, sqlgraph.NewFieldSpec(batchimageevent.FieldID, field.TypeInt64))
if ps := _d.mutation.predicates; len(ps) > 0 {
_spec.Predicate = func(selector *sql.Selector) {
for i := range ps {
ps[i](selector)
}
}
}
affected, err := sqlgraph.DeleteNodes(ctx, _d.driver, _spec)
if err != nil && sqlgraph.IsConstraintError(err) {
err = &ConstraintError{msg: err.Error(), wrap: err}
}
_d.mutation.done = true
return affected, err
}
// BatchImageEventDeleteOne is the builder for deleting a single BatchImageEvent entity.
type BatchImageEventDeleteOne struct {
_d *BatchImageEventDelete
}
// Where appends a list predicates to the BatchImageEventDelete builder.
func (_d *BatchImageEventDeleteOne) Where(ps ...predicate.BatchImageEvent) *BatchImageEventDeleteOne {
_d._d.mutation.Where(ps...)
return _d
}
// Exec executes the deletion query.
func (_d *BatchImageEventDeleteOne) Exec(ctx context.Context) error {
n, err := _d._d.Exec(ctx)
switch {
case err != nil:
return err
case n == 0:
return &NotFoundError{batchimageevent.Label}
default:
return nil
}
}
// ExecX is like Exec, but panics if an error occurs.
func (_d *BatchImageEventDeleteOne) ExecX(ctx context.Context) {
if err := _d.Exec(ctx); err != nil {
panic(err)
}
}
+564
View File
@@ -0,0 +1,564 @@
// Code generated by ent, DO NOT EDIT.
package ent
import (
"context"
"fmt"
"math"
"entgo.io/ent"
"entgo.io/ent/dialect"
"entgo.io/ent/dialect/sql"
"entgo.io/ent/dialect/sql/sqlgraph"
"entgo.io/ent/schema/field"
"github.com/Wei-Shaw/sub2api/ent/batchimageevent"
"github.com/Wei-Shaw/sub2api/ent/predicate"
)
// BatchImageEventQuery is the builder for querying BatchImageEvent entities.
type BatchImageEventQuery struct {
config
ctx *QueryContext
order []batchimageevent.OrderOption
inters []Interceptor
predicates []predicate.BatchImageEvent
modifiers []func(*sql.Selector)
// intermediate query (i.e. traversal path).
sql *sql.Selector
path func(context.Context) (*sql.Selector, error)
}
// Where adds a new predicate for the BatchImageEventQuery builder.
func (_q *BatchImageEventQuery) Where(ps ...predicate.BatchImageEvent) *BatchImageEventQuery {
_q.predicates = append(_q.predicates, ps...)
return _q
}
// Limit the number of records to be returned by this query.
func (_q *BatchImageEventQuery) Limit(limit int) *BatchImageEventQuery {
_q.ctx.Limit = &limit
return _q
}
// Offset to start from.
func (_q *BatchImageEventQuery) Offset(offset int) *BatchImageEventQuery {
_q.ctx.Offset = &offset
return _q
}
// Unique configures the query builder to filter duplicate records on query.
// By default, unique is set to true, and can be disabled using this method.
func (_q *BatchImageEventQuery) Unique(unique bool) *BatchImageEventQuery {
_q.ctx.Unique = &unique
return _q
}
// Order specifies how the records should be ordered.
func (_q *BatchImageEventQuery) Order(o ...batchimageevent.OrderOption) *BatchImageEventQuery {
_q.order = append(_q.order, o...)
return _q
}
// First returns the first BatchImageEvent entity from the query.
// Returns a *NotFoundError when no BatchImageEvent was found.
func (_q *BatchImageEventQuery) First(ctx context.Context) (*BatchImageEvent, error) {
nodes, err := _q.Limit(1).All(setContextOp(ctx, _q.ctx, ent.OpQueryFirst))
if err != nil {
return nil, err
}
if len(nodes) == 0 {
return nil, &NotFoundError{batchimageevent.Label}
}
return nodes[0], nil
}
// FirstX is like First, but panics if an error occurs.
func (_q *BatchImageEventQuery) FirstX(ctx context.Context) *BatchImageEvent {
node, err := _q.First(ctx)
if err != nil && !IsNotFound(err) {
panic(err)
}
return node
}
// FirstID returns the first BatchImageEvent ID from the query.
// Returns a *NotFoundError when no BatchImageEvent ID was found.
func (_q *BatchImageEventQuery) FirstID(ctx context.Context) (id int64, err error) {
var ids []int64
if ids, err = _q.Limit(1).IDs(setContextOp(ctx, _q.ctx, ent.OpQueryFirstID)); err != nil {
return
}
if len(ids) == 0 {
err = &NotFoundError{batchimageevent.Label}
return
}
return ids[0], nil
}
// FirstIDX is like FirstID, but panics if an error occurs.
func (_q *BatchImageEventQuery) FirstIDX(ctx context.Context) int64 {
id, err := _q.FirstID(ctx)
if err != nil && !IsNotFound(err) {
panic(err)
}
return id
}
// Only returns a single BatchImageEvent entity found by the query, ensuring it only returns one.
// Returns a *NotSingularError when more than one BatchImageEvent entity is found.
// Returns a *NotFoundError when no BatchImageEvent entities are found.
func (_q *BatchImageEventQuery) Only(ctx context.Context) (*BatchImageEvent, error) {
nodes, err := _q.Limit(2).All(setContextOp(ctx, _q.ctx, ent.OpQueryOnly))
if err != nil {
return nil, err
}
switch len(nodes) {
case 1:
return nodes[0], nil
case 0:
return nil, &NotFoundError{batchimageevent.Label}
default:
return nil, &NotSingularError{batchimageevent.Label}
}
}
// OnlyX is like Only, but panics if an error occurs.
func (_q *BatchImageEventQuery) OnlyX(ctx context.Context) *BatchImageEvent {
node, err := _q.Only(ctx)
if err != nil {
panic(err)
}
return node
}
// OnlyID is like Only, but returns the only BatchImageEvent ID in the query.
// Returns a *NotSingularError when more than one BatchImageEvent ID is found.
// Returns a *NotFoundError when no entities are found.
func (_q *BatchImageEventQuery) OnlyID(ctx context.Context) (id int64, err error) {
var ids []int64
if ids, err = _q.Limit(2).IDs(setContextOp(ctx, _q.ctx, ent.OpQueryOnlyID)); err != nil {
return
}
switch len(ids) {
case 1:
id = ids[0]
case 0:
err = &NotFoundError{batchimageevent.Label}
default:
err = &NotSingularError{batchimageevent.Label}
}
return
}
// OnlyIDX is like OnlyID, but panics if an error occurs.
func (_q *BatchImageEventQuery) OnlyIDX(ctx context.Context) int64 {
id, err := _q.OnlyID(ctx)
if err != nil {
panic(err)
}
return id
}
// All executes the query and returns a list of BatchImageEvents.
func (_q *BatchImageEventQuery) All(ctx context.Context) ([]*BatchImageEvent, error) {
ctx = setContextOp(ctx, _q.ctx, ent.OpQueryAll)
if err := _q.prepareQuery(ctx); err != nil {
return nil, err
}
qr := querierAll[[]*BatchImageEvent, *BatchImageEventQuery]()
return withInterceptors[[]*BatchImageEvent](ctx, _q, qr, _q.inters)
}
// AllX is like All, but panics if an error occurs.
func (_q *BatchImageEventQuery) AllX(ctx context.Context) []*BatchImageEvent {
nodes, err := _q.All(ctx)
if err != nil {
panic(err)
}
return nodes
}
// IDs executes the query and returns a list of BatchImageEvent IDs.
func (_q *BatchImageEventQuery) IDs(ctx context.Context) (ids []int64, err error) {
if _q.ctx.Unique == nil && _q.path != nil {
_q.Unique(true)
}
ctx = setContextOp(ctx, _q.ctx, ent.OpQueryIDs)
if err = _q.Select(batchimageevent.FieldID).Scan(ctx, &ids); err != nil {
return nil, err
}
return ids, nil
}
// IDsX is like IDs, but panics if an error occurs.
func (_q *BatchImageEventQuery) IDsX(ctx context.Context) []int64 {
ids, err := _q.IDs(ctx)
if err != nil {
panic(err)
}
return ids
}
// Count returns the count of the given query.
func (_q *BatchImageEventQuery) Count(ctx context.Context) (int, error) {
ctx = setContextOp(ctx, _q.ctx, ent.OpQueryCount)
if err := _q.prepareQuery(ctx); err != nil {
return 0, err
}
return withInterceptors[int](ctx, _q, querierCount[*BatchImageEventQuery](), _q.inters)
}
// CountX is like Count, but panics if an error occurs.
func (_q *BatchImageEventQuery) CountX(ctx context.Context) int {
count, err := _q.Count(ctx)
if err != nil {
panic(err)
}
return count
}
// Exist returns true if the query has elements in the graph.
func (_q *BatchImageEventQuery) Exist(ctx context.Context) (bool, error) {
ctx = setContextOp(ctx, _q.ctx, ent.OpQueryExist)
switch _, err := _q.FirstID(ctx); {
case IsNotFound(err):
return false, nil
case err != nil:
return false, fmt.Errorf("ent: check existence: %w", err)
default:
return true, nil
}
}
// ExistX is like Exist, but panics if an error occurs.
func (_q *BatchImageEventQuery) ExistX(ctx context.Context) bool {
exist, err := _q.Exist(ctx)
if err != nil {
panic(err)
}
return exist
}
// Clone returns a duplicate of the BatchImageEventQuery builder, including all associated steps. It can be
// used to prepare common query builders and use them differently after the clone is made.
func (_q *BatchImageEventQuery) Clone() *BatchImageEventQuery {
if _q == nil {
return nil
}
return &BatchImageEventQuery{
config: _q.config,
ctx: _q.ctx.Clone(),
order: append([]batchimageevent.OrderOption{}, _q.order...),
inters: append([]Interceptor{}, _q.inters...),
predicates: append([]predicate.BatchImageEvent{}, _q.predicates...),
// clone intermediate query.
sql: _q.sql.Clone(),
path: _q.path,
}
}
// GroupBy is used to group vertices by one or more fields/columns.
// It is often used with aggregate functions, like: count, max, mean, min, sum.
//
// Example:
//
// var v []struct {
// JobID string `json:"job_id,omitempty"`
// Count int `json:"count,omitempty"`
// }
//
// client.BatchImageEvent.Query().
// GroupBy(batchimageevent.FieldJobID).
// Aggregate(ent.Count()).
// Scan(ctx, &v)
func (_q *BatchImageEventQuery) GroupBy(field string, fields ...string) *BatchImageEventGroupBy {
_q.ctx.Fields = append([]string{field}, fields...)
grbuild := &BatchImageEventGroupBy{build: _q}
grbuild.flds = &_q.ctx.Fields
grbuild.label = batchimageevent.Label
grbuild.scan = grbuild.Scan
return grbuild
}
// Select allows the selection one or more fields/columns for the given query,
// instead of selecting all fields in the entity.
//
// Example:
//
// var v []struct {
// JobID string `json:"job_id,omitempty"`
// }
//
// client.BatchImageEvent.Query().
// Select(batchimageevent.FieldJobID).
// Scan(ctx, &v)
func (_q *BatchImageEventQuery) Select(fields ...string) *BatchImageEventSelect {
_q.ctx.Fields = append(_q.ctx.Fields, fields...)
sbuild := &BatchImageEventSelect{BatchImageEventQuery: _q}
sbuild.label = batchimageevent.Label
sbuild.flds, sbuild.scan = &_q.ctx.Fields, sbuild.Scan
return sbuild
}
// Aggregate returns a BatchImageEventSelect configured with the given aggregations.
func (_q *BatchImageEventQuery) Aggregate(fns ...AggregateFunc) *BatchImageEventSelect {
return _q.Select().Aggregate(fns...)
}
func (_q *BatchImageEventQuery) prepareQuery(ctx context.Context) error {
for _, inter := range _q.inters {
if inter == nil {
return fmt.Errorf("ent: uninitialized interceptor (forgotten import ent/runtime?)")
}
if trv, ok := inter.(Traverser); ok {
if err := trv.Traverse(ctx, _q); err != nil {
return err
}
}
}
for _, f := range _q.ctx.Fields {
if !batchimageevent.ValidColumn(f) {
return &ValidationError{Name: f, err: fmt.Errorf("ent: invalid field %q for query", f)}
}
}
if _q.path != nil {
prev, err := _q.path(ctx)
if err != nil {
return err
}
_q.sql = prev
}
return nil
}
func (_q *BatchImageEventQuery) sqlAll(ctx context.Context, hooks ...queryHook) ([]*BatchImageEvent, error) {
var (
nodes = []*BatchImageEvent{}
_spec = _q.querySpec()
)
_spec.ScanValues = func(columns []string) ([]any, error) {
return (*BatchImageEvent).scanValues(nil, columns)
}
_spec.Assign = func(columns []string, values []any) error {
node := &BatchImageEvent{config: _q.config}
nodes = append(nodes, node)
return node.assignValues(columns, values)
}
if len(_q.modifiers) > 0 {
_spec.Modifiers = _q.modifiers
}
for i := range hooks {
hooks[i](ctx, _spec)
}
if err := sqlgraph.QueryNodes(ctx, _q.driver, _spec); err != nil {
return nil, err
}
if len(nodes) == 0 {
return nodes, nil
}
return nodes, nil
}
func (_q *BatchImageEventQuery) sqlCount(ctx context.Context) (int, error) {
_spec := _q.querySpec()
if len(_q.modifiers) > 0 {
_spec.Modifiers = _q.modifiers
}
_spec.Node.Columns = _q.ctx.Fields
if len(_q.ctx.Fields) > 0 {
_spec.Unique = _q.ctx.Unique != nil && *_q.ctx.Unique
}
return sqlgraph.CountNodes(ctx, _q.driver, _spec)
}
func (_q *BatchImageEventQuery) querySpec() *sqlgraph.QuerySpec {
_spec := sqlgraph.NewQuerySpec(batchimageevent.Table, batchimageevent.Columns, sqlgraph.NewFieldSpec(batchimageevent.FieldID, field.TypeInt64))
_spec.From = _q.sql
if unique := _q.ctx.Unique; unique != nil {
_spec.Unique = *unique
} else if _q.path != nil {
_spec.Unique = true
}
if fields := _q.ctx.Fields; len(fields) > 0 {
_spec.Node.Columns = make([]string, 0, len(fields))
_spec.Node.Columns = append(_spec.Node.Columns, batchimageevent.FieldID)
for i := range fields {
if fields[i] != batchimageevent.FieldID {
_spec.Node.Columns = append(_spec.Node.Columns, fields[i])
}
}
}
if ps := _q.predicates; len(ps) > 0 {
_spec.Predicate = func(selector *sql.Selector) {
for i := range ps {
ps[i](selector)
}
}
}
if limit := _q.ctx.Limit; limit != nil {
_spec.Limit = *limit
}
if offset := _q.ctx.Offset; offset != nil {
_spec.Offset = *offset
}
if ps := _q.order; len(ps) > 0 {
_spec.Order = func(selector *sql.Selector) {
for i := range ps {
ps[i](selector)
}
}
}
return _spec
}
func (_q *BatchImageEventQuery) sqlQuery(ctx context.Context) *sql.Selector {
builder := sql.Dialect(_q.driver.Dialect())
t1 := builder.Table(batchimageevent.Table)
columns := _q.ctx.Fields
if len(columns) == 0 {
columns = batchimageevent.Columns
}
selector := builder.Select(t1.Columns(columns...)...).From(t1)
if _q.sql != nil {
selector = _q.sql
selector.Select(selector.Columns(columns...)...)
}
if _q.ctx.Unique != nil && *_q.ctx.Unique {
selector.Distinct()
}
for _, m := range _q.modifiers {
m(selector)
}
for _, p := range _q.predicates {
p(selector)
}
for _, p := range _q.order {
p(selector)
}
if offset := _q.ctx.Offset; offset != nil {
// limit is mandatory for offset clause. We start
// with default value, and override it below if needed.
selector.Offset(*offset).Limit(math.MaxInt32)
}
if limit := _q.ctx.Limit; limit != nil {
selector.Limit(*limit)
}
return selector
}
// ForUpdate locks the selected rows against concurrent updates, and prevent them from being
// updated, deleted or "selected ... for update" by other sessions, until the transaction is
// either committed or rolled-back.
func (_q *BatchImageEventQuery) ForUpdate(opts ...sql.LockOption) *BatchImageEventQuery {
if _q.driver.Dialect() == dialect.Postgres {
_q.Unique(false)
}
_q.modifiers = append(_q.modifiers, func(s *sql.Selector) {
s.ForUpdate(opts...)
})
return _q
}
// ForShare behaves similarly to ForUpdate, except that it acquires a shared mode lock
// on any rows that are read. Other sessions can read the rows, but cannot modify them
// until your transaction commits.
func (_q *BatchImageEventQuery) ForShare(opts ...sql.LockOption) *BatchImageEventQuery {
if _q.driver.Dialect() == dialect.Postgres {
_q.Unique(false)
}
_q.modifiers = append(_q.modifiers, func(s *sql.Selector) {
s.ForShare(opts...)
})
return _q
}
// BatchImageEventGroupBy is the group-by builder for BatchImageEvent entities.
type BatchImageEventGroupBy struct {
selector
build *BatchImageEventQuery
}
// Aggregate adds the given aggregation functions to the group-by query.
func (_g *BatchImageEventGroupBy) Aggregate(fns ...AggregateFunc) *BatchImageEventGroupBy {
_g.fns = append(_g.fns, fns...)
return _g
}
// Scan applies the selector query and scans the result into the given value.
func (_g *BatchImageEventGroupBy) Scan(ctx context.Context, v any) error {
ctx = setContextOp(ctx, _g.build.ctx, ent.OpQueryGroupBy)
if err := _g.build.prepareQuery(ctx); err != nil {
return err
}
return scanWithInterceptors[*BatchImageEventQuery, *BatchImageEventGroupBy](ctx, _g.build, _g, _g.build.inters, v)
}
func (_g *BatchImageEventGroupBy) sqlScan(ctx context.Context, root *BatchImageEventQuery, v any) error {
selector := root.sqlQuery(ctx).Select()
aggregation := make([]string, 0, len(_g.fns))
for _, fn := range _g.fns {
aggregation = append(aggregation, fn(selector))
}
if len(selector.SelectedColumns()) == 0 {
columns := make([]string, 0, len(*_g.flds)+len(_g.fns))
for _, f := range *_g.flds {
columns = append(columns, selector.C(f))
}
columns = append(columns, aggregation...)
selector.Select(columns...)
}
selector.GroupBy(selector.Columns(*_g.flds...)...)
if err := selector.Err(); err != nil {
return err
}
rows := &sql.Rows{}
query, args := selector.Query()
if err := _g.build.driver.Query(ctx, query, args, rows); err != nil {
return err
}
defer rows.Close()
return sql.ScanSlice(rows, v)
}
// BatchImageEventSelect is the builder for selecting fields of BatchImageEvent entities.
type BatchImageEventSelect struct {
*BatchImageEventQuery
selector
}
// Aggregate adds the given aggregation functions to the selector query.
func (_s *BatchImageEventSelect) Aggregate(fns ...AggregateFunc) *BatchImageEventSelect {
_s.fns = append(_s.fns, fns...)
return _s
}
// Scan applies the selector query and scans the result into the given value.
func (_s *BatchImageEventSelect) Scan(ctx context.Context, v any) error {
ctx = setContextOp(ctx, _s.ctx, ent.OpQuerySelect)
if err := _s.prepareQuery(ctx); err != nil {
return err
}
return scanWithInterceptors[*BatchImageEventQuery, *BatchImageEventSelect](ctx, _s.BatchImageEventQuery, _s, _s.inters, v)
}
func (_s *BatchImageEventSelect) sqlScan(ctx context.Context, root *BatchImageEventQuery, v any) error {
selector := root.sqlQuery(ctx)
aggregation := make([]string, 0, len(_s.fns))
for _, fn := range _s.fns {
aggregation = append(aggregation, fn(selector))
}
switch n := len(*_s.selector.flds); {
case n == 0 && len(aggregation) > 0:
selector.Select(aggregation...)
case n != 0 && len(aggregation) > 0:
selector.AppendSelect(aggregation...)
}
rows := &sql.Rows{}
query, args := selector.Query()
if err := _s.driver.Query(ctx, query, args, rows); err != nil {
return err
}
defer rows.Close()
return sql.ScanSlice(rows, v)
}
+377
View File
@@ -0,0 +1,377 @@
// Code generated by ent, DO NOT EDIT.
package ent
import (
"context"
"errors"
"fmt"
"entgo.io/ent/dialect/sql"
"entgo.io/ent/dialect/sql/sqlgraph"
"entgo.io/ent/schema/field"
"github.com/Wei-Shaw/sub2api/ent/batchimageevent"
"github.com/Wei-Shaw/sub2api/ent/predicate"
)
// BatchImageEventUpdate is the builder for updating BatchImageEvent entities.
type BatchImageEventUpdate struct {
config
hooks []Hook
mutation *BatchImageEventMutation
}
// Where appends a list predicates to the BatchImageEventUpdate builder.
func (_u *BatchImageEventUpdate) Where(ps ...predicate.BatchImageEvent) *BatchImageEventUpdate {
_u.mutation.Where(ps...)
return _u
}
// SetJobID sets the "job_id" field.
func (_u *BatchImageEventUpdate) SetJobID(v string) *BatchImageEventUpdate {
_u.mutation.SetJobID(v)
return _u
}
// SetNillableJobID sets the "job_id" field if the given value is not nil.
func (_u *BatchImageEventUpdate) SetNillableJobID(v *string) *BatchImageEventUpdate {
if v != nil {
_u.SetJobID(*v)
}
return _u
}
// SetEventType sets the "event_type" field.
func (_u *BatchImageEventUpdate) SetEventType(v string) *BatchImageEventUpdate {
_u.mutation.SetEventType(v)
return _u
}
// SetNillableEventType sets the "event_type" field if the given value is not nil.
func (_u *BatchImageEventUpdate) SetNillableEventType(v *string) *BatchImageEventUpdate {
if v != nil {
_u.SetEventType(*v)
}
return _u
}
// SetPayload sets the "payload" field.
func (_u *BatchImageEventUpdate) SetPayload(v map[string]interface{}) *BatchImageEventUpdate {
_u.mutation.SetPayload(v)
return _u
}
// ClearPayload clears the value of the "payload" field.
func (_u *BatchImageEventUpdate) ClearPayload() *BatchImageEventUpdate {
_u.mutation.ClearPayload()
return _u
}
// SetEventHash sets the "event_hash" field.
func (_u *BatchImageEventUpdate) SetEventHash(v string) *BatchImageEventUpdate {
_u.mutation.SetEventHash(v)
return _u
}
// SetNillableEventHash sets the "event_hash" field if the given value is not nil.
func (_u *BatchImageEventUpdate) SetNillableEventHash(v *string) *BatchImageEventUpdate {
if v != nil {
_u.SetEventHash(*v)
}
return _u
}
// ClearEventHash clears the value of the "event_hash" field.
func (_u *BatchImageEventUpdate) ClearEventHash() *BatchImageEventUpdate {
_u.mutation.ClearEventHash()
return _u
}
// Mutation returns the BatchImageEventMutation object of the builder.
func (_u *BatchImageEventUpdate) Mutation() *BatchImageEventMutation {
return _u.mutation
}
// Save executes the query and returns the number of nodes affected by the update operation.
func (_u *BatchImageEventUpdate) Save(ctx context.Context) (int, error) {
return withHooks(ctx, _u.sqlSave, _u.mutation, _u.hooks)
}
// SaveX is like Save, but panics if an error occurs.
func (_u *BatchImageEventUpdate) SaveX(ctx context.Context) int {
affected, err := _u.Save(ctx)
if err != nil {
panic(err)
}
return affected
}
// Exec executes the query.
func (_u *BatchImageEventUpdate) Exec(ctx context.Context) error {
_, err := _u.Save(ctx)
return err
}
// ExecX is like Exec, but panics if an error occurs.
func (_u *BatchImageEventUpdate) ExecX(ctx context.Context) {
if err := _u.Exec(ctx); err != nil {
panic(err)
}
}
// check runs all checks and user-defined validators on the builder.
func (_u *BatchImageEventUpdate) check() error {
if v, ok := _u.mutation.JobID(); ok {
if err := batchimageevent.JobIDValidator(v); err != nil {
return &ValidationError{Name: "job_id", err: fmt.Errorf(`ent: validator failed for field "BatchImageEvent.job_id": %w`, err)}
}
}
if v, ok := _u.mutation.EventType(); ok {
if err := batchimageevent.EventTypeValidator(v); err != nil {
return &ValidationError{Name: "event_type", err: fmt.Errorf(`ent: validator failed for field "BatchImageEvent.event_type": %w`, err)}
}
}
if v, ok := _u.mutation.EventHash(); ok {
if err := batchimageevent.EventHashValidator(v); err != nil {
return &ValidationError{Name: "event_hash", err: fmt.Errorf(`ent: validator failed for field "BatchImageEvent.event_hash": %w`, err)}
}
}
return nil
}
func (_u *BatchImageEventUpdate) sqlSave(ctx context.Context) (_node int, err error) {
if err := _u.check(); err != nil {
return _node, err
}
_spec := sqlgraph.NewUpdateSpec(batchimageevent.Table, batchimageevent.Columns, sqlgraph.NewFieldSpec(batchimageevent.FieldID, field.TypeInt64))
if ps := _u.mutation.predicates; len(ps) > 0 {
_spec.Predicate = func(selector *sql.Selector) {
for i := range ps {
ps[i](selector)
}
}
}
if value, ok := _u.mutation.JobID(); ok {
_spec.SetField(batchimageevent.FieldJobID, field.TypeString, value)
}
if value, ok := _u.mutation.EventType(); ok {
_spec.SetField(batchimageevent.FieldEventType, field.TypeString, value)
}
if value, ok := _u.mutation.Payload(); ok {
_spec.SetField(batchimageevent.FieldPayload, field.TypeJSON, value)
}
if _u.mutation.PayloadCleared() {
_spec.ClearField(batchimageevent.FieldPayload, field.TypeJSON)
}
if value, ok := _u.mutation.EventHash(); ok {
_spec.SetField(batchimageevent.FieldEventHash, field.TypeString, value)
}
if _u.mutation.EventHashCleared() {
_spec.ClearField(batchimageevent.FieldEventHash, field.TypeString)
}
if _node, err = sqlgraph.UpdateNodes(ctx, _u.driver, _spec); err != nil {
if _, ok := err.(*sqlgraph.NotFoundError); ok {
err = &NotFoundError{batchimageevent.Label}
} else if sqlgraph.IsConstraintError(err) {
err = &ConstraintError{msg: err.Error(), wrap: err}
}
return 0, err
}
_u.mutation.done = true
return _node, nil
}
// BatchImageEventUpdateOne is the builder for updating a single BatchImageEvent entity.
type BatchImageEventUpdateOne struct {
config
fields []string
hooks []Hook
mutation *BatchImageEventMutation
}
// SetJobID sets the "job_id" field.
func (_u *BatchImageEventUpdateOne) SetJobID(v string) *BatchImageEventUpdateOne {
_u.mutation.SetJobID(v)
return _u
}
// SetNillableJobID sets the "job_id" field if the given value is not nil.
func (_u *BatchImageEventUpdateOne) SetNillableJobID(v *string) *BatchImageEventUpdateOne {
if v != nil {
_u.SetJobID(*v)
}
return _u
}
// SetEventType sets the "event_type" field.
func (_u *BatchImageEventUpdateOne) SetEventType(v string) *BatchImageEventUpdateOne {
_u.mutation.SetEventType(v)
return _u
}
// SetNillableEventType sets the "event_type" field if the given value is not nil.
func (_u *BatchImageEventUpdateOne) SetNillableEventType(v *string) *BatchImageEventUpdateOne {
if v != nil {
_u.SetEventType(*v)
}
return _u
}
// SetPayload sets the "payload" field.
func (_u *BatchImageEventUpdateOne) SetPayload(v map[string]interface{}) *BatchImageEventUpdateOne {
_u.mutation.SetPayload(v)
return _u
}
// ClearPayload clears the value of the "payload" field.
func (_u *BatchImageEventUpdateOne) ClearPayload() *BatchImageEventUpdateOne {
_u.mutation.ClearPayload()
return _u
}
// SetEventHash sets the "event_hash" field.
func (_u *BatchImageEventUpdateOne) SetEventHash(v string) *BatchImageEventUpdateOne {
_u.mutation.SetEventHash(v)
return _u
}
// SetNillableEventHash sets the "event_hash" field if the given value is not nil.
func (_u *BatchImageEventUpdateOne) SetNillableEventHash(v *string) *BatchImageEventUpdateOne {
if v != nil {
_u.SetEventHash(*v)
}
return _u
}
// ClearEventHash clears the value of the "event_hash" field.
func (_u *BatchImageEventUpdateOne) ClearEventHash() *BatchImageEventUpdateOne {
_u.mutation.ClearEventHash()
return _u
}
// Mutation returns the BatchImageEventMutation object of the builder.
func (_u *BatchImageEventUpdateOne) Mutation() *BatchImageEventMutation {
return _u.mutation
}
// Where appends a list predicates to the BatchImageEventUpdate builder.
func (_u *BatchImageEventUpdateOne) Where(ps ...predicate.BatchImageEvent) *BatchImageEventUpdateOne {
_u.mutation.Where(ps...)
return _u
}
// Select allows selecting one or more fields (columns) of the returned entity.
// The default is selecting all fields defined in the entity schema.
func (_u *BatchImageEventUpdateOne) Select(field string, fields ...string) *BatchImageEventUpdateOne {
_u.fields = append([]string{field}, fields...)
return _u
}
// Save executes the query and returns the updated BatchImageEvent entity.
func (_u *BatchImageEventUpdateOne) Save(ctx context.Context) (*BatchImageEvent, error) {
return withHooks(ctx, _u.sqlSave, _u.mutation, _u.hooks)
}
// SaveX is like Save, but panics if an error occurs.
func (_u *BatchImageEventUpdateOne) SaveX(ctx context.Context) *BatchImageEvent {
node, err := _u.Save(ctx)
if err != nil {
panic(err)
}
return node
}
// Exec executes the query on the entity.
func (_u *BatchImageEventUpdateOne) Exec(ctx context.Context) error {
_, err := _u.Save(ctx)
return err
}
// ExecX is like Exec, but panics if an error occurs.
func (_u *BatchImageEventUpdateOne) ExecX(ctx context.Context) {
if err := _u.Exec(ctx); err != nil {
panic(err)
}
}
// check runs all checks and user-defined validators on the builder.
func (_u *BatchImageEventUpdateOne) check() error {
if v, ok := _u.mutation.JobID(); ok {
if err := batchimageevent.JobIDValidator(v); err != nil {
return &ValidationError{Name: "job_id", err: fmt.Errorf(`ent: validator failed for field "BatchImageEvent.job_id": %w`, err)}
}
}
if v, ok := _u.mutation.EventType(); ok {
if err := batchimageevent.EventTypeValidator(v); err != nil {
return &ValidationError{Name: "event_type", err: fmt.Errorf(`ent: validator failed for field "BatchImageEvent.event_type": %w`, err)}
}
}
if v, ok := _u.mutation.EventHash(); ok {
if err := batchimageevent.EventHashValidator(v); err != nil {
return &ValidationError{Name: "event_hash", err: fmt.Errorf(`ent: validator failed for field "BatchImageEvent.event_hash": %w`, err)}
}
}
return nil
}
func (_u *BatchImageEventUpdateOne) sqlSave(ctx context.Context) (_node *BatchImageEvent, err error) {
if err := _u.check(); err != nil {
return _node, err
}
_spec := sqlgraph.NewUpdateSpec(batchimageevent.Table, batchimageevent.Columns, sqlgraph.NewFieldSpec(batchimageevent.FieldID, field.TypeInt64))
id, ok := _u.mutation.ID()
if !ok {
return nil, &ValidationError{Name: "id", err: errors.New(`ent: missing "BatchImageEvent.id" for update`)}
}
_spec.Node.ID.Value = id
if fields := _u.fields; len(fields) > 0 {
_spec.Node.Columns = make([]string, 0, len(fields))
_spec.Node.Columns = append(_spec.Node.Columns, batchimageevent.FieldID)
for _, f := range fields {
if !batchimageevent.ValidColumn(f) {
return nil, &ValidationError{Name: f, err: fmt.Errorf("ent: invalid field %q for query", f)}
}
if f != batchimageevent.FieldID {
_spec.Node.Columns = append(_spec.Node.Columns, f)
}
}
}
if ps := _u.mutation.predicates; len(ps) > 0 {
_spec.Predicate = func(selector *sql.Selector) {
for i := range ps {
ps[i](selector)
}
}
}
if value, ok := _u.mutation.JobID(); ok {
_spec.SetField(batchimageevent.FieldJobID, field.TypeString, value)
}
if value, ok := _u.mutation.EventType(); ok {
_spec.SetField(batchimageevent.FieldEventType, field.TypeString, value)
}
if value, ok := _u.mutation.Payload(); ok {
_spec.SetField(batchimageevent.FieldPayload, field.TypeJSON, value)
}
if _u.mutation.PayloadCleared() {
_spec.ClearField(batchimageevent.FieldPayload, field.TypeJSON)
}
if value, ok := _u.mutation.EventHash(); ok {
_spec.SetField(batchimageevent.FieldEventHash, field.TypeString, value)
}
if _u.mutation.EventHashCleared() {
_spec.ClearField(batchimageevent.FieldEventHash, field.TypeString)
}
_node = &BatchImageEvent{config: _u.config}
_spec.Assign = _node.assignValues
_spec.ScanValues = _node.scanValues
if err = sqlgraph.UpdateNode(ctx, _u.driver, _spec); err != nil {
if _, ok := err.(*sqlgraph.NotFoundError); ok {
err = &NotFoundError{batchimageevent.Label}
} else if sqlgraph.IsConstraintError(err) {
err = &ConstraintError{msg: err.Error(), wrap: err}
}
return nil, err
}
_u.mutation.done = true
return _node, nil
}
+320
View File
@@ -0,0 +1,320 @@
// Code generated by ent, DO NOT EDIT.
package ent
import (
"fmt"
"strings"
"time"
"entgo.io/ent"
"entgo.io/ent/dialect/sql"
"github.com/Wei-Shaw/sub2api/ent/batchimageitem"
)
// BatchImageItem is the model entity for the BatchImageItem schema.
type BatchImageItem struct {
config `json:"-"`
// ID of the ent.
ID int64 `json:"id,omitempty"`
// JobID holds the value of the "job_id" field.
JobID string `json:"job_id,omitempty"`
// CustomID holds the value of the "custom_id" field.
CustomID string `json:"custom_id,omitempty"`
// Status holds the value of the "status" field.
Status string `json:"status,omitempty"`
// RequestHash holds the value of the "request_hash" field.
RequestHash *string `json:"request_hash,omitempty"`
// PromptPreview holds the value of the "prompt_preview" field.
PromptPreview *string `json:"prompt_preview,omitempty"`
// ProviderSourceObject holds the value of the "provider_source_object" field.
ProviderSourceObject *string `json:"provider_source_object,omitempty"`
// SourceLineNumber holds the value of the "source_line_number" field.
SourceLineNumber *int `json:"source_line_number,omitempty"`
// SourceByteOffset holds the value of the "source_byte_offset" field.
SourceByteOffset *int64 `json:"source_byte_offset,omitempty"`
// SourceByteLength holds the value of the "source_byte_length" field.
SourceByteLength *int64 `json:"source_byte_length,omitempty"`
// MimeType holds the value of the "mime_type" field.
MimeType *string `json:"mime_type,omitempty"`
// FileExtension holds the value of the "file_extension" field.
FileExtension *string `json:"file_extension,omitempty"`
// ImageCount holds the value of the "image_count" field.
ImageCount int `json:"image_count,omitempty"`
// ErrorCode holds the value of the "error_code" field.
ErrorCode *string `json:"error_code,omitempty"`
// ErrorMessage holds the value of the "error_message" field.
ErrorMessage *string `json:"error_message,omitempty"`
// BilledAmount holds the value of the "billed_amount" field.
BilledAmount *float64 `json:"billed_amount,omitempty"`
// CreatedAt holds the value of the "created_at" field.
CreatedAt time.Time `json:"created_at,omitempty"`
// IndexedAt holds the value of the "indexed_at" field.
IndexedAt *time.Time `json:"indexed_at,omitempty"`
selectValues sql.SelectValues
}
// scanValues returns the types for scanning values from sql.Rows.
func (*BatchImageItem) scanValues(columns []string) ([]any, error) {
values := make([]any, len(columns))
for i := range columns {
switch columns[i] {
case batchimageitem.FieldBilledAmount:
values[i] = new(sql.NullFloat64)
case batchimageitem.FieldID, batchimageitem.FieldSourceLineNumber, batchimageitem.FieldSourceByteOffset, batchimageitem.FieldSourceByteLength, batchimageitem.FieldImageCount:
values[i] = new(sql.NullInt64)
case batchimageitem.FieldJobID, batchimageitem.FieldCustomID, batchimageitem.FieldStatus, batchimageitem.FieldRequestHash, batchimageitem.FieldPromptPreview, batchimageitem.FieldProviderSourceObject, batchimageitem.FieldMimeType, batchimageitem.FieldFileExtension, batchimageitem.FieldErrorCode, batchimageitem.FieldErrorMessage:
values[i] = new(sql.NullString)
case batchimageitem.FieldCreatedAt, batchimageitem.FieldIndexedAt:
values[i] = new(sql.NullTime)
default:
values[i] = new(sql.UnknownType)
}
}
return values, nil
}
// assignValues assigns the values that were returned from sql.Rows (after scanning)
// to the BatchImageItem fields.
func (_m *BatchImageItem) assignValues(columns []string, values []any) error {
if m, n := len(values), len(columns); m < n {
return fmt.Errorf("mismatch number of scan values: %d != %d", m, n)
}
for i := range columns {
switch columns[i] {
case batchimageitem.FieldID:
value, ok := values[i].(*sql.NullInt64)
if !ok {
return fmt.Errorf("unexpected type %T for field id", value)
}
_m.ID = int64(value.Int64)
case batchimageitem.FieldJobID:
if value, ok := values[i].(*sql.NullString); !ok {
return fmt.Errorf("unexpected type %T for field job_id", values[i])
} else if value.Valid {
_m.JobID = value.String
}
case batchimageitem.FieldCustomID:
if value, ok := values[i].(*sql.NullString); !ok {
return fmt.Errorf("unexpected type %T for field custom_id", values[i])
} else if value.Valid {
_m.CustomID = value.String
}
case batchimageitem.FieldStatus:
if value, ok := values[i].(*sql.NullString); !ok {
return fmt.Errorf("unexpected type %T for field status", values[i])
} else if value.Valid {
_m.Status = value.String
}
case batchimageitem.FieldRequestHash:
if value, ok := values[i].(*sql.NullString); !ok {
return fmt.Errorf("unexpected type %T for field request_hash", values[i])
} else if value.Valid {
_m.RequestHash = new(string)
*_m.RequestHash = value.String
}
case batchimageitem.FieldPromptPreview:
if value, ok := values[i].(*sql.NullString); !ok {
return fmt.Errorf("unexpected type %T for field prompt_preview", values[i])
} else if value.Valid {
_m.PromptPreview = new(string)
*_m.PromptPreview = value.String
}
case batchimageitem.FieldProviderSourceObject:
if value, ok := values[i].(*sql.NullString); !ok {
return fmt.Errorf("unexpected type %T for field provider_source_object", values[i])
} else if value.Valid {
_m.ProviderSourceObject = new(string)
*_m.ProviderSourceObject = value.String
}
case batchimageitem.FieldSourceLineNumber:
if value, ok := values[i].(*sql.NullInt64); !ok {
return fmt.Errorf("unexpected type %T for field source_line_number", values[i])
} else if value.Valid {
_m.SourceLineNumber = new(int)
*_m.SourceLineNumber = int(value.Int64)
}
case batchimageitem.FieldSourceByteOffset:
if value, ok := values[i].(*sql.NullInt64); !ok {
return fmt.Errorf("unexpected type %T for field source_byte_offset", values[i])
} else if value.Valid {
_m.SourceByteOffset = new(int64)
*_m.SourceByteOffset = value.Int64
}
case batchimageitem.FieldSourceByteLength:
if value, ok := values[i].(*sql.NullInt64); !ok {
return fmt.Errorf("unexpected type %T for field source_byte_length", values[i])
} else if value.Valid {
_m.SourceByteLength = new(int64)
*_m.SourceByteLength = value.Int64
}
case batchimageitem.FieldMimeType:
if value, ok := values[i].(*sql.NullString); !ok {
return fmt.Errorf("unexpected type %T for field mime_type", values[i])
} else if value.Valid {
_m.MimeType = new(string)
*_m.MimeType = value.String
}
case batchimageitem.FieldFileExtension:
if value, ok := values[i].(*sql.NullString); !ok {
return fmt.Errorf("unexpected type %T for field file_extension", values[i])
} else if value.Valid {
_m.FileExtension = new(string)
*_m.FileExtension = value.String
}
case batchimageitem.FieldImageCount:
if value, ok := values[i].(*sql.NullInt64); !ok {
return fmt.Errorf("unexpected type %T for field image_count", values[i])
} else if value.Valid {
_m.ImageCount = int(value.Int64)
}
case batchimageitem.FieldErrorCode:
if value, ok := values[i].(*sql.NullString); !ok {
return fmt.Errorf("unexpected type %T for field error_code", values[i])
} else if value.Valid {
_m.ErrorCode = new(string)
*_m.ErrorCode = value.String
}
case batchimageitem.FieldErrorMessage:
if value, ok := values[i].(*sql.NullString); !ok {
return fmt.Errorf("unexpected type %T for field error_message", values[i])
} else if value.Valid {
_m.ErrorMessage = new(string)
*_m.ErrorMessage = value.String
}
case batchimageitem.FieldBilledAmount:
if value, ok := values[i].(*sql.NullFloat64); !ok {
return fmt.Errorf("unexpected type %T for field billed_amount", values[i])
} else if value.Valid {
_m.BilledAmount = new(float64)
*_m.BilledAmount = value.Float64
}
case batchimageitem.FieldCreatedAt:
if value, ok := values[i].(*sql.NullTime); !ok {
return fmt.Errorf("unexpected type %T for field created_at", values[i])
} else if value.Valid {
_m.CreatedAt = value.Time
}
case batchimageitem.FieldIndexedAt:
if value, ok := values[i].(*sql.NullTime); !ok {
return fmt.Errorf("unexpected type %T for field indexed_at", values[i])
} else if value.Valid {
_m.IndexedAt = new(time.Time)
*_m.IndexedAt = value.Time
}
default:
_m.selectValues.Set(columns[i], values[i])
}
}
return nil
}
// Value returns the ent.Value that was dynamically selected and assigned to the BatchImageItem.
// This includes values selected through modifiers, order, etc.
func (_m *BatchImageItem) Value(name string) (ent.Value, error) {
return _m.selectValues.Get(name)
}
// Update returns a builder for updating this BatchImageItem.
// Note that you need to call BatchImageItem.Unwrap() before calling this method if this BatchImageItem
// was returned from a transaction, and the transaction was committed or rolled back.
func (_m *BatchImageItem) Update() *BatchImageItemUpdateOne {
return NewBatchImageItemClient(_m.config).UpdateOne(_m)
}
// Unwrap unwraps the BatchImageItem entity that was returned from a transaction after it was closed,
// so that all future queries will be executed through the driver which created the transaction.
func (_m *BatchImageItem) Unwrap() *BatchImageItem {
_tx, ok := _m.config.driver.(*txDriver)
if !ok {
panic("ent: BatchImageItem is not a transactional entity")
}
_m.config.driver = _tx.drv
return _m
}
// String implements the fmt.Stringer.
func (_m *BatchImageItem) String() string {
var builder strings.Builder
builder.WriteString("BatchImageItem(")
builder.WriteString(fmt.Sprintf("id=%v, ", _m.ID))
builder.WriteString("job_id=")
builder.WriteString(_m.JobID)
builder.WriteString(", ")
builder.WriteString("custom_id=")
builder.WriteString(_m.CustomID)
builder.WriteString(", ")
builder.WriteString("status=")
builder.WriteString(_m.Status)
builder.WriteString(", ")
if v := _m.RequestHash; v != nil {
builder.WriteString("request_hash=")
builder.WriteString(*v)
}
builder.WriteString(", ")
if v := _m.PromptPreview; v != nil {
builder.WriteString("prompt_preview=")
builder.WriteString(*v)
}
builder.WriteString(", ")
if v := _m.ProviderSourceObject; v != nil {
builder.WriteString("provider_source_object=")
builder.WriteString(*v)
}
builder.WriteString(", ")
if v := _m.SourceLineNumber; v != nil {
builder.WriteString("source_line_number=")
builder.WriteString(fmt.Sprintf("%v", *v))
}
builder.WriteString(", ")
if v := _m.SourceByteOffset; v != nil {
builder.WriteString("source_byte_offset=")
builder.WriteString(fmt.Sprintf("%v", *v))
}
builder.WriteString(", ")
if v := _m.SourceByteLength; v != nil {
builder.WriteString("source_byte_length=")
builder.WriteString(fmt.Sprintf("%v", *v))
}
builder.WriteString(", ")
if v := _m.MimeType; v != nil {
builder.WriteString("mime_type=")
builder.WriteString(*v)
}
builder.WriteString(", ")
if v := _m.FileExtension; v != nil {
builder.WriteString("file_extension=")
builder.WriteString(*v)
}
builder.WriteString(", ")
builder.WriteString("image_count=")
builder.WriteString(fmt.Sprintf("%v", _m.ImageCount))
builder.WriteString(", ")
if v := _m.ErrorCode; v != nil {
builder.WriteString("error_code=")
builder.WriteString(*v)
}
builder.WriteString(", ")
if v := _m.ErrorMessage; v != nil {
builder.WriteString("error_message=")
builder.WriteString(*v)
}
builder.WriteString(", ")
if v := _m.BilledAmount; v != nil {
builder.WriteString("billed_amount=")
builder.WriteString(fmt.Sprintf("%v", *v))
}
builder.WriteString(", ")
builder.WriteString("created_at=")
builder.WriteString(_m.CreatedAt.Format(time.ANSIC))
builder.WriteString(", ")
if v := _m.IndexedAt; v != nil {
builder.WriteString("indexed_at=")
builder.WriteString(v.Format(time.ANSIC))
}
builder.WriteByte(')')
return builder.String()
}
// BatchImageItems is a parsable slice of BatchImageItem.
type BatchImageItems []*BatchImageItem
@@ -0,0 +1,200 @@
// Code generated by ent, DO NOT EDIT.
package batchimageitem
import (
"time"
"entgo.io/ent/dialect/sql"
)
const (
// Label holds the string label denoting the batchimageitem type in the database.
Label = "batch_image_item"
// FieldID holds the string denoting the id field in the database.
FieldID = "id"
// FieldJobID holds the string denoting the job_id field in the database.
FieldJobID = "job_id"
// FieldCustomID holds the string denoting the custom_id field in the database.
FieldCustomID = "custom_id"
// FieldStatus holds the string denoting the status field in the database.
FieldStatus = "status"
// FieldRequestHash holds the string denoting the request_hash field in the database.
FieldRequestHash = "request_hash"
// FieldPromptPreview holds the string denoting the prompt_preview field in the database.
FieldPromptPreview = "prompt_preview"
// FieldProviderSourceObject holds the string denoting the provider_source_object field in the database.
FieldProviderSourceObject = "provider_source_object"
// FieldSourceLineNumber holds the string denoting the source_line_number field in the database.
FieldSourceLineNumber = "source_line_number"
// FieldSourceByteOffset holds the string denoting the source_byte_offset field in the database.
FieldSourceByteOffset = "source_byte_offset"
// FieldSourceByteLength holds the string denoting the source_byte_length field in the database.
FieldSourceByteLength = "source_byte_length"
// FieldMimeType holds the string denoting the mime_type field in the database.
FieldMimeType = "mime_type"
// FieldFileExtension holds the string denoting the file_extension field in the database.
FieldFileExtension = "file_extension"
// FieldImageCount holds the string denoting the image_count field in the database.
FieldImageCount = "image_count"
// FieldErrorCode holds the string denoting the error_code field in the database.
FieldErrorCode = "error_code"
// FieldErrorMessage holds the string denoting the error_message field in the database.
FieldErrorMessage = "error_message"
// FieldBilledAmount holds the string denoting the billed_amount field in the database.
FieldBilledAmount = "billed_amount"
// FieldCreatedAt holds the string denoting the created_at field in the database.
FieldCreatedAt = "created_at"
// FieldIndexedAt holds the string denoting the indexed_at field in the database.
FieldIndexedAt = "indexed_at"
// Table holds the table name of the batchimageitem in the database.
Table = "batch_image_items"
)
// Columns holds all SQL columns for batchimageitem fields.
var Columns = []string{
FieldID,
FieldJobID,
FieldCustomID,
FieldStatus,
FieldRequestHash,
FieldPromptPreview,
FieldProviderSourceObject,
FieldSourceLineNumber,
FieldSourceByteOffset,
FieldSourceByteLength,
FieldMimeType,
FieldFileExtension,
FieldImageCount,
FieldErrorCode,
FieldErrorMessage,
FieldBilledAmount,
FieldCreatedAt,
FieldIndexedAt,
}
// ValidColumn reports if the column name is valid (part of the table columns).
func ValidColumn(column string) bool {
for i := range Columns {
if column == Columns[i] {
return true
}
}
return false
}
var (
// JobIDValidator is a validator for the "job_id" field. It is called by the builders before save.
JobIDValidator func(string) error
// CustomIDValidator is a validator for the "custom_id" field. It is called by the builders before save.
CustomIDValidator func(string) error
// StatusValidator is a validator for the "status" field. It is called by the builders before save.
StatusValidator func(string) error
// RequestHashValidator is a validator for the "request_hash" field. It is called by the builders before save.
RequestHashValidator func(string) error
// ProviderSourceObjectValidator is a validator for the "provider_source_object" field. It is called by the builders before save.
ProviderSourceObjectValidator func(string) error
// MimeTypeValidator is a validator for the "mime_type" field. It is called by the builders before save.
MimeTypeValidator func(string) error
// FileExtensionValidator is a validator for the "file_extension" field. It is called by the builders before save.
FileExtensionValidator func(string) error
// DefaultImageCount holds the default value on creation for the "image_count" field.
DefaultImageCount int
// ErrorCodeValidator is a validator for the "error_code" field. It is called by the builders before save.
ErrorCodeValidator func(string) error
// DefaultCreatedAt holds the default value on creation for the "created_at" field.
DefaultCreatedAt func() time.Time
)
// OrderOption defines the ordering options for the BatchImageItem queries.
type OrderOption func(*sql.Selector)
// ByID orders the results by the id field.
func ByID(opts ...sql.OrderTermOption) OrderOption {
return sql.OrderByField(FieldID, opts...).ToFunc()
}
// ByJobID orders the results by the job_id field.
func ByJobID(opts ...sql.OrderTermOption) OrderOption {
return sql.OrderByField(FieldJobID, opts...).ToFunc()
}
// ByCustomID orders the results by the custom_id field.
func ByCustomID(opts ...sql.OrderTermOption) OrderOption {
return sql.OrderByField(FieldCustomID, opts...).ToFunc()
}
// ByStatus orders the results by the status field.
func ByStatus(opts ...sql.OrderTermOption) OrderOption {
return sql.OrderByField(FieldStatus, opts...).ToFunc()
}
// ByRequestHash orders the results by the request_hash field.
func ByRequestHash(opts ...sql.OrderTermOption) OrderOption {
return sql.OrderByField(FieldRequestHash, opts...).ToFunc()
}
// ByPromptPreview orders the results by the prompt_preview field.
func ByPromptPreview(opts ...sql.OrderTermOption) OrderOption {
return sql.OrderByField(FieldPromptPreview, opts...).ToFunc()
}
// ByProviderSourceObject orders the results by the provider_source_object field.
func ByProviderSourceObject(opts ...sql.OrderTermOption) OrderOption {
return sql.OrderByField(FieldProviderSourceObject, opts...).ToFunc()
}
// BySourceLineNumber orders the results by the source_line_number field.
func BySourceLineNumber(opts ...sql.OrderTermOption) OrderOption {
return sql.OrderByField(FieldSourceLineNumber, opts...).ToFunc()
}
// BySourceByteOffset orders the results by the source_byte_offset field.
func BySourceByteOffset(opts ...sql.OrderTermOption) OrderOption {
return sql.OrderByField(FieldSourceByteOffset, opts...).ToFunc()
}
// BySourceByteLength orders the results by the source_byte_length field.
func BySourceByteLength(opts ...sql.OrderTermOption) OrderOption {
return sql.OrderByField(FieldSourceByteLength, opts...).ToFunc()
}
// ByMimeType orders the results by the mime_type field.
func ByMimeType(opts ...sql.OrderTermOption) OrderOption {
return sql.OrderByField(FieldMimeType, opts...).ToFunc()
}
// ByFileExtension orders the results by the file_extension field.
func ByFileExtension(opts ...sql.OrderTermOption) OrderOption {
return sql.OrderByField(FieldFileExtension, opts...).ToFunc()
}
// ByImageCount orders the results by the image_count field.
func ByImageCount(opts ...sql.OrderTermOption) OrderOption {
return sql.OrderByField(FieldImageCount, opts...).ToFunc()
}
// ByErrorCode orders the results by the error_code field.
func ByErrorCode(opts ...sql.OrderTermOption) OrderOption {
return sql.OrderByField(FieldErrorCode, opts...).ToFunc()
}
// ByErrorMessage orders the results by the error_message field.
func ByErrorMessage(opts ...sql.OrderTermOption) OrderOption {
return sql.OrderByField(FieldErrorMessage, opts...).ToFunc()
}
// ByBilledAmount orders the results by the billed_amount field.
func ByBilledAmount(opts ...sql.OrderTermOption) OrderOption {
return sql.OrderByField(FieldBilledAmount, opts...).ToFunc()
}
// ByCreatedAt orders the results by the created_at field.
func ByCreatedAt(opts ...sql.OrderTermOption) OrderOption {
return sql.OrderByField(FieldCreatedAt, opts...).ToFunc()
}
// ByIndexedAt orders the results by the indexed_at field.
func ByIndexedAt(opts ...sql.OrderTermOption) OrderOption {
return sql.OrderByField(FieldIndexedAt, opts...).ToFunc()
}
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
+88
View File
@@ -0,0 +1,88 @@
// Code generated by ent, DO NOT EDIT.
package ent
import (
"context"
"entgo.io/ent/dialect/sql"
"entgo.io/ent/dialect/sql/sqlgraph"
"entgo.io/ent/schema/field"
"github.com/Wei-Shaw/sub2api/ent/batchimageitem"
"github.com/Wei-Shaw/sub2api/ent/predicate"
)
// BatchImageItemDelete is the builder for deleting a BatchImageItem entity.
type BatchImageItemDelete struct {
config
hooks []Hook
mutation *BatchImageItemMutation
}
// Where appends a list predicates to the BatchImageItemDelete builder.
func (_d *BatchImageItemDelete) Where(ps ...predicate.BatchImageItem) *BatchImageItemDelete {
_d.mutation.Where(ps...)
return _d
}
// Exec executes the deletion query and returns how many vertices were deleted.
func (_d *BatchImageItemDelete) Exec(ctx context.Context) (int, error) {
return withHooks(ctx, _d.sqlExec, _d.mutation, _d.hooks)
}
// ExecX is like Exec, but panics if an error occurs.
func (_d *BatchImageItemDelete) ExecX(ctx context.Context) int {
n, err := _d.Exec(ctx)
if err != nil {
panic(err)
}
return n
}
func (_d *BatchImageItemDelete) sqlExec(ctx context.Context) (int, error) {
_spec := sqlgraph.NewDeleteSpec(batchimageitem.Table, sqlgraph.NewFieldSpec(batchimageitem.FieldID, field.TypeInt64))
if ps := _d.mutation.predicates; len(ps) > 0 {
_spec.Predicate = func(selector *sql.Selector) {
for i := range ps {
ps[i](selector)
}
}
}
affected, err := sqlgraph.DeleteNodes(ctx, _d.driver, _spec)
if err != nil && sqlgraph.IsConstraintError(err) {
err = &ConstraintError{msg: err.Error(), wrap: err}
}
_d.mutation.done = true
return affected, err
}
// BatchImageItemDeleteOne is the builder for deleting a single BatchImageItem entity.
type BatchImageItemDeleteOne struct {
_d *BatchImageItemDelete
}
// Where appends a list predicates to the BatchImageItemDelete builder.
func (_d *BatchImageItemDeleteOne) Where(ps ...predicate.BatchImageItem) *BatchImageItemDeleteOne {
_d._d.mutation.Where(ps...)
return _d
}
// Exec executes the deletion query.
func (_d *BatchImageItemDeleteOne) Exec(ctx context.Context) error {
n, err := _d._d.Exec(ctx)
switch {
case err != nil:
return err
case n == 0:
return &NotFoundError{batchimageitem.Label}
default:
return nil
}
}
// ExecX is like Exec, but panics if an error occurs.
func (_d *BatchImageItemDeleteOne) ExecX(ctx context.Context) {
if err := _d.Exec(ctx); err != nil {
panic(err)
}
}
+564
View File
@@ -0,0 +1,564 @@
// Code generated by ent, DO NOT EDIT.
package ent
import (
"context"
"fmt"
"math"
"entgo.io/ent"
"entgo.io/ent/dialect"
"entgo.io/ent/dialect/sql"
"entgo.io/ent/dialect/sql/sqlgraph"
"entgo.io/ent/schema/field"
"github.com/Wei-Shaw/sub2api/ent/batchimageitem"
"github.com/Wei-Shaw/sub2api/ent/predicate"
)
// BatchImageItemQuery is the builder for querying BatchImageItem entities.
type BatchImageItemQuery struct {
config
ctx *QueryContext
order []batchimageitem.OrderOption
inters []Interceptor
predicates []predicate.BatchImageItem
modifiers []func(*sql.Selector)
// intermediate query (i.e. traversal path).
sql *sql.Selector
path func(context.Context) (*sql.Selector, error)
}
// Where adds a new predicate for the BatchImageItemQuery builder.
func (_q *BatchImageItemQuery) Where(ps ...predicate.BatchImageItem) *BatchImageItemQuery {
_q.predicates = append(_q.predicates, ps...)
return _q
}
// Limit the number of records to be returned by this query.
func (_q *BatchImageItemQuery) Limit(limit int) *BatchImageItemQuery {
_q.ctx.Limit = &limit
return _q
}
// Offset to start from.
func (_q *BatchImageItemQuery) Offset(offset int) *BatchImageItemQuery {
_q.ctx.Offset = &offset
return _q
}
// Unique configures the query builder to filter duplicate records on query.
// By default, unique is set to true, and can be disabled using this method.
func (_q *BatchImageItemQuery) Unique(unique bool) *BatchImageItemQuery {
_q.ctx.Unique = &unique
return _q
}
// Order specifies how the records should be ordered.
func (_q *BatchImageItemQuery) Order(o ...batchimageitem.OrderOption) *BatchImageItemQuery {
_q.order = append(_q.order, o...)
return _q
}
// First returns the first BatchImageItem entity from the query.
// Returns a *NotFoundError when no BatchImageItem was found.
func (_q *BatchImageItemQuery) First(ctx context.Context) (*BatchImageItem, error) {
nodes, err := _q.Limit(1).All(setContextOp(ctx, _q.ctx, ent.OpQueryFirst))
if err != nil {
return nil, err
}
if len(nodes) == 0 {
return nil, &NotFoundError{batchimageitem.Label}
}
return nodes[0], nil
}
// FirstX is like First, but panics if an error occurs.
func (_q *BatchImageItemQuery) FirstX(ctx context.Context) *BatchImageItem {
node, err := _q.First(ctx)
if err != nil && !IsNotFound(err) {
panic(err)
}
return node
}
// FirstID returns the first BatchImageItem ID from the query.
// Returns a *NotFoundError when no BatchImageItem ID was found.
func (_q *BatchImageItemQuery) FirstID(ctx context.Context) (id int64, err error) {
var ids []int64
if ids, err = _q.Limit(1).IDs(setContextOp(ctx, _q.ctx, ent.OpQueryFirstID)); err != nil {
return
}
if len(ids) == 0 {
err = &NotFoundError{batchimageitem.Label}
return
}
return ids[0], nil
}
// FirstIDX is like FirstID, but panics if an error occurs.
func (_q *BatchImageItemQuery) FirstIDX(ctx context.Context) int64 {
id, err := _q.FirstID(ctx)
if err != nil && !IsNotFound(err) {
panic(err)
}
return id
}
// Only returns a single BatchImageItem entity found by the query, ensuring it only returns one.
// Returns a *NotSingularError when more than one BatchImageItem entity is found.
// Returns a *NotFoundError when no BatchImageItem entities are found.
func (_q *BatchImageItemQuery) Only(ctx context.Context) (*BatchImageItem, error) {
nodes, err := _q.Limit(2).All(setContextOp(ctx, _q.ctx, ent.OpQueryOnly))
if err != nil {
return nil, err
}
switch len(nodes) {
case 1:
return nodes[0], nil
case 0:
return nil, &NotFoundError{batchimageitem.Label}
default:
return nil, &NotSingularError{batchimageitem.Label}
}
}
// OnlyX is like Only, but panics if an error occurs.
func (_q *BatchImageItemQuery) OnlyX(ctx context.Context) *BatchImageItem {
node, err := _q.Only(ctx)
if err != nil {
panic(err)
}
return node
}
// OnlyID is like Only, but returns the only BatchImageItem ID in the query.
// Returns a *NotSingularError when more than one BatchImageItem ID is found.
// Returns a *NotFoundError when no entities are found.
func (_q *BatchImageItemQuery) OnlyID(ctx context.Context) (id int64, err error) {
var ids []int64
if ids, err = _q.Limit(2).IDs(setContextOp(ctx, _q.ctx, ent.OpQueryOnlyID)); err != nil {
return
}
switch len(ids) {
case 1:
id = ids[0]
case 0:
err = &NotFoundError{batchimageitem.Label}
default:
err = &NotSingularError{batchimageitem.Label}
}
return
}
// OnlyIDX is like OnlyID, but panics if an error occurs.
func (_q *BatchImageItemQuery) OnlyIDX(ctx context.Context) int64 {
id, err := _q.OnlyID(ctx)
if err != nil {
panic(err)
}
return id
}
// All executes the query and returns a list of BatchImageItems.
func (_q *BatchImageItemQuery) All(ctx context.Context) ([]*BatchImageItem, error) {
ctx = setContextOp(ctx, _q.ctx, ent.OpQueryAll)
if err := _q.prepareQuery(ctx); err != nil {
return nil, err
}
qr := querierAll[[]*BatchImageItem, *BatchImageItemQuery]()
return withInterceptors[[]*BatchImageItem](ctx, _q, qr, _q.inters)
}
// AllX is like All, but panics if an error occurs.
func (_q *BatchImageItemQuery) AllX(ctx context.Context) []*BatchImageItem {
nodes, err := _q.All(ctx)
if err != nil {
panic(err)
}
return nodes
}
// IDs executes the query and returns a list of BatchImageItem IDs.
func (_q *BatchImageItemQuery) IDs(ctx context.Context) (ids []int64, err error) {
if _q.ctx.Unique == nil && _q.path != nil {
_q.Unique(true)
}
ctx = setContextOp(ctx, _q.ctx, ent.OpQueryIDs)
if err = _q.Select(batchimageitem.FieldID).Scan(ctx, &ids); err != nil {
return nil, err
}
return ids, nil
}
// IDsX is like IDs, but panics if an error occurs.
func (_q *BatchImageItemQuery) IDsX(ctx context.Context) []int64 {
ids, err := _q.IDs(ctx)
if err != nil {
panic(err)
}
return ids
}
// Count returns the count of the given query.
func (_q *BatchImageItemQuery) Count(ctx context.Context) (int, error) {
ctx = setContextOp(ctx, _q.ctx, ent.OpQueryCount)
if err := _q.prepareQuery(ctx); err != nil {
return 0, err
}
return withInterceptors[int](ctx, _q, querierCount[*BatchImageItemQuery](), _q.inters)
}
// CountX is like Count, but panics if an error occurs.
func (_q *BatchImageItemQuery) CountX(ctx context.Context) int {
count, err := _q.Count(ctx)
if err != nil {
panic(err)
}
return count
}
// Exist returns true if the query has elements in the graph.
func (_q *BatchImageItemQuery) Exist(ctx context.Context) (bool, error) {
ctx = setContextOp(ctx, _q.ctx, ent.OpQueryExist)
switch _, err := _q.FirstID(ctx); {
case IsNotFound(err):
return false, nil
case err != nil:
return false, fmt.Errorf("ent: check existence: %w", err)
default:
return true, nil
}
}
// ExistX is like Exist, but panics if an error occurs.
func (_q *BatchImageItemQuery) ExistX(ctx context.Context) bool {
exist, err := _q.Exist(ctx)
if err != nil {
panic(err)
}
return exist
}
// Clone returns a duplicate of the BatchImageItemQuery builder, including all associated steps. It can be
// used to prepare common query builders and use them differently after the clone is made.
func (_q *BatchImageItemQuery) Clone() *BatchImageItemQuery {
if _q == nil {
return nil
}
return &BatchImageItemQuery{
config: _q.config,
ctx: _q.ctx.Clone(),
order: append([]batchimageitem.OrderOption{}, _q.order...),
inters: append([]Interceptor{}, _q.inters...),
predicates: append([]predicate.BatchImageItem{}, _q.predicates...),
// clone intermediate query.
sql: _q.sql.Clone(),
path: _q.path,
}
}
// GroupBy is used to group vertices by one or more fields/columns.
// It is often used with aggregate functions, like: count, max, mean, min, sum.
//
// Example:
//
// var v []struct {
// JobID string `json:"job_id,omitempty"`
// Count int `json:"count,omitempty"`
// }
//
// client.BatchImageItem.Query().
// GroupBy(batchimageitem.FieldJobID).
// Aggregate(ent.Count()).
// Scan(ctx, &v)
func (_q *BatchImageItemQuery) GroupBy(field string, fields ...string) *BatchImageItemGroupBy {
_q.ctx.Fields = append([]string{field}, fields...)
grbuild := &BatchImageItemGroupBy{build: _q}
grbuild.flds = &_q.ctx.Fields
grbuild.label = batchimageitem.Label
grbuild.scan = grbuild.Scan
return grbuild
}
// Select allows the selection one or more fields/columns for the given query,
// instead of selecting all fields in the entity.
//
// Example:
//
// var v []struct {
// JobID string `json:"job_id,omitempty"`
// }
//
// client.BatchImageItem.Query().
// Select(batchimageitem.FieldJobID).
// Scan(ctx, &v)
func (_q *BatchImageItemQuery) Select(fields ...string) *BatchImageItemSelect {
_q.ctx.Fields = append(_q.ctx.Fields, fields...)
sbuild := &BatchImageItemSelect{BatchImageItemQuery: _q}
sbuild.label = batchimageitem.Label
sbuild.flds, sbuild.scan = &_q.ctx.Fields, sbuild.Scan
return sbuild
}
// Aggregate returns a BatchImageItemSelect configured with the given aggregations.
func (_q *BatchImageItemQuery) Aggregate(fns ...AggregateFunc) *BatchImageItemSelect {
return _q.Select().Aggregate(fns...)
}
func (_q *BatchImageItemQuery) prepareQuery(ctx context.Context) error {
for _, inter := range _q.inters {
if inter == nil {
return fmt.Errorf("ent: uninitialized interceptor (forgotten import ent/runtime?)")
}
if trv, ok := inter.(Traverser); ok {
if err := trv.Traverse(ctx, _q); err != nil {
return err
}
}
}
for _, f := range _q.ctx.Fields {
if !batchimageitem.ValidColumn(f) {
return &ValidationError{Name: f, err: fmt.Errorf("ent: invalid field %q for query", f)}
}
}
if _q.path != nil {
prev, err := _q.path(ctx)
if err != nil {
return err
}
_q.sql = prev
}
return nil
}
func (_q *BatchImageItemQuery) sqlAll(ctx context.Context, hooks ...queryHook) ([]*BatchImageItem, error) {
var (
nodes = []*BatchImageItem{}
_spec = _q.querySpec()
)
_spec.ScanValues = func(columns []string) ([]any, error) {
return (*BatchImageItem).scanValues(nil, columns)
}
_spec.Assign = func(columns []string, values []any) error {
node := &BatchImageItem{config: _q.config}
nodes = append(nodes, node)
return node.assignValues(columns, values)
}
if len(_q.modifiers) > 0 {
_spec.Modifiers = _q.modifiers
}
for i := range hooks {
hooks[i](ctx, _spec)
}
if err := sqlgraph.QueryNodes(ctx, _q.driver, _spec); err != nil {
return nil, err
}
if len(nodes) == 0 {
return nodes, nil
}
return nodes, nil
}
func (_q *BatchImageItemQuery) sqlCount(ctx context.Context) (int, error) {
_spec := _q.querySpec()
if len(_q.modifiers) > 0 {
_spec.Modifiers = _q.modifiers
}
_spec.Node.Columns = _q.ctx.Fields
if len(_q.ctx.Fields) > 0 {
_spec.Unique = _q.ctx.Unique != nil && *_q.ctx.Unique
}
return sqlgraph.CountNodes(ctx, _q.driver, _spec)
}
func (_q *BatchImageItemQuery) querySpec() *sqlgraph.QuerySpec {
_spec := sqlgraph.NewQuerySpec(batchimageitem.Table, batchimageitem.Columns, sqlgraph.NewFieldSpec(batchimageitem.FieldID, field.TypeInt64))
_spec.From = _q.sql
if unique := _q.ctx.Unique; unique != nil {
_spec.Unique = *unique
} else if _q.path != nil {
_spec.Unique = true
}
if fields := _q.ctx.Fields; len(fields) > 0 {
_spec.Node.Columns = make([]string, 0, len(fields))
_spec.Node.Columns = append(_spec.Node.Columns, batchimageitem.FieldID)
for i := range fields {
if fields[i] != batchimageitem.FieldID {
_spec.Node.Columns = append(_spec.Node.Columns, fields[i])
}
}
}
if ps := _q.predicates; len(ps) > 0 {
_spec.Predicate = func(selector *sql.Selector) {
for i := range ps {
ps[i](selector)
}
}
}
if limit := _q.ctx.Limit; limit != nil {
_spec.Limit = *limit
}
if offset := _q.ctx.Offset; offset != nil {
_spec.Offset = *offset
}
if ps := _q.order; len(ps) > 0 {
_spec.Order = func(selector *sql.Selector) {
for i := range ps {
ps[i](selector)
}
}
}
return _spec
}
func (_q *BatchImageItemQuery) sqlQuery(ctx context.Context) *sql.Selector {
builder := sql.Dialect(_q.driver.Dialect())
t1 := builder.Table(batchimageitem.Table)
columns := _q.ctx.Fields
if len(columns) == 0 {
columns = batchimageitem.Columns
}
selector := builder.Select(t1.Columns(columns...)...).From(t1)
if _q.sql != nil {
selector = _q.sql
selector.Select(selector.Columns(columns...)...)
}
if _q.ctx.Unique != nil && *_q.ctx.Unique {
selector.Distinct()
}
for _, m := range _q.modifiers {
m(selector)
}
for _, p := range _q.predicates {
p(selector)
}
for _, p := range _q.order {
p(selector)
}
if offset := _q.ctx.Offset; offset != nil {
// limit is mandatory for offset clause. We start
// with default value, and override it below if needed.
selector.Offset(*offset).Limit(math.MaxInt32)
}
if limit := _q.ctx.Limit; limit != nil {
selector.Limit(*limit)
}
return selector
}
// ForUpdate locks the selected rows against concurrent updates, and prevent them from being
// updated, deleted or "selected ... for update" by other sessions, until the transaction is
// either committed or rolled-back.
func (_q *BatchImageItemQuery) ForUpdate(opts ...sql.LockOption) *BatchImageItemQuery {
if _q.driver.Dialect() == dialect.Postgres {
_q.Unique(false)
}
_q.modifiers = append(_q.modifiers, func(s *sql.Selector) {
s.ForUpdate(opts...)
})
return _q
}
// ForShare behaves similarly to ForUpdate, except that it acquires a shared mode lock
// on any rows that are read. Other sessions can read the rows, but cannot modify them
// until your transaction commits.
func (_q *BatchImageItemQuery) ForShare(opts ...sql.LockOption) *BatchImageItemQuery {
if _q.driver.Dialect() == dialect.Postgres {
_q.Unique(false)
}
_q.modifiers = append(_q.modifiers, func(s *sql.Selector) {
s.ForShare(opts...)
})
return _q
}
// BatchImageItemGroupBy is the group-by builder for BatchImageItem entities.
type BatchImageItemGroupBy struct {
selector
build *BatchImageItemQuery
}
// Aggregate adds the given aggregation functions to the group-by query.
func (_g *BatchImageItemGroupBy) Aggregate(fns ...AggregateFunc) *BatchImageItemGroupBy {
_g.fns = append(_g.fns, fns...)
return _g
}
// Scan applies the selector query and scans the result into the given value.
func (_g *BatchImageItemGroupBy) Scan(ctx context.Context, v any) error {
ctx = setContextOp(ctx, _g.build.ctx, ent.OpQueryGroupBy)
if err := _g.build.prepareQuery(ctx); err != nil {
return err
}
return scanWithInterceptors[*BatchImageItemQuery, *BatchImageItemGroupBy](ctx, _g.build, _g, _g.build.inters, v)
}
func (_g *BatchImageItemGroupBy) sqlScan(ctx context.Context, root *BatchImageItemQuery, v any) error {
selector := root.sqlQuery(ctx).Select()
aggregation := make([]string, 0, len(_g.fns))
for _, fn := range _g.fns {
aggregation = append(aggregation, fn(selector))
}
if len(selector.SelectedColumns()) == 0 {
columns := make([]string, 0, len(*_g.flds)+len(_g.fns))
for _, f := range *_g.flds {
columns = append(columns, selector.C(f))
}
columns = append(columns, aggregation...)
selector.Select(columns...)
}
selector.GroupBy(selector.Columns(*_g.flds...)...)
if err := selector.Err(); err != nil {
return err
}
rows := &sql.Rows{}
query, args := selector.Query()
if err := _g.build.driver.Query(ctx, query, args, rows); err != nil {
return err
}
defer rows.Close()
return sql.ScanSlice(rows, v)
}
// BatchImageItemSelect is the builder for selecting fields of BatchImageItem entities.
type BatchImageItemSelect struct {
*BatchImageItemQuery
selector
}
// Aggregate adds the given aggregation functions to the selector query.
func (_s *BatchImageItemSelect) Aggregate(fns ...AggregateFunc) *BatchImageItemSelect {
_s.fns = append(_s.fns, fns...)
return _s
}
// Scan applies the selector query and scans the result into the given value.
func (_s *BatchImageItemSelect) Scan(ctx context.Context, v any) error {
ctx = setContextOp(ctx, _s.ctx, ent.OpQuerySelect)
if err := _s.prepareQuery(ctx); err != nil {
return err
}
return scanWithInterceptors[*BatchImageItemQuery, *BatchImageItemSelect](ctx, _s.BatchImageItemQuery, _s, _s.inters, v)
}
func (_s *BatchImageItemSelect) sqlScan(ctx context.Context, root *BatchImageItemQuery, v any) error {
selector := root.sqlQuery(ctx)
aggregation := make([]string, 0, len(_s.fns))
for _, fn := range _s.fns {
aggregation = append(aggregation, fn(selector))
}
switch n := len(*_s.selector.flds); {
case n == 0 && len(aggregation) > 0:
selector.Select(aggregation...)
case n != 0 && len(aggregation) > 0:
selector.AppendSelect(aggregation...)
}
rows := &sql.Rows{}
query, args := selector.Query()
if err := _s.driver.Query(ctx, query, args, rows); err != nil {
return err
}
defer rows.Close()
return sql.ScanSlice(rows, v)
}
File diff suppressed because it is too large Load Diff
+609
View File
@@ -0,0 +1,609 @@
// Code generated by ent, DO NOT EDIT.
package ent
import (
"fmt"
"strings"
"time"
"entgo.io/ent"
"entgo.io/ent/dialect/sql"
"github.com/Wei-Shaw/sub2api/ent/batchimagejob"
)
// BatchImageJob is the model entity for the BatchImageJob schema.
type BatchImageJob struct {
config `json:"-"`
// ID of the ent.
ID int64 `json:"id,omitempty"`
// BatchID holds the value of the "batch_id" field.
BatchID string `json:"batch_id,omitempty"`
// UserID holds the value of the "user_id" field.
UserID int64 `json:"user_id,omitempty"`
// APIKeyID holds the value of the "api_key_id" field.
APIKeyID *int64 `json:"api_key_id,omitempty"`
// AccountID holds the value of the "account_id" field.
AccountID *int64 `json:"account_id,omitempty"`
// Provider holds the value of the "provider" field.
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.
ProviderJobName *string `json:"provider_job_name,omitempty"`
// ProviderInputRef holds the value of the "provider_input_ref" field.
ProviderInputRef *string `json:"provider_input_ref,omitempty"`
// ProviderOutputRef holds the value of the "provider_output_ref" field.
ProviderOutputRef *string `json:"provider_output_ref,omitempty"`
// GcsInputURI holds the value of the "gcs_input_uri" field.
GcsInputURI *string `json:"gcs_input_uri,omitempty"`
// GcsOutputURI holds the value of the "gcs_output_uri" field.
GcsOutputURI *string `json:"gcs_output_uri,omitempty"`
// ItemCount holds the value of the "item_count" field.
ItemCount int `json:"item_count,omitempty"`
// SuccessCount holds the value of the "success_count" field.
SuccessCount int `json:"success_count,omitempty"`
// FailCount holds the value of the "fail_count" field.
FailCount int `json:"fail_count,omitempty"`
// CancelledCount holds the value of the "cancelled_count" field.
CancelledCount int `json:"cancelled_count,omitempty"`
// EstimatedCost holds the value of the "estimated_cost" field.
EstimatedCost float64 `json:"estimated_cost,omitempty"`
// HoldAmount holds the value of the "hold_amount" field.
HoldAmount *float64 `json:"hold_amount,omitempty"`
// ActualCost holds the value of the "actual_cost" field.
ActualCost *float64 `json:"actual_cost,omitempty"`
// Currency holds the value of the "currency" field.
Currency string `json:"currency,omitempty"`
// HoldID holds the value of the "hold_id" field.
HoldID *string `json:"hold_id,omitempty"`
// IdempotencyKey holds the value of the "idempotency_key" field.
IdempotencyKey *string `json:"idempotency_key,omitempty"`
// RequestHash holds the value of the "request_hash" field.
RequestHash *string `json:"request_hash,omitempty"`
// ManifestHash holds the value of the "manifest_hash" field.
ManifestHash *string `json:"manifest_hash,omitempty"`
// RetryCount holds the value of the "retry_count" field.
RetryCount int `json:"retry_count,omitempty"`
// Version holds the value of the "version" field.
Version int `json:"version,omitempty"`
// OutputExpiresAt holds the value of the "output_expires_at" field.
OutputExpiresAt *time.Time `json:"output_expires_at,omitempty"`
// InputDeletedAt holds the value of the "input_deleted_at" field.
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.
LastErrorMessage *string `json:"last_error_message,omitempty"`
// CreatedAt holds the value of the "created_at" field.
CreatedAt time.Time `json:"created_at,omitempty"`
// UpdatedAt holds the value of the "updated_at" field.
UpdatedAt time.Time `json:"updated_at,omitempty"`
// SubmittedAt holds the value of the "submitted_at" field.
SubmittedAt *time.Time `json:"submitted_at,omitempty"`
// StartedAt holds the value of the "started_at" field.
StartedAt *time.Time `json:"started_at,omitempty"`
// FinishedAt holds the value of the "finished_at" field.
FinishedAt *time.Time `json:"finished_at,omitempty"`
// SettledAt holds the value of the "settled_at" field.
SettledAt *time.Time `json:"settled_at,omitempty"`
selectValues sql.SelectValues
}
// scanValues returns the types for scanning values from sql.Rows.
func (*BatchImageJob) scanValues(columns []string) ([]any, error) {
values := make([]any, len(columns))
for i := range columns {
switch columns[i] {
case batchimagejob.FieldEstimatedCost, batchimagejob.FieldHoldAmount, batchimagejob.FieldActualCost:
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.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.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)
}
}
return values, nil
}
// assignValues assigns the values that were returned from sql.Rows (after scanning)
// to the BatchImageJob fields.
func (_m *BatchImageJob) assignValues(columns []string, values []any) error {
if m, n := len(values), len(columns); m < n {
return fmt.Errorf("mismatch number of scan values: %d != %d", m, n)
}
for i := range columns {
switch columns[i] {
case batchimagejob.FieldID:
value, ok := values[i].(*sql.NullInt64)
if !ok {
return fmt.Errorf("unexpected type %T for field id", value)
}
_m.ID = int64(value.Int64)
case batchimagejob.FieldBatchID:
if value, ok := values[i].(*sql.NullString); !ok {
return fmt.Errorf("unexpected type %T for field batch_id", values[i])
} else if value.Valid {
_m.BatchID = value.String
}
case batchimagejob.FieldUserID:
if value, ok := values[i].(*sql.NullInt64); !ok {
return fmt.Errorf("unexpected type %T for field user_id", values[i])
} else if value.Valid {
_m.UserID = value.Int64
}
case batchimagejob.FieldAPIKeyID:
if value, ok := values[i].(*sql.NullInt64); !ok {
return fmt.Errorf("unexpected type %T for field api_key_id", values[i])
} else if value.Valid {
_m.APIKeyID = new(int64)
*_m.APIKeyID = value.Int64
}
case batchimagejob.FieldAccountID:
if value, ok := values[i].(*sql.NullInt64); !ok {
return fmt.Errorf("unexpected type %T for field account_id", values[i])
} else if value.Valid {
_m.AccountID = new(int64)
*_m.AccountID = value.Int64
}
case batchimagejob.FieldProvider:
if value, ok := values[i].(*sql.NullString); !ok {
return fmt.Errorf("unexpected type %T for field provider", values[i])
} else if value.Valid {
_m.Provider = value.String
}
case batchimagejob.FieldModel:
if value, ok := values[i].(*sql.NullString); !ok {
return fmt.Errorf("unexpected type %T for field model", values[i])
} 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])
} else if value.Valid {
_m.Status = value.String
}
case batchimagejob.FieldProviderJobName:
if value, ok := values[i].(*sql.NullString); !ok {
return fmt.Errorf("unexpected type %T for field provider_job_name", values[i])
} else if value.Valid {
_m.ProviderJobName = new(string)
*_m.ProviderJobName = value.String
}
case batchimagejob.FieldProviderInputRef:
if value, ok := values[i].(*sql.NullString); !ok {
return fmt.Errorf("unexpected type %T for field provider_input_ref", values[i])
} else if value.Valid {
_m.ProviderInputRef = new(string)
*_m.ProviderInputRef = value.String
}
case batchimagejob.FieldProviderOutputRef:
if value, ok := values[i].(*sql.NullString); !ok {
return fmt.Errorf("unexpected type %T for field provider_output_ref", values[i])
} else if value.Valid {
_m.ProviderOutputRef = new(string)
*_m.ProviderOutputRef = value.String
}
case batchimagejob.FieldGcsInputURI:
if value, ok := values[i].(*sql.NullString); !ok {
return fmt.Errorf("unexpected type %T for field gcs_input_uri", values[i])
} else if value.Valid {
_m.GcsInputURI = new(string)
*_m.GcsInputURI = value.String
}
case batchimagejob.FieldGcsOutputURI:
if value, ok := values[i].(*sql.NullString); !ok {
return fmt.Errorf("unexpected type %T for field gcs_output_uri", values[i])
} else if value.Valid {
_m.GcsOutputURI = new(string)
*_m.GcsOutputURI = value.String
}
case batchimagejob.FieldItemCount:
if value, ok := values[i].(*sql.NullInt64); !ok {
return fmt.Errorf("unexpected type %T for field item_count", values[i])
} else if value.Valid {
_m.ItemCount = int(value.Int64)
}
case batchimagejob.FieldSuccessCount:
if value, ok := values[i].(*sql.NullInt64); !ok {
return fmt.Errorf("unexpected type %T for field success_count", values[i])
} else if value.Valid {
_m.SuccessCount = int(value.Int64)
}
case batchimagejob.FieldFailCount:
if value, ok := values[i].(*sql.NullInt64); !ok {
return fmt.Errorf("unexpected type %T for field fail_count", values[i])
} else if value.Valid {
_m.FailCount = int(value.Int64)
}
case batchimagejob.FieldCancelledCount:
if value, ok := values[i].(*sql.NullInt64); !ok {
return fmt.Errorf("unexpected type %T for field cancelled_count", values[i])
} else if value.Valid {
_m.CancelledCount = int(value.Int64)
}
case batchimagejob.FieldEstimatedCost:
if value, ok := values[i].(*sql.NullFloat64); !ok {
return fmt.Errorf("unexpected type %T for field estimated_cost", values[i])
} else if value.Valid {
_m.EstimatedCost = value.Float64
}
case batchimagejob.FieldHoldAmount:
if value, ok := values[i].(*sql.NullFloat64); !ok {
return fmt.Errorf("unexpected type %T for field hold_amount", values[i])
} else if value.Valid {
_m.HoldAmount = new(float64)
*_m.HoldAmount = value.Float64
}
case batchimagejob.FieldActualCost:
if value, ok := values[i].(*sql.NullFloat64); !ok {
return fmt.Errorf("unexpected type %T for field actual_cost", values[i])
} else if value.Valid {
_m.ActualCost = new(float64)
*_m.ActualCost = value.Float64
}
case batchimagejob.FieldCurrency:
if value, ok := values[i].(*sql.NullString); !ok {
return fmt.Errorf("unexpected type %T for field currency", values[i])
} else if value.Valid {
_m.Currency = value.String
}
case batchimagejob.FieldHoldID:
if value, ok := values[i].(*sql.NullString); !ok {
return fmt.Errorf("unexpected type %T for field hold_id", values[i])
} else if value.Valid {
_m.HoldID = new(string)
*_m.HoldID = value.String
}
case batchimagejob.FieldIdempotencyKey:
if value, ok := values[i].(*sql.NullString); !ok {
return fmt.Errorf("unexpected type %T for field idempotency_key", values[i])
} else if value.Valid {
_m.IdempotencyKey = new(string)
*_m.IdempotencyKey = value.String
}
case batchimagejob.FieldRequestHash:
if value, ok := values[i].(*sql.NullString); !ok {
return fmt.Errorf("unexpected type %T for field request_hash", values[i])
} else if value.Valid {
_m.RequestHash = new(string)
*_m.RequestHash = value.String
}
case batchimagejob.FieldManifestHash:
if value, ok := values[i].(*sql.NullString); !ok {
return fmt.Errorf("unexpected type %T for field manifest_hash", values[i])
} else if value.Valid {
_m.ManifestHash = new(string)
*_m.ManifestHash = value.String
}
case batchimagejob.FieldRetryCount:
if value, ok := values[i].(*sql.NullInt64); !ok {
return fmt.Errorf("unexpected type %T for field retry_count", values[i])
} else if value.Valid {
_m.RetryCount = int(value.Int64)
}
case batchimagejob.FieldVersion:
if value, ok := values[i].(*sql.NullInt64); !ok {
return fmt.Errorf("unexpected type %T for field version", values[i])
} else if value.Valid {
_m.Version = int(value.Int64)
}
case batchimagejob.FieldOutputExpiresAt:
if value, ok := values[i].(*sql.NullTime); !ok {
return fmt.Errorf("unexpected type %T for field output_expires_at", values[i])
} else if value.Valid {
_m.OutputExpiresAt = new(time.Time)
*_m.OutputExpiresAt = value.Time
}
case batchimagejob.FieldInputDeletedAt:
if value, ok := values[i].(*sql.NullTime); !ok {
return fmt.Errorf("unexpected type %T for field input_deleted_at", values[i])
} else if value.Valid {
_m.InputDeletedAt = new(time.Time)
*_m.InputDeletedAt = value.Time
}
case batchimagejob.FieldOutputDeletedAt:
if value, ok := values[i].(*sql.NullTime); !ok {
return fmt.Errorf("unexpected type %T for field output_deleted_at", values[i])
} else if value.Valid {
_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])
} else if value.Valid {
_m.LastErrorCode = new(string)
*_m.LastErrorCode = value.String
}
case batchimagejob.FieldLastErrorMessage:
if value, ok := values[i].(*sql.NullString); !ok {
return fmt.Errorf("unexpected type %T for field last_error_message", values[i])
} else if value.Valid {
_m.LastErrorMessage = new(string)
*_m.LastErrorMessage = value.String
}
case batchimagejob.FieldCreatedAt:
if value, ok := values[i].(*sql.NullTime); !ok {
return fmt.Errorf("unexpected type %T for field created_at", values[i])
} else if value.Valid {
_m.CreatedAt = value.Time
}
case batchimagejob.FieldUpdatedAt:
if value, ok := values[i].(*sql.NullTime); !ok {
return fmt.Errorf("unexpected type %T for field updated_at", values[i])
} else if value.Valid {
_m.UpdatedAt = value.Time
}
case batchimagejob.FieldSubmittedAt:
if value, ok := values[i].(*sql.NullTime); !ok {
return fmt.Errorf("unexpected type %T for field submitted_at", values[i])
} else if value.Valid {
_m.SubmittedAt = new(time.Time)
*_m.SubmittedAt = value.Time
}
case batchimagejob.FieldStartedAt:
if value, ok := values[i].(*sql.NullTime); !ok {
return fmt.Errorf("unexpected type %T for field started_at", values[i])
} else if value.Valid {
_m.StartedAt = new(time.Time)
*_m.StartedAt = value.Time
}
case batchimagejob.FieldFinishedAt:
if value, ok := values[i].(*sql.NullTime); !ok {
return fmt.Errorf("unexpected type %T for field finished_at", values[i])
} else if value.Valid {
_m.FinishedAt = new(time.Time)
*_m.FinishedAt = value.Time
}
case batchimagejob.FieldSettledAt:
if value, ok := values[i].(*sql.NullTime); !ok {
return fmt.Errorf("unexpected type %T for field settled_at", values[i])
} else if value.Valid {
_m.SettledAt = new(time.Time)
*_m.SettledAt = value.Time
}
default:
_m.selectValues.Set(columns[i], values[i])
}
}
return nil
}
// Value returns the ent.Value that was dynamically selected and assigned to the BatchImageJob.
// This includes values selected through modifiers, order, etc.
func (_m *BatchImageJob) Value(name string) (ent.Value, error) {
return _m.selectValues.Get(name)
}
// Update returns a builder for updating this BatchImageJob.
// Note that you need to call BatchImageJob.Unwrap() before calling this method if this BatchImageJob
// was returned from a transaction, and the transaction was committed or rolled back.
func (_m *BatchImageJob) Update() *BatchImageJobUpdateOne {
return NewBatchImageJobClient(_m.config).UpdateOne(_m)
}
// Unwrap unwraps the BatchImageJob entity that was returned from a transaction after it was closed,
// so that all future queries will be executed through the driver which created the transaction.
func (_m *BatchImageJob) Unwrap() *BatchImageJob {
_tx, ok := _m.config.driver.(*txDriver)
if !ok {
panic("ent: BatchImageJob is not a transactional entity")
}
_m.config.driver = _tx.drv
return _m
}
// String implements the fmt.Stringer.
func (_m *BatchImageJob) String() string {
var builder strings.Builder
builder.WriteString("BatchImageJob(")
builder.WriteString(fmt.Sprintf("id=%v, ", _m.ID))
builder.WriteString("batch_id=")
builder.WriteString(_m.BatchID)
builder.WriteString(", ")
builder.WriteString("user_id=")
builder.WriteString(fmt.Sprintf("%v", _m.UserID))
builder.WriteString(", ")
if v := _m.APIKeyID; v != nil {
builder.WriteString("api_key_id=")
builder.WriteString(fmt.Sprintf("%v", *v))
}
builder.WriteString(", ")
if v := _m.AccountID; v != nil {
builder.WriteString("account_id=")
builder.WriteString(fmt.Sprintf("%v", *v))
}
builder.WriteString(", ")
builder.WriteString("provider=")
builder.WriteString(_m.Provider)
builder.WriteString(", ")
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(", ")
if v := _m.ProviderJobName; v != nil {
builder.WriteString("provider_job_name=")
builder.WriteString(*v)
}
builder.WriteString(", ")
if v := _m.ProviderInputRef; v != nil {
builder.WriteString("provider_input_ref=")
builder.WriteString(*v)
}
builder.WriteString(", ")
if v := _m.ProviderOutputRef; v != nil {
builder.WriteString("provider_output_ref=")
builder.WriteString(*v)
}
builder.WriteString(", ")
if v := _m.GcsInputURI; v != nil {
builder.WriteString("gcs_input_uri=")
builder.WriteString(*v)
}
builder.WriteString(", ")
if v := _m.GcsOutputURI; v != nil {
builder.WriteString("gcs_output_uri=")
builder.WriteString(*v)
}
builder.WriteString(", ")
builder.WriteString("item_count=")
builder.WriteString(fmt.Sprintf("%v", _m.ItemCount))
builder.WriteString(", ")
builder.WriteString("success_count=")
builder.WriteString(fmt.Sprintf("%v", _m.SuccessCount))
builder.WriteString(", ")
builder.WriteString("fail_count=")
builder.WriteString(fmt.Sprintf("%v", _m.FailCount))
builder.WriteString(", ")
builder.WriteString("cancelled_count=")
builder.WriteString(fmt.Sprintf("%v", _m.CancelledCount))
builder.WriteString(", ")
builder.WriteString("estimated_cost=")
builder.WriteString(fmt.Sprintf("%v", _m.EstimatedCost))
builder.WriteString(", ")
if v := _m.HoldAmount; v != nil {
builder.WriteString("hold_amount=")
builder.WriteString(fmt.Sprintf("%v", *v))
}
builder.WriteString(", ")
if v := _m.ActualCost; v != nil {
builder.WriteString("actual_cost=")
builder.WriteString(fmt.Sprintf("%v", *v))
}
builder.WriteString(", ")
builder.WriteString("currency=")
builder.WriteString(_m.Currency)
builder.WriteString(", ")
if v := _m.HoldID; v != nil {
builder.WriteString("hold_id=")
builder.WriteString(*v)
}
builder.WriteString(", ")
if v := _m.IdempotencyKey; v != nil {
builder.WriteString("idempotency_key=")
builder.WriteString(*v)
}
builder.WriteString(", ")
if v := _m.RequestHash; v != nil {
builder.WriteString("request_hash=")
builder.WriteString(*v)
}
builder.WriteString(", ")
if v := _m.ManifestHash; v != nil {
builder.WriteString("manifest_hash=")
builder.WriteString(*v)
}
builder.WriteString(", ")
builder.WriteString("retry_count=")
builder.WriteString(fmt.Sprintf("%v", _m.RetryCount))
builder.WriteString(", ")
builder.WriteString("version=")
builder.WriteString(fmt.Sprintf("%v", _m.Version))
builder.WriteString(", ")
if v := _m.OutputExpiresAt; v != nil {
builder.WriteString("output_expires_at=")
builder.WriteString(v.Format(time.ANSIC))
}
builder.WriteString(", ")
if v := _m.InputDeletedAt; v != nil {
builder.WriteString("input_deleted_at=")
builder.WriteString(v.Format(time.ANSIC))
}
builder.WriteString(", ")
if v := _m.OutputDeletedAt; v != nil {
builder.WriteString("output_deleted_at=")
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)
}
builder.WriteString(", ")
if v := _m.LastErrorMessage; v != nil {
builder.WriteString("last_error_message=")
builder.WriteString(*v)
}
builder.WriteString(", ")
builder.WriteString("created_at=")
builder.WriteString(_m.CreatedAt.Format(time.ANSIC))
builder.WriteString(", ")
builder.WriteString("updated_at=")
builder.WriteString(_m.UpdatedAt.Format(time.ANSIC))
builder.WriteString(", ")
if v := _m.SubmittedAt; v != nil {
builder.WriteString("submitted_at=")
builder.WriteString(v.Format(time.ANSIC))
}
builder.WriteString(", ")
if v := _m.StartedAt; v != nil {
builder.WriteString("started_at=")
builder.WriteString(v.Format(time.ANSIC))
}
builder.WriteString(", ")
if v := _m.FinishedAt; v != nil {
builder.WriteString("finished_at=")
builder.WriteString(v.Format(time.ANSIC))
}
builder.WriteString(", ")
if v := _m.SettledAt; v != nil {
builder.WriteString("settled_at=")
builder.WriteString(v.Format(time.ANSIC))
}
builder.WriteByte(')')
return builder.String()
}
// BatchImageJobs is a parsable slice of BatchImageJob.
type BatchImageJobs []*BatchImageJob
+420
View File
@@ -0,0 +1,420 @@
// Code generated by ent, DO NOT EDIT.
package batchimagejob
import (
"time"
"entgo.io/ent/dialect/sql"
)
const (
// Label holds the string label denoting the batchimagejob type in the database.
Label = "batch_image_job"
// FieldID holds the string denoting the id field in the database.
FieldID = "id"
// FieldBatchID holds the string denoting the batch_id field in the database.
FieldBatchID = "batch_id"
// FieldUserID holds the string denoting the user_id field in the database.
FieldUserID = "user_id"
// FieldAPIKeyID holds the string denoting the api_key_id field in the database.
FieldAPIKeyID = "api_key_id"
// FieldAccountID holds the string denoting the account_id field in the database.
FieldAccountID = "account_id"
// FieldProvider holds the string denoting the provider field in the database.
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.
FieldProviderJobName = "provider_job_name"
// FieldProviderInputRef holds the string denoting the provider_input_ref field in the database.
FieldProviderInputRef = "provider_input_ref"
// FieldProviderOutputRef holds the string denoting the provider_output_ref field in the database.
FieldProviderOutputRef = "provider_output_ref"
// FieldGcsInputURI holds the string denoting the gcs_input_uri field in the database.
FieldGcsInputURI = "gcs_input_uri"
// FieldGcsOutputURI holds the string denoting the gcs_output_uri field in the database.
FieldGcsOutputURI = "gcs_output_uri"
// FieldItemCount holds the string denoting the item_count field in the database.
FieldItemCount = "item_count"
// FieldSuccessCount holds the string denoting the success_count field in the database.
FieldSuccessCount = "success_count"
// FieldFailCount holds the string denoting the fail_count field in the database.
FieldFailCount = "fail_count"
// FieldCancelledCount holds the string denoting the cancelled_count field in the database.
FieldCancelledCount = "cancelled_count"
// FieldEstimatedCost holds the string denoting the estimated_cost field in the database.
FieldEstimatedCost = "estimated_cost"
// FieldHoldAmount holds the string denoting the hold_amount field in the database.
FieldHoldAmount = "hold_amount"
// FieldActualCost holds the string denoting the actual_cost field in the database.
FieldActualCost = "actual_cost"
// FieldCurrency holds the string denoting the currency field in the database.
FieldCurrency = "currency"
// FieldHoldID holds the string denoting the hold_id field in the database.
FieldHoldID = "hold_id"
// FieldIdempotencyKey holds the string denoting the idempotency_key field in the database.
FieldIdempotencyKey = "idempotency_key"
// FieldRequestHash holds the string denoting the request_hash field in the database.
FieldRequestHash = "request_hash"
// FieldManifestHash holds the string denoting the manifest_hash field in the database.
FieldManifestHash = "manifest_hash"
// FieldRetryCount holds the string denoting the retry_count field in the database.
FieldRetryCount = "retry_count"
// FieldVersion holds the string denoting the version field in the database.
FieldVersion = "version"
// FieldOutputExpiresAt holds the string denoting the output_expires_at field in the database.
FieldOutputExpiresAt = "output_expires_at"
// FieldInputDeletedAt holds the string denoting the input_deleted_at field in the database.
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.
FieldLastErrorMessage = "last_error_message"
// FieldCreatedAt holds the string denoting the created_at field in the database.
FieldCreatedAt = "created_at"
// FieldUpdatedAt holds the string denoting the updated_at field in the database.
FieldUpdatedAt = "updated_at"
// FieldSubmittedAt holds the string denoting the submitted_at field in the database.
FieldSubmittedAt = "submitted_at"
// FieldStartedAt holds the string denoting the started_at field in the database.
FieldStartedAt = "started_at"
// FieldFinishedAt holds the string denoting the finished_at field in the database.
FieldFinishedAt = "finished_at"
// FieldSettledAt holds the string denoting the settled_at field in the database.
FieldSettledAt = "settled_at"
// Table holds the table name of the batchimagejob in the database.
Table = "batch_image_jobs"
)
// Columns holds all SQL columns for batchimagejob fields.
var Columns = []string{
FieldID,
FieldBatchID,
FieldUserID,
FieldAPIKeyID,
FieldAccountID,
FieldProvider,
FieldModel,
FieldTaskName,
FieldStatus,
FieldProviderJobName,
FieldProviderInputRef,
FieldProviderOutputRef,
FieldGcsInputURI,
FieldGcsOutputURI,
FieldItemCount,
FieldSuccessCount,
FieldFailCount,
FieldCancelledCount,
FieldEstimatedCost,
FieldHoldAmount,
FieldActualCost,
FieldCurrency,
FieldHoldID,
FieldIdempotencyKey,
FieldRequestHash,
FieldManifestHash,
FieldRetryCount,
FieldVersion,
FieldOutputExpiresAt,
FieldInputDeletedAt,
FieldOutputDeletedAt,
FieldDownloadedAt,
FieldUserDeletedAt,
FieldLastErrorCode,
FieldLastErrorMessage,
FieldCreatedAt,
FieldUpdatedAt,
FieldSubmittedAt,
FieldStartedAt,
FieldFinishedAt,
FieldSettledAt,
}
// ValidColumn reports if the column name is valid (part of the table columns).
func ValidColumn(column string) bool {
for i := range Columns {
if column == Columns[i] {
return true
}
}
return false
}
var (
// BatchIDValidator is a validator for the "batch_id" field. It is called by the builders before save.
BatchIDValidator func(string) error
// ProviderValidator is a validator for the "provider" field. It is called by the builders before save.
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.
StatusValidator func(string) error
// ProviderJobNameValidator is a validator for the "provider_job_name" field. It is called by the builders before save.
ProviderJobNameValidator func(string) error
// ProviderInputRefValidator is a validator for the "provider_input_ref" field. It is called by the builders before save.
ProviderInputRefValidator func(string) error
// ProviderOutputRefValidator is a validator for the "provider_output_ref" field. It is called by the builders before save.
ProviderOutputRefValidator func(string) error
// GcsInputURIValidator is a validator for the "gcs_input_uri" field. It is called by the builders before save.
GcsInputURIValidator func(string) error
// GcsOutputURIValidator is a validator for the "gcs_output_uri" field. It is called by the builders before save.
GcsOutputURIValidator func(string) error
// DefaultSuccessCount holds the default value on creation for the "success_count" field.
DefaultSuccessCount int
// DefaultFailCount holds the default value on creation for the "fail_count" field.
DefaultFailCount int
// DefaultCancelledCount holds the default value on creation for the "cancelled_count" field.
DefaultCancelledCount int
// DefaultEstimatedCost holds the default value on creation for the "estimated_cost" field.
DefaultEstimatedCost float64
// DefaultCurrency holds the default value on creation for the "currency" field.
DefaultCurrency string
// CurrencyValidator is a validator for the "currency" field. It is called by the builders before save.
CurrencyValidator func(string) error
// HoldIDValidator is a validator for the "hold_id" field. It is called by the builders before save.
HoldIDValidator func(string) error
// IdempotencyKeyValidator is a validator for the "idempotency_key" field. It is called by the builders before save.
IdempotencyKeyValidator func(string) error
// RequestHashValidator is a validator for the "request_hash" field. It is called by the builders before save.
RequestHashValidator func(string) error
// ManifestHashValidator is a validator for the "manifest_hash" field. It is called by the builders before save.
ManifestHashValidator func(string) error
// DefaultRetryCount holds the default value on creation for the "retry_count" field.
DefaultRetryCount int
// DefaultVersion holds the default value on creation for the "version" field.
DefaultVersion int
// LastErrorCodeValidator is a validator for the "last_error_code" field. It is called by the builders before save.
LastErrorCodeValidator func(string) error
// DefaultCreatedAt holds the default value on creation for the "created_at" field.
DefaultCreatedAt func() time.Time
// DefaultUpdatedAt holds the default value on creation for the "updated_at" field.
DefaultUpdatedAt func() time.Time
// UpdateDefaultUpdatedAt holds the default value on update for the "updated_at" field.
UpdateDefaultUpdatedAt func() time.Time
)
// OrderOption defines the ordering options for the BatchImageJob queries.
type OrderOption func(*sql.Selector)
// ByID orders the results by the id field.
func ByID(opts ...sql.OrderTermOption) OrderOption {
return sql.OrderByField(FieldID, opts...).ToFunc()
}
// ByBatchID orders the results by the batch_id field.
func ByBatchID(opts ...sql.OrderTermOption) OrderOption {
return sql.OrderByField(FieldBatchID, opts...).ToFunc()
}
// ByUserID orders the results by the user_id field.
func ByUserID(opts ...sql.OrderTermOption) OrderOption {
return sql.OrderByField(FieldUserID, opts...).ToFunc()
}
// ByAPIKeyID orders the results by the api_key_id field.
func ByAPIKeyID(opts ...sql.OrderTermOption) OrderOption {
return sql.OrderByField(FieldAPIKeyID, opts...).ToFunc()
}
// ByAccountID orders the results by the account_id field.
func ByAccountID(opts ...sql.OrderTermOption) OrderOption {
return sql.OrderByField(FieldAccountID, opts...).ToFunc()
}
// ByProvider orders the results by the provider field.
func ByProvider(opts ...sql.OrderTermOption) OrderOption {
return sql.OrderByField(FieldProvider, opts...).ToFunc()
}
// ByModel orders the results by the model field.
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()
}
// ByProviderJobName orders the results by the provider_job_name field.
func ByProviderJobName(opts ...sql.OrderTermOption) OrderOption {
return sql.OrderByField(FieldProviderJobName, opts...).ToFunc()
}
// ByProviderInputRef orders the results by the provider_input_ref field.
func ByProviderInputRef(opts ...sql.OrderTermOption) OrderOption {
return sql.OrderByField(FieldProviderInputRef, opts...).ToFunc()
}
// ByProviderOutputRef orders the results by the provider_output_ref field.
func ByProviderOutputRef(opts ...sql.OrderTermOption) OrderOption {
return sql.OrderByField(FieldProviderOutputRef, opts...).ToFunc()
}
// ByGcsInputURI orders the results by the gcs_input_uri field.
func ByGcsInputURI(opts ...sql.OrderTermOption) OrderOption {
return sql.OrderByField(FieldGcsInputURI, opts...).ToFunc()
}
// ByGcsOutputURI orders the results by the gcs_output_uri field.
func ByGcsOutputURI(opts ...sql.OrderTermOption) OrderOption {
return sql.OrderByField(FieldGcsOutputURI, opts...).ToFunc()
}
// ByItemCount orders the results by the item_count field.
func ByItemCount(opts ...sql.OrderTermOption) OrderOption {
return sql.OrderByField(FieldItemCount, opts...).ToFunc()
}
// BySuccessCount orders the results by the success_count field.
func BySuccessCount(opts ...sql.OrderTermOption) OrderOption {
return sql.OrderByField(FieldSuccessCount, opts...).ToFunc()
}
// ByFailCount orders the results by the fail_count field.
func ByFailCount(opts ...sql.OrderTermOption) OrderOption {
return sql.OrderByField(FieldFailCount, opts...).ToFunc()
}
// ByCancelledCount orders the results by the cancelled_count field.
func ByCancelledCount(opts ...sql.OrderTermOption) OrderOption {
return sql.OrderByField(FieldCancelledCount, opts...).ToFunc()
}
// ByEstimatedCost orders the results by the estimated_cost field.
func ByEstimatedCost(opts ...sql.OrderTermOption) OrderOption {
return sql.OrderByField(FieldEstimatedCost, opts...).ToFunc()
}
// ByHoldAmount orders the results by the hold_amount field.
func ByHoldAmount(opts ...sql.OrderTermOption) OrderOption {
return sql.OrderByField(FieldHoldAmount, opts...).ToFunc()
}
// ByActualCost orders the results by the actual_cost field.
func ByActualCost(opts ...sql.OrderTermOption) OrderOption {
return sql.OrderByField(FieldActualCost, opts...).ToFunc()
}
// ByCurrency orders the results by the currency field.
func ByCurrency(opts ...sql.OrderTermOption) OrderOption {
return sql.OrderByField(FieldCurrency, opts...).ToFunc()
}
// ByHoldID orders the results by the hold_id field.
func ByHoldID(opts ...sql.OrderTermOption) OrderOption {
return sql.OrderByField(FieldHoldID, opts...).ToFunc()
}
// ByIdempotencyKey orders the results by the idempotency_key field.
func ByIdempotencyKey(opts ...sql.OrderTermOption) OrderOption {
return sql.OrderByField(FieldIdempotencyKey, opts...).ToFunc()
}
// ByRequestHash orders the results by the request_hash field.
func ByRequestHash(opts ...sql.OrderTermOption) OrderOption {
return sql.OrderByField(FieldRequestHash, opts...).ToFunc()
}
// ByManifestHash orders the results by the manifest_hash field.
func ByManifestHash(opts ...sql.OrderTermOption) OrderOption {
return sql.OrderByField(FieldManifestHash, opts...).ToFunc()
}
// ByRetryCount orders the results by the retry_count field.
func ByRetryCount(opts ...sql.OrderTermOption) OrderOption {
return sql.OrderByField(FieldRetryCount, opts...).ToFunc()
}
// ByVersion orders the results by the version field.
func ByVersion(opts ...sql.OrderTermOption) OrderOption {
return sql.OrderByField(FieldVersion, opts...).ToFunc()
}
// ByOutputExpiresAt orders the results by the output_expires_at field.
func ByOutputExpiresAt(opts ...sql.OrderTermOption) OrderOption {
return sql.OrderByField(FieldOutputExpiresAt, opts...).ToFunc()
}
// ByInputDeletedAt orders the results by the input_deleted_at field.
func ByInputDeletedAt(opts ...sql.OrderTermOption) OrderOption {
return sql.OrderByField(FieldInputDeletedAt, opts...).ToFunc()
}
// ByOutputDeletedAt orders the results by the output_deleted_at field.
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()
}
// ByLastErrorMessage orders the results by the last_error_message field.
func ByLastErrorMessage(opts ...sql.OrderTermOption) OrderOption {
return sql.OrderByField(FieldLastErrorMessage, opts...).ToFunc()
}
// ByCreatedAt orders the results by the created_at field.
func ByCreatedAt(opts ...sql.OrderTermOption) OrderOption {
return sql.OrderByField(FieldCreatedAt, opts...).ToFunc()
}
// ByUpdatedAt orders the results by the updated_at field.
func ByUpdatedAt(opts ...sql.OrderTermOption) OrderOption {
return sql.OrderByField(FieldUpdatedAt, opts...).ToFunc()
}
// BySubmittedAt orders the results by the submitted_at field.
func BySubmittedAt(opts ...sql.OrderTermOption) OrderOption {
return sql.OrderByField(FieldSubmittedAt, opts...).ToFunc()
}
// ByStartedAt orders the results by the started_at field.
func ByStartedAt(opts ...sql.OrderTermOption) OrderOption {
return sql.OrderByField(FieldStartedAt, opts...).ToFunc()
}
// ByFinishedAt orders the results by the finished_at field.
func ByFinishedAt(opts ...sql.OrderTermOption) OrderOption {
return sql.OrderByField(FieldFinishedAt, opts...).ToFunc()
}
// BySettledAt orders the results by the settled_at field.
func BySettledAt(opts ...sql.OrderTermOption) OrderOption {
return sql.OrderByField(FieldSettledAt, opts...).ToFunc()
}
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
+88
View File
@@ -0,0 +1,88 @@
// Code generated by ent, DO NOT EDIT.
package ent
import (
"context"
"entgo.io/ent/dialect/sql"
"entgo.io/ent/dialect/sql/sqlgraph"
"entgo.io/ent/schema/field"
"github.com/Wei-Shaw/sub2api/ent/batchimagejob"
"github.com/Wei-Shaw/sub2api/ent/predicate"
)
// BatchImageJobDelete is the builder for deleting a BatchImageJob entity.
type BatchImageJobDelete struct {
config
hooks []Hook
mutation *BatchImageJobMutation
}
// Where appends a list predicates to the BatchImageJobDelete builder.
func (_d *BatchImageJobDelete) Where(ps ...predicate.BatchImageJob) *BatchImageJobDelete {
_d.mutation.Where(ps...)
return _d
}
// Exec executes the deletion query and returns how many vertices were deleted.
func (_d *BatchImageJobDelete) Exec(ctx context.Context) (int, error) {
return withHooks(ctx, _d.sqlExec, _d.mutation, _d.hooks)
}
// ExecX is like Exec, but panics if an error occurs.
func (_d *BatchImageJobDelete) ExecX(ctx context.Context) int {
n, err := _d.Exec(ctx)
if err != nil {
panic(err)
}
return n
}
func (_d *BatchImageJobDelete) sqlExec(ctx context.Context) (int, error) {
_spec := sqlgraph.NewDeleteSpec(batchimagejob.Table, sqlgraph.NewFieldSpec(batchimagejob.FieldID, field.TypeInt64))
if ps := _d.mutation.predicates; len(ps) > 0 {
_spec.Predicate = func(selector *sql.Selector) {
for i := range ps {
ps[i](selector)
}
}
}
affected, err := sqlgraph.DeleteNodes(ctx, _d.driver, _spec)
if err != nil && sqlgraph.IsConstraintError(err) {
err = &ConstraintError{msg: err.Error(), wrap: err}
}
_d.mutation.done = true
return affected, err
}
// BatchImageJobDeleteOne is the builder for deleting a single BatchImageJob entity.
type BatchImageJobDeleteOne struct {
_d *BatchImageJobDelete
}
// Where appends a list predicates to the BatchImageJobDelete builder.
func (_d *BatchImageJobDeleteOne) Where(ps ...predicate.BatchImageJob) *BatchImageJobDeleteOne {
_d._d.mutation.Where(ps...)
return _d
}
// Exec executes the deletion query.
func (_d *BatchImageJobDeleteOne) Exec(ctx context.Context) error {
n, err := _d._d.Exec(ctx)
switch {
case err != nil:
return err
case n == 0:
return &NotFoundError{batchimagejob.Label}
default:
return nil
}
}
// ExecX is like Exec, but panics if an error occurs.
func (_d *BatchImageJobDeleteOne) ExecX(ctx context.Context) {
if err := _d.Exec(ctx); err != nil {
panic(err)
}
}
+564
View File
@@ -0,0 +1,564 @@
// Code generated by ent, DO NOT EDIT.
package ent
import (
"context"
"fmt"
"math"
"entgo.io/ent"
"entgo.io/ent/dialect"
"entgo.io/ent/dialect/sql"
"entgo.io/ent/dialect/sql/sqlgraph"
"entgo.io/ent/schema/field"
"github.com/Wei-Shaw/sub2api/ent/batchimagejob"
"github.com/Wei-Shaw/sub2api/ent/predicate"
)
// BatchImageJobQuery is the builder for querying BatchImageJob entities.
type BatchImageJobQuery struct {
config
ctx *QueryContext
order []batchimagejob.OrderOption
inters []Interceptor
predicates []predicate.BatchImageJob
modifiers []func(*sql.Selector)
// intermediate query (i.e. traversal path).
sql *sql.Selector
path func(context.Context) (*sql.Selector, error)
}
// Where adds a new predicate for the BatchImageJobQuery builder.
func (_q *BatchImageJobQuery) Where(ps ...predicate.BatchImageJob) *BatchImageJobQuery {
_q.predicates = append(_q.predicates, ps...)
return _q
}
// Limit the number of records to be returned by this query.
func (_q *BatchImageJobQuery) Limit(limit int) *BatchImageJobQuery {
_q.ctx.Limit = &limit
return _q
}
// Offset to start from.
func (_q *BatchImageJobQuery) Offset(offset int) *BatchImageJobQuery {
_q.ctx.Offset = &offset
return _q
}
// Unique configures the query builder to filter duplicate records on query.
// By default, unique is set to true, and can be disabled using this method.
func (_q *BatchImageJobQuery) Unique(unique bool) *BatchImageJobQuery {
_q.ctx.Unique = &unique
return _q
}
// Order specifies how the records should be ordered.
func (_q *BatchImageJobQuery) Order(o ...batchimagejob.OrderOption) *BatchImageJobQuery {
_q.order = append(_q.order, o...)
return _q
}
// First returns the first BatchImageJob entity from the query.
// Returns a *NotFoundError when no BatchImageJob was found.
func (_q *BatchImageJobQuery) First(ctx context.Context) (*BatchImageJob, error) {
nodes, err := _q.Limit(1).All(setContextOp(ctx, _q.ctx, ent.OpQueryFirst))
if err != nil {
return nil, err
}
if len(nodes) == 0 {
return nil, &NotFoundError{batchimagejob.Label}
}
return nodes[0], nil
}
// FirstX is like First, but panics if an error occurs.
func (_q *BatchImageJobQuery) FirstX(ctx context.Context) *BatchImageJob {
node, err := _q.First(ctx)
if err != nil && !IsNotFound(err) {
panic(err)
}
return node
}
// FirstID returns the first BatchImageJob ID from the query.
// Returns a *NotFoundError when no BatchImageJob ID was found.
func (_q *BatchImageJobQuery) FirstID(ctx context.Context) (id int64, err error) {
var ids []int64
if ids, err = _q.Limit(1).IDs(setContextOp(ctx, _q.ctx, ent.OpQueryFirstID)); err != nil {
return
}
if len(ids) == 0 {
err = &NotFoundError{batchimagejob.Label}
return
}
return ids[0], nil
}
// FirstIDX is like FirstID, but panics if an error occurs.
func (_q *BatchImageJobQuery) FirstIDX(ctx context.Context) int64 {
id, err := _q.FirstID(ctx)
if err != nil && !IsNotFound(err) {
panic(err)
}
return id
}
// Only returns a single BatchImageJob entity found by the query, ensuring it only returns one.
// Returns a *NotSingularError when more than one BatchImageJob entity is found.
// Returns a *NotFoundError when no BatchImageJob entities are found.
func (_q *BatchImageJobQuery) Only(ctx context.Context) (*BatchImageJob, error) {
nodes, err := _q.Limit(2).All(setContextOp(ctx, _q.ctx, ent.OpQueryOnly))
if err != nil {
return nil, err
}
switch len(nodes) {
case 1:
return nodes[0], nil
case 0:
return nil, &NotFoundError{batchimagejob.Label}
default:
return nil, &NotSingularError{batchimagejob.Label}
}
}
// OnlyX is like Only, but panics if an error occurs.
func (_q *BatchImageJobQuery) OnlyX(ctx context.Context) *BatchImageJob {
node, err := _q.Only(ctx)
if err != nil {
panic(err)
}
return node
}
// OnlyID is like Only, but returns the only BatchImageJob ID in the query.
// Returns a *NotSingularError when more than one BatchImageJob ID is found.
// Returns a *NotFoundError when no entities are found.
func (_q *BatchImageJobQuery) OnlyID(ctx context.Context) (id int64, err error) {
var ids []int64
if ids, err = _q.Limit(2).IDs(setContextOp(ctx, _q.ctx, ent.OpQueryOnlyID)); err != nil {
return
}
switch len(ids) {
case 1:
id = ids[0]
case 0:
err = &NotFoundError{batchimagejob.Label}
default:
err = &NotSingularError{batchimagejob.Label}
}
return
}
// OnlyIDX is like OnlyID, but panics if an error occurs.
func (_q *BatchImageJobQuery) OnlyIDX(ctx context.Context) int64 {
id, err := _q.OnlyID(ctx)
if err != nil {
panic(err)
}
return id
}
// All executes the query and returns a list of BatchImageJobs.
func (_q *BatchImageJobQuery) All(ctx context.Context) ([]*BatchImageJob, error) {
ctx = setContextOp(ctx, _q.ctx, ent.OpQueryAll)
if err := _q.prepareQuery(ctx); err != nil {
return nil, err
}
qr := querierAll[[]*BatchImageJob, *BatchImageJobQuery]()
return withInterceptors[[]*BatchImageJob](ctx, _q, qr, _q.inters)
}
// AllX is like All, but panics if an error occurs.
func (_q *BatchImageJobQuery) AllX(ctx context.Context) []*BatchImageJob {
nodes, err := _q.All(ctx)
if err != nil {
panic(err)
}
return nodes
}
// IDs executes the query and returns a list of BatchImageJob IDs.
func (_q *BatchImageJobQuery) IDs(ctx context.Context) (ids []int64, err error) {
if _q.ctx.Unique == nil && _q.path != nil {
_q.Unique(true)
}
ctx = setContextOp(ctx, _q.ctx, ent.OpQueryIDs)
if err = _q.Select(batchimagejob.FieldID).Scan(ctx, &ids); err != nil {
return nil, err
}
return ids, nil
}
// IDsX is like IDs, but panics if an error occurs.
func (_q *BatchImageJobQuery) IDsX(ctx context.Context) []int64 {
ids, err := _q.IDs(ctx)
if err != nil {
panic(err)
}
return ids
}
// Count returns the count of the given query.
func (_q *BatchImageJobQuery) Count(ctx context.Context) (int, error) {
ctx = setContextOp(ctx, _q.ctx, ent.OpQueryCount)
if err := _q.prepareQuery(ctx); err != nil {
return 0, err
}
return withInterceptors[int](ctx, _q, querierCount[*BatchImageJobQuery](), _q.inters)
}
// CountX is like Count, but panics if an error occurs.
func (_q *BatchImageJobQuery) CountX(ctx context.Context) int {
count, err := _q.Count(ctx)
if err != nil {
panic(err)
}
return count
}
// Exist returns true if the query has elements in the graph.
func (_q *BatchImageJobQuery) Exist(ctx context.Context) (bool, error) {
ctx = setContextOp(ctx, _q.ctx, ent.OpQueryExist)
switch _, err := _q.FirstID(ctx); {
case IsNotFound(err):
return false, nil
case err != nil:
return false, fmt.Errorf("ent: check existence: %w", err)
default:
return true, nil
}
}
// ExistX is like Exist, but panics if an error occurs.
func (_q *BatchImageJobQuery) ExistX(ctx context.Context) bool {
exist, err := _q.Exist(ctx)
if err != nil {
panic(err)
}
return exist
}
// Clone returns a duplicate of the BatchImageJobQuery builder, including all associated steps. It can be
// used to prepare common query builders and use them differently after the clone is made.
func (_q *BatchImageJobQuery) Clone() *BatchImageJobQuery {
if _q == nil {
return nil
}
return &BatchImageJobQuery{
config: _q.config,
ctx: _q.ctx.Clone(),
order: append([]batchimagejob.OrderOption{}, _q.order...),
inters: append([]Interceptor{}, _q.inters...),
predicates: append([]predicate.BatchImageJob{}, _q.predicates...),
// clone intermediate query.
sql: _q.sql.Clone(),
path: _q.path,
}
}
// GroupBy is used to group vertices by one or more fields/columns.
// It is often used with aggregate functions, like: count, max, mean, min, sum.
//
// Example:
//
// var v []struct {
// BatchID string `json:"batch_id,omitempty"`
// Count int `json:"count,omitempty"`
// }
//
// client.BatchImageJob.Query().
// GroupBy(batchimagejob.FieldBatchID).
// Aggregate(ent.Count()).
// Scan(ctx, &v)
func (_q *BatchImageJobQuery) GroupBy(field string, fields ...string) *BatchImageJobGroupBy {
_q.ctx.Fields = append([]string{field}, fields...)
grbuild := &BatchImageJobGroupBy{build: _q}
grbuild.flds = &_q.ctx.Fields
grbuild.label = batchimagejob.Label
grbuild.scan = grbuild.Scan
return grbuild
}
// Select allows the selection one or more fields/columns for the given query,
// instead of selecting all fields in the entity.
//
// Example:
//
// var v []struct {
// BatchID string `json:"batch_id,omitempty"`
// }
//
// client.BatchImageJob.Query().
// Select(batchimagejob.FieldBatchID).
// Scan(ctx, &v)
func (_q *BatchImageJobQuery) Select(fields ...string) *BatchImageJobSelect {
_q.ctx.Fields = append(_q.ctx.Fields, fields...)
sbuild := &BatchImageJobSelect{BatchImageJobQuery: _q}
sbuild.label = batchimagejob.Label
sbuild.flds, sbuild.scan = &_q.ctx.Fields, sbuild.Scan
return sbuild
}
// Aggregate returns a BatchImageJobSelect configured with the given aggregations.
func (_q *BatchImageJobQuery) Aggregate(fns ...AggregateFunc) *BatchImageJobSelect {
return _q.Select().Aggregate(fns...)
}
func (_q *BatchImageJobQuery) prepareQuery(ctx context.Context) error {
for _, inter := range _q.inters {
if inter == nil {
return fmt.Errorf("ent: uninitialized interceptor (forgotten import ent/runtime?)")
}
if trv, ok := inter.(Traverser); ok {
if err := trv.Traverse(ctx, _q); err != nil {
return err
}
}
}
for _, f := range _q.ctx.Fields {
if !batchimagejob.ValidColumn(f) {
return &ValidationError{Name: f, err: fmt.Errorf("ent: invalid field %q for query", f)}
}
}
if _q.path != nil {
prev, err := _q.path(ctx)
if err != nil {
return err
}
_q.sql = prev
}
return nil
}
func (_q *BatchImageJobQuery) sqlAll(ctx context.Context, hooks ...queryHook) ([]*BatchImageJob, error) {
var (
nodes = []*BatchImageJob{}
_spec = _q.querySpec()
)
_spec.ScanValues = func(columns []string) ([]any, error) {
return (*BatchImageJob).scanValues(nil, columns)
}
_spec.Assign = func(columns []string, values []any) error {
node := &BatchImageJob{config: _q.config}
nodes = append(nodes, node)
return node.assignValues(columns, values)
}
if len(_q.modifiers) > 0 {
_spec.Modifiers = _q.modifiers
}
for i := range hooks {
hooks[i](ctx, _spec)
}
if err := sqlgraph.QueryNodes(ctx, _q.driver, _spec); err != nil {
return nil, err
}
if len(nodes) == 0 {
return nodes, nil
}
return nodes, nil
}
func (_q *BatchImageJobQuery) sqlCount(ctx context.Context) (int, error) {
_spec := _q.querySpec()
if len(_q.modifiers) > 0 {
_spec.Modifiers = _q.modifiers
}
_spec.Node.Columns = _q.ctx.Fields
if len(_q.ctx.Fields) > 0 {
_spec.Unique = _q.ctx.Unique != nil && *_q.ctx.Unique
}
return sqlgraph.CountNodes(ctx, _q.driver, _spec)
}
func (_q *BatchImageJobQuery) querySpec() *sqlgraph.QuerySpec {
_spec := sqlgraph.NewQuerySpec(batchimagejob.Table, batchimagejob.Columns, sqlgraph.NewFieldSpec(batchimagejob.FieldID, field.TypeInt64))
_spec.From = _q.sql
if unique := _q.ctx.Unique; unique != nil {
_spec.Unique = *unique
} else if _q.path != nil {
_spec.Unique = true
}
if fields := _q.ctx.Fields; len(fields) > 0 {
_spec.Node.Columns = make([]string, 0, len(fields))
_spec.Node.Columns = append(_spec.Node.Columns, batchimagejob.FieldID)
for i := range fields {
if fields[i] != batchimagejob.FieldID {
_spec.Node.Columns = append(_spec.Node.Columns, fields[i])
}
}
}
if ps := _q.predicates; len(ps) > 0 {
_spec.Predicate = func(selector *sql.Selector) {
for i := range ps {
ps[i](selector)
}
}
}
if limit := _q.ctx.Limit; limit != nil {
_spec.Limit = *limit
}
if offset := _q.ctx.Offset; offset != nil {
_spec.Offset = *offset
}
if ps := _q.order; len(ps) > 0 {
_spec.Order = func(selector *sql.Selector) {
for i := range ps {
ps[i](selector)
}
}
}
return _spec
}
func (_q *BatchImageJobQuery) sqlQuery(ctx context.Context) *sql.Selector {
builder := sql.Dialect(_q.driver.Dialect())
t1 := builder.Table(batchimagejob.Table)
columns := _q.ctx.Fields
if len(columns) == 0 {
columns = batchimagejob.Columns
}
selector := builder.Select(t1.Columns(columns...)...).From(t1)
if _q.sql != nil {
selector = _q.sql
selector.Select(selector.Columns(columns...)...)
}
if _q.ctx.Unique != nil && *_q.ctx.Unique {
selector.Distinct()
}
for _, m := range _q.modifiers {
m(selector)
}
for _, p := range _q.predicates {
p(selector)
}
for _, p := range _q.order {
p(selector)
}
if offset := _q.ctx.Offset; offset != nil {
// limit is mandatory for offset clause. We start
// with default value, and override it below if needed.
selector.Offset(*offset).Limit(math.MaxInt32)
}
if limit := _q.ctx.Limit; limit != nil {
selector.Limit(*limit)
}
return selector
}
// ForUpdate locks the selected rows against concurrent updates, and prevent them from being
// updated, deleted or "selected ... for update" by other sessions, until the transaction is
// either committed or rolled-back.
func (_q *BatchImageJobQuery) ForUpdate(opts ...sql.LockOption) *BatchImageJobQuery {
if _q.driver.Dialect() == dialect.Postgres {
_q.Unique(false)
}
_q.modifiers = append(_q.modifiers, func(s *sql.Selector) {
s.ForUpdate(opts...)
})
return _q
}
// ForShare behaves similarly to ForUpdate, except that it acquires a shared mode lock
// on any rows that are read. Other sessions can read the rows, but cannot modify them
// until your transaction commits.
func (_q *BatchImageJobQuery) ForShare(opts ...sql.LockOption) *BatchImageJobQuery {
if _q.driver.Dialect() == dialect.Postgres {
_q.Unique(false)
}
_q.modifiers = append(_q.modifiers, func(s *sql.Selector) {
s.ForShare(opts...)
})
return _q
}
// BatchImageJobGroupBy is the group-by builder for BatchImageJob entities.
type BatchImageJobGroupBy struct {
selector
build *BatchImageJobQuery
}
// Aggregate adds the given aggregation functions to the group-by query.
func (_g *BatchImageJobGroupBy) Aggregate(fns ...AggregateFunc) *BatchImageJobGroupBy {
_g.fns = append(_g.fns, fns...)
return _g
}
// Scan applies the selector query and scans the result into the given value.
func (_g *BatchImageJobGroupBy) Scan(ctx context.Context, v any) error {
ctx = setContextOp(ctx, _g.build.ctx, ent.OpQueryGroupBy)
if err := _g.build.prepareQuery(ctx); err != nil {
return err
}
return scanWithInterceptors[*BatchImageJobQuery, *BatchImageJobGroupBy](ctx, _g.build, _g, _g.build.inters, v)
}
func (_g *BatchImageJobGroupBy) sqlScan(ctx context.Context, root *BatchImageJobQuery, v any) error {
selector := root.sqlQuery(ctx).Select()
aggregation := make([]string, 0, len(_g.fns))
for _, fn := range _g.fns {
aggregation = append(aggregation, fn(selector))
}
if len(selector.SelectedColumns()) == 0 {
columns := make([]string, 0, len(*_g.flds)+len(_g.fns))
for _, f := range *_g.flds {
columns = append(columns, selector.C(f))
}
columns = append(columns, aggregation...)
selector.Select(columns...)
}
selector.GroupBy(selector.Columns(*_g.flds...)...)
if err := selector.Err(); err != nil {
return err
}
rows := &sql.Rows{}
query, args := selector.Query()
if err := _g.build.driver.Query(ctx, query, args, rows); err != nil {
return err
}
defer rows.Close()
return sql.ScanSlice(rows, v)
}
// BatchImageJobSelect is the builder for selecting fields of BatchImageJob entities.
type BatchImageJobSelect struct {
*BatchImageJobQuery
selector
}
// Aggregate adds the given aggregation functions to the selector query.
func (_s *BatchImageJobSelect) Aggregate(fns ...AggregateFunc) *BatchImageJobSelect {
_s.fns = append(_s.fns, fns...)
return _s
}
// Scan applies the selector query and scans the result into the given value.
func (_s *BatchImageJobSelect) Scan(ctx context.Context, v any) error {
ctx = setContextOp(ctx, _s.ctx, ent.OpQuerySelect)
if err := _s.prepareQuery(ctx); err != nil {
return err
}
return scanWithInterceptors[*BatchImageJobQuery, *BatchImageJobSelect](ctx, _s.BatchImageJobQuery, _s, _s.inters, v)
}
func (_s *BatchImageJobSelect) sqlScan(ctx context.Context, root *BatchImageJobQuery, v any) error {
selector := root.sqlQuery(ctx)
aggregation := make([]string, 0, len(_s.fns))
for _, fn := range _s.fns {
aggregation = append(aggregation, fn(selector))
}
switch n := len(*_s.selector.flds); {
case n == 0 && len(aggregation) > 0:
selector.Select(aggregation...)
case n != 0 && len(aggregation) > 0:
selector.AppendSelect(aggregation...)
}
rows := &sql.Rows{}
query, args := selector.Query()
if err := _s.driver.Query(ctx, query, args, rows); err != nil {
return err
}
defer rows.Close()
return sql.ScanSlice(rows, v)
}
File diff suppressed because it is too large Load Diff
+457 -32
View File
@@ -22,6 +22,9 @@ import (
"github.com/Wei-Shaw/sub2api/ent/apikey"
"github.com/Wei-Shaw/sub2api/ent/authidentity"
"github.com/Wei-Shaw/sub2api/ent/authidentitychannel"
"github.com/Wei-Shaw/sub2api/ent/batchimageevent"
"github.com/Wei-Shaw/sub2api/ent/batchimageitem"
"github.com/Wei-Shaw/sub2api/ent/batchimagejob"
"github.com/Wei-Shaw/sub2api/ent/channelmonitor"
"github.com/Wei-Shaw/sub2api/ent/channelmonitordailyrollup"
"github.com/Wei-Shaw/sub2api/ent/channelmonitorhistory"
@@ -73,6 +76,12 @@ type Client struct {
AuthIdentity *AuthIdentityClient
// AuthIdentityChannel is the client for interacting with the AuthIdentityChannel builders.
AuthIdentityChannel *AuthIdentityChannelClient
// BatchImageEvent is the client for interacting with the BatchImageEvent builders.
BatchImageEvent *BatchImageEventClient
// BatchImageItem is the client for interacting with the BatchImageItem builders.
BatchImageItem *BatchImageItemClient
// BatchImageJob is the client for interacting with the BatchImageJob builders.
BatchImageJob *BatchImageJobClient
// ChannelMonitor is the client for interacting with the ChannelMonitor builders.
ChannelMonitor *ChannelMonitorClient
// ChannelMonitorDailyRollup is the client for interacting with the ChannelMonitorDailyRollup builders.
@@ -147,6 +156,9 @@ func (c *Client) init() {
c.AnnouncementRead = NewAnnouncementReadClient(c.config)
c.AuthIdentity = NewAuthIdentityClient(c.config)
c.AuthIdentityChannel = NewAuthIdentityChannelClient(c.config)
c.BatchImageEvent = NewBatchImageEventClient(c.config)
c.BatchImageItem = NewBatchImageItemClient(c.config)
c.BatchImageJob = NewBatchImageJobClient(c.config)
c.ChannelMonitor = NewChannelMonitorClient(c.config)
c.ChannelMonitorDailyRollup = NewChannelMonitorDailyRollupClient(c.config)
c.ChannelMonitorHistory = NewChannelMonitorHistoryClient(c.config)
@@ -274,6 +286,9 @@ func (c *Client) Tx(ctx context.Context) (*Tx, error) {
AnnouncementRead: NewAnnouncementReadClient(cfg),
AuthIdentity: NewAuthIdentityClient(cfg),
AuthIdentityChannel: NewAuthIdentityChannelClient(cfg),
BatchImageEvent: NewBatchImageEventClient(cfg),
BatchImageItem: NewBatchImageItemClient(cfg),
BatchImageJob: NewBatchImageJobClient(cfg),
ChannelMonitor: NewChannelMonitorClient(cfg),
ChannelMonitorDailyRollup: NewChannelMonitorDailyRollupClient(cfg),
ChannelMonitorHistory: NewChannelMonitorHistoryClient(cfg),
@@ -328,6 +343,9 @@ func (c *Client) BeginTx(ctx context.Context, opts *sql.TxOptions) (*Tx, error)
AnnouncementRead: NewAnnouncementReadClient(cfg),
AuthIdentity: NewAuthIdentityClient(cfg),
AuthIdentityChannel: NewAuthIdentityChannelClient(cfg),
BatchImageEvent: NewBatchImageEventClient(cfg),
BatchImageItem: NewBatchImageItemClient(cfg),
BatchImageJob: NewBatchImageJobClient(cfg),
ChannelMonitor: NewChannelMonitorClient(cfg),
ChannelMonitorDailyRollup: NewChannelMonitorDailyRollupClient(cfg),
ChannelMonitorHistory: NewChannelMonitorHistoryClient(cfg),
@@ -386,14 +404,15 @@ func (c *Client) Close() error {
func (c *Client) Use(hooks ...Hook) {
for _, n := range []interface{ Use(...Hook) }{
c.APIKey, c.Account, c.AccountGroup, c.Announcement, c.AnnouncementRead,
c.AuthIdentity, c.AuthIdentityChannel, c.ChannelMonitor,
c.ChannelMonitorDailyRollup, c.ChannelMonitorHistory,
c.ChannelMonitorRequestTemplate, c.ErrorPassthroughRule, c.Group,
c.IdempotencyRecord, c.IdentityAdoptionDecision, c.PaymentAuditLog,
c.PaymentOrder, c.PaymentProviderInstance, c.PendingAuthSession, c.PromoCode,
c.PromoCodeUsage, c.Proxy, c.RedeemCode, c.SecuritySecret, c.Setting,
c.SubscriptionPlan, c.TLSFingerprintProfile, c.UsageCleanupTask, c.UsageLog,
c.User, c.UserAllowedGroup, c.UserAttributeDefinition, c.UserAttributeValue,
c.AuthIdentity, c.AuthIdentityChannel, c.BatchImageEvent, c.BatchImageItem,
c.BatchImageJob, c.ChannelMonitor, c.ChannelMonitorDailyRollup,
c.ChannelMonitorHistory, c.ChannelMonitorRequestTemplate,
c.ErrorPassthroughRule, c.Group, c.IdempotencyRecord,
c.IdentityAdoptionDecision, c.PaymentAuditLog, c.PaymentOrder,
c.PaymentProviderInstance, c.PendingAuthSession, c.PromoCode, c.PromoCodeUsage,
c.Proxy, c.RedeemCode, c.SecuritySecret, c.Setting, c.SubscriptionPlan,
c.TLSFingerprintProfile, c.UsageCleanupTask, c.UsageLog, c.User,
c.UserAllowedGroup, c.UserAttributeDefinition, c.UserAttributeValue,
c.UserPlatformQuota, c.UserSubscription,
} {
n.Use(hooks...)
@@ -405,14 +424,15 @@ func (c *Client) Use(hooks ...Hook) {
func (c *Client) Intercept(interceptors ...Interceptor) {
for _, n := range []interface{ Intercept(...Interceptor) }{
c.APIKey, c.Account, c.AccountGroup, c.Announcement, c.AnnouncementRead,
c.AuthIdentity, c.AuthIdentityChannel, c.ChannelMonitor,
c.ChannelMonitorDailyRollup, c.ChannelMonitorHistory,
c.ChannelMonitorRequestTemplate, c.ErrorPassthroughRule, c.Group,
c.IdempotencyRecord, c.IdentityAdoptionDecision, c.PaymentAuditLog,
c.PaymentOrder, c.PaymentProviderInstance, c.PendingAuthSession, c.PromoCode,
c.PromoCodeUsage, c.Proxy, c.RedeemCode, c.SecuritySecret, c.Setting,
c.SubscriptionPlan, c.TLSFingerprintProfile, c.UsageCleanupTask, c.UsageLog,
c.User, c.UserAllowedGroup, c.UserAttributeDefinition, c.UserAttributeValue,
c.AuthIdentity, c.AuthIdentityChannel, c.BatchImageEvent, c.BatchImageItem,
c.BatchImageJob, c.ChannelMonitor, c.ChannelMonitorDailyRollup,
c.ChannelMonitorHistory, c.ChannelMonitorRequestTemplate,
c.ErrorPassthroughRule, c.Group, c.IdempotencyRecord,
c.IdentityAdoptionDecision, c.PaymentAuditLog, c.PaymentOrder,
c.PaymentProviderInstance, c.PendingAuthSession, c.PromoCode, c.PromoCodeUsage,
c.Proxy, c.RedeemCode, c.SecuritySecret, c.Setting, c.SubscriptionPlan,
c.TLSFingerprintProfile, c.UsageCleanupTask, c.UsageLog, c.User,
c.UserAllowedGroup, c.UserAttributeDefinition, c.UserAttributeValue,
c.UserPlatformQuota, c.UserSubscription,
} {
n.Intercept(interceptors...)
@@ -436,6 +456,12 @@ func (c *Client) Mutate(ctx context.Context, m Mutation) (Value, error) {
return c.AuthIdentity.mutate(ctx, m)
case *AuthIdentityChannelMutation:
return c.AuthIdentityChannel.mutate(ctx, m)
case *BatchImageEventMutation:
return c.BatchImageEvent.mutate(ctx, m)
case *BatchImageItemMutation:
return c.BatchImageItem.mutate(ctx, m)
case *BatchImageJobMutation:
return c.BatchImageJob.mutate(ctx, m)
case *ChannelMonitorMutation:
return c.ChannelMonitor.mutate(ctx, m)
case *ChannelMonitorDailyRollupMutation:
@@ -1671,6 +1697,405 @@ func (c *AuthIdentityChannelClient) mutate(ctx context.Context, m *AuthIdentityC
}
}
// BatchImageEventClient is a client for the BatchImageEvent schema.
type BatchImageEventClient struct {
config
}
// NewBatchImageEventClient returns a client for the BatchImageEvent from the given config.
func NewBatchImageEventClient(c config) *BatchImageEventClient {
return &BatchImageEventClient{config: c}
}
// Use adds a list of mutation hooks to the hooks stack.
// A call to `Use(f, g, h)` equals to `batchimageevent.Hooks(f(g(h())))`.
func (c *BatchImageEventClient) Use(hooks ...Hook) {
c.hooks.BatchImageEvent = append(c.hooks.BatchImageEvent, hooks...)
}
// Intercept adds a list of query interceptors to the interceptors stack.
// A call to `Intercept(f, g, h)` equals to `batchimageevent.Intercept(f(g(h())))`.
func (c *BatchImageEventClient) Intercept(interceptors ...Interceptor) {
c.inters.BatchImageEvent = append(c.inters.BatchImageEvent, interceptors...)
}
// Create returns a builder for creating a BatchImageEvent entity.
func (c *BatchImageEventClient) Create() *BatchImageEventCreate {
mutation := newBatchImageEventMutation(c.config, OpCreate)
return &BatchImageEventCreate{config: c.config, hooks: c.Hooks(), mutation: mutation}
}
// CreateBulk returns a builder for creating a bulk of BatchImageEvent entities.
func (c *BatchImageEventClient) CreateBulk(builders ...*BatchImageEventCreate) *BatchImageEventCreateBulk {
return &BatchImageEventCreateBulk{config: c.config, builders: builders}
}
// MapCreateBulk creates a bulk creation builder from the given slice. For each item in the slice, the function creates
// a builder and applies setFunc on it.
func (c *BatchImageEventClient) MapCreateBulk(slice any, setFunc func(*BatchImageEventCreate, int)) *BatchImageEventCreateBulk {
rv := reflect.ValueOf(slice)
if rv.Kind() != reflect.Slice {
return &BatchImageEventCreateBulk{err: fmt.Errorf("calling to BatchImageEventClient.MapCreateBulk with wrong type %T, need slice", slice)}
}
builders := make([]*BatchImageEventCreate, rv.Len())
for i := 0; i < rv.Len(); i++ {
builders[i] = c.Create()
setFunc(builders[i], i)
}
return &BatchImageEventCreateBulk{config: c.config, builders: builders}
}
// Update returns an update builder for BatchImageEvent.
func (c *BatchImageEventClient) Update() *BatchImageEventUpdate {
mutation := newBatchImageEventMutation(c.config, OpUpdate)
return &BatchImageEventUpdate{config: c.config, hooks: c.Hooks(), mutation: mutation}
}
// UpdateOne returns an update builder for the given entity.
func (c *BatchImageEventClient) UpdateOne(_m *BatchImageEvent) *BatchImageEventUpdateOne {
mutation := newBatchImageEventMutation(c.config, OpUpdateOne, withBatchImageEvent(_m))
return &BatchImageEventUpdateOne{config: c.config, hooks: c.Hooks(), mutation: mutation}
}
// UpdateOneID returns an update builder for the given id.
func (c *BatchImageEventClient) UpdateOneID(id int64) *BatchImageEventUpdateOne {
mutation := newBatchImageEventMutation(c.config, OpUpdateOne, withBatchImageEventID(id))
return &BatchImageEventUpdateOne{config: c.config, hooks: c.Hooks(), mutation: mutation}
}
// Delete returns a delete builder for BatchImageEvent.
func (c *BatchImageEventClient) Delete() *BatchImageEventDelete {
mutation := newBatchImageEventMutation(c.config, OpDelete)
return &BatchImageEventDelete{config: c.config, hooks: c.Hooks(), mutation: mutation}
}
// DeleteOne returns a builder for deleting the given entity.
func (c *BatchImageEventClient) DeleteOne(_m *BatchImageEvent) *BatchImageEventDeleteOne {
return c.DeleteOneID(_m.ID)
}
// DeleteOneID returns a builder for deleting the given entity by its id.
func (c *BatchImageEventClient) DeleteOneID(id int64) *BatchImageEventDeleteOne {
builder := c.Delete().Where(batchimageevent.ID(id))
builder.mutation.id = &id
builder.mutation.op = OpDeleteOne
return &BatchImageEventDeleteOne{builder}
}
// Query returns a query builder for BatchImageEvent.
func (c *BatchImageEventClient) Query() *BatchImageEventQuery {
return &BatchImageEventQuery{
config: c.config,
ctx: &QueryContext{Type: TypeBatchImageEvent},
inters: c.Interceptors(),
}
}
// Get returns a BatchImageEvent entity by its id.
func (c *BatchImageEventClient) Get(ctx context.Context, id int64) (*BatchImageEvent, error) {
return c.Query().Where(batchimageevent.ID(id)).Only(ctx)
}
// GetX is like Get, but panics if an error occurs.
func (c *BatchImageEventClient) GetX(ctx context.Context, id int64) *BatchImageEvent {
obj, err := c.Get(ctx, id)
if err != nil {
panic(err)
}
return obj
}
// Hooks returns the client hooks.
func (c *BatchImageEventClient) Hooks() []Hook {
return c.hooks.BatchImageEvent
}
// Interceptors returns the client interceptors.
func (c *BatchImageEventClient) Interceptors() []Interceptor {
return c.inters.BatchImageEvent
}
func (c *BatchImageEventClient) mutate(ctx context.Context, m *BatchImageEventMutation) (Value, error) {
switch m.Op() {
case OpCreate:
return (&BatchImageEventCreate{config: c.config, hooks: c.Hooks(), mutation: m}).Save(ctx)
case OpUpdate:
return (&BatchImageEventUpdate{config: c.config, hooks: c.Hooks(), mutation: m}).Save(ctx)
case OpUpdateOne:
return (&BatchImageEventUpdateOne{config: c.config, hooks: c.Hooks(), mutation: m}).Save(ctx)
case OpDelete, OpDeleteOne:
return (&BatchImageEventDelete{config: c.config, hooks: c.Hooks(), mutation: m}).Exec(ctx)
default:
return nil, fmt.Errorf("ent: unknown BatchImageEvent mutation op: %q", m.Op())
}
}
// BatchImageItemClient is a client for the BatchImageItem schema.
type BatchImageItemClient struct {
config
}
// NewBatchImageItemClient returns a client for the BatchImageItem from the given config.
func NewBatchImageItemClient(c config) *BatchImageItemClient {
return &BatchImageItemClient{config: c}
}
// Use adds a list of mutation hooks to the hooks stack.
// A call to `Use(f, g, h)` equals to `batchimageitem.Hooks(f(g(h())))`.
func (c *BatchImageItemClient) Use(hooks ...Hook) {
c.hooks.BatchImageItem = append(c.hooks.BatchImageItem, hooks...)
}
// Intercept adds a list of query interceptors to the interceptors stack.
// A call to `Intercept(f, g, h)` equals to `batchimageitem.Intercept(f(g(h())))`.
func (c *BatchImageItemClient) Intercept(interceptors ...Interceptor) {
c.inters.BatchImageItem = append(c.inters.BatchImageItem, interceptors...)
}
// Create returns a builder for creating a BatchImageItem entity.
func (c *BatchImageItemClient) Create() *BatchImageItemCreate {
mutation := newBatchImageItemMutation(c.config, OpCreate)
return &BatchImageItemCreate{config: c.config, hooks: c.Hooks(), mutation: mutation}
}
// CreateBulk returns a builder for creating a bulk of BatchImageItem entities.
func (c *BatchImageItemClient) CreateBulk(builders ...*BatchImageItemCreate) *BatchImageItemCreateBulk {
return &BatchImageItemCreateBulk{config: c.config, builders: builders}
}
// MapCreateBulk creates a bulk creation builder from the given slice. For each item in the slice, the function creates
// a builder and applies setFunc on it.
func (c *BatchImageItemClient) MapCreateBulk(slice any, setFunc func(*BatchImageItemCreate, int)) *BatchImageItemCreateBulk {
rv := reflect.ValueOf(slice)
if rv.Kind() != reflect.Slice {
return &BatchImageItemCreateBulk{err: fmt.Errorf("calling to BatchImageItemClient.MapCreateBulk with wrong type %T, need slice", slice)}
}
builders := make([]*BatchImageItemCreate, rv.Len())
for i := 0; i < rv.Len(); i++ {
builders[i] = c.Create()
setFunc(builders[i], i)
}
return &BatchImageItemCreateBulk{config: c.config, builders: builders}
}
// Update returns an update builder for BatchImageItem.
func (c *BatchImageItemClient) Update() *BatchImageItemUpdate {
mutation := newBatchImageItemMutation(c.config, OpUpdate)
return &BatchImageItemUpdate{config: c.config, hooks: c.Hooks(), mutation: mutation}
}
// UpdateOne returns an update builder for the given entity.
func (c *BatchImageItemClient) UpdateOne(_m *BatchImageItem) *BatchImageItemUpdateOne {
mutation := newBatchImageItemMutation(c.config, OpUpdateOne, withBatchImageItem(_m))
return &BatchImageItemUpdateOne{config: c.config, hooks: c.Hooks(), mutation: mutation}
}
// UpdateOneID returns an update builder for the given id.
func (c *BatchImageItemClient) UpdateOneID(id int64) *BatchImageItemUpdateOne {
mutation := newBatchImageItemMutation(c.config, OpUpdateOne, withBatchImageItemID(id))
return &BatchImageItemUpdateOne{config: c.config, hooks: c.Hooks(), mutation: mutation}
}
// Delete returns a delete builder for BatchImageItem.
func (c *BatchImageItemClient) Delete() *BatchImageItemDelete {
mutation := newBatchImageItemMutation(c.config, OpDelete)
return &BatchImageItemDelete{config: c.config, hooks: c.Hooks(), mutation: mutation}
}
// DeleteOne returns a builder for deleting the given entity.
func (c *BatchImageItemClient) DeleteOne(_m *BatchImageItem) *BatchImageItemDeleteOne {
return c.DeleteOneID(_m.ID)
}
// DeleteOneID returns a builder for deleting the given entity by its id.
func (c *BatchImageItemClient) DeleteOneID(id int64) *BatchImageItemDeleteOne {
builder := c.Delete().Where(batchimageitem.ID(id))
builder.mutation.id = &id
builder.mutation.op = OpDeleteOne
return &BatchImageItemDeleteOne{builder}
}
// Query returns a query builder for BatchImageItem.
func (c *BatchImageItemClient) Query() *BatchImageItemQuery {
return &BatchImageItemQuery{
config: c.config,
ctx: &QueryContext{Type: TypeBatchImageItem},
inters: c.Interceptors(),
}
}
// Get returns a BatchImageItem entity by its id.
func (c *BatchImageItemClient) Get(ctx context.Context, id int64) (*BatchImageItem, error) {
return c.Query().Where(batchimageitem.ID(id)).Only(ctx)
}
// GetX is like Get, but panics if an error occurs.
func (c *BatchImageItemClient) GetX(ctx context.Context, id int64) *BatchImageItem {
obj, err := c.Get(ctx, id)
if err != nil {
panic(err)
}
return obj
}
// Hooks returns the client hooks.
func (c *BatchImageItemClient) Hooks() []Hook {
return c.hooks.BatchImageItem
}
// Interceptors returns the client interceptors.
func (c *BatchImageItemClient) Interceptors() []Interceptor {
return c.inters.BatchImageItem
}
func (c *BatchImageItemClient) mutate(ctx context.Context, m *BatchImageItemMutation) (Value, error) {
switch m.Op() {
case OpCreate:
return (&BatchImageItemCreate{config: c.config, hooks: c.Hooks(), mutation: m}).Save(ctx)
case OpUpdate:
return (&BatchImageItemUpdate{config: c.config, hooks: c.Hooks(), mutation: m}).Save(ctx)
case OpUpdateOne:
return (&BatchImageItemUpdateOne{config: c.config, hooks: c.Hooks(), mutation: m}).Save(ctx)
case OpDelete, OpDeleteOne:
return (&BatchImageItemDelete{config: c.config, hooks: c.Hooks(), mutation: m}).Exec(ctx)
default:
return nil, fmt.Errorf("ent: unknown BatchImageItem mutation op: %q", m.Op())
}
}
// BatchImageJobClient is a client for the BatchImageJob schema.
type BatchImageJobClient struct {
config
}
// NewBatchImageJobClient returns a client for the BatchImageJob from the given config.
func NewBatchImageJobClient(c config) *BatchImageJobClient {
return &BatchImageJobClient{config: c}
}
// Use adds a list of mutation hooks to the hooks stack.
// A call to `Use(f, g, h)` equals to `batchimagejob.Hooks(f(g(h())))`.
func (c *BatchImageJobClient) Use(hooks ...Hook) {
c.hooks.BatchImageJob = append(c.hooks.BatchImageJob, hooks...)
}
// Intercept adds a list of query interceptors to the interceptors stack.
// A call to `Intercept(f, g, h)` equals to `batchimagejob.Intercept(f(g(h())))`.
func (c *BatchImageJobClient) Intercept(interceptors ...Interceptor) {
c.inters.BatchImageJob = append(c.inters.BatchImageJob, interceptors...)
}
// Create returns a builder for creating a BatchImageJob entity.
func (c *BatchImageJobClient) Create() *BatchImageJobCreate {
mutation := newBatchImageJobMutation(c.config, OpCreate)
return &BatchImageJobCreate{config: c.config, hooks: c.Hooks(), mutation: mutation}
}
// CreateBulk returns a builder for creating a bulk of BatchImageJob entities.
func (c *BatchImageJobClient) CreateBulk(builders ...*BatchImageJobCreate) *BatchImageJobCreateBulk {
return &BatchImageJobCreateBulk{config: c.config, builders: builders}
}
// MapCreateBulk creates a bulk creation builder from the given slice. For each item in the slice, the function creates
// a builder and applies setFunc on it.
func (c *BatchImageJobClient) MapCreateBulk(slice any, setFunc func(*BatchImageJobCreate, int)) *BatchImageJobCreateBulk {
rv := reflect.ValueOf(slice)
if rv.Kind() != reflect.Slice {
return &BatchImageJobCreateBulk{err: fmt.Errorf("calling to BatchImageJobClient.MapCreateBulk with wrong type %T, need slice", slice)}
}
builders := make([]*BatchImageJobCreate, rv.Len())
for i := 0; i < rv.Len(); i++ {
builders[i] = c.Create()
setFunc(builders[i], i)
}
return &BatchImageJobCreateBulk{config: c.config, builders: builders}
}
// Update returns an update builder for BatchImageJob.
func (c *BatchImageJobClient) Update() *BatchImageJobUpdate {
mutation := newBatchImageJobMutation(c.config, OpUpdate)
return &BatchImageJobUpdate{config: c.config, hooks: c.Hooks(), mutation: mutation}
}
// UpdateOne returns an update builder for the given entity.
func (c *BatchImageJobClient) UpdateOne(_m *BatchImageJob) *BatchImageJobUpdateOne {
mutation := newBatchImageJobMutation(c.config, OpUpdateOne, withBatchImageJob(_m))
return &BatchImageJobUpdateOne{config: c.config, hooks: c.Hooks(), mutation: mutation}
}
// UpdateOneID returns an update builder for the given id.
func (c *BatchImageJobClient) UpdateOneID(id int64) *BatchImageJobUpdateOne {
mutation := newBatchImageJobMutation(c.config, OpUpdateOne, withBatchImageJobID(id))
return &BatchImageJobUpdateOne{config: c.config, hooks: c.Hooks(), mutation: mutation}
}
// Delete returns a delete builder for BatchImageJob.
func (c *BatchImageJobClient) Delete() *BatchImageJobDelete {
mutation := newBatchImageJobMutation(c.config, OpDelete)
return &BatchImageJobDelete{config: c.config, hooks: c.Hooks(), mutation: mutation}
}
// DeleteOne returns a builder for deleting the given entity.
func (c *BatchImageJobClient) DeleteOne(_m *BatchImageJob) *BatchImageJobDeleteOne {
return c.DeleteOneID(_m.ID)
}
// DeleteOneID returns a builder for deleting the given entity by its id.
func (c *BatchImageJobClient) DeleteOneID(id int64) *BatchImageJobDeleteOne {
builder := c.Delete().Where(batchimagejob.ID(id))
builder.mutation.id = &id
builder.mutation.op = OpDeleteOne
return &BatchImageJobDeleteOne{builder}
}
// Query returns a query builder for BatchImageJob.
func (c *BatchImageJobClient) Query() *BatchImageJobQuery {
return &BatchImageJobQuery{
config: c.config,
ctx: &QueryContext{Type: TypeBatchImageJob},
inters: c.Interceptors(),
}
}
// Get returns a BatchImageJob entity by its id.
func (c *BatchImageJobClient) Get(ctx context.Context, id int64) (*BatchImageJob, error) {
return c.Query().Where(batchimagejob.ID(id)).Only(ctx)
}
// GetX is like Get, but panics if an error occurs.
func (c *BatchImageJobClient) GetX(ctx context.Context, id int64) *BatchImageJob {
obj, err := c.Get(ctx, id)
if err != nil {
panic(err)
}
return obj
}
// Hooks returns the client hooks.
func (c *BatchImageJobClient) Hooks() []Hook {
return c.hooks.BatchImageJob
}
// Interceptors returns the client interceptors.
func (c *BatchImageJobClient) Interceptors() []Interceptor {
return c.inters.BatchImageJob
}
func (c *BatchImageJobClient) mutate(ctx context.Context, m *BatchImageJobMutation) (Value, error) {
switch m.Op() {
case OpCreate:
return (&BatchImageJobCreate{config: c.config, hooks: c.Hooks(), mutation: m}).Save(ctx)
case OpUpdate:
return (&BatchImageJobUpdate{config: c.config, hooks: c.Hooks(), mutation: m}).Save(ctx)
case OpUpdateOne:
return (&BatchImageJobUpdateOne{config: c.config, hooks: c.Hooks(), mutation: m}).Save(ctx)
case OpDelete, OpDeleteOne:
return (&BatchImageJobDelete{config: c.config, hooks: c.Hooks(), mutation: m}).Exec(ctx)
default:
return nil, fmt.Errorf("ent: unknown BatchImageJob mutation op: %q", m.Op())
}
}
// ChannelMonitorClient is a client for the ChannelMonitor schema.
type ChannelMonitorClient struct {
config
@@ -6242,25 +6667,25 @@ func (c *UserSubscriptionClient) mutate(ctx context.Context, m *UserSubscription
type (
hooks struct {
APIKey, Account, AccountGroup, Announcement, AnnouncementRead, AuthIdentity,
AuthIdentityChannel, ChannelMonitor, ChannelMonitorDailyRollup,
ChannelMonitorHistory, ChannelMonitorRequestTemplate, ErrorPassthroughRule,
Group, IdempotencyRecord, IdentityAdoptionDecision, PaymentAuditLog,
PaymentOrder, PaymentProviderInstance, PendingAuthSession, PromoCode,
PromoCodeUsage, Proxy, RedeemCode, SecuritySecret, Setting, SubscriptionPlan,
TLSFingerprintProfile, UsageCleanupTask, UsageLog, User, UserAllowedGroup,
UserAttributeDefinition, UserAttributeValue, UserPlatformQuota,
UserSubscription []ent.Hook
AuthIdentityChannel, BatchImageEvent, BatchImageItem, BatchImageJob,
ChannelMonitor, ChannelMonitorDailyRollup, ChannelMonitorHistory,
ChannelMonitorRequestTemplate, ErrorPassthroughRule, Group, IdempotencyRecord,
IdentityAdoptionDecision, PaymentAuditLog, PaymentOrder,
PaymentProviderInstance, PendingAuthSession, PromoCode, PromoCodeUsage, Proxy,
RedeemCode, SecuritySecret, Setting, SubscriptionPlan, TLSFingerprintProfile,
UsageCleanupTask, UsageLog, User, UserAllowedGroup, UserAttributeDefinition,
UserAttributeValue, UserPlatformQuota, UserSubscription []ent.Hook
}
inters struct {
APIKey, Account, AccountGroup, Announcement, AnnouncementRead, AuthIdentity,
AuthIdentityChannel, ChannelMonitor, ChannelMonitorDailyRollup,
ChannelMonitorHistory, ChannelMonitorRequestTemplate, ErrorPassthroughRule,
Group, IdempotencyRecord, IdentityAdoptionDecision, PaymentAuditLog,
PaymentOrder, PaymentProviderInstance, PendingAuthSession, PromoCode,
PromoCodeUsage, Proxy, RedeemCode, SecuritySecret, Setting, SubscriptionPlan,
TLSFingerprintProfile, UsageCleanupTask, UsageLog, User, UserAllowedGroup,
UserAttributeDefinition, UserAttributeValue, UserPlatformQuota,
UserSubscription []ent.Interceptor
AuthIdentityChannel, BatchImageEvent, BatchImageItem, BatchImageJob,
ChannelMonitor, ChannelMonitorDailyRollup, ChannelMonitorHistory,
ChannelMonitorRequestTemplate, ErrorPassthroughRule, Group, IdempotencyRecord,
IdentityAdoptionDecision, PaymentAuditLog, PaymentOrder,
PaymentProviderInstance, PendingAuthSession, PromoCode, PromoCodeUsage, Proxy,
RedeemCode, SecuritySecret, Setting, SubscriptionPlan, TLSFingerprintProfile,
UsageCleanupTask, UsageLog, User, UserAllowedGroup, UserAttributeDefinition,
UserAttributeValue, UserPlatformQuota, UserSubscription []ent.Interceptor
}
)
+6
View File
@@ -19,6 +19,9 @@ import (
"github.com/Wei-Shaw/sub2api/ent/apikey"
"github.com/Wei-Shaw/sub2api/ent/authidentity"
"github.com/Wei-Shaw/sub2api/ent/authidentitychannel"
"github.com/Wei-Shaw/sub2api/ent/batchimageevent"
"github.com/Wei-Shaw/sub2api/ent/batchimageitem"
"github.com/Wei-Shaw/sub2api/ent/batchimagejob"
"github.com/Wei-Shaw/sub2api/ent/channelmonitor"
"github.com/Wei-Shaw/sub2api/ent/channelmonitordailyrollup"
"github.com/Wei-Shaw/sub2api/ent/channelmonitorhistory"
@@ -114,6 +117,9 @@ func checkColumn(t, c string) error {
announcementread.Table: announcementread.ValidColumn,
authidentity.Table: authidentity.ValidColumn,
authidentitychannel.Table: authidentitychannel.ValidColumn,
batchimageevent.Table: batchimageevent.ValidColumn,
batchimageitem.Table: batchimageitem.ValidColumn,
batchimagejob.Table: batchimagejob.ValidColumn,
channelmonitor.Table: channelmonitor.ValidColumn,
channelmonitordailyrollup.Table: channelmonitordailyrollup.ValidColumn,
channelmonitorhistory.Table: channelmonitorhistory.ValidColumn,
+37 -4
View File
@@ -33,9 +33,9 @@ type Group struct {
RateMultiplier float64 `json:"rate_multiplier,omitempty"`
// 是否启用高峰时段倍率
PeakRateEnabled bool `json:"peak_rate_enabled,omitempty"`
// 高峰开始时间 HH:MM(含),如 14:00;空表示未配置
// 高峰开始时间 HH:MM(含),如 14:00;空表示未配置;不支持跨天
PeakStart string `json:"peak_start,omitempty"`
// 高峰结束时间 HH:MM(不含),如 18:00
// 高峰结束时间 HH:MM(不含),必须大于 peak_start;不支持跨天,如 22:00-02:00
PeakEnd string `json:"peak_end,omitempty"`
// 高峰时段叠加倍率,仅在 peak_rate_enabled 且处于 [peak_start, peak_end) 时乘入文本倍率
PeakRateMultiplier float64 `json:"peak_rate_multiplier,omitempty"`
@@ -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(", ")
+30
View File
@@ -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()
+105
View File
@@ -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))
+235
View File
@@ -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) {
+142
View File
@@ -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)
}
+36
View File
@@ -93,6 +93,42 @@ func (f AuthIdentityChannelFunc) Mutate(ctx context.Context, m ent.Mutation) (en
return nil, fmt.Errorf("unexpected mutation type %T. expect *ent.AuthIdentityChannelMutation", m)
}
// The BatchImageEventFunc type is an adapter to allow the use of ordinary
// function as BatchImageEvent mutator.
type BatchImageEventFunc func(context.Context, *ent.BatchImageEventMutation) (ent.Value, error)
// Mutate calls f(ctx, m).
func (f BatchImageEventFunc) Mutate(ctx context.Context, m ent.Mutation) (ent.Value, error) {
if mv, ok := m.(*ent.BatchImageEventMutation); ok {
return f(ctx, mv)
}
return nil, fmt.Errorf("unexpected mutation type %T. expect *ent.BatchImageEventMutation", m)
}
// The BatchImageItemFunc type is an adapter to allow the use of ordinary
// function as BatchImageItem mutator.
type BatchImageItemFunc func(context.Context, *ent.BatchImageItemMutation) (ent.Value, error)
// Mutate calls f(ctx, m).
func (f BatchImageItemFunc) Mutate(ctx context.Context, m ent.Mutation) (ent.Value, error) {
if mv, ok := m.(*ent.BatchImageItemMutation); ok {
return f(ctx, mv)
}
return nil, fmt.Errorf("unexpected mutation type %T. expect *ent.BatchImageItemMutation", m)
}
// The BatchImageJobFunc type is an adapter to allow the use of ordinary
// function as BatchImageJob mutator.
type BatchImageJobFunc func(context.Context, *ent.BatchImageJobMutation) (ent.Value, error)
// Mutate calls f(ctx, m).
func (f BatchImageJobFunc) Mutate(ctx context.Context, m ent.Mutation) (ent.Value, error) {
if mv, ok := m.(*ent.BatchImageJobMutation); ok {
return f(ctx, mv)
}
return nil, fmt.Errorf("unexpected mutation type %T. expect *ent.BatchImageJobMutation", m)
}
// The ChannelMonitorFunc type is an adapter to allow the use of ordinary
// function as ChannelMonitor mutator.
type ChannelMonitorFunc func(context.Context, *ent.ChannelMonitorMutation) (ent.Value, error)
+90
View File
@@ -15,6 +15,9 @@ import (
"github.com/Wei-Shaw/sub2api/ent/apikey"
"github.com/Wei-Shaw/sub2api/ent/authidentity"
"github.com/Wei-Shaw/sub2api/ent/authidentitychannel"
"github.com/Wei-Shaw/sub2api/ent/batchimageevent"
"github.com/Wei-Shaw/sub2api/ent/batchimageitem"
"github.com/Wei-Shaw/sub2api/ent/batchimagejob"
"github.com/Wei-Shaw/sub2api/ent/channelmonitor"
"github.com/Wei-Shaw/sub2api/ent/channelmonitordailyrollup"
"github.com/Wei-Shaw/sub2api/ent/channelmonitorhistory"
@@ -291,6 +294,87 @@ func (f TraverseAuthIdentityChannel) Traverse(ctx context.Context, q ent.Query)
return fmt.Errorf("unexpected query type %T. expect *ent.AuthIdentityChannelQuery", q)
}
// The BatchImageEventFunc type is an adapter to allow the use of ordinary function as a Querier.
type BatchImageEventFunc func(context.Context, *ent.BatchImageEventQuery) (ent.Value, error)
// Query calls f(ctx, q).
func (f BatchImageEventFunc) Query(ctx context.Context, q ent.Query) (ent.Value, error) {
if q, ok := q.(*ent.BatchImageEventQuery); ok {
return f(ctx, q)
}
return nil, fmt.Errorf("unexpected query type %T. expect *ent.BatchImageEventQuery", q)
}
// The TraverseBatchImageEvent type is an adapter to allow the use of ordinary function as Traverser.
type TraverseBatchImageEvent func(context.Context, *ent.BatchImageEventQuery) error
// Intercept is a dummy implementation of Intercept that returns the next Querier in the pipeline.
func (f TraverseBatchImageEvent) Intercept(next ent.Querier) ent.Querier {
return next
}
// Traverse calls f(ctx, q).
func (f TraverseBatchImageEvent) Traverse(ctx context.Context, q ent.Query) error {
if q, ok := q.(*ent.BatchImageEventQuery); ok {
return f(ctx, q)
}
return fmt.Errorf("unexpected query type %T. expect *ent.BatchImageEventQuery", q)
}
// The BatchImageItemFunc type is an adapter to allow the use of ordinary function as a Querier.
type BatchImageItemFunc func(context.Context, *ent.BatchImageItemQuery) (ent.Value, error)
// Query calls f(ctx, q).
func (f BatchImageItemFunc) Query(ctx context.Context, q ent.Query) (ent.Value, error) {
if q, ok := q.(*ent.BatchImageItemQuery); ok {
return f(ctx, q)
}
return nil, fmt.Errorf("unexpected query type %T. expect *ent.BatchImageItemQuery", q)
}
// The TraverseBatchImageItem type is an adapter to allow the use of ordinary function as Traverser.
type TraverseBatchImageItem func(context.Context, *ent.BatchImageItemQuery) error
// Intercept is a dummy implementation of Intercept that returns the next Querier in the pipeline.
func (f TraverseBatchImageItem) Intercept(next ent.Querier) ent.Querier {
return next
}
// Traverse calls f(ctx, q).
func (f TraverseBatchImageItem) Traverse(ctx context.Context, q ent.Query) error {
if q, ok := q.(*ent.BatchImageItemQuery); ok {
return f(ctx, q)
}
return fmt.Errorf("unexpected query type %T. expect *ent.BatchImageItemQuery", q)
}
// The BatchImageJobFunc type is an adapter to allow the use of ordinary function as a Querier.
type BatchImageJobFunc func(context.Context, *ent.BatchImageJobQuery) (ent.Value, error)
// Query calls f(ctx, q).
func (f BatchImageJobFunc) Query(ctx context.Context, q ent.Query) (ent.Value, error) {
if q, ok := q.(*ent.BatchImageJobQuery); ok {
return f(ctx, q)
}
return nil, fmt.Errorf("unexpected query type %T. expect *ent.BatchImageJobQuery", q)
}
// The TraverseBatchImageJob type is an adapter to allow the use of ordinary function as Traverser.
type TraverseBatchImageJob func(context.Context, *ent.BatchImageJobQuery) error
// Intercept is a dummy implementation of Intercept that returns the next Querier in the pipeline.
func (f TraverseBatchImageJob) Intercept(next ent.Querier) ent.Querier {
return next
}
// Traverse calls f(ctx, q).
func (f TraverseBatchImageJob) Traverse(ctx context.Context, q ent.Query) error {
if q, ok := q.(*ent.BatchImageJobQuery); ok {
return f(ctx, q)
}
return fmt.Errorf("unexpected query type %T. expect *ent.BatchImageJobQuery", q)
}
// The ChannelMonitorFunc type is an adapter to allow the use of ordinary function as a Querier.
type ChannelMonitorFunc func(context.Context, *ent.ChannelMonitorQuery) (ent.Value, error)
@@ -1064,6 +1148,12 @@ func NewQuery(q ent.Query) (Query, error) {
return &query[*ent.AuthIdentityQuery, predicate.AuthIdentity, authidentity.OrderOption]{typ: ent.TypeAuthIdentity, tq: q}, nil
case *ent.AuthIdentityChannelQuery:
return &query[*ent.AuthIdentityChannelQuery, predicate.AuthIdentityChannel, authidentitychannel.OrderOption]{typ: ent.TypeAuthIdentityChannel, tq: q}, nil
case *ent.BatchImageEventQuery:
return &query[*ent.BatchImageEventQuery, predicate.BatchImageEvent, batchimageevent.OrderOption]{typ: ent.TypeBatchImageEvent, tq: q}, nil
case *ent.BatchImageItemQuery:
return &query[*ent.BatchImageItemQuery, predicate.BatchImageItem, batchimageitem.OrderOption]{typ: ent.TypeBatchImageItem, tq: q}, nil
case *ent.BatchImageJobQuery:
return &query[*ent.BatchImageJobQuery, predicate.BatchImageJob, batchimagejob.OrderOption]{typ: ent.TypeBatchImageJob, tq: q}, nil
case *ent.ChannelMonitorQuery:
return &query[*ent.ChannelMonitorQuery, predicate.ChannelMonitor, channelmonitor.OrderOption]{typ: ent.TypeChannelMonitor, tq: q}, nil
case *ent.ChannelMonitorDailyRollupQuery:
+200 -2
View File
@@ -435,6 +435,188 @@ var (
},
},
}
// BatchImageEventsColumns holds the columns for the "batch_image_events" table.
BatchImageEventsColumns = []*schema.Column{
{Name: "id", Type: field.TypeInt64, Increment: true},
{Name: "job_id", Type: field.TypeString, Size: 64},
{Name: "event_type", Type: field.TypeString, Size: 64},
{Name: "payload", Type: field.TypeJSON, Nullable: true, SchemaType: map[string]string{"postgres": "jsonb"}},
{Name: "event_hash", Type: field.TypeString, Nullable: true, Size: 128},
{Name: "created_at", Type: field.TypeTime, SchemaType: map[string]string{"postgres": "timestamptz"}},
}
// BatchImageEventsTable holds the schema information for the "batch_image_events" table.
BatchImageEventsTable = &schema.Table{
Name: "batch_image_events",
Columns: BatchImageEventsColumns,
PrimaryKey: []*schema.Column{BatchImageEventsColumns[0]},
Indexes: []*schema.Index{
{
Name: "batchimageevent_job_id_created_at",
Unique: false,
Columns: []*schema.Column{BatchImageEventsColumns[1], BatchImageEventsColumns[5]},
},
{
Name: "batchimageevent_event_type",
Unique: false,
Columns: []*schema.Column{BatchImageEventsColumns[2]},
},
{
Name: "batchimageevent_job_id_event_hash",
Unique: true,
Columns: []*schema.Column{BatchImageEventsColumns[1], BatchImageEventsColumns[4]},
Annotation: &entsql.IndexAnnotation{
Where: "event_hash IS NOT NULL AND event_hash <> ''",
},
},
},
}
// BatchImageItemsColumns holds the columns for the "batch_image_items" table.
BatchImageItemsColumns = []*schema.Column{
{Name: "id", Type: field.TypeInt64, Increment: true},
{Name: "job_id", Type: field.TypeString, Size: 64},
{Name: "custom_id", Type: field.TypeString, Size: 255},
{Name: "status", Type: field.TypeString, Size: 32},
{Name: "request_hash", Type: field.TypeString, Nullable: true, Size: 128},
{Name: "prompt_preview", Type: field.TypeString, Nullable: true, SchemaType: map[string]string{"postgres": "text"}},
{Name: "provider_source_object", Type: field.TypeString, Nullable: true, Size: 1024},
{Name: "source_line_number", Type: field.TypeInt, Nullable: true},
{Name: "source_byte_offset", Type: field.TypeInt64, Nullable: true},
{Name: "source_byte_length", Type: field.TypeInt64, Nullable: true},
{Name: "mime_type", Type: field.TypeString, Nullable: true, Size: 128},
{Name: "file_extension", Type: field.TypeString, Nullable: true, Size: 32},
{Name: "image_count", Type: field.TypeInt, Default: 0},
{Name: "error_code", Type: field.TypeString, Nullable: true, Size: 128},
{Name: "error_message", Type: field.TypeString, Nullable: true, SchemaType: map[string]string{"postgres": "text"}},
{Name: "billed_amount", Type: field.TypeFloat64, Nullable: true, SchemaType: map[string]string{"postgres": "decimal(20,10)"}},
{Name: "created_at", Type: field.TypeTime, SchemaType: map[string]string{"postgres": "timestamptz"}},
{Name: "indexed_at", Type: field.TypeTime, Nullable: true, SchemaType: map[string]string{"postgres": "timestamptz"}},
}
// BatchImageItemsTable holds the schema information for the "batch_image_items" table.
BatchImageItemsTable = &schema.Table{
Name: "batch_image_items",
Columns: BatchImageItemsColumns,
PrimaryKey: []*schema.Column{BatchImageItemsColumns[0]},
Indexes: []*schema.Index{
{
Name: "batchimageitem_job_id_custom_id",
Unique: true,
Columns: []*schema.Column{BatchImageItemsColumns[1], BatchImageItemsColumns[2]},
},
{
Name: "batchimageitem_job_id_status",
Unique: false,
Columns: []*schema.Column{BatchImageItemsColumns[1], BatchImageItemsColumns[3]},
},
{
Name: "batchimageitem_provider_source_object",
Unique: false,
Columns: []*schema.Column{BatchImageItemsColumns[6]},
},
},
}
// BatchImageJobsColumns holds the columns for the "batch_image_jobs" table.
BatchImageJobsColumns = []*schema.Column{
{Name: "id", Type: field.TypeInt64, Increment: true},
{Name: "batch_id", Type: field.TypeString, Size: 64},
{Name: "user_id", Type: field.TypeInt64},
{Name: "api_key_id", Type: field.TypeInt64, Nullable: true},
{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},
{Name: "provider_output_ref", Type: field.TypeString, Nullable: true, Size: 1024},
{Name: "gcs_input_uri", Type: field.TypeString, Nullable: true, Size: 1024},
{Name: "gcs_output_uri", Type: field.TypeString, Nullable: true, Size: 1024},
{Name: "item_count", Type: field.TypeInt},
{Name: "success_count", Type: field.TypeInt, Default: 0},
{Name: "fail_count", Type: field.TypeInt, Default: 0},
{Name: "cancelled_count", Type: field.TypeInt, Default: 0},
{Name: "estimated_cost", Type: field.TypeFloat64, Default: 0, SchemaType: map[string]string{"postgres": "decimal(20,10)"}},
{Name: "hold_amount", Type: field.TypeFloat64, Nullable: true, SchemaType: map[string]string{"postgres": "decimal(20,10)"}},
{Name: "actual_cost", Type: field.TypeFloat64, Nullable: true, SchemaType: map[string]string{"postgres": "decimal(20,10)"}},
{Name: "currency", Type: field.TypeString, Size: 16, Default: "USD"},
{Name: "hold_id", Type: field.TypeString, Nullable: true, Size: 128},
{Name: "idempotency_key", Type: field.TypeString, Nullable: true, Size: 255},
{Name: "request_hash", Type: field.TypeString, Nullable: true, Size: 128},
{Name: "manifest_hash", Type: field.TypeString, Nullable: true, Size: 128},
{Name: "retry_count", Type: field.TypeInt, Default: 0},
{Name: "version", Type: field.TypeInt, Default: 0},
{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"}},
{Name: "updated_at", Type: field.TypeTime, SchemaType: map[string]string{"postgres": "timestamptz"}},
{Name: "submitted_at", Type: field.TypeTime, Nullable: true, SchemaType: map[string]string{"postgres": "timestamptz"}},
{Name: "started_at", Type: field.TypeTime, Nullable: true, SchemaType: map[string]string{"postgres": "timestamptz"}},
{Name: "finished_at", Type: field.TypeTime, Nullable: true, SchemaType: map[string]string{"postgres": "timestamptz"}},
{Name: "settled_at", Type: field.TypeTime, Nullable: true, SchemaType: map[string]string{"postgres": "timestamptz"}},
}
// BatchImageJobsTable holds the schema information for the "batch_image_jobs" table.
BatchImageJobsTable = &schema.Table{
Name: "batch_image_jobs",
Columns: BatchImageJobsColumns,
PrimaryKey: []*schema.Column{BatchImageJobsColumns[0]},
Indexes: []*schema.Index{
{
Name: "batchimagejob_batch_id",
Unique: true,
Columns: []*schema.Column{BatchImageJobsColumns[1]},
},
{
Name: "batchimagejob_user_id_created_at",
Unique: false,
Columns: []*schema.Column{BatchImageJobsColumns[2], BatchImageJobsColumns[35]},
},
{
Name: "batchimagejob_status",
Unique: false,
Columns: []*schema.Column{BatchImageJobsColumns[8]},
},
{
Name: "batchimagejob_provider_status",
Unique: false,
Columns: []*schema.Column{BatchImageJobsColumns[5], BatchImageJobsColumns[8]},
},
{
Name: "batchimagejob_idempotency_key",
Unique: false,
Columns: []*schema.Column{BatchImageJobsColumns[23]},
Annotation: &entsql.IndexAnnotation{
Where: "idempotency_key IS NOT NULL AND idempotency_key <> ''",
},
},
{
Name: "batchimagejob_manifest_hash",
Unique: true,
Columns: []*schema.Column{BatchImageJobsColumns[25]},
Annotation: &entsql.IndexAnnotation{
Where: "manifest_hash IS NOT NULL AND manifest_hash <> ''",
},
},
{
Name: "batchimagejob_output_expires_at",
Unique: false,
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]},
},
},
}
// ChannelMonitorsColumns holds the columns for the "channel_monitors" table.
ChannelMonitorsColumns = []*schema.Column{
{Name: "id", Type: field.TypeInt64, Increment: true},
@@ -670,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},
@@ -725,7 +910,7 @@ var (
{
Name: "group_sort_order",
Unique: false,
Columns: []*schema.Column{GroupsColumns[32]},
Columns: []*schema.Column{GroupsColumns[35]},
},
},
}
@@ -1500,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: ""},
@@ -1526,7 +1712,7 @@ var (
{
Name: "user_status",
Unique: false,
Columns: []*schema.Column{UsersColumns[9]},
Columns: []*schema.Column{UsersColumns[10]},
},
{
Name: "user_deleted_at",
@@ -1799,6 +1985,9 @@ var (
AnnouncementReadsTable,
AuthIdentitiesTable,
AuthIdentityChannelsTable,
BatchImageEventsTable,
BatchImageItemsTable,
BatchImageJobsTable,
ChannelMonitorsTable,
ChannelMonitorDailyRollupsTable,
ChannelMonitorHistoriesTable,
@@ -1862,6 +2051,15 @@ func init() {
AuthIdentityChannelsTable.Annotation = &entsql.Annotation{
Table: "auth_identity_channels",
}
BatchImageEventsTable.Annotation = &entsql.Annotation{
Table: "batch_image_events",
}
BatchImageItemsTable.Annotation = &entsql.Annotation{
Table: "batch_image_items",
}
BatchImageJobsTable.Annotation = &entsql.Annotation{
Table: "batch_image_jobs",
}
ChannelMonitorsTable.ForeignKeys[0].RefTable = ChannelMonitorRequestTemplatesTable
ChannelMonitorsTable.Annotation = &entsql.Annotation{
Table: "channel_monitors",
+5793 -2
View File
File diff suppressed because it is too large Load Diff
+9
View File
@@ -27,6 +27,15 @@ type AuthIdentity func(*sql.Selector)
// AuthIdentityChannel is the predicate function for authidentitychannel builders.
type AuthIdentityChannel func(*sql.Selector)
// BatchImageEvent is the predicate function for batchimageevent builders.
type BatchImageEvent func(*sql.Selector)
// BatchImageItem is the predicate function for batchimageitem builders.
type BatchImageItem func(*sql.Selector)
// BatchImageJob is the predicate function for batchimagejob builders.
type BatchImageJob func(*sql.Selector)
// ChannelMonitor is the predicate function for channelmonitor builders.
type ChannelMonitor func(*sql.Selector)
+210 -25
View File
@@ -12,6 +12,9 @@ import (
"github.com/Wei-Shaw/sub2api/ent/apikey"
"github.com/Wei-Shaw/sub2api/ent/authidentity"
"github.com/Wei-Shaw/sub2api/ent/authidentitychannel"
"github.com/Wei-Shaw/sub2api/ent/batchimageevent"
"github.com/Wei-Shaw/sub2api/ent/batchimageitem"
"github.com/Wei-Shaw/sub2api/ent/batchimagejob"
"github.com/Wei-Shaw/sub2api/ent/channelmonitor"
"github.com/Wei-Shaw/sub2api/ent/channelmonitordailyrollup"
"github.com/Wei-Shaw/sub2api/ent/channelmonitorhistory"
@@ -432,6 +435,172 @@ func init() {
authidentitychannelDescMetadata := authidentitychannelFields[6].Descriptor()
// authidentitychannel.DefaultMetadata holds the default value on creation for the metadata field.
authidentitychannel.DefaultMetadata = authidentitychannelDescMetadata.Default.(func() map[string]interface{})
batchimageeventFields := schema.BatchImageEvent{}.Fields()
_ = batchimageeventFields
// batchimageeventDescJobID is the schema descriptor for job_id field.
batchimageeventDescJobID := batchimageeventFields[0].Descriptor()
// batchimageevent.JobIDValidator is a validator for the "job_id" field. It is called by the builders before save.
batchimageevent.JobIDValidator = batchimageeventDescJobID.Validators[0].(func(string) error)
// batchimageeventDescEventType is the schema descriptor for event_type field.
batchimageeventDescEventType := batchimageeventFields[1].Descriptor()
// batchimageevent.EventTypeValidator is a validator for the "event_type" field. It is called by the builders before save.
batchimageevent.EventTypeValidator = batchimageeventDescEventType.Validators[0].(func(string) error)
// batchimageeventDescEventHash is the schema descriptor for event_hash field.
batchimageeventDescEventHash := batchimageeventFields[3].Descriptor()
// batchimageevent.EventHashValidator is a validator for the "event_hash" field. It is called by the builders before save.
batchimageevent.EventHashValidator = batchimageeventDescEventHash.Validators[0].(func(string) error)
// batchimageeventDescCreatedAt is the schema descriptor for created_at field.
batchimageeventDescCreatedAt := batchimageeventFields[4].Descriptor()
// batchimageevent.DefaultCreatedAt holds the default value on creation for the created_at field.
batchimageevent.DefaultCreatedAt = batchimageeventDescCreatedAt.Default.(func() time.Time)
batchimageitemFields := schema.BatchImageItem{}.Fields()
_ = batchimageitemFields
// batchimageitemDescJobID is the schema descriptor for job_id field.
batchimageitemDescJobID := batchimageitemFields[0].Descriptor()
// batchimageitem.JobIDValidator is a validator for the "job_id" field. It is called by the builders before save.
batchimageitem.JobIDValidator = batchimageitemDescJobID.Validators[0].(func(string) error)
// batchimageitemDescCustomID is the schema descriptor for custom_id field.
batchimageitemDescCustomID := batchimageitemFields[1].Descriptor()
// batchimageitem.CustomIDValidator is a validator for the "custom_id" field. It is called by the builders before save.
batchimageitem.CustomIDValidator = batchimageitemDescCustomID.Validators[0].(func(string) error)
// batchimageitemDescStatus is the schema descriptor for status field.
batchimageitemDescStatus := batchimageitemFields[2].Descriptor()
// batchimageitem.StatusValidator is a validator for the "status" field. It is called by the builders before save.
batchimageitem.StatusValidator = batchimageitemDescStatus.Validators[0].(func(string) error)
// batchimageitemDescRequestHash is the schema descriptor for request_hash field.
batchimageitemDescRequestHash := batchimageitemFields[3].Descriptor()
// batchimageitem.RequestHashValidator is a validator for the "request_hash" field. It is called by the builders before save.
batchimageitem.RequestHashValidator = batchimageitemDescRequestHash.Validators[0].(func(string) error)
// batchimageitemDescProviderSourceObject is the schema descriptor for provider_source_object field.
batchimageitemDescProviderSourceObject := batchimageitemFields[5].Descriptor()
// batchimageitem.ProviderSourceObjectValidator is a validator for the "provider_source_object" field. It is called by the builders before save.
batchimageitem.ProviderSourceObjectValidator = batchimageitemDescProviderSourceObject.Validators[0].(func(string) error)
// batchimageitemDescMimeType is the schema descriptor for mime_type field.
batchimageitemDescMimeType := batchimageitemFields[9].Descriptor()
// batchimageitem.MimeTypeValidator is a validator for the "mime_type" field. It is called by the builders before save.
batchimageitem.MimeTypeValidator = batchimageitemDescMimeType.Validators[0].(func(string) error)
// batchimageitemDescFileExtension is the schema descriptor for file_extension field.
batchimageitemDescFileExtension := batchimageitemFields[10].Descriptor()
// batchimageitem.FileExtensionValidator is a validator for the "file_extension" field. It is called by the builders before save.
batchimageitem.FileExtensionValidator = batchimageitemDescFileExtension.Validators[0].(func(string) error)
// batchimageitemDescImageCount is the schema descriptor for image_count field.
batchimageitemDescImageCount := batchimageitemFields[11].Descriptor()
// batchimageitem.DefaultImageCount holds the default value on creation for the image_count field.
batchimageitem.DefaultImageCount = batchimageitemDescImageCount.Default.(int)
// batchimageitemDescErrorCode is the schema descriptor for error_code field.
batchimageitemDescErrorCode := batchimageitemFields[12].Descriptor()
// batchimageitem.ErrorCodeValidator is a validator for the "error_code" field. It is called by the builders before save.
batchimageitem.ErrorCodeValidator = batchimageitemDescErrorCode.Validators[0].(func(string) error)
// batchimageitemDescCreatedAt is the schema descriptor for created_at field.
batchimageitemDescCreatedAt := batchimageitemFields[15].Descriptor()
// batchimageitem.DefaultCreatedAt holds the default value on creation for the created_at field.
batchimageitem.DefaultCreatedAt = batchimageitemDescCreatedAt.Default.(func() time.Time)
batchimagejobFields := schema.BatchImageJob{}.Fields()
_ = batchimagejobFields
// batchimagejobDescBatchID is the schema descriptor for batch_id field.
batchimagejobDescBatchID := batchimagejobFields[0].Descriptor()
// batchimagejob.BatchIDValidator is a validator for the "batch_id" field. It is called by the builders before save.
batchimagejob.BatchIDValidator = batchimagejobDescBatchID.Validators[0].(func(string) error)
// batchimagejobDescProvider is the schema descriptor for provider field.
batchimagejobDescProvider := batchimagejobFields[4].Descriptor()
// batchimagejob.ProviderValidator is a validator for the "provider" field. It is called by the builders before save.
batchimagejob.ProviderValidator = batchimagejobDescProvider.Validators[0].(func(string) error)
// batchimagejobDescModel is the schema descriptor for model field.
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[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[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[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[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[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[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[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[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[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[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[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[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[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[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[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[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[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[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[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[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.
batchimagejob.UpdateDefaultUpdatedAt = batchimagejobDescUpdatedAt.UpdateDefault.(func() time.Time)
channelmonitorMixin := schema.ChannelMonitor{}.Mixin()
channelmonitorMixinFields0 := channelmonitorMixin[0].Fields()
_ = channelmonitorMixinFields0
@@ -846,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()
@@ -1860,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()
+43
View File
@@ -0,0 +1,43 @@
package schema
import (
"time"
"entgo.io/ent"
"entgo.io/ent/dialect"
"entgo.io/ent/dialect/entsql"
"entgo.io/ent/schema"
"entgo.io/ent/schema/field"
"entgo.io/ent/schema/index"
)
// BatchImageEvent records append-only operational events for batch image jobs.
type BatchImageEvent struct {
ent.Schema
}
func (BatchImageEvent) Annotations() []schema.Annotation {
return []schema.Annotation{
entsql.Annotation{Table: "batch_image_events"},
}
}
func (BatchImageEvent) Fields() []ent.Field {
return []ent.Field{
field.String("job_id").MaxLen(64),
field.String("event_type").MaxLen(64),
field.JSON("payload", map[string]any{}).
Optional().
SchemaType(map[string]string{dialect.Postgres: "jsonb"}),
field.String("event_hash").Optional().Nillable().MaxLen(128),
field.Time("created_at").Immutable().Default(time.Now).SchemaType(map[string]string{dialect.Postgres: "timestamptz"}),
}
}
func (BatchImageEvent) Indexes() []ent.Index {
return []ent.Index{
index.Fields("job_id", "created_at"),
index.Fields("event_type"),
index.Fields("job_id", "event_hash").Unique().Annotations(entsql.IndexWhere("event_hash IS NOT NULL AND event_hash <> ''")),
}
}
+53
View File
@@ -0,0 +1,53 @@
package schema
import (
"time"
"entgo.io/ent"
"entgo.io/ent/dialect"
"entgo.io/ent/dialect/entsql"
"entgo.io/ent/schema"
"entgo.io/ent/schema/field"
"entgo.io/ent/schema/index"
)
// BatchImageItem holds indexed output rows for a batch image job.
type BatchImageItem struct {
ent.Schema
}
func (BatchImageItem) Annotations() []schema.Annotation {
return []schema.Annotation{
entsql.Annotation{Table: "batch_image_items"},
}
}
func (BatchImageItem) Fields() []ent.Field {
return []ent.Field{
field.String("job_id").MaxLen(64),
field.String("custom_id").MaxLen(255),
field.String("status").MaxLen(32),
field.String("request_hash").Optional().Nillable().MaxLen(128),
field.String("prompt_preview").Optional().Nillable().SchemaType(map[string]string{dialect.Postgres: "text"}),
field.String("provider_source_object").Optional().Nillable().MaxLen(1024),
field.Int("source_line_number").Optional().Nillable(),
field.Int64("source_byte_offset").Optional().Nillable(),
field.Int64("source_byte_length").Optional().Nillable(),
field.String("mime_type").Optional().Nillable().MaxLen(128),
field.String("file_extension").Optional().Nillable().MaxLen(32),
field.Int("image_count").Default(0),
field.String("error_code").Optional().Nillable().MaxLen(128),
field.String("error_message").Optional().Nillable().SchemaType(map[string]string{dialect.Postgres: "text"}),
field.Float("billed_amount").Optional().Nillable().SchemaType(map[string]string{dialect.Postgres: "decimal(20,10)"}),
field.Time("created_at").Immutable().Default(time.Now).SchemaType(map[string]string{dialect.Postgres: "timestamptz"}),
field.Time("indexed_at").Optional().Nillable().SchemaType(map[string]string{dialect.Postgres: "timestamptz"}),
}
}
func (BatchImageItem) Indexes() []ent.Index {
return []ent.Index{
index.Fields("job_id", "custom_id").Unique(),
index.Fields("job_id", "status"),
index.Fields("provider_source_object"),
}
}
+86
View File
@@ -0,0 +1,86 @@
package schema
import (
"time"
"entgo.io/ent"
"entgo.io/ent/dialect"
"entgo.io/ent/dialect/entsql"
"entgo.io/ent/schema"
"entgo.io/ent/schema/field"
"entgo.io/ent/schema/index"
)
// BatchImageJob holds the schema definition for asynchronous image batch jobs.
//
// 删除策略:账务源保留
// 这张表是批量生图任务的账务和状态源;用户侧删除仅通过 user_deleted_at
// 从列表隐藏,输出清理通过 output_deleted 状态和删除时间字段表达。
type BatchImageJob struct {
ent.Schema
}
func (BatchImageJob) Annotations() []schema.Annotation {
return []schema.Annotation{
entsql.Annotation{Table: "batch_image_jobs"},
}
}
func (BatchImageJob) Fields() []ent.Field {
return []ent.Field{
field.String("batch_id").MaxLen(64).Immutable(),
field.Int64("user_id"),
field.Int64("api_key_id").Optional().Nillable(),
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),
field.String("provider_output_ref").Optional().Nillable().MaxLen(1024),
field.String("gcs_input_uri").Optional().Nillable().MaxLen(1024),
field.String("gcs_output_uri").Optional().Nillable().MaxLen(1024),
field.Int("item_count"),
field.Int("success_count").Default(0),
field.Int("fail_count").Default(0),
field.Int("cancelled_count").Default(0),
field.Float("estimated_cost").SchemaType(map[string]string{dialect.Postgres: "decimal(20,10)"}).Default(0),
field.Float("hold_amount").Optional().Nillable().SchemaType(map[string]string{dialect.Postgres: "decimal(20,10)"}),
field.Float("actual_cost").Optional().Nillable().SchemaType(map[string]string{dialect.Postgres: "decimal(20,10)"}),
field.String("currency").MaxLen(16).Default("USD"),
field.String("hold_id").Optional().Nillable().MaxLen(128),
field.String("idempotency_key").Optional().Nillable().MaxLen(255),
field.String("request_hash").Optional().Nillable().MaxLen(128),
field.String("manifest_hash").Optional().Nillable().MaxLen(128),
field.Int("retry_count").Default(0),
field.Int("version").Default(0),
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"}),
field.Time("updated_at").Default(time.Now).UpdateDefault(time.Now).SchemaType(map[string]string{dialect.Postgres: "timestamptz"}),
field.Time("submitted_at").Optional().Nillable().SchemaType(map[string]string{dialect.Postgres: "timestamptz"}),
field.Time("started_at").Optional().Nillable().SchemaType(map[string]string{dialect.Postgres: "timestamptz"}),
field.Time("finished_at").Optional().Nillable().SchemaType(map[string]string{dialect.Postgres: "timestamptz"}),
field.Time("settled_at").Optional().Nillable().SchemaType(map[string]string{dialect.Postgres: "timestamptz"}),
}
}
func (BatchImageJob) Indexes() []ent.Index {
return []ent.Index{
index.Fields("batch_id").Unique(),
index.Fields("user_id", "created_at"),
index.Fields("status"),
index.Fields("provider", "status"),
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"),
}
}
+11
View File
@@ -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").
+3
View File
@@ -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").
+9
View File
@@ -28,6 +28,12 @@ type Tx struct {
AuthIdentity *AuthIdentityClient
// AuthIdentityChannel is the client for interacting with the AuthIdentityChannel builders.
AuthIdentityChannel *AuthIdentityChannelClient
// BatchImageEvent is the client for interacting with the BatchImageEvent builders.
BatchImageEvent *BatchImageEventClient
// BatchImageItem is the client for interacting with the BatchImageItem builders.
BatchImageItem *BatchImageItemClient
// BatchImageJob is the client for interacting with the BatchImageJob builders.
BatchImageJob *BatchImageJobClient
// ChannelMonitor is the client for interacting with the ChannelMonitor builders.
ChannelMonitor *ChannelMonitorClient
// ChannelMonitorDailyRollup is the client for interacting with the ChannelMonitorDailyRollup builders.
@@ -222,6 +228,9 @@ func (tx *Tx) init() {
tx.AnnouncementRead = NewAnnouncementReadClient(tx.config)
tx.AuthIdentity = NewAuthIdentityClient(tx.config)
tx.AuthIdentityChannel = NewAuthIdentityChannelClient(tx.config)
tx.BatchImageEvent = NewBatchImageEventClient(tx.config)
tx.BatchImageItem = NewBatchImageItemClient(tx.config)
tx.BatchImageJob = NewBatchImageJobClient(tx.config)
tx.ChannelMonitor = NewChannelMonitorClient(tx.config)
tx.ChannelMonitorDailyRollup = NewChannelMonitorDailyRollupClient(tx.config)
tx.ChannelMonitorHistory = NewChannelMonitorHistoryClient(tx.config)
+12 -1
View File
@@ -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(", ")
+10
View File
@@ -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()
+45
View File
@@ -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))
+85
View File
@@ -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) {
+54
View File
@@ -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)
}
+154 -1
View File
@@ -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 (
@@ -93,6 +93,7 @@ type Config struct {
Gemini GeminiConfig `mapstructure:"gemini"`
Update UpdateConfig `mapstructure:"update"`
Idempotency IdempotencyConfig `mapstructure:"idempotency"`
BatchImage BatchImageConfig `mapstructure:"batch_image"`
}
type LogConfig struct {
@@ -175,6 +176,56 @@ type IdempotencyConfig struct {
CleanupBatchSize int `mapstructure:"cleanup_batch_size"`
}
type BatchImageConfig struct {
Enabled bool `mapstructure:"enabled"`
MaxItemsPerJobDefault int `mapstructure:"max_items_per_job_default"`
MaxItemsPerJobTrial int `mapstructure:"max_items_per_job_trial"`
MaxOutputImagesPerJob int `mapstructure:"max_output_images_per_job"`
MaxOutputImagesPerItem int `mapstructure:"max_output_images_per_item"`
MaxPromptCharsPerItem int `mapstructure:"max_prompt_chars_per_item"`
MaxReferenceImagesPerJob int `mapstructure:"max_reference_images_per_job"`
MaxReferenceInlineBytesPerJob int `mapstructure:"max_reference_inline_bytes_per_job"`
DefaultResponseMimeType string `mapstructure:"default_response_mime_type"`
DefaultImageSize string `mapstructure:"default_image_size"`
MaxDownloadItemsZip int `mapstructure:"max_download_items_zip"`
MaxDownloadBytesPerRequest int64 `mapstructure:"max_download_bytes_per_request"`
MaxDownloadDurationSeconds int `mapstructure:"max_download_duration_seconds"`
MaxDownloadConcurrencyPerUser int `mapstructure:"max_download_concurrency_per_user"`
InputRetentionAfterTerminalHours int `mapstructure:"input_retention_after_terminal_hours"`
OutputRetentionAfterTerminalHours int `mapstructure:"output_retention_after_terminal_hours"`
OutputRetentionMaxDays int `mapstructure:"output_retention_max_days"`
CleanupIntervalMinutes int `mapstructure:"cleanup_interval_minutes"`
CleanupBatchSize int `mapstructure:"cleanup_batch_size"`
QueueEnabled bool `mapstructure:"queue_enabled"`
QueueReadyKey string `mapstructure:"queue_ready_key"`
QueueDelayedKey string `mapstructure:"queue_delayed_key"`
QueueActiveKey string `mapstructure:"queue_active_key"`
InflightKeyPrefix string `mapstructure:"inflight_key_prefix"`
LockKeyPrefix string `mapstructure:"lock_key_prefix"`
IdempotencyKeyPrefix string `mapstructure:"idempotency_key_prefix"`
InflightTTLSeconds int `mapstructure:"inflight_ttl_seconds"`
JobLockTTLSeconds int `mapstructure:"job_lock_ttl_seconds"`
DefaultRequeueDelaySeconds int `mapstructure:"default_requeue_delay_seconds"`
ErrorRetryDelaySeconds int `mapstructure:"error_retry_delay_seconds"`
LockConflictDelaySeconds int `mapstructure:"lock_conflict_delay_seconds"`
StaleActiveAfterSeconds int `mapstructure:"stale_active_after_seconds"`
DelayedMoverIntervalSeconds int `mapstructure:"delayed_mover_interval_seconds"`
RecoveryIntervalSeconds int `mapstructure:"recovery_interval_seconds"`
DelayedMoveLimit int `mapstructure:"delayed_move_limit"`
RecoverLimit int `mapstructure:"recover_limit"`
VertexEnabled bool `mapstructure:"vertex_enabled"`
VertexProjectID string `mapstructure:"vertex_project_id"`
VertexLocation string `mapstructure:"vertex_location"`
// VertexManagedGCSBucket is a server-owned bucket for batch JSONL input/output.
// Disable Cloud Storage soft delete on this bucket to avoid retaining deleted batch objects.
VertexManagedGCSBucket string `mapstructure:"vertex_managed_gcs_bucket"`
VertexManagedGCSPrefix string `mapstructure:"vertex_managed_gcs_prefix"`
VertexInputRetentionHours int `mapstructure:"vertex_input_retention_hours"`
VertexOutputRetentionHours int `mapstructure:"vertex_output_retention_hours"`
VertexBatchPredictionBaseURL string `mapstructure:"vertex_batch_prediction_base_url"`
VertexGCSBaseURL string `mapstructure:"vertex_gcs_base_url"`
}
type LinuxDoConnectConfig struct {
Enabled bool `mapstructure:"enabled"`
ClientID string `mapstructure:"client_id"`
@@ -1732,6 +1783,53 @@ func setDefaults() {
viper.SetDefault("redis.min_idle_conns", 128)
viper.SetDefault("redis.enable_tls", false)
// Batch Image queue
viper.SetDefault("batch_image.enabled", false)
viper.SetDefault("batch_image.max_items_per_job_default", 200)
viper.SetDefault("batch_image.max_items_per_job_trial", 50)
viper.SetDefault("batch_image.max_output_images_per_job", 200)
viper.SetDefault("batch_image.max_output_images_per_item", 4)
viper.SetDefault("batch_image.max_prompt_chars_per_item", 8000)
viper.SetDefault("batch_image.max_reference_images_per_job", 1000)
viper.SetDefault("batch_image.max_reference_inline_bytes_per_job", 134217728)
viper.SetDefault("batch_image.default_response_mime_type", "image/png")
viper.SetDefault("batch_image.default_image_size", "1K")
viper.SetDefault("batch_image.max_download_items_zip", 200)
viper.SetDefault("batch_image.max_download_bytes_per_request", 536870912)
viper.SetDefault("batch_image.max_download_duration_seconds", 600)
viper.SetDefault("batch_image.max_download_concurrency_per_user", 1)
viper.SetDefault("batch_image.input_retention_after_terminal_hours", 24)
viper.SetDefault("batch_image.output_retention_after_terminal_hours", 72)
viper.SetDefault("batch_image.output_retention_max_days", 7)
viper.SetDefault("batch_image.cleanup_interval_minutes", 30)
viper.SetDefault("batch_image.cleanup_batch_size", 100)
viper.SetDefault("batch_image.queue_enabled", false)
viper.SetDefault("batch_image.queue_ready_key", "batch_image:queue:ready")
viper.SetDefault("batch_image.queue_delayed_key", "batch_image:queue:delayed")
viper.SetDefault("batch_image.queue_active_key", "batch_image:queue:active")
viper.SetDefault("batch_image.inflight_key_prefix", "batch_image:queue:inflight:")
viper.SetDefault("batch_image.lock_key_prefix", "batch_image:queue:lock:")
viper.SetDefault("batch_image.idempotency_key_prefix", "batch_image:queue:idem:")
viper.SetDefault("batch_image.inflight_ttl_seconds", 604800)
viper.SetDefault("batch_image.job_lock_ttl_seconds", 300)
viper.SetDefault("batch_image.default_requeue_delay_seconds", 30)
viper.SetDefault("batch_image.error_retry_delay_seconds", 60)
viper.SetDefault("batch_image.lock_conflict_delay_seconds", 5)
viper.SetDefault("batch_image.stale_active_after_seconds", 600)
viper.SetDefault("batch_image.delayed_mover_interval_seconds", 5)
viper.SetDefault("batch_image.recovery_interval_seconds", 300)
viper.SetDefault("batch_image.delayed_move_limit", 100)
viper.SetDefault("batch_image.recover_limit", 100)
viper.SetDefault("batch_image.vertex_enabled", false)
viper.SetDefault("batch_image.vertex_project_id", "")
viper.SetDefault("batch_image.vertex_location", "global")
viper.SetDefault("batch_image.vertex_managed_gcs_bucket", "")
viper.SetDefault("batch_image.vertex_managed_gcs_prefix", "batch-image/{env}/{batch_id}")
viper.SetDefault("batch_image.vertex_input_retention_hours", 24)
viper.SetDefault("batch_image.vertex_output_retention_hours", 72)
viper.SetDefault("batch_image.vertex_batch_prediction_base_url", "")
viper.SetDefault("batch_image.vertex_gcs_base_url", "")
// Ops (vNext)
viper.SetDefault("ops.enabled", true)
viper.SetDefault("ops.use_preaggregated_tables", true)
@@ -2333,6 +2431,61 @@ func (c *Config) Validate() error {
if c.Redis.MinIdleConns > c.Redis.PoolSize {
return fmt.Errorf("redis.min_idle_conns cannot exceed redis.pool_size")
}
if c.BatchImage.QueueEnabled {
if strings.TrimSpace(c.BatchImage.QueueReadyKey) == "" {
return fmt.Errorf("batch_image.queue_ready_key must not be empty")
}
if strings.TrimSpace(c.BatchImage.QueueDelayedKey) == "" {
return fmt.Errorf("batch_image.queue_delayed_key must not be empty")
}
if strings.TrimSpace(c.BatchImage.QueueActiveKey) == "" {
return fmt.Errorf("batch_image.queue_active_key must not be empty")
}
if strings.TrimSpace(c.BatchImage.InflightKeyPrefix) == "" {
return fmt.Errorf("batch_image.inflight_key_prefix must not be empty")
}
if strings.TrimSpace(c.BatchImage.LockKeyPrefix) == "" {
return fmt.Errorf("batch_image.lock_key_prefix must not be empty")
}
if c.BatchImage.InflightTTLSeconds <= 0 {
return fmt.Errorf("batch_image.inflight_ttl_seconds must be positive")
}
if c.BatchImage.JobLockTTLSeconds <= 0 {
return fmt.Errorf("batch_image.job_lock_ttl_seconds must be positive")
}
if c.BatchImage.StaleActiveAfterSeconds <= 0 {
return fmt.Errorf("batch_image.stale_active_after_seconds must be positive")
}
if c.BatchImage.DelayedMoveLimit <= 0 {
return fmt.Errorf("batch_image.delayed_move_limit must be positive")
}
if c.BatchImage.RecoverLimit <= 0 {
return fmt.Errorf("batch_image.recover_limit must be positive")
}
}
if c.BatchImage.VertexEnabled {
if strings.TrimSpace(c.BatchImage.VertexManagedGCSBucket) == "" {
return fmt.Errorf("batch_image.vertex_managed_gcs_bucket must not be empty when vertex is enabled")
}
if strings.Contains(c.BatchImage.VertexManagedGCSBucket, "://") {
return fmt.Errorf("batch_image.vertex_managed_gcs_bucket must be a bucket name, not a URI")
}
if strings.TrimSpace(c.BatchImage.VertexLocation) == "" {
return fmt.Errorf("batch_image.vertex_location must not be empty when vertex is enabled")
}
if strings.TrimSpace(c.BatchImage.VertexManagedGCSPrefix) == "" {
return fmt.Errorf("batch_image.vertex_managed_gcs_prefix must not be empty when vertex is enabled")
}
if !strings.Contains(c.BatchImage.VertexManagedGCSPrefix, "{batch_id}") {
return fmt.Errorf("batch_image.vertex_managed_gcs_prefix must contain {batch_id}")
}
if c.BatchImage.VertexInputRetentionHours <= 0 {
return fmt.Errorf("batch_image.vertex_input_retention_hours must be positive")
}
if c.BatchImage.VertexOutputRetentionHours <= 0 {
return fmt.Errorf("batch_image.vertex_output_retention_hours must be positive")
}
}
if c.Dashboard.Enabled {
if c.Dashboard.StatsFreshTTLSeconds <= 0 {
return fmt.Errorf("dashboard_cache.stats_fresh_ttl_seconds must be positive")
+8
View File
@@ -270,6 +270,14 @@ func TestLoadDefaultIdempotencyConfig(t *testing.T) {
}
}
func TestLoadDefaultBatchImageQueueDisabled(t *testing.T) {
resetViperWithJWTSecret(t)
cfg, err := Load()
require.NoError(t, err)
require.False(t, cfg.BatchImage.QueueEnabled)
}
func TestLoadIdempotencyConfigFromEnv(t *testing.T) {
resetViperWithJWTSecret(t)
t.Setenv("IDEMPOTENCY_OBSERVE_ONLY", "false")
@@ -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,
@@ -0,0 +1,257 @@
package handler
import (
"errors"
"io"
"net/http"
"strconv"
"strings"
infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors"
"github.com/Wei-Shaw/sub2api/internal/server/middleware"
"github.com/Wei-Shaw/sub2api/internal/service"
"github.com/gin-gonic/gin"
)
type BatchImageHandler struct {
service *service.BatchImagePublicService
download *service.BatchImageDownloadService
cleanup *service.BatchImageCleanupService
}
func NewBatchImageHandler(service *service.BatchImagePublicService, download *service.BatchImageDownloadService, cleanup *service.BatchImageCleanupService) *BatchImageHandler {
return &BatchImageHandler{service: service, download: download, cleanup: cleanup}
}
func (h *BatchImageHandler) Submit(c *gin.Context) {
var req service.BatchImageSubmitRequest
if err := c.ShouldBindJSON(&req); err != nil {
batchImageError(c, service.ErrBatchImageInvalidItems)
return
}
owner, ok := batchImageOwnerFromContext(c)
if !ok {
batchImageError(c, infraerrors.New(http.StatusUnauthorized, "API_KEY_REQUIRED", "API key is required"))
return
}
got, err := h.service.Submit(c.Request.Context(), owner, req, c.GetHeader("Idempotency-Key"))
if err != nil {
batchImageError(c, err)
return
}
c.JSON(http.StatusOK, got)
}
func (h *BatchImageHandler) Get(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.Get(c.Request.Context(), owner, c.Param("id"))
if err != nil {
batchImageError(c, err)
return
}
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 {
batchImageError(c, infraerrors.New(http.StatusUnauthorized, "API_KEY_REQUIRED", "API key is required"))
return
}
limit, _ := strconv.Atoi(c.Query("limit"))
got, err := h.service.ListItems(c.Request.Context(), owner, c.Param("id"), service.BatchImageItemsQuery{
Status: c.Query("status"),
Limit: limit,
Cursor: c.Query("cursor"),
})
if err != nil {
batchImageError(c, err)
return
}
c.JSON(http.StatusOK, got)
}
func (h *BatchImageHandler) Cancel(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.Cancel(c.Request.Context(), owner, c.Param("id"))
if err != nil {
batchImageError(c, err)
return
}
c.JSON(http.StatusOK, got)
}
func (h *BatchImageHandler) ItemContent(c *gin.Context) {
owner, ok := batchImageOwnerFromContext(c)
if !ok {
batchImageError(c, infraerrors.New(http.StatusUnauthorized, "API_KEY_REQUIRED", "API key is required"))
return
}
imageIndex := 0
if raw := c.Query("image_index"); raw != "" {
parsed, err := strconv.Atoi(raw)
if err != nil {
batchImageError(c, service.ErrBatchImageItemImageIndexOutOfRange)
return
}
imageIndex = parsed
}
stream, err := h.download.OpenItemContent(c.Request.Context(), owner, c.Param("id"), c.Param("custom_id"), imageIndex)
if err != nil {
batchImageError(c, err)
return
}
defer func() { _ = stream.Reader.Close() }()
c.Header("Content-Type", stream.ContentType)
c.Header("Content-Disposition", service.BatchImageContentDispositionAttachment(stream.Filename))
c.Header("Cache-Control", "private, max-age=300")
c.Header("X-Content-Type-Options", "nosniff")
if stream.ContentLength != nil && *stream.ContentLength >= 0 {
c.Header("Content-Length", strconv.FormatInt(*stream.ContentLength, 10))
}
c.Status(http.StatusOK)
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) {
owner, ok := batchImageOwnerFromContext(c)
if !ok {
batchImageError(c, infraerrors.New(http.StatusUnauthorized, "API_KEY_REQUIRED", "API key is required"))
return
}
maxItems, _ := strconv.Atoi(c.Query("max_items"))
c.Header("Content-Type", "application/zip")
c.Header("Content-Disposition", service.BatchImageContentDispositionAttachment(c.Param("id")+".zip"))
c.Header("Cache-Control", "private, no-store")
c.Header("X-Content-Type-Options", "nosniff")
result, err := h.download.StreamZip(c.Request.Context(), owner, c.Param("id"), service.BatchImageZipOptions{
Status: c.Query("status"),
MaxItems: maxItems,
IncludeManifest: true,
}, c.Writer)
if err != nil {
if result == nil || !c.Writer.Written() {
batchImageError(c, err)
}
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) {
owner, ok := batchImageOwnerFromContext(c)
if !ok {
batchImageError(c, infraerrors.New(http.StatusUnauthorized, "API_KEY_REQUIRED", "API key is required"))
return
}
got, err := h.cleanup.DeleteOutputsForOwner(c.Request.Context(), owner, c.Param("id"))
if err != nil {
batchImageError(c, err)
return
}
c.JSON(http.StatusOK, got)
}
func batchImageOwnerFromContext(c *gin.Context) (service.BatchImageOwner, bool) {
apiKey, ok := middleware.GetAPIKeyFromContext(c)
if !ok || apiKey == nil || apiKey.ID <= 0 || apiKey.UserID <= 0 {
return service.BatchImageOwner{}, false
}
return service.BatchImageOwner{
UserID: apiKey.UserID,
APIKeyID: apiKey.ID,
GroupID: apiKey.GroupID,
}, true
}
func batchImageError(c *gin.Context, err error) {
status := infraerrors.Code(err)
code := infraerrors.Reason(err)
message := infraerrors.Message(err)
if err == nil {
status = http.StatusInternalServerError
code = "INTERNAL_ERROR"
message = "internal error"
}
if status == 0 || (status == http.StatusInternalServerError && strings.TrimSpace(code) == "") {
status = http.StatusInternalServerError
code = "INTERNAL_ERROR"
message = "internal error"
}
if errors.Is(err, service.ErrBatchImageJobNotFound) {
status = http.StatusNotFound
code = "BATCH_IMAGE_NOT_FOUND"
message = "batch image job not found"
}
c.JSON(status, gin.H{
"error": gin.H{
"type": "invalid_request_error",
"code": code,
"message": message,
},
})
}
+4
View File
@@ -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,
@@ -180,8 +181,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,
+7 -3
View File
@@ -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"`
@@ -99,9 +100,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"`
+1
View File
@@ -58,6 +58,7 @@ type Handlers struct {
Payment *PaymentHandler
PaymentWebhook *PaymentWebhookHandler
AvailableChannel *AvailableChannelHandler
BatchImage *BatchImageHandler
}
// BuildInfo contains build-time information
+3
View File
@@ -115,6 +115,7 @@ func ProvideHandlers(
paymentHandler *PaymentHandler,
paymentWebhookHandler *PaymentWebhookHandler,
availableChannelHandler *AvailableChannelHandler,
batchImageHandler *BatchImageHandler,
_ *service.IdempotencyCoordinator,
_ *service.IdempotencyCleanupService,
) *Handlers {
@@ -135,6 +136,7 @@ func ProvideHandlers(
Payment: paymentHandler,
PaymentWebhook: paymentWebhookHandler,
AvailableChannel: availableChannelHandler,
BatchImage: batchImageHandler,
}
}
@@ -156,6 +158,7 @@ var ProviderSet = wire.NewSet(
NewPaymentHandler,
NewPaymentWebhookHandler,
NewAvailableChannelHandler,
NewBatchImageHandler,
// Admin handlers
admin.NewDashboardHandler,
@@ -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,
@@ -0,0 +1,102 @@
package repository
import (
"context"
"sync"
"time"
"github.com/Wei-Shaw/sub2api/internal/config"
"github.com/Wei-Shaw/sub2api/internal/service"
"github.com/redis/go-redis/v9"
)
const (
defaultBatchImageDownloadActivePrefix = "batch_image:download:active:"
defaultBatchImageDownloadActiveTTL = 10 * time.Minute
defaultBatchImageDownloadConcurrency = 2
)
var batchImageDownloadAcquireScript = redis.NewScript(`
local current = tonumber(redis.call("GET", KEYS[1]) or "0")
local max = tonumber(ARGV[1])
if current >= max then
return 0
end
redis.call("INCR", KEYS[1])
redis.call("EXPIRE", KEYS[1], ARGV[2])
return 1
`)
var batchImageDownloadReleaseScript = redis.NewScript(`
local current = tonumber(redis.call("GET", KEYS[1]) or "0")
if current <= 1 then
redis.call("DEL", KEYS[1])
return 0
end
return redis.call("DECR", KEYS[1])
`)
type batchImageDownloadLimiter struct {
rdb *redis.Client
activePrefix string
maxActive int
ttl time.Duration
}
func NewBatchImageDownloadLimiter(rdb *redis.Client, cfg *config.Config) service.BatchImageDownloadLimiter {
maxActive := defaultBatchImageDownloadConcurrency
ttl := defaultBatchImageDownloadActiveTTL
if cfg != nil {
if cfg.BatchImage.MaxDownloadConcurrencyPerUser > 0 {
maxActive = cfg.BatchImage.MaxDownloadConcurrencyPerUser
}
if cfg.BatchImage.MaxDownloadDurationSeconds > 0 {
ttl = time.Duration(cfg.BatchImage.MaxDownloadDurationSeconds) * time.Second
}
}
return &batchImageDownloadLimiter{
rdb: rdb,
activePrefix: defaultBatchImageDownloadActivePrefix,
maxActive: maxActive,
ttl: ttl,
}
}
func (l *batchImageDownloadLimiter) Acquire(ctx context.Context, userID string, kind string) (service.BatchImageDownloadPermit, error) {
if l == nil || l.rdb == nil {
return nil, service.ErrBatchImageDownloadLimited
}
key := l.activeKey(userID)
ok, err := batchImageDownloadAcquireScript.Run(ctx, l.rdb, []string{key}, l.maxActive, int(l.ttl.Seconds())).Int()
if err != nil {
return nil, err
}
if ok != 1 {
return nil, service.ErrBatchImageDownloadLimited
}
return &batchImageDownloadPermit{rdb: l.rdb, key: key}, nil
}
func (l *batchImageDownloadLimiter) activeKey(userID string) string {
return l.activePrefix + userID
}
type batchImageDownloadPermit struct {
rdb *redis.Client
key string
once sync.Once
err error
}
func (p *batchImageDownloadPermit) Release(ctx context.Context) error {
if p == nil || p.rdb == nil || p.key == "" {
return nil
}
p.once.Do(func() {
_, p.err = batchImageDownloadReleaseScript.Run(ctx, p.rdb, []string{p.key}).Result()
})
return p.err
}
var _ service.BatchImageDownloadLimiter = (*batchImageDownloadLimiter)(nil)
var _ service.BatchImageDownloadPermit = (*batchImageDownloadPermit)(nil)
@@ -0,0 +1,43 @@
//go:build unit
package repository
import (
"context"
"testing"
"time"
"github.com/Wei-Shaw/sub2api/internal/service"
"github.com/alicebob/miniredis/v2"
"github.com/redis/go-redis/v9"
"github.com/stretchr/testify/require"
)
func TestBatchImageDownloadLimiter_AcquireDenyReleaseAndTTL(t *testing.T) {
ctx := context.Background()
mr := miniredis.RunT(t)
rdb := redis.NewClient(&redis.Options{Addr: mr.Addr()})
t.Cleanup(func() { _ = rdb.Close() })
limiter := &batchImageDownloadLimiter{
rdb: rdb,
activePrefix: defaultBatchImageDownloadActivePrefix,
maxActive: 1,
ttl: time.Minute,
}
permit, err := limiter.Acquire(ctx, "11", "zip")
require.NoError(t, err)
require.NotNil(t, permit)
require.True(t, mr.TTL(limiter.activeKey("11")) > 0)
_, err = limiter.Acquire(ctx, "11", "zip")
require.ErrorIs(t, err, service.ErrBatchImageDownloadLimited)
require.NoError(t, permit.Release(ctx))
require.NoError(t, permit.Release(ctx))
require.False(t, mr.Exists(limiter.activeKey("11")))
permit, err = limiter.Acquire(ctx, "11", "zip")
require.NoError(t, err)
require.NotNil(t, permit)
}
@@ -0,0 +1,280 @@
package repository
import (
"context"
"crypto/rand"
"encoding/hex"
"errors"
"time"
"github.com/Wei-Shaw/sub2api/internal/config"
"github.com/Wei-Shaw/sub2api/internal/service"
"github.com/redis/go-redis/v9"
)
const (
defaultBatchImageReadyKey = "batch_image:queue:ready"
defaultBatchImageDelayedKey = "batch_image:queue:delayed"
defaultBatchImageActiveKey = "batch_image:queue:active"
defaultBatchImageInflightPrefix = "batch_image:queue:inflight:"
defaultBatchImageLockPrefix = "batch_image:queue:lock:"
defaultBatchImageInflightTTL = 7 * 24 * time.Hour
defaultBatchImageJobLockTTL = 5 * time.Minute
)
var batchImageMoveDueDelayedScript = redis.NewScript(`
local jobs = redis.call("ZRANGEBYSCORE", KEYS[1], "-inf", ARGV[1], "LIMIT", 0, ARGV[2])
for _, job in ipairs(jobs) do
redis.call("ZREM", KEYS[1], job)
redis.call("LPUSH", KEYS[2], job)
end
return #jobs
`)
var batchImageRecoverStaleActiveScript = redis.NewScript(`
local jobs = redis.call("ZRANGEBYSCORE", KEYS[1], "-inf", ARGV[1], "LIMIT", 0, ARGV[2])
for _, job in ipairs(jobs) do
redis.call("ZREM", KEYS[1], job)
redis.call("LPUSH", KEYS[2], job)
end
return #jobs
`)
var batchImageReleaseLockScript = redis.NewScript(`
if redis.call("GET", KEYS[1]) == ARGV[1] then
return redis.call("DEL", KEYS[1])
end
return 0
`)
type batchImageQueue struct {
rdb *redis.Client
readyKey string
delayedKey string
activeKey string
inflightPrefix string
lockPrefix string
inflightTTL time.Duration
lockTTL time.Duration
}
func NewBatchImageQueue(rdb *redis.Client, cfg *config.Config) service.BatchImageQueue {
return newBatchImageQueueWithOptions(rdb, batchImageQueueOptionsFromConfig(cfg))
}
type batchImageQueueOptions struct {
ReadyKey string
DelayedKey string
ActiveKey string
InflightPrefix string
LockPrefix string
InflightTTL time.Duration
LockTTL time.Duration
}
func newBatchImageQueueWithOptions(rdb *redis.Client, opts batchImageQueueOptions) *batchImageQueue {
opts = normalizeBatchImageQueueOptions(opts)
return &batchImageQueue{
rdb: rdb,
readyKey: opts.ReadyKey,
delayedKey: opts.DelayedKey,
activeKey: opts.ActiveKey,
inflightPrefix: opts.InflightPrefix,
lockPrefix: opts.LockPrefix,
inflightTTL: opts.InflightTTL,
lockTTL: opts.LockTTL,
}
}
func batchImageQueueOptionsFromConfig(cfg *config.Config) batchImageQueueOptions {
if cfg == nil {
return batchImageQueueOptions{}
}
return batchImageQueueOptions{
ReadyKey: cfg.BatchImage.QueueReadyKey,
DelayedKey: cfg.BatchImage.QueueDelayedKey,
ActiveKey: cfg.BatchImage.QueueActiveKey,
InflightPrefix: cfg.BatchImage.InflightKeyPrefix,
LockPrefix: cfg.BatchImage.LockKeyPrefix,
InflightTTL: time.Duration(cfg.BatchImage.InflightTTLSeconds) * time.Second,
LockTTL: time.Duration(cfg.BatchImage.JobLockTTLSeconds) * time.Second,
}
}
func normalizeBatchImageQueueOptions(opts batchImageQueueOptions) batchImageQueueOptions {
if opts.ReadyKey == "" {
opts.ReadyKey = defaultBatchImageReadyKey
}
if opts.DelayedKey == "" {
opts.DelayedKey = defaultBatchImageDelayedKey
}
if opts.ActiveKey == "" {
opts.ActiveKey = defaultBatchImageActiveKey
}
if opts.InflightPrefix == "" {
opts.InflightPrefix = defaultBatchImageInflightPrefix
}
if opts.LockPrefix == "" {
opts.LockPrefix = defaultBatchImageLockPrefix
}
if opts.InflightTTL <= 0 {
opts.InflightTTL = defaultBatchImageInflightTTL
}
if opts.LockTTL <= 0 {
opts.LockTTL = defaultBatchImageJobLockTTL
}
return opts
}
func (q *batchImageQueue) Enqueue(ctx context.Context, batchID string) error {
if !service.IsValidBatchImageID(batchID) {
return service.ErrInvalidBatchImageQueuePayload
}
ok, err := q.rdb.SetNX(ctx, q.inflightKey(batchID), batchID, q.inflightTTL).Result()
if err != nil {
return err
}
if !ok {
return service.ErrBatchImageAlreadyQueued
}
if err := q.rdb.LPush(ctx, q.readyKey, batchID).Err(); err != nil {
_ = q.rdb.Del(ctx, q.inflightKey(batchID)).Err()
return err
}
return nil
}
func (q *batchImageQueue) Reserve(ctx context.Context, blockTimeout time.Duration) (service.ReservedBatchImageJob, error) {
result, err := q.rdb.BRPop(ctx, blockTimeout, q.readyKey).Result()
if errors.Is(err, redis.Nil) {
return service.ReservedBatchImageJob{}, service.ErrBatchImageQueueEmpty
}
if err != nil {
return service.ReservedBatchImageJob{}, err
}
if len(result) != 2 || !service.IsValidBatchImageID(result[1]) {
return service.ReservedBatchImageJob{}, service.ErrInvalidBatchImageQueuePayload
}
batchID := result[1]
if err := q.rdb.ZAdd(ctx, q.activeKey, redis.Z{
Score: float64(time.Now().UnixMilli()),
Member: batchID,
}).Err(); err != nil {
return service.ReservedBatchImageJob{}, err
}
return service.ReservedBatchImageJob{BatchID: batchID}, nil
}
func (q *batchImageQueue) RequeueAfter(ctx context.Context, batchID string, delay time.Duration) error {
if !service.IsValidBatchImageID(batchID) {
return service.ErrInvalidBatchImageQueuePayload
}
pipe := q.rdb.TxPipeline()
pipe.ZRem(ctx, q.activeKey, batchID)
pipe.ZRem(ctx, q.delayedKey, batchID)
if delay <= 0 {
pipe.LPush(ctx, q.readyKey, batchID)
} else {
pipe.ZAdd(ctx, q.delayedKey, redis.Z{
Score: float64(time.Now().Add(delay).UnixMilli()),
Member: batchID,
})
}
_, err := pipe.Exec(ctx)
return err
}
func (q *batchImageQueue) Ack(ctx context.Context, batchID string) error {
if !service.IsValidBatchImageID(batchID) {
return service.ErrInvalidBatchImageQueuePayload
}
pipe := q.rdb.TxPipeline()
pipe.ZRem(ctx, q.activeKey, batchID)
pipe.ZRem(ctx, q.delayedKey, batchID)
pipe.Del(ctx, q.inflightKey(batchID))
_, err := pipe.Exec(ctx)
return err
}
func (q *batchImageQueue) Heartbeat(ctx context.Context, batchID string) error {
if !service.IsValidBatchImageID(batchID) {
return service.ErrInvalidBatchImageQueuePayload
}
return q.rdb.ZAdd(ctx, q.activeKey, redis.Z{
Score: float64(time.Now().UnixMilli()),
Member: batchID,
}).Err()
}
func (q *batchImageQueue) MoveDueDelayedToReady(ctx context.Context, limit int) (int, error) {
if limit <= 0 {
limit = 100
}
return batchImageMoveDueDelayedScript.Run(ctx, q.rdb, []string{q.delayedKey, q.readyKey}, time.Now().UnixMilli(), limit).Int()
}
func (q *batchImageQueue) RecoverStaleActive(ctx context.Context, staleAfter time.Duration, limit int) (int, error) {
if staleAfter <= 0 {
return 0, service.ErrInvalidBatchImageQueuePayload
}
if limit <= 0 {
limit = 100
}
cutoff := time.Now().Add(-staleAfter).UnixMilli()
return batchImageRecoverStaleActiveScript.Run(ctx, q.rdb, []string{q.activeKey, q.readyKey}, cutoff, limit).Int()
}
func (q *batchImageQueue) TryAcquireJobLock(ctx context.Context, batchID string, ttl time.Duration) (service.BatchImageJobLock, bool, error) {
if !service.IsValidBatchImageID(batchID) {
return nil, false, service.ErrInvalidBatchImageQueuePayload
}
if ttl <= 0 {
ttl = q.lockTTL
}
token, err := newBatchImageLockToken()
if err != nil {
return nil, false, err
}
key := q.lockKey(batchID)
ok, err := q.rdb.SetNX(ctx, key, token, ttl).Result()
if err != nil {
return nil, false, err
}
if !ok {
return nil, false, nil
}
return &batchImageRedisJobLock{rdb: q.rdb, key: key, token: token}, true, nil
}
func (q *batchImageQueue) inflightKey(batchID string) string {
return q.inflightPrefix + batchID
}
func (q *batchImageQueue) lockKey(batchID string) string {
return q.lockPrefix + batchID
}
type batchImageRedisJobLock struct {
rdb *redis.Client
key string
token string
}
func (l *batchImageRedisJobLock) Release(ctx context.Context) error {
if l == nil || l.rdb == nil || l.key == "" || l.token == "" {
return nil
}
return batchImageReleaseLockScript.Run(ctx, l.rdb, []string{l.key}, l.token).Err()
}
func newBatchImageLockToken() (string, error) {
var b [16]byte
if _, err := rand.Read(b[:]); err != nil {
return "", err
}
return hex.EncodeToString(b[:]), nil
}
var _ service.BatchImageQueue = (*batchImageQueue)(nil)
@@ -0,0 +1,123 @@
//go:build unit
package repository
import (
"context"
"errors"
"testing"
"time"
"github.com/Wei-Shaw/sub2api/internal/service"
"github.com/alicebob/miniredis/v2"
"github.com/redis/go-redis/v9"
"github.com/stretchr/testify/require"
)
func TestBatchImageQueue_DuplicateEnqueueReturnsAlreadyQueued(t *testing.T) {
ctx := context.Background()
queue, _ := newBatchImageQueueTest(t)
batchID := "imgbatch_duplicate"
require.NoError(t, queue.Enqueue(ctx, batchID))
err := queue.Enqueue(ctx, batchID)
require.Error(t, err)
require.True(t, errors.Is(err, service.ErrBatchImageAlreadyQueued))
}
func TestBatchImageQueue_RequeueAfterMovesJobFromActiveToDelayed(t *testing.T) {
ctx := context.Background()
queue, _ := newBatchImageQueueTest(t)
batchID := "imgbatch_requeue_after"
require.NoError(t, queue.rdb.ZAdd(ctx, queue.activeKey, redis.Z{
Score: float64(time.Now().UnixMilli()),
Member: batchID,
}).Err())
require.NoError(t, queue.RequeueAfter(ctx, batchID, time.Minute))
require.ErrorIs(t, queue.rdb.ZScore(ctx, queue.activeKey, batchID).Err(), redis.Nil)
score, err := queue.rdb.ZScore(ctx, queue.delayedKey, batchID).Result()
require.NoError(t, err)
require.Greater(t, score, float64(time.Now().UnixMilli()))
}
func TestBatchImageQueue_MoveDueDelayedToReadyMovesDueJobs(t *testing.T) {
ctx := context.Background()
queue, _ := newBatchImageQueueTest(t)
dueBatchID := "imgbatch_due"
futureBatchID := "imgbatch_future"
now := time.Now()
require.NoError(t, queue.rdb.ZAdd(ctx, queue.delayedKey,
redis.Z{Score: float64(now.Add(-time.Second).UnixMilli()), Member: dueBatchID},
redis.Z{Score: float64(now.Add(time.Hour).UnixMilli()), Member: futureBatchID},
).Err())
moved, err := queue.MoveDueDelayedToReady(ctx, 10)
require.NoError(t, err)
require.Equal(t, 1, moved)
require.ErrorIs(t, queue.rdb.ZScore(ctx, queue.delayedKey, dueBatchID).Err(), redis.Nil)
require.NoError(t, queue.rdb.ZScore(ctx, queue.delayedKey, futureBatchID).Err())
reserved, err := queue.Reserve(ctx, time.Millisecond)
require.NoError(t, err)
require.Equal(t, dueBatchID, reserved.BatchID)
}
func TestBatchImageQueue_RecoverStaleActiveMovesStaleJobsToReady(t *testing.T) {
ctx := context.Background()
queue, _ := newBatchImageQueueTest(t)
staleBatchID := "imgbatch_stale"
recentBatchID := "imgbatch_recent"
now := time.Now()
require.NoError(t, queue.rdb.ZAdd(ctx, queue.activeKey,
redis.Z{Score: float64(now.Add(-time.Hour).UnixMilli()), Member: staleBatchID},
redis.Z{Score: float64(now.UnixMilli()), Member: recentBatchID},
).Err())
moved, err := queue.RecoverStaleActive(ctx, 10*time.Minute, 10)
require.NoError(t, err)
require.Equal(t, 1, moved)
require.ErrorIs(t, queue.rdb.ZScore(ctx, queue.activeKey, staleBatchID).Err(), redis.Nil)
require.NoError(t, queue.rdb.ZScore(ctx, queue.activeKey, recentBatchID).Err())
reserved, err := queue.Reserve(ctx, time.Millisecond)
require.NoError(t, err)
require.Equal(t, staleBatchID, reserved.BatchID)
}
func TestBatchImageQueue_JobLockReleaseOnlyDeletesMatchingToken(t *testing.T) {
ctx := context.Background()
queue, _ := newBatchImageQueueTest(t)
batchID := "imgbatch_lock"
lock, ok, err := queue.TryAcquireJobLock(ctx, batchID, time.Minute)
require.NoError(t, err)
require.True(t, ok)
require.NoError(t, queue.rdb.Set(ctx, queue.lockKey(batchID), "other-token", time.Minute).Err())
require.NoError(t, lock.Release(ctx))
got, err := queue.rdb.Get(ctx, queue.lockKey(batchID)).Result()
require.NoError(t, err)
require.Equal(t, "other-token", got)
require.NoError(t, queue.rdb.Del(ctx, queue.lockKey(batchID)).Err())
lock, ok, err = queue.TryAcquireJobLock(ctx, batchID, time.Minute)
require.NoError(t, err)
require.True(t, ok)
require.NoError(t, lock.Release(ctx))
require.ErrorIs(t, queue.rdb.Get(ctx, queue.lockKey(batchID)).Err(), redis.Nil)
}
func newBatchImageQueueTest(t *testing.T) (*batchImageQueue, *miniredis.Miniredis) {
t.Helper()
mr := miniredis.RunT(t)
rdb := redis.NewClient(&redis.Options{Addr: mr.Addr()})
t.Cleanup(func() {
_ = rdb.Close()
})
queue := newBatchImageQueueWithOptions(rdb, batchImageQueueOptions{
InflightTTL: time.Hour,
LockTTL: time.Minute,
})
return queue, mr
}
@@ -0,0 +1,941 @@
package repository
import (
"context"
"database/sql"
"encoding/json"
"strconv"
"time"
"github.com/Wei-Shaw/sub2api/internal/service"
)
type batchImageSQLExecutor interface {
ExecContext(ctx context.Context, query string, args ...any) (sql.Result, error)
QueryContext(ctx context.Context, query string, args ...any) (*sql.Rows, error)
QueryRowContext(ctx context.Context, query string, args ...any) *sql.Row
}
type batchImageRepository struct {
db *sql.DB
sql batchImageSQLExecutor
}
func NewBatchImageRepository(db *sql.DB) service.BatchImageRepository {
return &batchImageRepository{db: db, sql: db}
}
func (r *batchImageRepository) CreateBatchImageJob(ctx context.Context, params service.CreateBatchImageJobParams) (*service.BatchImageJob, error) {
if !service.IsSupportedBatchImageProvider(params.Provider) {
return nil, service.ErrBatchImageInvalidProvider
}
if params.BatchID == "" {
batchID, err := service.NewBatchImageID()
if err != nil {
return nil, err
}
params.BatchID = batchID
}
if params.Status == "" {
params.Status = service.BatchImageJobStatusCreated
}
if params.Currency == "" {
params.Currency = "USD"
}
job, err := createBatchImageJobWithSQL(ctx, r.sql, params)
if err != nil {
return nil, translatePersistenceError(err, nil, service.ErrBatchImageJobExists)
}
return job, nil
}
func (r *batchImageRepository) GetBatchImageJobByBatchID(ctx context.Context, batchID string) (*service.BatchImageJob, error) {
job, err := scanBatchImageJob(r.sql.QueryRowContext(ctx, batchImageJobSelectSQL+" WHERE batch_id = $1", batchID))
if err != nil {
return nil, translatePersistenceError(err, service.ErrBatchImageJobNotFound, nil)
}
return job, nil
}
func (r *batchImageRepository) GetBatchImageJobByIdempotencyKey(ctx context.Context, userID, apiKeyID int64, key string) (*service.BatchImageJob, error) {
job, err := scanBatchImageJob(r.sql.QueryRowContext(ctx, batchImageJobSelectSQL+`
WHERE user_id = $1 AND api_key_id = $2 AND idempotency_key = $3
ORDER BY id DESC LIMIT 1`, userID, apiKeyID, key))
if err != nil {
return nil, translatePersistenceError(err, service.ErrBatchImageJobNotFound, nil)
}
return job, nil
}
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 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 func() { _ = 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 {
return nil, translatePersistenceError(err, service.ErrBatchImageJobNotFound, nil)
}
return job, nil
}
func (r *batchImageRepository) TransitionBatchImageJobStatus(ctx context.Context, batchID, toStatus string, opts service.BatchImageTransitionOptions) error {
if r.db == nil {
return r.transitionBatchImageJobStatusWithSQL(ctx, r.sql, batchID, toStatus, opts)
}
tx, err := r.db.BeginTx(ctx, nil)
if err != nil {
return err
}
defer func() {
_ = tx.Rollback()
}()
if err := r.transitionBatchImageJobStatusWithSQL(ctx, tx, batchID, toStatus, opts); err != nil {
return err
}
return tx.Commit()
}
func (r *batchImageRepository) UpdateBatchImageJobProviderOutputRef(ctx context.Context, batchID, providerOutputRef string) error {
res, err := r.sql.ExecContext(ctx, `
UPDATE batch_image_jobs
SET provider_output_ref = $2, updated_at = $3
WHERE batch_id = $1`, batchID, providerOutputRef, time.Now())
if err != nil {
return translatePersistenceError(err, service.ErrBatchImageJobNotFound, nil)
}
if affected, err := res.RowsAffected(); err == nil && affected == 0 {
return service.ErrBatchImageJobNotFound
}
return nil
}
func (r *batchImageRepository) UpdateBatchImageJobProviderSubmit(ctx context.Context, params service.UpdateBatchImageJobProviderSubmitParams) error {
if r.db == nil {
return r.updateBatchImageJobProviderSubmitWithSQL(ctx, r.sql, params)
}
tx, err := r.db.BeginTx(ctx, nil)
if err != nil {
return err
}
defer func() {
_ = tx.Rollback()
}()
if err := r.updateBatchImageJobProviderSubmitWithSQL(ctx, tx, params); err != nil {
return err
}
return tx.Commit()
}
func (r *batchImageRepository) updateBatchImageJobProviderSubmitWithSQL(ctx context.Context, sqlq batchImageSQLExecutor, params service.UpdateBatchImageJobProviderSubmitParams) error {
var current string
if err := sqlq.QueryRowContext(ctx, `SELECT status FROM batch_image_jobs WHERE batch_id = $1 FOR UPDATE`, params.BatchID).Scan(&current); err != nil {
return translatePersistenceError(err, service.ErrBatchImageJobNotFound, nil)
}
if !service.CanTransitionBatchImageJob(current, service.BatchImageJobStatusSubmitted) {
return service.ErrBatchImageInvalidTransition
}
now := time.Now()
if _, err := sqlq.ExecContext(ctx, `
UPDATE batch_image_jobs
SET status = 'submitted',
provider_job_name = $2,
provider_input_ref = NULLIF($3, ''),
provider_output_ref = NULLIF($4, ''),
gcs_input_uri = NULLIF($5, ''),
gcs_output_uri = NULLIF($6, ''),
submitted_at = CASE WHEN submitted_at IS NULL THEN $7 ELSE submitted_at END,
updated_at = $7,
version = version + 1
WHERE batch_id = $1`, params.BatchID, params.ProviderJobName, params.ProviderInputRef, params.ProviderOutputRef, params.GCSInputURI, params.GCSOutputURI, now); err != nil {
return err
}
return appendBatchImageEventWithSQL(ctx, sqlq, params.BatchID, "provider_submitted", params.EventPayload)
}
func (r *batchImageRepository) RecordBatchImageJobSubmitFailure(ctx context.Context, batchID, code, message string, markFailed bool) error {
now := time.Now()
statusSQL := "status"
if markFailed {
statusSQL = "'failed'"
}
_, err := r.sql.ExecContext(ctx, `
UPDATE batch_image_jobs
SET status = `+statusSQL+`,
last_error_code = $2,
last_error_message = $3,
finished_at = CASE WHEN `+statusSQL+` = 'failed' AND finished_at IS NULL THEN $4 ELSE finished_at END,
updated_at = $4,
version = version + 1
WHERE batch_id = $1`, batchID, code, message, now)
if err != nil {
return err
}
eventType := "submit_failed"
if !markFailed {
eventType = "queue_failed"
}
return appendBatchImageEventWithSQL(ctx, r.sql, batchID, eventType, map[string]any{"error_code": code})
}
func (r *batchImageRepository) MarkBatchImageJobSettled(ctx context.Context, params service.MarkBatchImageJobSettledParams) error {
if r.db == nil {
return r.markBatchImageJobSettledWithSQL(ctx, r.sql, params)
}
tx, err := r.db.BeginTx(ctx, nil)
if err != nil {
return err
}
defer func() {
_ = tx.Rollback()
}()
if err := r.markBatchImageJobSettledWithSQL(ctx, tx, params); err != nil {
return err
}
return tx.Commit()
}
func (r *batchImageRepository) markBatchImageJobSettledWithSQL(ctx context.Context, sqlq batchImageSQLExecutor, params service.MarkBatchImageJobSettledParams) error {
now := time.Now()
if params.Now != nil {
now = *params.Now
}
outputExpiresAt := params.OutputExpiresAt
res, err := sqlq.ExecContext(ctx, `
UPDATE batch_image_jobs
SET status = 'completed',
actual_cost = $2,
manifest_hash = $3,
settled_at = CASE WHEN settled_at IS NULL THEN $4 ELSE settled_at END,
finished_at = CASE WHEN finished_at IS NULL THEN $4 ELSE finished_at END,
output_expires_at = CASE WHEN output_expires_at IS NULL THEN $5 ELSE output_expires_at END,
updated_at = $4,
version = version + 1
WHERE batch_id = $1
AND status = 'settling'
AND (manifest_hash IS NULL OR manifest_hash = '' OR manifest_hash = $3)`, params.BatchID, params.ActualCost, params.ManifestHash, now, outputExpiresAt)
if err != nil {
return err
}
affected, err := res.RowsAffected()
if err != nil {
return err
}
if affected == 0 {
job, getErr := scanBatchImageJob(sqlq.QueryRowContext(ctx, batchImageJobSelectSQL+" WHERE batch_id = $1", params.BatchID))
if getErr != nil {
return translatePersistenceError(getErr, service.ErrBatchImageJobNotFound, nil)
}
if job.Status != service.BatchImageJobStatusSettling {
if job.Status == service.BatchImageJobStatusCompleted {
return service.ErrBatchImageAlreadySettled
}
return service.ErrBatchImageSettlementInvalidStatus
}
return service.ErrBatchImageSettlementManifestConflict
}
return appendBatchImageEventWithSQL(ctx, sqlq, params.BatchID, "settlement_completed", params.EventPayload)
}
func (r *batchImageRepository) SetBatchImageJobSettlementFailed(ctx context.Context, batchID, code, message string) (int, error) {
var retryCount int
err := r.sql.QueryRowContext(ctx, `
UPDATE batch_image_jobs
SET last_error_code = $2,
last_error_message = $3,
retry_count = retry_count + 1,
updated_at = $4
WHERE batch_id = $1
RETURNING retry_count`, batchID, code, message, time.Now()).Scan(&retryCount)
if err != nil {
return 0, translatePersistenceError(err, service.ErrBatchImageJobNotFound, nil)
}
return retryCount, appendBatchImageEventWithSQL(ctx, r.sql, batchID, "settlement_failed", map[string]any{
"error_code": code,
})
}
func (r *batchImageRepository) transitionBatchImageJobStatusWithSQL(ctx context.Context, sqlq batchImageSQLExecutor, batchID, toStatus string, opts service.BatchImageTransitionOptions) error {
var current string
if err := sqlq.QueryRowContext(ctx, `SELECT status FROM batch_image_jobs WHERE batch_id = $1 FOR UPDATE`, batchID).Scan(&current); err != nil {
return translatePersistenceError(err, service.ErrBatchImageJobNotFound, nil)
}
if !service.CanTransitionBatchImageJob(current, toStatus) {
return service.ErrBatchImageInvalidTransition
}
now := time.Now()
if opts.Now != nil {
now = *opts.Now
}
if _, err := sqlq.ExecContext(ctx, `
UPDATE batch_image_jobs
SET
status = $2::varchar,
version = version + 1,
updated_at = $3,
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
}
if opts.EventType != "" {
return appendBatchImageEventWithSQL(ctx, sqlq, batchID, opts.EventType, opts.EventPayload)
}
return nil
}
func (r *batchImageRepository) CreateBatchImageItem(ctx context.Context, params service.CreateBatchImageItemParams) (*service.BatchImageItem, error) {
item, err := createBatchImageItemWithSQL(ctx, r.sql, params)
if err != nil {
return nil, translatePersistenceError(err, nil, service.ErrBatchImageItemExists)
}
return item, nil
}
func (r *batchImageRepository) BulkCreateBatchImageItems(ctx context.Context, params []service.CreateBatchImageItemParams) error {
if len(params) == 0 {
return nil
}
if r.db == nil {
for _, param := range params {
if _, err := createBatchImageItemWithSQL(ctx, r.sql, param); err != nil {
return translatePersistenceError(err, nil, service.ErrBatchImageItemExists)
}
}
return nil
}
tx, err := r.db.BeginTx(ctx, nil)
if err != nil {
return err
}
defer func() {
_ = tx.Rollback()
}()
for _, param := range params {
if _, err := createBatchImageItemWithSQL(ctx, tx, param); err != nil {
return translatePersistenceError(err, nil, service.ErrBatchImageItemExists)
}
}
return tx.Commit()
}
func (r *batchImageRepository) ReplaceBatchImageItemsForJob(ctx context.Context, batchID string, items []service.CreateBatchImageItemParams, counts service.BatchImageCounts) error {
if r.db == nil {
return r.replaceBatchImageItemsForJobWithSQL(ctx, r.sql, batchID, items, counts)
}
tx, err := r.db.BeginTx(ctx, nil)
if err != nil {
return err
}
defer func() {
_ = tx.Rollback()
}()
if err := r.replaceBatchImageItemsForJobWithSQL(ctx, tx, batchID, items, counts); err != nil {
return err
}
return tx.Commit()
}
func (r *batchImageRepository) replaceBatchImageItemsForJobWithSQL(ctx context.Context, sqlq batchImageSQLExecutor, batchID string, items []service.CreateBatchImageItemParams, counts service.BatchImageCounts) error {
var id int64
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, `
UPDATE batch_image_jobs
SET success_count = $2,
fail_count = $3,
updated_at = $4
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 func() { _ = 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 {
limit = 100
}
if filter.Offset < 0 {
filter.Offset = 0
}
query := batchImageItemSelectSQL + " WHERE job_id = $1"
args := []any{batchID}
if filter.Status != "" {
query += " AND status = $2"
args = append(args, filter.Status)
}
query += " ORDER BY id ASC 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 func() { _ = rows.Close() }()
var items []*service.BatchImageItem
for rows.Next() {
item, err := scanBatchImageItem(rows)
if err != nil {
return nil, err
}
items = append(items, item)
}
if err := rows.Err(); err != nil {
return nil, err
}
return items, nil
}
func (r *batchImageRepository) ListBatchImageItemsForOwner(ctx context.Context, userID, apiKeyID int64, batchID string, filter service.BatchImageItemFilter) ([]*service.BatchImageItem, error) {
if _, err := r.GetBatchImageJobByBatchIDForOwner(ctx, userID, apiKeyID, batchID); err != nil {
return nil, err
}
return r.ListBatchImageItems(ctx, batchID, filter)
}
func (r *batchImageRepository) GetBatchImageJobForDownload(ctx context.Context, userID, apiKeyID int64, batchID string) (*service.BatchImageJob, error) {
return r.GetBatchImageJobByBatchIDForOwner(ctx, userID, apiKeyID, batchID)
}
func (r *batchImageRepository) GetBatchImageItemForDownload(ctx context.Context, batchID, customID string) (*service.BatchImageItem, error) {
item, err := scanBatchImageItem(r.sql.QueryRowContext(ctx, batchImageItemSelectSQL+`
WHERE job_id = $1 AND custom_id = $2`, batchID, customID))
if err != nil {
return nil, translatePersistenceError(err, service.ErrBatchImageItemNotFound, nil)
}
return item, nil
}
func (r *batchImageRepository) ListBatchImageItemsForDownload(ctx context.Context, batchID string, status string, limit int) ([]*service.BatchImageItem, error) {
return r.ListBatchImageItems(ctx, batchID, service.BatchImageItemFilter{Status: status, Limit: limit})
}
func (r *batchImageRepository) ListBatchImageJobsDueForInputCleanup(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 input_deleted_at IS NULL
AND provider_input_ref IS NOT NULL
AND status IN ('completed', 'failed', 'cancelled', 'output_deleted')
AND COALESCE(finished_at, settled_at, updated_at, created_at) <= $1
ORDER BY id ASC
LIMIT $2`, cutoff, limit)
if err != nil {
return nil, err
}
defer func() { _ = rows.Close() }()
return scanBatchImageJobs(rows)
}
func (r *batchImageRepository) ListBatchImageJobsDueForOutputCleanup(ctx context.Context, now time.Time, limit int) ([]*service.BatchImageJob, error) {
if limit <= 0 || limit > 1000 {
limit = 100
}
rows, err := r.sql.QueryContext(ctx, batchImageJobSelectSQL+`
WHERE output_deleted_at IS NULL
AND provider_output_ref IS NOT NULL
AND status = 'completed'
AND output_expires_at IS NOT NULL
AND output_expires_at <= $1
ORDER BY output_expires_at ASC, id ASC
LIMIT $2`, now, limit)
if err != nil {
return nil, err
}
defer func() { _ = rows.Close() }()
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 func() { _ = 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
SET input_deleted_at = CASE WHEN input_deleted_at IS NULL THEN $2 ELSE input_deleted_at END,
updated_at = $2,
version = version + 1
WHERE batch_id = $1`, batchID, deletedAt)
if err != nil {
return err
}
if affected, err := res.RowsAffected(); err == nil && affected == 0 {
return service.ErrBatchImageJobNotFound
}
return appendBatchImageEventWithSQL(ctx, r.sql, batchID, "input_cleanup_completed", map[string]any{
"batch_id": batchID,
"cleanup_target": "input",
"deleted_at": deletedAt.UTC().Format(time.RFC3339),
})
}
func (r *batchImageRepository) MarkBatchImageOutputDeleted(ctx context.Context, batchID string, deletedAt time.Time) error {
res, err := r.sql.ExecContext(ctx, `
UPDATE batch_image_jobs
SET status = CASE WHEN status = 'completed' THEN 'output_deleted' ELSE status END,
output_deleted_at = CASE WHEN output_deleted_at IS NULL THEN $2 ELSE output_deleted_at END,
finished_at = CASE WHEN status = 'completed' AND finished_at IS NULL THEN $2 ELSE finished_at END,
updated_at = $2,
version = version + 1
WHERE batch_id = $1`, batchID, deletedAt)
if err != nil {
return err
}
if affected, err := res.RowsAffected(); err == nil && affected == 0 {
return service.ErrBatchImageJobNotFound
}
return appendBatchImageEventWithSQL(ctx, r.sql, batchID, "output_cleanup_completed", map[string]any{
"batch_id": batchID,
"cleanup_target": "output",
"deleted_at": deletedAt.UTC().Format(time.RFC3339),
})
}
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
SET output_expires_at = CASE WHEN output_expires_at IS NULL THEN $2 ELSE output_expires_at END,
updated_at = $3
WHERE batch_id = $1`, batchID, expiresAt, time.Now())
if err != nil {
return err
}
if affected, err := res.RowsAffected(); err == nil && affected == 0 {
return service.ErrBatchImageJobNotFound
}
return nil
}
func (r *batchImageRepository) RecordBatchImageCleanupFailure(ctx context.Context, batchID, code, message string) error {
_, err := r.sql.ExecContext(ctx, `
UPDATE batch_image_jobs
SET last_error_code = $2,
last_error_message = $3,
retry_count = retry_count + 1,
updated_at = $4
WHERE batch_id = $1`, batchID, code, message, time.Now())
if err != nil {
return err
}
return appendBatchImageEventWithSQL(ctx, r.sql, batchID, "output_cleanup_failed", map[string]any{"error_code": code})
}
func (r *batchImageRepository) AppendBatchImageEvent(ctx context.Context, batchID, eventType string, payload any) error {
return appendBatchImageEventWithSQL(ctx, r.sql, batchID, eventType, payload)
}
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, 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,
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, $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.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.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,
))
}
func createBatchImageItemWithSQL(ctx context.Context, sqlq batchImageSQLExecutor, params service.CreateBatchImageItemParams) (*service.BatchImageItem, error) {
return scanBatchImageItem(sqlq.QueryRowContext(ctx, `
INSERT INTO batch_image_items (
job_id, custom_id, status, request_hash, prompt_preview, provider_source_object,
source_line_number, source_byte_offset, source_byte_length,
mime_type, file_extension, image_count,
error_code, error_message, billed_amount, indexed_at
) VALUES (
$1, $2, $3, $4, $5, $6,
$7, $8, $9,
$10, $11, $12,
$13, $14, $15, $16
)
RETURNING `+batchImageItemColumns,
params.JobID, params.CustomID, params.Status, params.RequestHash, params.PromptPreview, params.ProviderSourceObject,
params.SourceLineNumber, params.SourceByteOffset, params.SourceByteLength,
params.MimeType, params.FileExtension, params.ImageCount,
params.ErrorCode, params.ErrorMessage, params.BilledAmount, params.IndexedAt,
))
}
func appendBatchImageEventWithSQL(ctx context.Context, sqlq batchImageSQLExecutor, batchID, eventType string, payload any) error {
var payloadArg any
if payload != nil {
payloadBytes, err := json.Marshal(payload)
if err != nil {
return err
}
payloadArg = string(payloadBytes)
}
_, err := sqlq.ExecContext(ctx, `
INSERT INTO batch_image_events (job_id, event_type, payload)
VALUES ($1, $2, $3)`, batchID, eventType, payloadArg)
return err
}
type rowScanner interface {
Scan(dest ...any) error
}
const batchImageJobColumns = `
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,
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, downloaded_at, user_deleted_at,
last_error_code, last_error_message,
created_at, updated_at, submitted_at, started_at, finished_at, settled_at`
const batchImageJobSelectSQL = `SELECT ` + batchImageJobColumns + ` FROM batch_image_jobs`
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, 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.TaskName, &parentBatchID, &job.Status,
&providerJobName, &providerInputRef, &providerOutputRef, &gcsInputURI, &gcsOutputURI,
&job.ItemCount, &job.SuccessCount, &job.FailCount, &job.CancelledCount,
&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, &downloadedAt, &userDeletedAt,
&lastErrorCode, &lastErrorMessage,
&job.CreatedAt, &job.UpdatedAt, &submittedAt, &startedAt, &finishedAt, &settledAt,
)
if err != nil {
return nil, err
}
job.APIKeyID = batchImageNullInt64Ptr(apiKeyID)
job.AccountID = batchImageNullInt64Ptr(accountID)
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)
job.ActualCost = batchImageNullFloat64Ptr(actualCost)
job.HoldID = batchImageNullStringPtr(holdID)
job.IdempotencyKey = batchImageNullStringPtr(idempotencyKey)
job.RequestHash = batchImageNullStringPtr(requestHash)
job.ManifestHash = batchImageNullStringPtr(manifestHash)
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)
job.StartedAt = batchImageNullTimePtr(startedAt)
job.FinishedAt = batchImageNullTimePtr(finishedAt)
job.SettledAt = batchImageNullTimePtr(settledAt)
return &job, nil
}
func scanBatchImageJobs(rows *sql.Rows) ([]*service.BatchImageJob, error) {
var jobs []*service.BatchImageJob
for rows.Next() {
job, err := scanBatchImageJob(rows)
if err != nil {
return nil, err
}
jobs = append(jobs, job)
}
if err := rows.Err(); err != nil {
return nil, err
}
return jobs, nil
}
const batchImageItemColumns = `
id, job_id, custom_id, status, request_hash, prompt_preview, provider_source_object,
source_line_number, source_byte_offset, source_byte_length,
mime_type, file_extension, image_count,
error_code, error_message, billed_amount,
created_at, indexed_at`
const batchImageItemSelectSQL = `SELECT ` + batchImageItemColumns + ` FROM batch_image_items`
func scanBatchImageItem(row rowScanner) (*service.BatchImageItem, error) {
var item service.BatchImageItem
var requestHash, promptPreview, providerSourceObject sql.NullString
var sourceLineNumber sql.NullInt64
var sourceByteOffset, sourceByteLength sql.NullInt64
var mimeType, fileExtension, errorCode, errorMessage sql.NullString
var billedAmount sql.NullFloat64
var indexedAt sql.NullTime
err := row.Scan(
&item.ID, &item.JobID, &item.CustomID, &item.Status, &requestHash, &promptPreview, &providerSourceObject,
&sourceLineNumber, &sourceByteOffset, &sourceByteLength,
&mimeType, &fileExtension, &item.ImageCount,
&errorCode, &errorMessage, &billedAmount,
&item.CreatedAt, &indexedAt,
)
if err != nil {
return nil, err
}
item.RequestHash = batchImageNullStringPtr(requestHash)
item.PromptPreview = batchImageNullStringPtr(promptPreview)
item.ProviderSourceObject = batchImageNullStringPtr(providerSourceObject)
item.SourceLineNumber = batchImageNullIntPtr(sourceLineNumber)
item.SourceByteOffset = batchImageNullInt64Ptr(sourceByteOffset)
item.SourceByteLength = batchImageNullInt64Ptr(sourceByteLength)
item.MimeType = batchImageNullStringPtr(mimeType)
item.FileExtension = batchImageNullStringPtr(fileExtension)
item.ErrorCode = batchImageNullStringPtr(errorCode)
item.ErrorMessage = batchImageNullStringPtr(errorMessage)
item.BilledAmount = batchImageNullFloat64Ptr(billedAmount)
item.IndexedAt = batchImageNullTimePtr(indexedAt)
return &item, nil
}
func batchImageNullStringPtr(v sql.NullString) *string {
if !v.Valid {
return nil
}
return &v.String
}
func batchImageNullInt64Ptr(v sql.NullInt64) *int64 {
if !v.Valid {
return nil
}
return &v.Int64
}
func batchImageNullIntPtr(v sql.NullInt64) *int {
if !v.Valid {
return nil
}
i := int(v.Int64)
return &i
}
func batchImageNullFloat64Ptr(v sql.NullFloat64) *float64 {
if !v.Valid {
return nil
}
return &v.Float64
}
func batchImageNullTimePtr(v sql.NullTime) *time.Time {
if !v.Valid {
return nil
}
return &v.Time
}
var _ service.BatchImageRepository = (*batchImageRepository)(nil)
@@ -0,0 +1,371 @@
//go:build integration
package repository
import (
"context"
"crypto/sha1"
"encoding/hex"
"errors"
"regexp"
"strings"
"testing"
"time"
"github.com/Wei-Shaw/sub2api/internal/service"
"github.com/stretchr/testify/require"
)
func newBatchImageRepositoryWithSQL(sqlq batchImageSQLExecutor) *batchImageRepository {
return &batchImageRepository{sql: sqlq}
}
func TestBatchImageRepository_CreateJobAndDuplicates(t *testing.T) {
ctx := context.Background()
tx := testTx(t)
repo := newBatchImageRepositoryWithSQL(tx)
batchID := batchImageTestID(t, "create")
job, err := repo.CreateBatchImageJob(ctx, service.CreateBatchImageJobParams{
BatchID: batchID,
UserID: 1001,
Provider: service.BatchImageProviderGeminiAPI,
Model: "gemini-2.5-flash-image",
ItemCount: 2,
EstimatedCost: 0.02,
})
require.NoError(t, err)
require.Equal(t, batchID, job.BatchID)
require.Equal(t, service.BatchImageJobStatusCreated, job.Status)
require.Equal(t, "USD", job.Currency)
_, err = repo.CreateBatchImageJob(ctx, service.CreateBatchImageJobParams{
BatchID: batchID,
UserID: 1001,
Provider: service.BatchImageProviderGeminiAPI,
Model: "gemini-2.5-flash-image",
ItemCount: 1,
})
require.Error(t, err)
require.True(t, errors.Is(err, service.ErrBatchImageJobExists))
}
func TestBatchImageRepository_InvalidProvider(t *testing.T) {
tx := testTx(t)
repo := newBatchImageRepositoryWithSQL(tx)
_, err := repo.CreateBatchImageJob(context.Background(), service.CreateBatchImageJobParams{
BatchID: batchImageTestID(t, "provider"),
UserID: 1001,
Provider: "unknown",
Model: "gemini-2.5-flash-image",
ItemCount: 1,
})
require.Error(t, err)
require.True(t, errors.Is(err, service.ErrBatchImageInvalidProvider))
}
func TestBatchImageRepository_TransitionIncrementsVersionAndEvents(t *testing.T) {
ctx := context.Background()
tx := testTx(t)
repo := newBatchImageRepositoryWithSQL(tx)
batchID := batchImageTestID(t, "transition")
now := time.Date(2026, 7, 3, 8, 0, 0, 0, time.UTC)
_, err := repo.CreateBatchImageJob(ctx, service.CreateBatchImageJobParams{
BatchID: batchID,
UserID: 1001,
Provider: service.BatchImageProviderVertex,
Model: "gemini-2.5-flash-image",
ItemCount: 1,
})
require.NoError(t, err)
err = repo.TransitionBatchImageJobStatus(ctx, batchID, service.BatchImageJobStatusUploading, service.BatchImageTransitionOptions{
EventType: "status_changed",
EventPayload: map[string]any{"to": service.BatchImageJobStatusUploading},
Now: &now,
})
require.NoError(t, err)
job, err := repo.GetBatchImageJobByBatchID(ctx, batchID)
require.NoError(t, err)
require.Equal(t, service.BatchImageJobStatusUploading, job.Status)
require.Equal(t, 1, job.Version)
var eventCount int
err = tx.QueryRowContext(ctx, `SELECT COUNT(*) FROM batch_image_events WHERE job_id = $1 AND event_type = 'status_changed'`, batchID).Scan(&eventCount)
require.NoError(t, err)
require.Equal(t, 1, eventCount)
}
func TestBatchImageRepository_InvalidTransition(t *testing.T) {
ctx := context.Background()
tx := testTx(t)
repo := newBatchImageRepositoryWithSQL(tx)
batchID := batchImageTestID(t, "invalid-transition")
_, err := repo.CreateBatchImageJob(ctx, service.CreateBatchImageJobParams{
BatchID: batchID,
UserID: 1001,
Provider: service.BatchImageProviderGeminiAPI,
Model: "gemini-2.5-flash-image",
ItemCount: 1,
})
require.NoError(t, err)
err = repo.TransitionBatchImageJobStatus(ctx, batchID, service.BatchImageJobStatusRunning, service.BatchImageTransitionOptions{})
require.Error(t, err)
require.True(t, errors.Is(err, service.ErrBatchImageInvalidTransition))
}
func TestBatchImageRepository_TerminalStatusCannotMoveBack(t *testing.T) {
ctx := context.Background()
tx := testTx(t)
repo := newBatchImageRepositoryWithSQL(tx)
batchID := batchImageTestID(t, "terminal")
_, err := repo.CreateBatchImageJob(ctx, service.CreateBatchImageJobParams{
BatchID: batchID,
UserID: 1001,
Provider: service.BatchImageProviderGeminiAPI,
Model: "gemini-2.5-flash-image",
Status: service.BatchImageJobStatusCompleted,
ItemCount: 1,
})
require.NoError(t, err)
err = repo.TransitionBatchImageJobStatus(ctx, batchID, service.BatchImageJobStatusRunning, service.BatchImageTransitionOptions{})
require.Error(t, err)
require.True(t, errors.Is(err, service.ErrBatchImageInvalidTransition))
}
func TestBatchImageRepository_ItemCustomIDUniqueness(t *testing.T) {
ctx := context.Background()
tx := testTx(t)
repo := newBatchImageRepositoryWithSQL(tx)
firstBatchID := batchImageTestID(t, "items-a")
secondBatchID := batchImageTestID(t, "items-b")
for _, batchID := range []string{firstBatchID, secondBatchID} {
_, err := repo.CreateBatchImageJob(ctx, service.CreateBatchImageJobParams{
BatchID: batchID,
UserID: 1001,
Provider: service.BatchImageProviderGeminiAPI,
Model: "gemini-2.5-flash-image",
ItemCount: 1,
})
require.NoError(t, err)
}
_, err := repo.CreateBatchImageItem(ctx, service.CreateBatchImageItemParams{
JobID: firstBatchID,
CustomID: "line-1",
Status: service.BatchImageItemStatusSuccess,
ImageCount: 1,
})
require.NoError(t, err)
_, err = tx.ExecContext(ctx, `SAVEPOINT batch_image_duplicate_item`)
require.NoError(t, err)
_, err = repo.CreateBatchImageItem(ctx, service.CreateBatchImageItemParams{
JobID: firstBatchID,
CustomID: "line-1",
Status: service.BatchImageItemStatusFailed,
})
require.Error(t, err)
require.True(t, errors.Is(err, service.ErrBatchImageItemExists))
_, rollbackErr := tx.ExecContext(ctx, `ROLLBACK TO SAVEPOINT batch_image_duplicate_item`)
require.NoError(t, rollbackErr)
_, err = repo.CreateBatchImageItem(ctx, service.CreateBatchImageItemParams{
JobID: secondBatchID,
CustomID: "line-1",
Status: service.BatchImageItemStatusSuccess,
ImageCount: 1,
})
require.NoError(t, err)
items, err := repo.ListBatchImageItems(ctx, firstBatchID, service.BatchImageItemFilter{})
require.NoError(t, err)
require.Len(t, items, 1)
}
func TestBatchImageRepository_ReplaceBatchImageItemsForJob(t *testing.T) {
ctx := context.Background()
tx := testTx(t)
repo := newBatchImageRepositoryWithSQL(tx)
batchID := batchImageTestID(t, "replace-items")
lineOne := 1
lineTwo := 2
_, err := repo.CreateBatchImageJob(ctx, service.CreateBatchImageJobParams{
BatchID: batchID,
UserID: 1001,
Provider: service.BatchImageProviderGeminiAPI,
Model: "gemini-2.5-flash-image",
ItemCount: 2,
})
require.NoError(t, err)
err = repo.ReplaceBatchImageItemsForJob(ctx, batchID, []service.CreateBatchImageItemParams{
{CustomID: "old", Status: service.BatchImageItemStatusSuccess, SourceLineNumber: &lineOne, ImageCount: 1},
}, service.BatchImageCounts{SuccessCount: 1})
require.NoError(t, err)
err = repo.ReplaceBatchImageItemsForJob(ctx, batchID, []service.CreateBatchImageItemParams{
{CustomID: "new-ok", Status: service.BatchImageItemStatusSuccess, SourceLineNumber: &lineOne, ImageCount: 1},
{CustomID: "new-fail", Status: service.BatchImageItemStatusFailed, SourceLineNumber: &lineTwo, ErrorCode: batchImageTestStringPtr("SAFETY_BLOCKED")},
}, service.BatchImageCounts{SuccessCount: 1, FailCount: 1})
require.NoError(t, err)
items, err := repo.ListBatchImageItems(ctx, batchID, service.BatchImageItemFilter{})
require.NoError(t, err)
require.Len(t, items, 2)
require.Equal(t, "new-ok", items[0].CustomID)
require.Equal(t, "new-fail", items[1].CustomID)
job, err := repo.GetBatchImageJobByBatchID(ctx, batchID)
require.NoError(t, err)
require.Equal(t, 1, job.SuccessCount)
require.Equal(t, 1, job.FailCount)
}
func TestBatchImageRepository_MarkBatchImageJobSettled(t *testing.T) {
ctx := context.Background()
tx := testTx(t)
repo := newBatchImageRepositoryWithSQL(tx)
batchID := batchImageTestID(t, "settled")
apiKeyID := int64(2001)
accountID := int64(3001)
providerJob := "providers/job"
outputRef := "files/output"
now := time.Date(2026, 7, 4, 10, 0, 0, 0, time.UTC)
_, err := repo.CreateBatchImageJob(ctx, service.CreateBatchImageJobParams{
BatchID: batchID,
UserID: 1001,
APIKeyID: &apiKeyID,
AccountID: &accountID,
Provider: service.BatchImageProviderGeminiAPI,
Model: "gemini-image",
Status: service.BatchImageJobStatusSettling,
ProviderJobName: &providerJob,
ProviderOutputRef: &outputRef,
ItemCount: 3,
SuccessCount: 2,
FailCount: 1,
})
require.NoError(t, err)
err = repo.MarkBatchImageJobSettled(ctx, service.MarkBatchImageJobSettledParams{
BatchID: batchID,
ActualCost: 0.5,
ManifestHash: "manifest-hash",
EventPayload: map[string]any{"request_id": "batch_image_settlement:" + batchID},
Now: &now,
})
require.NoError(t, err)
job, err := repo.GetBatchImageJobByBatchID(ctx, batchID)
require.NoError(t, err)
require.Equal(t, service.BatchImageJobStatusCompleted, job.Status)
require.NotNil(t, job.ActualCost)
require.Equal(t, 0.5, *job.ActualCost)
require.Equal(t, "manifest-hash", batchImageDerefTest(job.ManifestHash))
require.NotNil(t, job.SettledAt)
require.Equal(t, now, *job.SettledAt)
var eventCount int
err = tx.QueryRowContext(ctx, `SELECT COUNT(*) FROM batch_image_events WHERE job_id = $1 AND event_type = 'settlement_completed'`, batchID).Scan(&eventCount)
require.NoError(t, err)
require.Equal(t, 1, eventCount)
}
func TestBatchImageRepository_SetBatchImageJobSettlementFailed(t *testing.T) {
ctx := context.Background()
tx := testTx(t)
repo := newBatchImageRepositoryWithSQL(tx)
batchID := batchImageTestID(t, "settlement-failed")
_, err := repo.CreateBatchImageJob(ctx, service.CreateBatchImageJobParams{
BatchID: batchID,
UserID: 1001,
Provider: service.BatchImageProviderGeminiAPI,
Model: "gemini-image",
Status: service.BatchImageJobStatusSettling,
ItemCount: 1,
SuccessCount: 1,
})
require.NoError(t, err)
retryCount, err := repo.SetBatchImageJobSettlementFailed(ctx, batchID, "SETTLEMENT_BILLING_FAILED", "temporary")
require.NoError(t, err)
require.Equal(t, 1, retryCount)
job, err := repo.GetBatchImageJobByBatchID(ctx, batchID)
require.NoError(t, err)
require.Equal(t, service.BatchImageJobStatusSettling, job.Status)
require.Equal(t, "SETTLEMENT_BILLING_FAILED", batchImageDerefTest(job.LastErrorCode))
require.Equal(t, "temporary", batchImageDerefTest(job.LastErrorMessage))
require.Equal(t, 1, job.RetryCount)
}
func TestBatchImageRepository_AppendEvent(t *testing.T) {
ctx := context.Background()
tx := testTx(t)
repo := newBatchImageRepositoryWithSQL(tx)
batchID := batchImageTestID(t, "event")
_, err := repo.CreateBatchImageJob(ctx, service.CreateBatchImageJobParams{
BatchID: batchID,
UserID: 1001,
Provider: service.BatchImageProviderVertex,
Model: "gemini-2.5-flash-image",
ItemCount: 1,
})
require.NoError(t, err)
err = repo.AppendBatchImageEvent(ctx, batchID, "job_created", map[string]any{"batch_id": batchID})
require.NoError(t, err)
var payload string
err = tx.QueryRowContext(ctx, `SELECT payload::text FROM batch_image_events WHERE job_id = $1 AND event_type = 'job_created'`, batchID).Scan(&payload)
require.NoError(t, err)
require.Contains(t, payload, batchID)
}
func batchImageTestID(t *testing.T, prefix string) string {
t.Helper()
safePrefix := batchImageSafeTestIDSegment(prefix, 20)
sum := sha1.Sum([]byte(t.Name()))
return "imgbatch_" + safePrefix + "_" + hex.EncodeToString(sum[:])[:16]
}
func batchImageSafeTestIDSegment(v string, maxLen int) string {
v = strings.ToLower(strings.TrimSpace(v))
v = regexp.MustCompile(`[^a-z0-9_-]+`).ReplaceAllString(v, "-")
v = strings.Trim(v, "-_")
if v == "" {
v = "job"
}
if len(v) > maxLen {
v = v[:maxLen]
v = strings.Trim(v, "-_")
}
if v == "" {
return "job"
}
return v
}
func batchImageTestStringPtr(v string) *string {
return &v
}
func batchImageDerefTest(v *string) string {
if v == nil {
return ""
}
return *v
}
@@ -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).
@@ -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 迁移文件应用到指定的数据库。
@@ -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, `
@@ -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())
}
+3
View File
@@ -77,6 +77,7 @@ var ProviderSet = wire.NewSet(
NewAnnouncementReadRepository,
NewUsageLogRepository,
NewUsageBillingRepository,
NewBatchImageRepository,
NewIdempotencyRepository,
NewUsageCleanupRepository,
NewDashboardAggregationRepository,
@@ -115,6 +116,8 @@ var ProviderSet = wire.NewSet(
NewRedeemCache,
NewUpdateCache,
NewGeminiTokenCache,
NewBatchImageQueue,
NewBatchImageDownloadLimiter,
NewLeaderLockCache,
ProvideSchedulerCache,
NewSchedulerOutboxRepository,
+7 -3
View File
@@ -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,
@@ -361,6 +362,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,
@@ -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 {
@@ -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
}
@@ -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)
@@ -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)))
+10
View File
@@ -164,6 +164,16 @@ 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)
}
+48 -6
View File
@@ -215,9 +215,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
@@ -261,9 +264,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
@@ -1857,6 +1863,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 {
@@ -1892,6 +1912,7 @@ func (s *adminServiceImpl) CreateGroup(ctx context.Context, input *CreateGroupIn
}
allowImageGeneration := input.AllowImageGeneration || defaultAllowImageGenerationForPlatform(platform)
allowBatchImageGeneration := input.AllowBatchImageGeneration && allowImageGeneration && platform == PlatformGemini
// 如果指定了复制账号的源分组,先获取账号 ID 列表
var accountIDsToCopy []int64
@@ -1937,8 +1958,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,
@@ -2123,6 +2147,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.Platform != PlatformGemini {
group.AllowBatchImageGeneration = false
}
if input.ImageRateIndependent != nil {
group.ImageRateIndependent = *input.ImageRateIndependent
}
@@ -2132,6 +2162,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
}
@@ -232,6 +232,46 @@ 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)
}
func TestAdminService_CreateGroup_DisablesBatchImageForNonGeminiPlatform(t *testing.T) {
repo := &groupRepoStubForAdmin{}
svc := &adminServiceImpl{groupRepo: repo}
group, err := svc.CreateGroup(context.Background(), &CreateGroupInput{
Name: "openai-image",
Description: "OpenAI image group",
Platform: PlatformOpenAI,
RateMultiplier: 1.0,
AllowImageGeneration: true,
AllowBatchImageGeneration: true,
})
require.NoError(t, err)
require.NotNil(t, group)
require.NotNil(t, repo.created)
require.True(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 +366,53 @@ 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_DisablesBatchImageWhenPlatformChangesFromGemini(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}
group, err := svc.UpdateGroup(context.Background(), 1, &UpdateGroupInput{
Platform: PlatformOpenAI,
})
require.NoError(t, err)
require.NotNil(t, group)
require.NotNil(t, repo.updated)
require.Equal(t, PlatformOpenAI, repo.updated.Platform)
require.False(t, repo.updated.AllowBatchImageGeneration)
require.False(t, group.AllowBatchImageGeneration)
}
func TestAdminService_UpdateGroup_ClearsDescriptionWhenEmptyString(t *testing.T) {
existingGroup := &Group{
ID: 1,
@@ -384,6 +471,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,
@@ -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"`
@@ -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,
+406
View File
@@ -0,0 +1,406 @@
package service
import (
"context"
"crypto/rand"
"encoding/hex"
"net/http"
"time"
infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors"
)
const (
BatchImageProviderGeminiAPI = "gemini_api"
BatchImageProviderVertex = "vertex"
)
const (
BatchImageJobStatusCreated = "created"
BatchImageJobStatusUploading = "uploading"
BatchImageJobStatusSubmitted = "submitted"
BatchImageJobStatusRunning = "running"
BatchImageJobStatusIndexing = "indexing"
BatchImageJobStatusSettling = "settling"
BatchImageJobStatusCompleted = "completed"
BatchImageJobStatusFailed = "failed"
BatchImageJobStatusCancelled = "cancelled"
BatchImageJobStatusOutputDeleted = "output_deleted"
)
const (
BatchImageItemStatusPending = "pending"
BatchImageItemStatusSuccess = "success"
BatchImageItemStatusFailed = "failed"
BatchImageItemStatusCancelled = "cancelled"
)
var (
ErrBatchImageJobNotFound = infraerrors.New(http.StatusNotFound, "BATCH_IMAGE_JOB_NOT_FOUND", "batch image job not found")
ErrBatchImageJobExists = infraerrors.New(http.StatusConflict, "BATCH_IMAGE_JOB_EXISTS", "batch image job already exists")
ErrBatchImageItemExists = infraerrors.New(http.StatusConflict, "BATCH_IMAGE_ITEM_EXISTS", "batch image item already exists")
ErrBatchImageInvalidTransition = infraerrors.New(http.StatusBadRequest, "BATCH_IMAGE_INVALID_TRANSITION", "invalid batch image job status transition")
ErrBatchImageInvalidProvider = infraerrors.New(http.StatusBadRequest, "BATCH_IMAGE_INVALID_PROVIDER", "invalid batch image provider")
ErrBatchImageMissingProviderJobName = infraerrors.New(http.StatusBadRequest, "BATCH_IMAGE_MISSING_PROVIDER_JOB_NAME", "batch image provider job name is missing")
ErrBatchImageMissingAccountID = infraerrors.New(http.StatusBadRequest, "BATCH_IMAGE_MISSING_ACCOUNT_ID", "batch image account id is missing")
ErrBatchImageUnsupportedProvider = infraerrors.New(http.StatusBadRequest, "BATCH_IMAGE_UNSUPPORTED_PROVIDER", "unsupported batch image provider")
ErrBatchImageIndexOutputMissing = infraerrors.New(http.StatusBadGateway, "BATCH_IMAGE_INDEX_OUTPUT_MISSING", "batch image provider output is missing")
ErrBatchImageIndexParseFailed = infraerrors.New(http.StatusBadGateway, "BATCH_IMAGE_INDEX_PARSE_FAILED", "batch image provider output parse failed")
ErrBatchImageIndexNoResultLines = infraerrors.New(http.StatusBadGateway, "BATCH_IMAGE_INDEX_NO_RESULT_LINES", "batch image provider output has no result lines")
ErrBatchImageDuplicateCustomID = infraerrors.New(http.StatusBadGateway, "DUPLICATE_CUSTOM_ID_IN_OUTPUT", "batch image provider output contains duplicate custom id")
ErrBatchImageSettlementInvalidStatus = infraerrors.New(http.StatusBadRequest, "BATCH_IMAGE_SETTLEMENT_INVALID_STATUS", "batch image job is not ready for settlement")
ErrBatchImageSettlementManifestConflict = infraerrors.New(http.StatusConflict, "BATCH_IMAGE_SETTLEMENT_MANIFEST_CONFLICT", "batch image settlement manifest hash conflict")
ErrBatchImageSettlementPricingMissing = infraerrors.New(http.StatusBadRequest, "BATCH_IMAGE_SETTLEMENT_PRICING_MISSING", "batch image settlement pricing is missing")
ErrBatchImageSettlementBillingFailed = infraerrors.New(http.StatusBadGateway, "BATCH_IMAGE_SETTLEMENT_BILLING_FAILED", "batch image settlement billing failed")
ErrBatchImageAlreadySettled = infraerrors.New(http.StatusConflict, "BATCH_IMAGE_ALREADY_SETTLED", "batch image job is already settled")
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")
ErrBatchImageDuplicateCustomIDInRequest = infraerrors.New(http.StatusBadRequest, "BATCH_IMAGE_DUPLICATE_CUSTOM_ID", "batch image custom ids must be unique")
ErrBatchImagePromptTooLong = infraerrors.New(http.StatusBadRequest, "BATCH_IMAGE_PROMPT_TOO_LONG", "batch image prompt is too long")
ErrBatchImageInvalidReferenceImage = infraerrors.New(http.StatusBadRequest, "BATCH_IMAGE_INVALID_REFERENCE_IMAGE", "batch image reference image is invalid")
ErrBatchImageTooManyReferenceImages = infraerrors.New(http.StatusBadRequest, "BATCH_IMAGE_TOO_MANY_REFERENCE_IMAGES", "too many batch image reference images for this model")
ErrBatchImageReferenceImagesTooLarge = infraerrors.New(http.StatusBadRequest, "BATCH_IMAGE_REFERENCE_IMAGES_TOO_LARGE", "batch image reference images are too large")
ErrBatchImageTooManyOutputImages = infraerrors.New(http.StatusBadRequest, "BATCH_IMAGE_TOO_MANY_OUTPUT_IMAGES", "too many batch image output images")
ErrBatchImageProviderSubmitFailed = infraerrors.New(http.StatusBadGateway, "BATCH_IMAGE_PROVIDER_SUBMIT_FAILED", "batch image provider submit failed")
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")
ErrBatchImageItemNotFound = infraerrors.New(http.StatusNotFound, "BATCH_IMAGE_ITEM_NOT_FOUND", "batch image item not found")
ErrBatchImageItemFailed = infraerrors.New(http.StatusConflict, "BATCH_IMAGE_ITEM_FAILED", "batch image item did not succeed")
ErrBatchImageResultMissing = infraerrors.New(http.StatusInternalServerError, "BATCH_IMAGE_RESULT_MISSING", "batch image result is missing")
ErrBatchImageDownloadLimited = infraerrors.New(http.StatusTooManyRequests, "BATCH_IMAGE_DOWNLOAD_LIMITED", "too many batch image downloads")
ErrBatchImageDownloadFailed = infraerrors.New(http.StatusInternalServerError, "BATCH_IMAGE_DOWNLOAD_FAILED", "batch image download failed")
ErrBatchImageDownloadTooLarge = infraerrors.New(http.StatusBadRequest, "BATCH_IMAGE_DOWNLOAD_TOO_LARGE", "batch image download is too large")
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")
)
type BatchImageJob struct {
ID int64
BatchID string
UserID int64
APIKeyID *int64
AccountID *int64
Provider string
Model string
TaskName string
ParentBatchID *string
Status string
ProviderJobName *string
ProviderInputRef *string
ProviderOutputRef *string
GCSInputURI *string
GCSOutputURI *string
ItemCount int
SuccessCount int
FailCount int
CancelledCount int
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
ManifestHash *string
RetryCount int
Version int
OutputExpiresAt *time.Time
InputDeletedAt *time.Time
OutputDeletedAt *time.Time
DownloadedAt *time.Time
UserDeletedAt *time.Time
LastErrorCode *string
LastErrorMessage *string
CreatedAt time.Time
UpdatedAt time.Time
SubmittedAt *time.Time
StartedAt *time.Time
FinishedAt *time.Time
SettledAt *time.Time
}
type CreateBatchImageJobParams struct {
BatchID string
UserID int64
APIKeyID *int64
AccountID *int64
Provider string
Model string
TaskName string
ParentBatchID *string
Status string
ProviderJobName *string
ProviderInputRef *string
ProviderOutputRef *string
GCSInputURI *string
GCSOutputURI *string
ItemCount int
SuccessCount int
FailCount int
CancelledCount int
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
ManifestHash *string
RetryCount int
OutputExpiresAt *time.Time
}
type BatchImageItem struct {
ID int64
JobID string
CustomID string
Status string
RequestHash *string
PromptPreview *string
ProviderSourceObject *string
SourceLineNumber *int
SourceByteOffset *int64
SourceByteLength *int64
MimeType *string
FileExtension *string
ImageCount int
ErrorCode *string
ErrorMessage *string
BilledAmount *float64
CreatedAt time.Time
IndexedAt *time.Time
}
type CreateBatchImageItemParams struct {
JobID string
CustomID string
Status string
RequestHash *string
PromptPreview *string
ProviderSourceObject *string
SourceLineNumber *int
SourceByteOffset *int64
SourceByteLength *int64
MimeType *string
FileExtension *string
ImageCount int
ErrorCode *string
ErrorMessage *string
BilledAmount *float64
IndexedAt *time.Time
}
type BatchImageItemFilter struct {
Status string
Limit int
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
}
type UpdateBatchImageJobProviderSubmitParams struct {
BatchID string
ProviderJobName string
ProviderInputRef string
ProviderOutputRef string
GCSInputURI string
GCSOutputURI string
EventPayload any
}
type BatchImageTransitionOptions struct {
EventType string
EventPayload any
ErrorCode *string
ErrorMessage *string
Now *time.Time
}
type MarkBatchImageJobSettledParams struct {
BatchID string
ActualCost float64
ManifestHash string
EventPayload any
Now *time.Time
OutputExpiresAt *time.Time
}
type BatchImageEvent struct {
ID int64
JobID string
EventType string
Payload []byte
EventHash *string
CreatedAt time.Time
}
type BatchImageRepository interface {
CreateBatchImageJob(ctx context.Context, params CreateBatchImageJobParams) (*BatchImageJob, error)
GetBatchImageJobByBatchID(ctx context.Context, batchID string) (*BatchImageJob, error)
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
RecordBatchImageJobSubmitFailure(ctx context.Context, batchID, code, message string, markFailed bool) error
MarkBatchImageJobSettled(ctx context.Context, params MarkBatchImageJobSettledParams) error
SetBatchImageJobSettlementFailed(ctx context.Context, batchID, code, message string) (int, error)
CreateBatchImageItem(ctx context.Context, params CreateBatchImageItemParams) (*BatchImageItem, error)
BulkCreateBatchImageItems(ctx context.Context, params []CreateBatchImageItemParams) error
ReplaceBatchImageItemsForJob(ctx context.Context, batchID string, items []CreateBatchImageItemParams, counts BatchImageCounts) error
ListBatchImageItems(ctx context.Context, batchID string, filter BatchImageItemFilter) ([]*BatchImageItem, error)
ListBatchImageItemsForOwner(ctx context.Context, userID, apiKeyID int64, batchID string, filter BatchImageItemFilter) ([]*BatchImageItem, error)
GetBatchImageJobForDownload(ctx context.Context, userID, apiKeyID int64, batchID string) (*BatchImageJob, error)
GetBatchImageItemForDownload(ctx context.Context, batchID, customID string) (*BatchImageItem, error)
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
}
func NewBatchImageID() (string, error) {
var b [16]byte
if _, err := rand.Read(b[:]); err != nil {
return "", err
}
return "imgbatch_" + hex.EncodeToString(b[:]), nil
}
func IsSupportedBatchImageProvider(provider string) bool {
switch provider {
case BatchImageProviderGeminiAPI, BatchImageProviderVertex:
return true
default:
return false
}
}
func IsTerminalBatchImageJobStatus(status string) bool {
switch status {
case BatchImageJobStatusCompleted, BatchImageJobStatusFailed, BatchImageJobStatusCancelled, BatchImageJobStatusOutputDeleted:
return true
default:
return false
}
}
func CanTransitionBatchImageJob(from, to string) bool {
if from == "" || to == "" {
return false
}
if IsTerminalBatchImageJobStatus(from) {
return to == BatchImageJobStatusOutputDeleted &&
from != BatchImageJobStatusOutputDeleted &&
(from == BatchImageJobStatusCompleted || from == BatchImageJobStatusFailed || from == BatchImageJobStatusCancelled)
}
if to == BatchImageJobStatusFailed {
return true
}
allowed := map[string]map[string]struct{}{
BatchImageJobStatusCreated: {
BatchImageJobStatusUploading: {},
BatchImageJobStatusSubmitted: {},
BatchImageJobStatusCancelled: {},
},
BatchImageJobStatusUploading: {
BatchImageJobStatusSubmitted: {},
BatchImageJobStatusCancelled: {},
},
BatchImageJobStatusSubmitted: {
BatchImageJobStatusRunning: {},
BatchImageJobStatusIndexing: {},
BatchImageJobStatusFailed: {},
BatchImageJobStatusCancelled: {},
},
BatchImageJobStatusRunning: {
BatchImageJobStatusRunning: {},
BatchImageJobStatusIndexing: {},
BatchImageJobStatusFailed: {},
BatchImageJobStatusCancelled: {},
},
BatchImageJobStatusIndexing: {
BatchImageJobStatusSettling: {},
BatchImageJobStatusFailed: {},
},
BatchImageJobStatusSettling: {
BatchImageJobStatusCompleted: {},
},
}
_, ok := allowed[from][to]
return ok
}
@@ -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
}
@@ -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
}
@@ -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)
}
@@ -0,0 +1,294 @@
package service
import (
"context"
"errors"
"strings"
"sync"
"time"
"github.com/Wei-Shaw/sub2api/internal/config"
infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors"
)
const (
defaultBatchImageInputRetentionAfterTerminal = 24 * time.Hour
defaultBatchImageOutputRetentionAfterTerminal = 72 * time.Hour
defaultBatchImageCleanupInterval = 30 * time.Minute
defaultBatchImageCleanupBatchSize = 100
)
type BatchImageCleanupService struct {
Repo BatchImageRepository
ProviderRegistry *BatchImageProviderRegistry
AccountResolver BatchImageAccountResolver
Config *config.Config
cancel context.CancelFunc
done chan struct{}
mu sync.Mutex
}
func NewBatchImageCleanupService(repo BatchImageRepository, accountRepo AccountRepository, cfg *config.Config) *BatchImageCleanupService {
return &BatchImageCleanupService{
Repo: repo,
ProviderRegistry: NewBatchImageProviderRegistryFromConfig(cfg),
AccountResolver: &BatchImageAccountRepositoryResolver{Repo: accountRepo},
Config: cfg,
}
}
func (s *BatchImageCleanupService) DeleteOutputsForOwner(ctx context.Context, owner BatchImageOwner, batchID string) (*BatchImagePublicBatch, error) {
job, err := s.Repo.GetBatchImageJobByBatchIDForOwner(ctx, owner.UserID, owner.APIKeyID, batchID)
if err != nil {
return nil, err
}
if job.Status == BatchImageJobStatusOutputDeleted || job.OutputDeletedAt != nil {
return BatchImageJobToPublic(job), nil
}
if job.Status != BatchImageJobStatusCompleted {
return nil, ErrBatchImageOutputDeleteNotReady
}
_ = s.Repo.AppendBatchImageEvent(ctx, job.BatchID, "manual_output_delete_requested", map[string]any{
"batch_id": job.BatchID,
"cleanup_target": "output",
"reason": "manual",
})
if err := s.cleanupJob(ctx, job, CleanupTargetOutput, "manual"); err != nil {
return nil, err
}
updated, err := s.Repo.GetBatchImageJobByBatchIDForOwner(ctx, owner.UserID, owner.APIKeyID, batchID)
if err != nil {
return nil, err
}
return BatchImageJobToPublic(updated), nil
}
func (s *BatchImageCleanupService) CleanupInput(ctx context.Context, batchID string) error {
job, err := s.Repo.GetBatchImageJobByBatchID(ctx, batchID)
if err != nil {
return err
}
return s.cleanupJob(ctx, job, CleanupTargetInput, "ttl")
}
func (s *BatchImageCleanupService) CleanupOutput(ctx context.Context, batchID string, reason string) error {
job, err := s.Repo.GetBatchImageJobByBatchID(ctx, batchID)
if err != nil {
return err
}
return s.cleanupJob(ctx, job, CleanupTargetOutput, reason)
}
func (s *BatchImageCleanupService) RunOnce(ctx context.Context, now time.Time) (BatchImageCleanupRunResult, error) {
if s == nil || s.Repo == nil {
return BatchImageCleanupRunResult{}, ErrBatchImageCleanupFailed
}
if now.IsZero() {
now = time.Now()
}
limit := s.cleanupBatchSize()
result := BatchImageCleanupRunResult{}
inputCutoff := now.Add(-s.inputRetentionAfterTerminal())
inputJobs, err := s.Repo.ListBatchImageJobsDueForInputCleanup(ctx, inputCutoff, limit)
if err != nil {
return result, err
}
for _, job := range inputJobs {
if job == nil {
continue
}
if err := s.cleanupJob(ctx, job, CleanupTargetInput, "ttl"); err != nil {
result.Failures++
continue
}
result.InputCleaned++
}
outputJobs, err := s.Repo.ListBatchImageJobsDueForOutputCleanup(ctx, now, limit)
if err != nil {
return result, err
}
for _, job := range outputJobs {
if job == nil {
continue
}
if err := s.cleanupJob(ctx, job, CleanupTargetOutput, "expired"); err != nil {
result.Failures++
continue
}
result.OutputCleaned++
}
return result, nil
}
func (s *BatchImageCleanupService) Start() {
if s == nil || s.Repo == nil || s.Config == nil || !s.Config.BatchImage.Enabled || s.cleanupInterval() <= 0 {
return
}
s.mu.Lock()
defer s.mu.Unlock()
if s.cancel != nil {
return
}
ctx, cancel := context.WithCancel(context.Background())
s.cancel = cancel
s.done = make(chan struct{})
go func() {
defer close(s.done)
ticker := time.NewTicker(s.cleanupInterval())
defer ticker.Stop()
for {
_, _ = s.RunOnce(ctx, time.Now())
select {
case <-ctx.Done():
return
case <-ticker.C:
}
}
}()
}
func (s *BatchImageCleanupService) Stop() {
if s == nil {
return
}
s.mu.Lock()
cancel := s.cancel
done := s.done
s.cancel = nil
s.done = nil
s.mu.Unlock()
if cancel != nil {
cancel()
}
if done != nil {
<-done
}
}
func (s *BatchImageCleanupService) cleanupJob(ctx context.Context, job *BatchImageJob, target CleanupTarget, reason string) error {
if job == nil {
return ErrBatchImageJobNotFound
}
switch target {
case CleanupTargetInput:
if job.InputDeletedAt != nil {
return nil
}
if !IsTerminalBatchImageJobStatus(job.Status) {
return ErrBatchImageCleanupFailed
}
_ = s.Repo.AppendBatchImageEvent(ctx, job.BatchID, "input_cleanup_started", cleanupEventPayload(job.BatchID, target, reason, nil))
case CleanupTargetOutput:
if job.OutputDeletedAt != nil || job.Status == BatchImageJobStatusOutputDeleted {
return nil
}
if job.Status != BatchImageJobStatusCompleted && job.Status != BatchImageJobStatusFailed && job.Status != BatchImageJobStatusCancelled {
return ErrBatchImageOutputDeleteNotReady
}
_ = s.Repo.AppendBatchImageEvent(ctx, job.BatchID, "output_cleanup_started", cleanupEventPayload(job.BatchID, target, reason, nil))
default:
return ErrUnsupportedCleanupTarget
}
if err := s.callProviderCleanup(ctx, job, target); err != nil {
code := cleanupFailureCode(err)
msg := sanitizeBatchImagePublicMessage(err.Error())
_ = s.Repo.RecordBatchImageCleanupFailure(ctx, job.BatchID, code, msg)
event := string(target) + "_cleanup_failed"
_ = s.Repo.AppendBatchImageEvent(ctx, job.BatchID, event, map[string]any{"batch_id": job.BatchID, "cleanup_target": string(target), "reason": reason, "error_code": code})
if errors.Is(err, ErrBatchImageProviderUnsafeCleanupPath) {
return ErrBatchImageCleanupUnsafePath
}
return ErrBatchImageProviderCleanupFailed
}
deletedAt := time.Now()
if target == CleanupTargetInput {
return s.Repo.MarkBatchImageInputDeleted(ctx, job.BatchID, deletedAt)
}
return s.Repo.MarkBatchImageOutputDeleted(ctx, job.BatchID, deletedAt)
}
func (s *BatchImageCleanupService) callProviderCleanup(ctx context.Context, job *BatchImageJob, target CleanupTarget) error {
if s == nil || s.ProviderRegistry == nil || s.AccountResolver == nil {
return ErrBatchImageCleanupFailed
}
provider, ok := s.ProviderRegistry.Get(job.Provider)
if !ok || provider == nil {
return ErrBatchImageUnsupportedProvider
}
if job.AccountID == nil || *job.AccountID <= 0 {
return ErrBatchImageMissingAccountID
}
account, err := s.AccountResolver.ResolveBatchImageAccount(ctx, *job.AccountID)
if err != nil {
return err
}
if err := provider.Cleanup(ctx, job, account, target); err != nil {
if cleanupErrorIsNotFound(err) {
return nil
}
return err
}
return nil
}
func (s *BatchImageCleanupService) inputRetentionAfterTerminal() time.Duration {
if s != nil && s.Config != nil && s.Config.BatchImage.InputRetentionAfterTerminalHours > 0 {
return time.Duration(s.Config.BatchImage.InputRetentionAfterTerminalHours) * time.Hour
}
return defaultBatchImageInputRetentionAfterTerminal
}
func (s *BatchImageCleanupService) cleanupInterval() time.Duration {
if s != nil && s.Config != nil && s.Config.BatchImage.CleanupIntervalMinutes > 0 {
return time.Duration(s.Config.BatchImage.CleanupIntervalMinutes) * time.Minute
}
return defaultBatchImageCleanupInterval
}
func (s *BatchImageCleanupService) cleanupBatchSize() int {
if s != nil && s.Config != nil && s.Config.BatchImage.CleanupBatchSize > 0 {
return s.Config.BatchImage.CleanupBatchSize
}
return defaultBatchImageCleanupBatchSize
}
type BatchImageCleanupRunResult struct {
InputCleaned int
OutputCleaned int
Failures int
}
func cleanupEventPayload(batchID string, target CleanupTarget, reason string, deletedAt *time.Time) map[string]any {
payload := map[string]any{
"batch_id": batchID,
"cleanup_target": string(target),
"reason": reason,
}
if deletedAt != nil {
payload["deleted_at"] = deletedAt.UTC().Format(time.RFC3339)
}
return payload
}
func cleanupErrorIsNotFound(err error) bool {
if err == nil {
return false
}
reason := strings.ToUpper(infraerrors.Reason(err))
msg := strings.ToUpper(err.Error())
return strings.Contains(reason, "NOT_FOUND") || strings.Contains(msg, "NOT FOUND") || strings.Contains(msg, "404")
}
func cleanupFailureCode(err error) string {
if errors.Is(err, ErrBatchImageProviderUnsafeCleanupPath) {
return "BATCH_IMAGE_CLEANUP_UNSAFE_PATH"
}
reason := strings.TrimSpace(infraerrors.Reason(err))
if reason != "" {
return reason
}
return "BATCH_IMAGE_PROVIDER_CLEANUP_FAILED"
}
@@ -0,0 +1,231 @@
//go:build unit
package service
import (
"context"
"encoding/json"
"errors"
"strings"
"testing"
"time"
"github.com/Wei-Shaw/sub2api/internal/config"
infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors"
"github.com/stretchr/testify/require"
)
func TestBatchImageCleanupService_DeleteOutputsForOwner(t *testing.T) {
ctx := context.Background()
t.Run("deletes completed output and returns public dto", func(t *testing.T) {
svc, repo, provider := newTestBatchImageCleanupService()
got, err := svc.DeleteOutputsForOwner(ctx, testBatchImageOwner(), "imgbatch_cleanup")
require.NoError(t, err)
require.Equal(t, "output_deleted", got.Status)
require.NotNil(t, got.OutputDeletedAt)
require.Equal(t, []CleanupTarget{CleanupTargetOutput}, provider.cleanupTargets)
require.NotNil(t, repo.jobs["imgbatch_cleanup"].OutputDeletedAt)
require.Equal(t, BatchImageJobStatusOutputDeleted, repo.jobs["imgbatch_cleanup"].Status)
body := mustJSON(t, got)
requireBatchImagePublicJSONHasNoInternals(t, body)
})
t.Run("repeated delete is idempotent", func(t *testing.T) {
svc, repo, provider := newTestBatchImageCleanupService()
deletedAt := time.Now()
repo.jobs["imgbatch_cleanup"].Status = BatchImageJobStatusOutputDeleted
repo.jobs["imgbatch_cleanup"].OutputDeletedAt = &deletedAt
got, err := svc.DeleteOutputsForOwner(ctx, testBatchImageOwner(), "imgbatch_cleanup")
require.NoError(t, err)
require.Equal(t, "output_deleted", got.Status)
require.Empty(t, provider.cleanupTargets)
})
t.Run("not completed returns not ready", func(t *testing.T) {
svc, repo, _ := newTestBatchImageCleanupService()
repo.jobs["imgbatch_cleanup"].Status = BatchImageJobStatusRunning
_, err := svc.DeleteOutputsForOwner(ctx, testBatchImageOwner(), "imgbatch_cleanup")
require.ErrorIs(t, err, ErrBatchImageOutputDeleteNotReady)
})
t.Run("non owner returns not found", func(t *testing.T) {
svc, _, _ := newTestBatchImageCleanupService()
_, err := svc.DeleteOutputsForOwner(ctx, BatchImageOwner{UserID: 11, APIKeyID: 999}, "imgbatch_cleanup")
require.ErrorIs(t, err, ErrBatchImageJobNotFound)
})
t.Run("provider not found is success", func(t *testing.T) {
svc, repo, provider := newTestBatchImageCleanupService()
provider.cleanupErr = infraerrors.New(404, "PROVIDER_NOT_FOUND", "provider file not found: gs://hidden")
got, err := svc.DeleteOutputsForOwner(ctx, testBatchImageOwner(), "imgbatch_cleanup")
require.NoError(t, err)
require.Equal(t, "output_deleted", got.Status)
require.NotNil(t, repo.jobs["imgbatch_cleanup"].OutputDeletedAt)
})
t.Run("provider transient error is sanitized and records failure", func(t *testing.T) {
svc, repo, provider := newTestBatchImageCleanupService()
provider.cleanupErr = errors.New("temporary cleanup failed for gs://secret-output")
_, err := svc.DeleteOutputsForOwner(ctx, testBatchImageOwner(), "imgbatch_cleanup")
require.ErrorIs(t, err, ErrBatchImageProviderCleanupFailed)
require.Equal(t, "BATCH_IMAGE_PROVIDER_CLEANUP_FAILED", infraerrors.Reason(err))
require.NotContains(t, infraerrors.Message(err), "gs://")
require.Equal(t, "BATCH_IMAGE_PROVIDER_CLEANUP_FAILED", batchImageDerefString(repo.jobs["imgbatch_cleanup"].LastErrorCode))
require.Equal(t, "upstream provider operation failed", batchImageDerefString(repo.jobs["imgbatch_cleanup"].LastErrorMessage))
})
t.Run("unsafe cleanup path is not swallowed", func(t *testing.T) {
svc, repo, provider := newTestBatchImageCleanupService()
provider.cleanupErr = ErrBatchImageProviderUnsafeCleanupPath
_, err := svc.DeleteOutputsForOwner(ctx, testBatchImageOwner(), "imgbatch_cleanup")
require.ErrorIs(t, err, ErrBatchImageCleanupUnsafePath)
require.Equal(t, "BATCH_IMAGE_CLEANUP_UNSAFE_PATH", batchImageDerefString(repo.jobs["imgbatch_cleanup"].LastErrorCode))
})
}
func TestBatchImageCleanupService_InputOutputAndWorker(t *testing.T) {
ctx := context.Background()
now := time.Now()
t.Run("input cleanup marks input only", func(t *testing.T) {
svc, repo, provider := newTestBatchImageCleanupService()
err := svc.CleanupInput(ctx, "imgbatch_cleanup")
require.NoError(t, err)
require.Equal(t, []CleanupTarget{CleanupTargetInput}, provider.cleanupTargets)
require.NotNil(t, repo.jobs["imgbatch_cleanup"].InputDeletedAt)
require.Equal(t, BatchImageJobStatusCompleted, repo.jobs["imgbatch_cleanup"].Status)
err = svc.CleanupInput(ctx, "imgbatch_cleanup")
require.NoError(t, err)
require.Len(t, provider.cleanupTargets, 1)
})
t.Run("output cleanup for failed job keeps status", func(t *testing.T) {
svc, repo, _ := newTestBatchImageCleanupService()
repo.jobs["imgbatch_cleanup"].Status = BatchImageJobStatusFailed
err := svc.CleanupOutput(ctx, "imgbatch_cleanup", "ttl")
require.NoError(t, err)
require.Equal(t, BatchImageJobStatusFailed, repo.jobs["imgbatch_cleanup"].Status)
require.NotNil(t, repo.jobs["imgbatch_cleanup"].OutputDeletedAt)
})
t.Run("worker processes due jobs and continues after failure", func(t *testing.T) {
svc, repo, provider := newTestBatchImageCleanupService()
provider.cleanupErr = nil
old := now.Add(-48 * time.Hour)
expired := now.Add(-time.Minute)
future := now.Add(time.Hour)
repo.jobs["imgbatch_cleanup"].FinishedAt = &old
repo.jobs["imgbatch_cleanup"].OutputExpiresAt = &expired
repo.jobs["imgbatch_running"] = cleanupTestJob("imgbatch_running", BatchImageJobStatusRunning)
repo.jobs["imgbatch_running"].FinishedAt = &old
repo.jobs["imgbatch_running"].OutputExpiresAt = &expired
repo.jobs["imgbatch_future"] = cleanupTestJob("imgbatch_future", BatchImageJobStatusCompleted)
repo.jobs["imgbatch_future"].FinishedAt = &old
repo.jobs["imgbatch_future"].OutputExpiresAt = &future
result, err := svc.RunOnce(ctx, now)
require.NoError(t, err)
require.Equal(t, 2, result.InputCleaned)
require.Equal(t, 1, result.OutputCleaned)
require.Equal(t, BatchImageJobStatusRunning, repo.jobs["imgbatch_running"].Status)
require.Nil(t, repo.jobs["imgbatch_future"].OutputDeletedAt)
require.NotContains(t, strings.Join(repo.events["imgbatch_running"], ","), "cleanup")
})
}
func TestBatchImageSettlementOutputExpiration(t *testing.T) {
repo := newFakeBatchImageRepository()
job := testSettlingBatchImageJob("imgbatch_expire")
repo.jobs[job.BatchID] = job
billing := &fakeBatchImageBillingRepo{}
svc := &BatchImageSettlementService{
Repo: repo,
BillingRepo: billing,
Pricing: &fakeBatchImagePricingResolver{unitPrice: 0.25},
Config: &config.Config{BatchImage: config.BatchImageConfig{OutputRetentionAfterTerminalHours: 5}},
}
_, err := svc.Settle(context.Background(), job.BatchID)
require.NoError(t, err)
require.NotNil(t, repo.jobs[job.BatchID].OutputExpiresAt)
require.WithinDuration(t, time.Now().Add(5*time.Hour), *repo.jobs[job.BatchID].OutputExpiresAt, time.Minute)
existing := time.Now().Add(time.Hour)
second := testSettlingBatchImageJob("imgbatch_keep_expire")
second.OutputExpiresAt = &existing
repo.jobs[second.BatchID] = second
_, err = svc.Settle(context.Background(), second.BatchID)
require.NoError(t, err)
require.Equal(t, existing, *repo.jobs[second.BatchID].OutputExpiresAt)
}
func TestBatchImageDownloadAfterOutputDeletedReturnsGone(t *testing.T) {
svc, repo, _ := newTestBatchImageDownloadService()
now := time.Now()
repo.jobs["imgbatch_download"].Status = BatchImageJobStatusOutputDeleted
repo.jobs["imgbatch_download"].OutputDeletedAt = &now
stream, err := svc.OpenItemContent(context.Background(), testBatchImageOwner(), "imgbatch_download", "cover/../001", 0)
require.Nil(t, stream)
require.ErrorIs(t, err, ErrBatchImageOutputDeleted)
var out strings.Builder
result, err := svc.StreamZip(context.Background(), testBatchImageOwner(), "imgbatch_download", BatchImageZipOptions{}, &out)
require.Nil(t, result)
require.ErrorIs(t, err, ErrBatchImageOutputDeleted)
}
func newTestBatchImageCleanupService() (*BatchImageCleanupService, *fakeBatchImageRepository, *publicBatchImageProvider) {
repo := newFakeBatchImageRepository()
repo.jobs["imgbatch_cleanup"] = cleanupTestJob("imgbatch_cleanup", BatchImageJobStatusCompleted)
provider := &publicBatchImageProvider{name: BatchImageProviderGeminiAPI}
accountID := int64(101)
svc := &BatchImageCleanupService{
Repo: repo,
ProviderRegistry: NewBatchImageProviderRegistry(provider),
AccountResolver: &fakeBatchImageAccountResolver{account: &Account{ID: accountID, Platform: PlatformGemini, Type: AccountTypeAPIKey, Status: StatusActive, Schedulable: true}},
Config: &config.Config{BatchImage: config.BatchImageConfig{CleanupBatchSize: 10, InputRetentionAfterTerminalHours: 24}},
}
return svc, repo, provider
}
func cleanupTestJob(batchID, status string) *BatchImageJob {
apiKeyID := int64(22)
accountID := int64(101)
now := time.Now().Add(-48 * time.Hour)
return &BatchImageJob{
BatchID: batchID,
UserID: 11,
APIKeyID: &apiKeyID,
AccountID: &accountID,
Provider: BatchImageProviderGeminiAPI,
Model: "gemini-2.5-flash-image",
Status: status,
ProviderJobName: batchImageStringPtr("providers/internal/job"),
ProviderInputRef: batchImageStringPtr("files/internal/input"),
ProviderOutputRef: batchImageStringPtr("files/internal/output"),
ItemCount: 1,
SuccessCount: 1,
CreatedAt: now,
UpdatedAt: now,
FinishedAt: &now,
SettledAt: &now,
}
}
func mustJSON(t *testing.T, v any) string {
t.Helper()
b, err := json.Marshal(v)
require.NoError(t, err)
return string(b)
}
@@ -0,0 +1,659 @@
package service
import (
"archive/zip"
"bufio"
"context"
"encoding/base64"
"encoding/json"
"errors"
"fmt"
"io"
"net/http"
"path/filepath"
"sort"
"strings"
"sync"
"time"
"unicode"
"github.com/Wei-Shaw/sub2api/internal/config"
infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors"
)
const (
defaultBatchImageZipMaxItems = 200
defaultBatchImageZipMaxBytes = 512 * 1024 * 1024
defaultBatchImageDownloadDuration = 10 * time.Minute
defaultBatchImageDownloadConcurrency = 1
batchImageDownloadScannerMaxLineBytes = 16 * 1024 * 1024
)
var errBatchImageDownloadSizeExceeded = errors.New("batch image download size limit exceeded")
type BatchImageDownloadLimiter interface {
Acquire(ctx context.Context, userID string, kind string) (BatchImageDownloadPermit, error)
}
type BatchImageDownloadPermit interface {
Release(ctx context.Context) error
}
type BatchImageContentStream struct {
Reader io.ReadCloser
ContentType string
Filename string
ContentLength *int64
}
type BatchImageZipOptions struct {
Status string
MaxItems int
IncludeManifest bool
}
type BatchImageZipResult struct {
FileCount int
ErrorCount int
}
type BatchImageLineImages struct {
CustomID string
Images []BatchImageInlineImage
ErrorCode string
ErrorMessage string
}
type BatchImageInlineImage struct {
MimeType string
Extension string
Base64Data string
}
type BatchImageDownloadService struct {
Repo BatchImageRepository
ProviderRegistry *BatchImageProviderRegistry
AccountResolver BatchImageAccountResolver
Limiter BatchImageDownloadLimiter
Config *config.Config
}
type batchImageDownloadLimitWriter struct {
w io.Writer
limit int64
written int64
}
func (w *batchImageDownloadLimitWriter) Write(p []byte) (int, error) {
if w == nil || w.w == nil {
return 0, io.ErrClosedPipe
}
if w.limit > 0 && w.written+int64(len(p)) > w.limit {
return 0, errBatchImageDownloadSizeExceeded
}
n, err := w.w.Write(p)
w.written += int64(n)
return n, err
}
func NewBatchImageDownloadService(repo BatchImageRepository, accountRepo AccountRepository, limiter BatchImageDownloadLimiter, cfg *config.Config) *BatchImageDownloadService {
return &BatchImageDownloadService{
Repo: repo,
ProviderRegistry: NewBatchImageProviderRegistryFromConfig(cfg),
AccountResolver: &BatchImageAccountRepositoryResolver{Repo: accountRepo},
Limiter: limiter,
Config: cfg,
}
}
func (s *BatchImageDownloadService) OpenItemContent(ctx context.Context, owner BatchImageOwner, batchID string, customID string, imageIndex int) (*BatchImageContentStream, error) {
if imageIndex < 0 {
return nil, ErrBatchImageItemImageIndexOutOfRange
}
job, err := s.getCompletedJob(ctx, owner, batchID)
if err != nil {
return nil, err
}
item, err := s.Repo.GetBatchImageItemForDownload(ctx, job.BatchID, customID)
if err != nil {
return nil, err
}
if item.Status != BatchImageItemStatusSuccess {
return nil, ErrBatchImageItemFailed
}
if imageIndex >= item.ImageCount {
return nil, ErrBatchImageItemImageIndexOutOfRange
}
permit, err := s.acquirePermit(ctx, owner.UserID, "item")
if err != nil {
return nil, err
}
releasePermit := true
defer func() {
if releasePermit && permit != nil {
_ = permit.Release(ctx)
}
}()
provider, account, err := s.providerAndAccount(ctx, job)
if err != nil {
return nil, err
}
r, _, err := provider.OpenResult(ctx, job, account)
if err != nil {
return nil, ErrBatchImageResultMissing.WithCause(err)
}
defer func() { _ = r.Close() }()
line, err := findBatchImageLineImages(r, item.CustomID)
if err != nil {
return nil, err
}
if imageIndex >= len(line.Images) {
return nil, ErrBatchImageItemImageIndexOutOfRange
}
image := line.Images[imageIndex]
if strings.TrimSpace(image.Base64Data) == "" {
return nil, ErrBatchImageResultMissing
}
contentType := strings.TrimSpace(image.MimeType)
if contentType == "" {
contentType = "application/octet-stream"
}
extension := strings.TrimSpace(image.Extension)
if extension == "" {
extension = batchImageFileExtension(contentType)
}
if extension == "" {
extension = "bin"
}
reader := base64.NewDecoder(base64.StdEncoding, strings.NewReader(image.Base64Data))
releasePermit = false
return &BatchImageContentStream{
Reader: &batchImagePermitReadCloser{Reader: reader, permit: permit},
ContentType: contentType,
Filename: BatchImageSafeDownloadFilename(item.CustomID, extension),
}, nil
}
func (s *BatchImageDownloadService) StreamZip(ctx context.Context, owner BatchImageOwner, batchID string, opts BatchImageZipOptions, w io.Writer) (*BatchImageZipResult, error) {
job, err := s.getCompletedJob(ctx, owner, batchID)
if err != nil {
return nil, err
}
maxItems := opts.MaxItems
if maxItems <= 0 {
maxItems = s.maxZipItems()
}
if job.SuccessCount > maxItems {
return nil, ErrBatchImageZipTooManyItems
}
successItems, err := s.Repo.ListBatchImageItemsForDownload(ctx, job.BatchID, BatchImageItemStatusSuccess, maxItems+1)
if err != nil {
return nil, err
}
if len(successItems) > maxItems {
return nil, ErrBatchImageZipTooManyItems
}
failedItems, err := s.Repo.ListBatchImageItemsForDownload(ctx, job.BatchID, BatchImageItemStatusFailed, maxItems)
if err != nil {
return nil, err
}
permit, err := s.acquirePermit(ctx, owner.UserID, "zip")
if err != nil {
return nil, err
}
if permit != nil {
defer func() { _ = permit.Release(ctx) }()
}
provider, account, err := s.providerAndAccount(ctx, job)
if err != nil {
return nil, err
}
r, _, err := provider.OpenResult(ctx, job, account)
if err != nil {
return nil, ErrBatchImageResultMissing.WithCause(err)
}
defer func() { _ = r.Close() }()
streamCtx := ctx
cancel := func() {}
if d := s.maxDownloadDuration(); d > 0 {
streamCtx, cancel = context.WithTimeout(ctx, d)
}
defer cancel()
limitedWriter := &batchImageDownloadLimitWriter{w: w, limit: s.maxDownloadBytes()}
zipWriter := zip.NewWriter(limitedWriter)
result, manifestFiles, zipErrors, err := s.writeZipImages(streamCtx, zipWriter, r, successItems)
if err != nil {
_ = zipWriter.Close()
if errors.Is(err, errBatchImageDownloadSizeExceeded) {
return result, ErrBatchImageDownloadTooLarge.WithCause(err)
}
return result, ErrBatchImageDownloadFailed.WithCause(err)
}
zipErrors = append(zipErrors, batchImageZipErrorsFromItems(failedItems)...)
if err := writeBatchImageZipJSON(zipWriter, "manifest.json", batchImageZipManifest{
BatchID: job.BatchID,
Model: job.Model,
ItemCount: job.ItemCount,
SuccessCount: job.SuccessCount,
FailCount: job.FailCount,
Files: manifestFiles,
}); err != nil {
_ = zipWriter.Close()
if errors.Is(err, errBatchImageDownloadSizeExceeded) {
return result, ErrBatchImageDownloadTooLarge.WithCause(err)
}
return result, ErrBatchImageDownloadFailed.WithCause(err)
}
if err := writeBatchImageZipJSON(zipWriter, "errors.json", zipErrors); err != nil {
_ = zipWriter.Close()
if errors.Is(err, errBatchImageDownloadSizeExceeded) {
return result, ErrBatchImageDownloadTooLarge.WithCause(err)
}
return result, ErrBatchImageDownloadFailed.WithCause(err)
}
result.ErrorCount = len(zipErrors)
if err := zipWriter.Close(); err != nil {
if errors.Is(err, errBatchImageDownloadSizeExceeded) {
return result, ErrBatchImageDownloadTooLarge.WithCause(err)
}
return result, ErrBatchImageDownloadFailed.WithCause(err)
}
return result, nil
}
func (s *BatchImageDownloadService) writeZipImages(ctx context.Context, zipWriter *zip.Writer, resultReader io.Reader, successItems []*BatchImageItem) (*BatchImageZipResult, []batchImageZipManifestFile, []batchImageZipError, error) {
successByID := make(map[string]*BatchImageItem, len(successItems))
missing := make(map[string]struct{}, len(successItems))
for _, item := range successItems {
if item == nil {
continue
}
successByID[item.CustomID] = item
missing[item.CustomID] = struct{}{}
}
scanner := bufio.NewScanner(resultReader)
scanner.Buffer(make([]byte, 0, 64*1024), batchImageDownloadScannerMaxLineBytes)
result := &BatchImageZipResult{}
var manifestFiles []batchImageZipManifestFile
var zipErrors []batchImageZipError
for scanner.Scan() {
if err := ctx.Err(); err != nil {
return result, manifestFiles, zipErrors, err
}
line := strings.TrimSpace(scanner.Text())
if line == "" {
continue
}
images, err := ExtractBatchImagePartsFromResultLine([]byte(line))
if err != nil {
return result, manifestFiles, zipErrors, err
}
item := successByID[images.CustomID]
if item == nil {
continue
}
delete(missing, images.CustomID)
if len(images.Images) == 0 {
zipErrors = append(zipErrors, batchImageZipError{CustomID: images.CustomID, Code: "EMPTY_IMAGE_OUTPUT", Message: "provider response contained no image output"})
continue
}
for idx, image := range images.Images {
extension := image.Extension
if extension == "" {
extension = "bin"
}
filename := batchImageZipImageFilename(item.CustomID, idx, extension)
entry, err := zipWriter.CreateHeader(&zip.FileHeader{Name: filename, Method: zip.Deflate})
if err != nil {
return result, manifestFiles, zipErrors, err
}
decoder := base64.NewDecoder(base64.StdEncoding, strings.NewReader(image.Base64Data))
if _, err := io.Copy(entry, decoder); err != nil {
zipErrors = append(zipErrors, batchImageZipError{CustomID: item.CustomID, Code: "IMAGE_DECODE_FAILED", Message: "image data could not be decoded"})
continue
}
result.FileCount++
manifestFiles = append(manifestFiles, batchImageZipManifestFile{
CustomID: item.CustomID,
Filename: filename,
MimeType: image.MimeType,
ImageIndex: idx,
})
}
}
if err := scanner.Err(); err != nil {
return result, manifestFiles, zipErrors, err
}
missingIDs := make([]string, 0, len(missing))
for customID := range missing {
missingIDs = append(missingIDs, customID)
}
sort.Strings(missingIDs)
for _, customID := range missingIDs {
zipErrors = append(zipErrors, batchImageZipError{CustomID: customID, Code: "RESULT_MISSING", Message: "provider result was not found for item"})
}
return result, manifestFiles, zipErrors, nil
}
func (s *BatchImageDownloadService) getCompletedJob(ctx context.Context, owner BatchImageOwner, batchID string) (*BatchImageJob, error) {
if s == nil || s.Repo == nil {
return nil, ErrBatchImageDownloadFailed
}
job, err := s.Repo.GetBatchImageJobForDownload(ctx, owner.UserID, owner.APIKeyID, batchID)
if err != nil {
return nil, err
}
switch job.Status {
case BatchImageJobStatusCompleted:
return job, nil
case BatchImageJobStatusOutputDeleted:
return nil, ErrBatchImageOutputDeleted
default:
return nil, ErrBatchImageNotReady
}
}
func (s *BatchImageDownloadService) providerAndAccount(ctx context.Context, job *BatchImageJob) (BatchImageProvider, *Account, error) {
if s == nil || s.ProviderRegistry == nil || s.AccountResolver == nil || job == nil {
return nil, nil, ErrBatchImageDownloadFailed
}
provider, ok := s.ProviderRegistry.Get(job.Provider)
if !ok || provider == nil {
return nil, nil, ErrBatchImageUnsupportedProvider
}
if job.AccountID == nil || *job.AccountID <= 0 {
return nil, nil, ErrBatchImageMissingAccountID
}
account, err := s.AccountResolver.ResolveBatchImageAccount(ctx, *job.AccountID)
if err != nil {
return nil, nil, ErrBatchImageDownloadFailed
}
if !provider.SupportsAccount(account) {
return nil, nil, ErrBatchImageProviderUnsupportedAccount
}
return provider, account, nil
}
func (s *BatchImageDownloadService) acquirePermit(ctx context.Context, userID int64, kind string) (BatchImageDownloadPermit, error) {
if s == nil || s.Limiter == nil {
return nil, nil
}
permit, err := s.Limiter.Acquire(ctx, fmt.Sprintf("%d", userID), kind)
if err != nil {
if infraerrors.Code(err) == http.StatusTooManyRequests {
return nil, ErrBatchImageDownloadLimited
}
return nil, ErrBatchImageDownloadLimited.WithCause(err)
}
return permit, nil
}
func (s *BatchImageDownloadService) maxZipItems() int {
if s != nil && s.Config != nil && s.Config.BatchImage.MaxDownloadItemsZip > 0 {
return s.Config.BatchImage.MaxDownloadItemsZip
}
return defaultBatchImageZipMaxItems
}
func (s *BatchImageDownloadService) maxDownloadBytes() int64 {
if s != nil && s.Config != nil && s.Config.BatchImage.MaxDownloadBytesPerRequest > 0 {
return s.Config.BatchImage.MaxDownloadBytesPerRequest
}
return defaultBatchImageZipMaxBytes
}
func (s *BatchImageDownloadService) maxDownloadDuration() time.Duration {
if s != nil && s.Config != nil && s.Config.BatchImage.MaxDownloadDurationSeconds > 0 {
return time.Duration(s.Config.BatchImage.MaxDownloadDurationSeconds) * time.Second
}
return defaultBatchImageDownloadDuration
}
func ExtractBatchImagePartsFromResultLine(line []byte) (*BatchImageLineImages, error) {
var obj map[string]any
if err := json.Unmarshal(line, &obj); err != nil {
return nil, ErrBatchImageIndexParseFailed.WithCause(err)
}
customID := batchImageFirstNonEmptyString(
batchImageMapString(obj, "key"),
batchImageMapString(obj, "custom_id"),
batchImageMapString(obj, "customId"),
batchImageNestedString(obj, "request", "key"),
)
if customID == "" {
return nil, ErrBatchImageIndexParseFailed.WithCause(fmt.Errorf("missing custom id"))
}
out := &BatchImageLineImages{CustomID: customID}
out.Images = append(out.Images, extractBatchImageInlineImages(batchImageNestedAny(obj, "response", "candidates"))...)
out.Images = append(out.Images, extractBatchImageInlineImages(obj["candidates"])...)
if len(out.Images) > 0 {
return out, nil
}
if code, message, ok := batchImageFailureFromProviderFields(obj); ok {
out.ErrorCode = code
out.ErrorMessage = truncateBatchImageMessage(message, batchImageMaxErrorMessageLength)
return out, nil
}
if _, hasResponse := obj["response"]; hasResponse || batchImageHasCandidates(obj) {
out.ErrorCode = "EMPTY_IMAGE_OUTPUT"
out.ErrorMessage = "provider response contained no image output"
return out, nil
}
out.ErrorCode = "PROVIDER_ITEM_FAILED"
out.ErrorMessage = "provider result line contained no image output"
return out, nil
}
func extractBatchImageInlineImages(raw any) []BatchImageInlineImage {
candidates, ok := raw.([]any)
if !ok {
return nil
}
var images []BatchImageInlineImage
for _, candidateRaw := range candidates {
candidate, ok := candidateRaw.(map[string]any)
if !ok {
continue
}
parts, ok := batchImageNestedAny(candidate, "content", "parts").([]any)
if !ok {
continue
}
for _, partRaw := range parts {
part, ok := partRaw.(map[string]any)
if !ok {
continue
}
inline, ok := firstMap(part["inlineData"], part["inline_data"])
if !ok {
continue
}
data := strings.TrimSpace(batchImageMapString(inline, "data"))
mime := strings.TrimSpace(batchImageFirstNonEmptyString(batchImageMapString(inline, "mimeType"), batchImageMapString(inline, "mime_type")))
if data == "" || !strings.HasPrefix(strings.ToLower(mime), "image/") {
continue
}
images = append(images, BatchImageInlineImage{
MimeType: mime,
Extension: batchImageFileExtension(mime),
Base64Data: data,
})
}
}
return images
}
func findBatchImageLineImages(r io.Reader, customID string) (*BatchImageLineImages, error) {
scanner := bufio.NewScanner(r)
scanner.Buffer(make([]byte, 0, 64*1024), batchImageDownloadScannerMaxLineBytes)
for scanner.Scan() {
line := strings.TrimSpace(scanner.Text())
if line == "" {
continue
}
parsed, err := ExtractBatchImagePartsFromResultLine([]byte(line))
if err != nil {
return nil, err
}
if parsed.CustomID != customID {
continue
}
if len(parsed.Images) == 0 {
if parsed.ErrorCode != "" {
return nil, ErrBatchImageItemFailed
}
return nil, ErrBatchImageResultMissing
}
return parsed, nil
}
if err := scanner.Err(); err != nil {
return nil, ErrBatchImageDownloadFailed.WithCause(err)
}
return nil, ErrBatchImageResultMissing
}
func BatchImageSafeDownloadFilename(customID, extension string) string {
base := sanitizeBatchImageFilenameBase(customID)
extension = sanitizeBatchImageFilenameExtension(extension)
if extension == "" {
extension = "bin"
}
return base + "." + extension
}
func BatchImageContentDispositionAttachment(filename string) string {
filename = strings.ReplaceAll(filename, "\\", "_")
filename = strings.ReplaceAll(filename, `"`, "_")
filename = sanitizeBatchImageFilenameBase(strings.TrimSuffix(filename, filepath.Ext(filename))) + filepath.Ext(filename)
return `attachment; filename="` + filename + `"`
}
func sanitizeBatchImageFilenameBase(value string) string {
value = strings.TrimSpace(value)
if value == "" {
return "image"
}
var b strings.Builder
for _, r := range value {
switch {
case r == '/' || r == '\\' || r == ':' || r == 0:
_ = b.WriteByte('_')
case unicode.IsControl(r):
_ = b.WriteByte('_')
case unicode.IsLetter(r) || unicode.IsDigit(r) || r == '_' || r == '-' || r == '.':
_, _ = b.WriteRune(r)
default:
_ = b.WriteByte('_')
}
}
out := strings.Trim(b.String(), ". ")
for strings.Contains(out, "..") {
out = strings.ReplaceAll(out, "..", "_")
}
out = strings.Trim(out, ". ")
if out == "" {
out = "image"
}
if len(out) > 120 {
out = strings.TrimRight(out[:120], ". ")
}
if out == "" {
out = "image"
}
return out
}
func sanitizeBatchImageFilenameExtension(extension string) string {
extension = strings.TrimPrefix(strings.TrimSpace(strings.ToLower(extension)), ".")
var b strings.Builder
for _, r := range extension {
if unicode.IsLetter(r) || unicode.IsDigit(r) {
_, _ = b.WriteRune(r)
}
}
out := b.String()
if len(out) > 12 {
out = out[:12]
}
return out
}
func batchImageZipImageFilename(customID string, imageIndex int, extension string) string {
base := sanitizeBatchImageFilenameBase(customID)
if imageIndex > 0 {
base = fmt.Sprintf("%s_%d", base, imageIndex+1)
}
return "images/" + BatchImageSafeDownloadFilename(base, extension)
}
func writeBatchImageZipJSON(zipWriter *zip.Writer, name string, value any) error {
entry, err := zipWriter.CreateHeader(&zip.FileHeader{Name: name, Method: zip.Deflate})
if err != nil {
return err
}
encoder := json.NewEncoder(entry)
encoder.SetIndent("", " ")
return encoder.Encode(value)
}
type batchImageZipManifest struct {
BatchID string `json:"batch_id"`
Model string `json:"model"`
ItemCount int `json:"item_count"`
SuccessCount int `json:"success_count"`
FailCount int `json:"fail_count"`
Files []batchImageZipManifestFile `json:"files"`
}
type batchImageZipManifestFile struct {
CustomID string `json:"custom_id"`
Filename string `json:"filename"`
MimeType string `json:"mime_type"`
ImageIndex int `json:"image_index"`
}
type batchImageZipError struct {
CustomID string `json:"custom_id"`
Code string `json:"code"`
Message string `json:"message"`
}
func batchImageZipErrorsFromItems(items []*BatchImageItem) []batchImageZipError {
out := make([]batchImageZipError, 0, len(items))
for _, item := range items {
if item == nil {
continue
}
out = append(out, batchImageZipError{
CustomID: item.CustomID,
Code: batchImageDerefString(item.ErrorCode),
Message: sanitizeBatchImagePublicMessage(batchImageDerefString(item.ErrorMessage)),
})
}
return out
}
type batchImagePermitReadCloser struct {
io.Reader
permit BatchImageDownloadPermit
once sync.Once
err error
}
func (r *batchImagePermitReadCloser) Close() error {
r.once.Do(func() {
if r.permit != nil {
r.err = r.permit.Release(context.Background())
}
})
return r.err
}
@@ -0,0 +1,299 @@
//go:build unit
package service
import (
"archive/zip"
"bytes"
"context"
"encoding/json"
"io"
"strings"
"testing"
"time"
"github.com/Wei-Shaw/sub2api/internal/config"
"github.com/stretchr/testify/require"
)
func TestBatchImageDownloadService_OpenItemContent(t *testing.T) {
ctx := context.Background()
t.Run("streams image bytes with safe headers data", func(t *testing.T) {
svc, _, limiter := newTestBatchImageDownloadService()
stream, err := svc.OpenItemContent(ctx, testBatchImageOwner(), "imgbatch_download", "cover/../001", 1)
require.NoError(t, err)
defer stream.Reader.Close()
body, err := io.ReadAll(stream.Reader)
require.NoError(t, err)
require.Equal(t, []byte("second"), body)
require.Equal(t, "image/jpeg", stream.ContentType)
require.Equal(t, "cover___001.jpg", stream.Filename)
require.Equal(t, 1, limiter.acquireCount)
require.Zero(t, limiter.releaseCount)
require.NoError(t, stream.Reader.Close())
require.Equal(t, 1, limiter.releaseCount)
})
tests := []struct {
name string
mutate func(*fakeBatchImageRepository)
id string
item string
index int
want error
}{
{name: "non_owner", id: "imgbatch_download", item: "cover/../001", mutate: func(r *fakeBatchImageRepository) {
v := int64(999)
r.jobs["imgbatch_download"].APIKeyID = &v
}, want: ErrBatchImageJobNotFound},
{name: "not_completed", id: "imgbatch_download", item: "cover/../001", mutate: func(r *fakeBatchImageRepository) {
r.jobs["imgbatch_download"].Status = BatchImageJobStatusRunning
}, want: ErrBatchImageNotReady},
{name: "output_deleted", id: "imgbatch_download", item: "cover/../001", mutate: func(r *fakeBatchImageRepository) {
r.jobs["imgbatch_download"].Status = BatchImageJobStatusOutputDeleted
}, want: ErrBatchImageOutputDeleted},
{name: "missing_item", id: "imgbatch_download", item: "missing", want: ErrBatchImageItemNotFound},
{name: "failed_item", id: "imgbatch_download", item: "bad", want: ErrBatchImageItemFailed},
{name: "out_of_range", id: "imgbatch_download", item: "cover/../001", index: 2, want: ErrBatchImageItemImageIndexOutOfRange},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
svc, repo, _ := newTestBatchImageDownloadService()
if tt.mutate != nil {
tt.mutate(repo)
}
got, err := svc.OpenItemContent(ctx, testBatchImageOwner(), tt.id, tt.item, tt.index)
require.Nil(t, got)
require.ErrorIs(t, err, tt.want)
require.NotContains(t, err.Error(), batchImageDownloadTestBase64)
require.NotContains(t, err.Error(), "providers/")
require.NotContains(t, err.Error(), "gs://")
})
}
}
func TestBatchImageDownloadService_StreamZip(t *testing.T) {
ctx := context.Background()
t.Run("streams zip with images manifest and errors", func(t *testing.T) {
svc, _, limiter := newTestBatchImageDownloadService()
var buf bytes.Buffer
result, err := svc.StreamZip(ctx, testBatchImageOwner(), "imgbatch_download", BatchImageZipOptions{}, &buf)
require.NoError(t, err)
require.Equal(t, 3, result.FileCount)
require.Equal(t, 1, limiter.acquireCount)
require.Equal(t, 1, limiter.releaseCount)
files := readZipFiles(t, buf.Bytes())
require.Equal(t, []byte("first"), files["images/cover___001.png"])
require.Equal(t, []byte("second"), files["images/cover___001_2.jpg"])
require.Equal(t, []byte("third"), files["images/ok_2.webp"])
require.Contains(t, files, "manifest.json")
require.Contains(t, files, "errors.json")
zipText := string(bytes.Join(mapValues(files), []byte("\n")))
require.NotContains(t, zipText, batchImageDownloadTestBase64)
require.NotContains(t, zipText, "provider_job_name")
require.NotContains(t, zipText, "provider_input_ref")
require.NotContains(t, zipText, "gcs_output_uri")
require.NotContains(t, zipText, "account_id")
require.NotContains(t, zipText, "providers/")
require.NotContains(t, zipText, "gs://")
var manifest struct {
Files []struct {
CustomID string `json:"custom_id"`
Filename string `json:"filename"`
MimeType string `json:"mime_type"`
ImageIndex int `json:"image_index"`
} `json:"files"`
}
require.NoError(t, json.Unmarshal(files["manifest.json"], &manifest))
require.Len(t, manifest.Files, 3)
require.Equal(t, "images/cover___001_2.jpg", manifest.Files[1].Filename)
require.Equal(t, 1, manifest.Files[1].ImageIndex)
var errorsJSON []map[string]string
require.NoError(t, json.Unmarshal(files["errors.json"], &errorsJSON))
require.Len(t, errorsJSON, 1)
require.Equal(t, "bad", errorsJSON[0]["custom_id"])
require.Equal(t, "SAFETY_BLOCKED", errorsJSON[0]["code"])
})
t.Run("limiter denial returns public limit error", func(t *testing.T) {
svc, _, limiter := newTestBatchImageDownloadService()
limiter.deny = true
var buf bytes.Buffer
result, err := svc.StreamZip(ctx, testBatchImageOwner(), "imgbatch_download", BatchImageZipOptions{}, &buf)
require.Nil(t, result)
require.ErrorIs(t, err, ErrBatchImageDownloadLimited)
require.Empty(t, buf.Bytes())
})
t.Run("rejects too many zip items before opening output", func(t *testing.T) {
svc, repo, _ := newTestBatchImageDownloadService()
repo.jobs["imgbatch_download"].SuccessCount = 3
svc.Config.BatchImage.MaxDownloadItemsZip = 1
var buf bytes.Buffer
result, err := svc.StreamZip(ctx, testBatchImageOwner(), "imgbatch_download", BatchImageZipOptions{}, &buf)
require.Nil(t, result)
require.ErrorIs(t, err, ErrBatchImageZipTooManyItems)
require.Empty(t, buf.Bytes())
})
}
func TestExtractBatchImagePartsFromResultLine(t *testing.T) {
tests := []struct {
name string
line string
wantID string
wantMime string
wantError string
}{
{name: "inlineData_mimeType_response", line: `{"key":"a","response":{"candidates":[{"content":{"parts":[{"inlineData":{"mimeType":"image/png","data":"` + batchImageDownloadTestBase64 + `"}}]}}]}}`, wantID: "a", wantMime: "image/png"},
{name: "inline_data_mime_type_top_level", line: `{"custom_id":"b","candidates":[{"content":{"parts":[{"inline_data":{"mime_type":"image/jpeg","data":"` + batchImageDownloadTestBase64 + `"}}]}}]}`, wantID: "b", wantMime: "image/jpeg"},
{name: "status_failure", line: `{"key":"c","status":{"code":"INVALID_ARGUMENT","message":"bad prompt"}}`, wantID: "c", wantError: "INVALID_ARGUMENT"},
{name: "error_failure", line: `{"key":"d","error":{"code":"SAFETY","message":"blocked"}}`, wantID: "d", wantError: "SAFETY_BLOCKED"},
{name: "empty_output", line: `{"key":"e","response":{"candidates":[{"content":{"parts":[{"text":"none"}]}}]}}`, wantID: "e", wantError: "EMPTY_IMAGE_OUTPUT"},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got, err := ExtractBatchImagePartsFromResultLine([]byte(tt.line))
require.NoError(t, err)
require.Equal(t, tt.wantID, got.CustomID)
if tt.wantMime != "" {
require.Len(t, got.Images, 1)
require.Equal(t, tt.wantMime, got.Images[0].MimeType)
require.NotEmpty(t, got.Images[0].Base64Data)
}
if tt.wantError != "" {
require.Equal(t, tt.wantError, got.ErrorCode)
}
})
}
_, err := ExtractBatchImagePartsFromResultLine([]byte(`{"response":{"candidates":[{"content":{"parts":[{"inlineData":{"mimeType":"image/png","data":"` + batchImageDownloadTestBase64 + `"}}]}}]}}`))
require.Error(t, err)
require.NotContains(t, err.Error(), batchImageDownloadTestBase64)
}
func TestBatchImageDownloadFilenames(t *testing.T) {
require.Equal(t, "___secret_name.png", BatchImageSafeDownloadFilename("../../secret\nname", "png"))
require.Equal(t, `attachment; filename="cover_001.png"`, BatchImageContentDispositionAttachment(`cover"001.png`))
}
func newTestBatchImageDownloadService() (*BatchImageDownloadService, *fakeBatchImageRepository, *fakeBatchImageDownloadLimiter) {
repo := newFakeBatchImageRepository()
apiKeyID := int64(22)
accountID := int64(101)
repo.jobs["imgbatch_download"] = &BatchImageJob{
BatchID: "imgbatch_download",
UserID: 11,
APIKeyID: &apiKeyID,
AccountID: &accountID,
Provider: BatchImageProviderGeminiAPI,
Model: "gemini-2.5-flash-image",
Status: BatchImageJobStatusCompleted,
ProviderJobName: batchImageStringPtr("providers/internal/job"),
ProviderOutputRef: batchImageStringPtr("gs://bucket/internal/output.jsonl"),
ItemCount: 3,
SuccessCount: 2,
FailCount: 1,
CreatedAt: time.Now(),
}
mime := "image/png"
ext := "png"
webp := "image/webp"
webpExt := "webp"
code := "SAFETY_BLOCKED"
msg := "blocked in gs://bucket/internal/output.jsonl"
repo.items["imgbatch_download"] = []CreateBatchImageItemParams{
{JobID: "imgbatch_download", CustomID: "cover/../001", Status: BatchImageItemStatusSuccess, MimeType: &mime, FileExtension: &ext, ImageCount: 2},
{JobID: "imgbatch_download", CustomID: "bad", Status: BatchImageItemStatusFailed, ErrorCode: &code, ErrorMessage: &msg},
{JobID: "imgbatch_download", CustomID: "ok_2", Status: BatchImageItemStatusSuccess, MimeType: &webp, FileExtension: &webpExt, ImageCount: 1},
}
provider := &publicBatchImageProvider{name: BatchImageProviderGeminiAPI, result: batchImageDownloadResultJSONL()}
limiter := &fakeBatchImageDownloadLimiter{}
svc := &BatchImageDownloadService{
Repo: repo,
ProviderRegistry: NewBatchImageProviderRegistry(provider),
AccountResolver: &fakeBatchImageAccountResolver{account: &Account{ID: accountID, Platform: PlatformGemini, Type: AccountTypeAPIKey, Status: StatusActive, Schedulable: true}},
Limiter: limiter,
Config: &config.Config{BatchImage: config.BatchImageConfig{MaxDownloadItemsZip: 10, MaxDownloadDurationSeconds: 60}},
}
return svc, repo, limiter
}
const batchImageDownloadTestBase64 = "Zmlyc3Q="
func batchImageDownloadResultJSONL() string {
return strings.Join([]string{
`{"key":"cover/../001","response":{"candidates":[{"content":{"parts":[{"inlineData":{"mimeType":"image/png","data":"Zmlyc3Q="}},{"inlineData":{"mimeType":"image/jpeg","data":"c2Vjb25k"}}]}}]}}`,
`{"key":"bad","error":{"code":"SAFETY","message":"blocked"}}`,
`{"key":"ok_2","candidates":[{"content":{"parts":[{"inline_data":{"mime_type":"image/webp","data":"dGhpcmQ="}}]}}]}`,
}, "\n") + "\n"
}
func readZipFiles(t *testing.T, data []byte) map[string][]byte {
t.Helper()
reader, err := zip.NewReader(bytes.NewReader(data), int64(len(data)))
require.NoError(t, err)
out := make(map[string][]byte, len(reader.File))
for _, file := range reader.File {
rc, err := file.Open()
require.NoError(t, err)
body, err := io.ReadAll(rc)
require.NoError(t, err)
require.NoError(t, rc.Close())
out[file.Name] = body
}
return out
}
func mapValues(in map[string][]byte) [][]byte {
out := make([][]byte, 0, len(in))
for _, value := range in {
out = append(out, value)
}
return out
}
type fakeBatchImageDownloadLimiter struct {
acquireCount int
releaseCount int
deny bool
}
func (l *fakeBatchImageDownloadLimiter) Acquire(context.Context, string, string) (BatchImageDownloadPermit, error) {
l.acquireCount++
if l.deny {
return nil, ErrBatchImageDownloadLimited
}
return &fakeBatchImageDownloadPermit{release: func() { l.releaseCount++ }}, nil
}
type fakeBatchImageDownloadPermit struct {
once bool
release func()
}
func (p *fakeBatchImageDownloadPermit) Release(context.Context) error {
if p.once {
return nil
}
p.once = true
if p.release != nil {
p.release()
}
return nil
}
var _ BatchImageDownloadLimiter = (*fakeBatchImageDownloadLimiter)(nil)
var _ BatchImageDownloadPermit = (*fakeBatchImageDownloadPermit)(nil)
@@ -0,0 +1,265 @@
//go:build unit
package service
import (
"archive/zip"
"bytes"
"context"
"encoding/json"
"io"
"strings"
"testing"
"time"
"github.com/Wei-Shaw/sub2api/internal/config"
"github.com/stretchr/testify/require"
)
func TestBatchImageMVPFlow(t *testing.T) {
ctx := context.Background()
repo := newFakeBatchImageRepository()
queue := &publicBatchImageQueue{}
provider := &batchImageSmokeProvider{
name: BatchImageProviderGeminiAPI,
states: []BatchProviderInternalState{
BatchProviderStateRunning,
BatchProviderStateSucceeded,
},
result: batchImageSmokeResultJSONL(),
}
accountID := int64(101)
accountRepo := &publicBatchImageAccountRepo{accounts: []Account{testBatchImageAccount(accountID, AccountTypeAPIKey)}}
cfg := &config.Config{BatchImage: config.BatchImageConfig{
Enabled: true,
MaxItemsPerJobDefault: 10,
MaxPromptCharsPerItem: 8000,
DefaultResponseMimeType: "image/png",
DefaultImageSize: "1K",
MaxDownloadItemsZip: 10,
MaxDownloadDurationSeconds: 60,
OutputRetentionAfterTerminalHours: 72,
}}
registry := NewBatchImageProviderRegistry(provider)
billing := &fakeBatchImageBillingRepo{}
pricing := &fakeBatchImagePricingResolver{unitPrice: 0.25}
owner := testBatchImageOwner()
publicSvc := &BatchImagePublicService{
Repo: repo,
AccountRepo: accountRepo,
Queue: queue,
ProviderRegistry: registry,
Pricing: pricing,
BillingRepo: billing,
Config: cfg,
}
processor := &BatchImagePipelineProcessor{
ProviderProcessor: &BatchImageProviderProcessor{
Repo: repo,
ProviderRegistry: registry,
AccountResolver: &fakeBatchImageAccountResolver{account: &accountRepo.accounts[0]},
BillingRepo: billing,
},
SettlementService: &BatchImageSettlementService{
Repo: repo,
BillingRepo: billing,
Pricing: pricing,
Config: cfg,
},
}
downloadSvc := &BatchImageDownloadService{
Repo: repo,
ProviderRegistry: registry,
AccountResolver: &fakeBatchImageAccountResolver{account: &accountRepo.accounts[0]},
Limiter: &fakeBatchImageDownloadLimiter{},
Config: cfg,
}
cleanupSvc := &BatchImageCleanupService{
Repo: repo,
ProviderRegistry: registry,
AccountResolver: &fakeBatchImageAccountResolver{account: &accountRepo.accounts[0]},
Config: cfg,
}
submitted, err := publicSvc.Submit(ctx, owner, validBatchImageSubmitRequest(), "")
require.NoError(t, err)
require.Equal(t, "image.batch", submitted.Object)
require.True(t, strings.HasPrefix(submitted.ID, "imgbatch_"))
require.Equal(t, "queued", submitted.Status)
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)
require.NoError(t, err)
require.False(t, firstProcess.Terminal)
require.Equal(t, BatchImageJobStatusRunning, repo.jobs[submitted.ID].Status)
indexProcess, err := processor.Process(ctx, submitted.ID)
require.NoError(t, err)
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])
settleProcess, err := processor.Process(ctx, submitted.ID)
require.NoError(t, err)
require.True(t, settleProcess.Terminal)
job := repo.jobs[submitted.ID]
require.Equal(t, BatchImageJobStatusCompleted, job.Status)
require.NotNil(t, job.OutputExpiresAt)
require.Equal(t, 1, job.SuccessCount)
require.Equal(t, 1, job.FailCount)
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.captures, 1)
status, err := publicSvc.Get(ctx, owner, submitted.ID)
require.NoError(t, err)
require.Equal(t, "completed", status.Status)
require.Equal(t, 1, status.SuccessCount)
require.Equal(t, 1, status.FailCount)
require.NotNil(t, status.ActualCost)
requireBatchImagePublicJSONHasNoInternals(t, mustMarshalBatchImageSmokeJSON(t, status))
items, err := publicSvc.ListItems(ctx, owner, submitted.ID, BatchImageItemsQuery{Limit: 100})
require.NoError(t, err)
require.False(t, items.HasMore)
require.Len(t, items.Data, 2)
require.Equal(t, "cover_001", items.Data[0].CustomID)
require.Equal(t, "succeeded", items.Data[0].Status)
require.Equal(t, "cover_002", items.Data[1].CustomID)
require.Equal(t, "failed", items.Data[1].Status)
require.NotNil(t, items.Data[1].Error)
require.Nil(t, repo.items[submitted.ID][1].BilledAmount)
requireBatchImagePublicJSONHasNoInternals(t, mustMarshalBatchImageSmokeJSON(t, items))
stream, err := downloadSvc.OpenItemContent(ctx, owner, submitted.ID, "cover_001", 0)
require.NoError(t, err)
body, err := io.ReadAll(stream.Reader)
require.NoError(t, err)
require.NoError(t, stream.Reader.Close())
require.Equal(t, []byte("smoke-png"), body)
require.Equal(t, "image/png", stream.ContentType)
require.Equal(t, "cover_001.png", stream.Filename)
var zipBuf bytes.Buffer
zipResult, err := downloadSvc.StreamZip(ctx, owner, submitted.ID, BatchImageZipOptions{}, &zipBuf)
require.NoError(t, err)
require.Equal(t, 1, zipResult.FileCount)
require.Equal(t, 1, zipResult.ErrorCount)
zipFiles := readZipFiles(t, zipBuf.Bytes())
require.Equal(t, []byte("smoke-png"), zipFiles["images/cover_001.png"])
require.Contains(t, zipFiles, "manifest.json")
require.Contains(t, zipFiles, "errors.json")
requireBatchImagePublicJSONHasNoInternals(t, string(bytes.Join(mapValues(zipFiles), []byte("\n"))))
zipReader, err := zip.NewReader(bytes.NewReader(zipBuf.Bytes()), int64(zipBuf.Len()))
require.NoError(t, err)
require.ElementsMatch(t, []string{"images/cover_001.png", "manifest.json", "errors.json"}, batchImageSmokeZipNames(zipReader))
deleted, err := cleanupSvc.DeleteOutputsForOwner(ctx, owner, submitted.ID)
require.NoError(t, err)
require.Equal(t, "output_deleted", deleted.Status)
require.Equal(t, []CleanupTarget{CleanupTargetOutput}, provider.cleanupTargets)
requireBatchImagePublicJSONHasNoInternals(t, mustMarshalBatchImageSmokeJSON(t, deleted))
deletedAgain, err := cleanupSvc.DeleteOutputsForOwner(ctx, owner, submitted.ID)
require.NoError(t, err)
require.Equal(t, "output_deleted", deletedAgain.Status)
require.Equal(t, []CleanupTarget{CleanupTargetOutput}, provider.cleanupTargets)
stream, err = downloadSvc.OpenItemContent(ctx, owner, submitted.ID, "cover_001", 0)
require.Nil(t, stream)
require.ErrorIs(t, err, ErrBatchImageOutputDeleted)
var afterDelete bytes.Buffer
zipResult, err = downloadSvc.StreamZip(ctx, owner, submitted.ID, BatchImageZipOptions{}, &afterDelete)
require.Nil(t, zipResult)
require.ErrorIs(t, err, ErrBatchImageOutputDeleted)
require.Empty(t, afterDelete.Bytes())
}
func batchImageSmokeResultJSONL() string {
return strings.Join([]string{
`{"key":"cover_001","response":{"candidates":[{"content":{"parts":[{"inlineData":{"mimeType":"image/png","data":"c21va2UtcG5n"}}]}}]}}`,
`{"key":"cover_002","status":{"code":3,"message":"blocked by safety policy"}}`,
}, "\n") + "\n"
}
func mustMarshalBatchImageSmokeJSON(t *testing.T, value any) string {
t.Helper()
body, err := json.Marshal(value)
require.NoError(t, err)
return string(body)
}
func batchImageSmokeZipNames(reader *zip.Reader) []string {
names := make([]string, 0, len(reader.File))
for _, file := range reader.File {
names = append(names, file.Name)
}
return names
}
type batchImageSmokeProvider struct {
name string
states []BatchProviderInternalState
submits []BatchImageInput
result string
cleanupTargets []CleanupTarget
}
func (p *batchImageSmokeProvider) Name() string { return p.name }
func (p *batchImageSmokeProvider) SupportsAccount(account *Account) bool {
return account != nil && account.IsSchedulable()
}
func (p *batchImageSmokeProvider) Submit(_ context.Context, _ *BatchImageJob, _ *Account, input BatchImageInput) (*BatchProviderJob, error) {
p.submits = append(p.submits, input)
return &BatchProviderJob{
ProviderJobName: "providers/fake-provider-job/raw-id",
ProviderInputRef: "files/fake-provider-job/input.jsonl",
ProviderOutputRef: "files/fake-provider-job/output.jsonl",
}, nil
}
func (p *batchImageSmokeProvider) Get(context.Context, *BatchImageJob, *Account) (*BatchProviderStatus, error) {
state := BatchProviderStateSucceeded
if len(p.states) > 0 {
state = p.states[0]
p.states = p.states[1:]
}
return &BatchProviderStatus{
RawState: strings.ToUpper(string(state)),
InternalState: state,
Done: state == BatchProviderStateSucceeded,
ProviderOutputRef: "files/fake-provider-job/output.jsonl",
}, nil
}
func (p *batchImageSmokeProvider) Cancel(context.Context, *BatchImageJob, *Account) error {
return nil
}
func (p *batchImageSmokeProvider) OpenResult(context.Context, *BatchImageJob, *Account) (io.ReadCloser, string, error) {
return io.NopCloser(strings.NewReader(p.result)), "application/jsonl", nil
}
func (p *batchImageSmokeProvider) Cleanup(_ context.Context, _ *BatchImageJob, _ *Account, target CleanupTarget) error {
p.cleanupTargets = append(p.cleanupTargets, target)
return nil
}
var _ BatchImageProvider = (*batchImageSmokeProvider)(nil)
@@ -0,0 +1,596 @@
package service
import (
"bufio"
"context"
"encoding/json"
"errors"
"fmt"
"io"
"net/http"
"strconv"
"strings"
"time"
infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors"
"github.com/Wei-Shaw/sub2api/internal/pkg/logger"
"go.uber.org/zap"
)
const (
BatchImageParsedStatusSucceeded = "succeeded"
BatchImageParsedStatusFailed = "failed"
defaultBatchImageProcessorRequeue = 30 * time.Second
batchImageProviderErrorRequeue = time.Minute
batchImageMaxErrorMessageLength = 1000
)
type BatchImageAccountResolver interface {
ResolveBatchImageAccount(ctx context.Context, accountID int64) (*Account, error)
}
type BatchImageAccountLookup interface {
GetByID(ctx context.Context, id int64) (*Account, error)
}
type BatchImageAccountRepositoryResolver struct {
Repo BatchImageAccountLookup
}
func (r *BatchImageAccountRepositoryResolver) ResolveBatchImageAccount(ctx context.Context, accountID int64) (*Account, error) {
if r == nil || r.Repo == nil {
return nil, ErrAccountNotFound
}
return r.Repo.GetByID(ctx, accountID)
}
type BatchImageProviderProcessor struct {
Repo BatchImageRepository
ProviderRegistry *BatchImageProviderRegistry
AccountResolver BatchImageAccountResolver
Indexer *BatchImageResultIndexer
BillingRepo UsageBillingRepository
AuthCache APIKeyAuthCacheInvalidator
DefaultRequeue time.Duration
}
func (p *BatchImageProviderProcessor) Process(ctx context.Context, batchID string) (BatchImageProcessResult, error) {
if p == nil || p.Repo == nil || p.ProviderRegistry == nil || p.AccountResolver == nil {
return BatchImageProcessResult{}, infraerrors.New(http.StatusInternalServerError, "BATCH_IMAGE_PROCESSOR_NOT_CONFIGURED", "batch image processor is not configured")
}
job, err := p.Repo.GetBatchImageJobByBatchID(ctx, batchID)
if err != nil {
return BatchImageProcessResult{}, err
}
if isBatchImageProcessorDoneStatus(job.Status) {
if err := p.releaseTerminalHold(ctx, job); err != nil {
return BatchImageProcessResult{}, err
}
return BatchImageProcessResult{Terminal: true}, nil
}
provider, ok := p.ProviderRegistry.Get(job.Provider)
if !ok || provider == nil {
return BatchImageProcessResult{}, ErrBatchImageUnsupportedProvider
}
if job.AccountID == nil || *job.AccountID <= 0 {
return BatchImageProcessResult{}, ErrBatchImageMissingAccountID
}
account, err := p.AccountResolver.ResolveBatchImageAccount(ctx, *job.AccountID)
if err != nil {
return BatchImageProcessResult{}, err
}
if !provider.SupportsAccount(account) {
return BatchImageProcessResult{}, ErrBatchImageProviderUnsupportedAccount
}
if strings.TrimSpace(batchImageDerefString(job.ProviderJobName)) == "" {
return BatchImageProcessResult{}, ErrBatchImageMissingProviderJobName
}
if job.Status == BatchImageJobStatusIndexing {
return p.indexAndSettle(ctx, job, provider, account)
}
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 {
return BatchImageProcessResult{RequeueAfter: p.requeueDelay(0)}, nil
}
if err := p.persistProviderOutputRef(ctx, job, status.ProviderOutputRef); err != nil {
return BatchImageProcessResult{}, err
}
switch status.InternalState {
case BatchProviderStateQueued:
return BatchImageProcessResult{RequeueAfter: p.requeueDelay(status.SuggestedRequeueAfter)}, nil
case BatchProviderStateRunning:
if job.Status != BatchImageJobStatusRunning {
if err := p.Repo.TransitionBatchImageJobStatus(ctx, job.BatchID, BatchImageJobStatusRunning, BatchImageTransitionOptions{
EventType: "provider_status_checked",
EventPayload: map[string]any{"provider_state": status.RawState},
}); err != nil {
return BatchImageProcessResult{}, err
}
job.Status = BatchImageJobStatusRunning
}
return BatchImageProcessResult{RequeueAfter: p.requeueDelay(status.SuggestedRequeueAfter)}, nil
case BatchProviderStateSucceeded:
if job.Status != BatchImageJobStatusIndexing {
if err := p.Repo.TransitionBatchImageJobStatus(ctx, job.BatchID, BatchImageJobStatusIndexing, BatchImageTransitionOptions{
EventType: "indexing_started",
EventPayload: map[string]any{"provider_state": status.RawState},
}); err != nil {
return BatchImageProcessResult{}, err
}
job.Status = BatchImageJobStatusIndexing
}
return p.indexAndSettle(ctx, job, provider, account)
case BatchProviderStateFailed, BatchProviderStateExpired:
code := strings.TrimSpace(status.ErrorCode)
if code == "" && status.InternalState == BatchProviderStateExpired {
code = "PROVIDER_BATCH_EXPIRED"
}
if code == "" {
code = "PROVIDER_BATCH_FAILED"
}
msg := truncateBatchImageMessage(status.ErrorMessage, batchImageMaxErrorMessageLength)
if err := p.Repo.TransitionBatchImageJobStatus(ctx, job.BatchID, BatchImageJobStatusFailed, BatchImageTransitionOptions{
EventType: "job_failed",
EventPayload: map[string]any{"provider_state": status.RawState, "error_code": code},
ErrorCode: batchImageStringPtr(code),
ErrorMessage: batchImageOptionalStringPtr(msg),
}); 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{
EventType: "job_failed",
EventPayload: map[string]any{"provider_state": status.RawState, "error_code": "PROVIDER_BATCH_CANCELLED"},
}); 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
}
}
func (p *BatchImageProviderProcessor) indexAndSettle(ctx context.Context, job *BatchImageJob, provider BatchImageProvider, account *Account) (BatchImageProcessResult, error) {
indexer := p.Indexer
if indexer == nil {
indexer = &BatchImageResultIndexer{Repo: p.Repo}
}
if indexer.Repo == nil {
indexer.Repo = p.Repo
}
result, err := indexer.Index(ctx, job, provider, account)
if err != nil {
if errors.Is(err, ErrBatchImageIndexOutputMissing) {
return BatchImageProcessResult{}, err
}
code := "INDEX_PARSE_FAILED"
if errors.Is(err, ErrBatchImageDuplicateCustomID) {
code = "DUPLICATE_CUSTOM_ID_IN_OUTPUT"
}
msg := truncateBatchImageMessage(err.Error(), batchImageMaxErrorMessageLength)
transitionErr := p.Repo.TransitionBatchImageJobStatus(ctx, job.BatchID, BatchImageJobStatusFailed, BatchImageTransitionOptions{
EventType: "indexing_failed",
EventPayload: map[string]any{"error_code": code},
ErrorCode: batchImageStringPtr(code),
ErrorMessage: batchImageOptionalStringPtr(msg),
})
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
}
if err := p.Repo.TransitionBatchImageJobStatus(ctx, job.BatchID, BatchImageJobStatusSettling, BatchImageTransitionOptions{
EventType: "indexing_completed",
EventPayload: map[string]any{
"success_count": result.SuccessCount,
"fail_count": result.FailCount,
"total_count": result.TotalCount,
},
}); err != nil {
return BatchImageProcessResult{}, err
}
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 {
ref = strings.TrimSpace(ref)
if ref == "" || job == nil || batchImageDerefString(job.ProviderOutputRef) == ref {
return nil
}
if err := p.Repo.UpdateBatchImageJobProviderOutputRef(ctx, job.BatchID, ref); err != nil {
return err
}
job.ProviderOutputRef = &ref
return nil
}
func (p *BatchImageProviderProcessor) requeueDelay(suggested time.Duration) time.Duration {
if suggested > 0 {
return suggested
}
if p != nil && p.DefaultRequeue > 0 {
return p.DefaultRequeue
}
return defaultBatchImageProcessorRequeue
}
func isBatchImageProcessorDoneStatus(status string) bool {
if status == BatchImageJobStatusSettling {
return true
}
return IsTerminalBatchImageJobStatus(status)
}
type BatchImageIndexResult struct {
SuccessCount int
FailCount int
TotalCount int
}
type BatchImageResultIndexer struct {
Repo BatchImageRepository
}
func (i *BatchImageResultIndexer) Index(ctx context.Context, job *BatchImageJob, provider BatchImageProvider, account *Account) (*BatchImageIndexResult, error) {
if i == nil || i.Repo == nil || job == nil || provider == nil {
return nil, ErrBatchImageIndexOutputMissing
}
r, _, err := provider.OpenResult(ctx, job, account)
if err != nil {
return nil, ErrBatchImageIndexOutputMissing.WithCause(err)
}
defer func() { _ = r.Close() }()
scanner := bufio.NewScanner(r)
scanner.Buffer(make([]byte, 0, 64*1024), 16*1024*1024)
seen := make(map[string]int)
var items []CreateBatchImageItemParams
result := &BatchImageIndexResult{}
lineNumber := 0
now := time.Now()
sourceObject := batchImageDerefString(job.ProviderOutputRef)
if sourceObject == "" {
sourceObject = batchImageDerefString(job.ProviderJobName)
}
for scanner.Scan() {
lineNumber++
line := strings.TrimSpace(scanner.Text())
if line == "" {
continue
}
parsed, err := ParseBatchImageResultLine([]byte(line), lineNumber)
if err != nil {
return nil, err
}
if firstLine, ok := seen[parsed.CustomID]; ok {
return nil, ErrBatchImageDuplicateCustomID.WithCause(fmt.Errorf("custom id %q duplicated at lines %d and %d", parsed.CustomID, firstLine, lineNumber))
}
seen[parsed.CustomID] = lineNumber
lineNo := parsed.SourceLineNumber
item := CreateBatchImageItemParams{
JobID: job.BatchID,
CustomID: parsed.CustomID,
Status: BatchImageItemStatusFailed,
ProviderSourceObject: batchImageOptionalStringPtr(sourceObject),
SourceLineNumber: &lineNo,
ImageCount: parsed.ImageCount,
IndexedAt: &now,
}
if parsed.Status == BatchImageParsedStatusSucceeded {
item.Status = BatchImageItemStatusSuccess
item.MimeType = batchImageOptionalStringPtr(parsed.MimeType)
item.FileExtension = batchImageOptionalStringPtr(parsed.FileExtension)
result.SuccessCount++
} else {
item.ErrorCode = batchImageOptionalStringPtr(parsed.ErrorCode)
item.ErrorMessage = batchImageOptionalStringPtr(parsed.ErrorMessage)
result.FailCount++
}
items = append(items, item)
result.TotalCount++
}
if err := scanner.Err(); err != nil {
if errors.Is(err, io.ErrUnexpectedEOF) {
return nil, ErrBatchImageIndexParseFailed.WithCause(err)
}
return nil, err
}
if result.TotalCount == 0 {
return nil, ErrBatchImageIndexNoResultLines
}
if err := i.Repo.ReplaceBatchImageItemsForJob(ctx, job.BatchID, items, BatchImageCounts{
SuccessCount: result.SuccessCount,
FailCount: result.FailCount,
}); err != nil {
return nil, err
}
return result, nil
}
type ParsedBatchImageResult struct {
CustomID string
Status string
MimeType string
FileExtension string
ImageCount int
ErrorCode string
ErrorMessage string
SourceLineNumber int
}
func ParseBatchImageResultLine(line []byte, lineNumber int) (*ParsedBatchImageResult, error) {
var obj map[string]any
if err := json.Unmarshal(line, &obj); err != nil {
return nil, ErrBatchImageIndexParseFailed.WithCause(fmt.Errorf("line %d: %w", lineNumber, err))
}
customID := batchImageFirstNonEmptyString(
batchImageMapString(obj, "key"),
batchImageMapString(obj, "custom_id"),
batchImageMapString(obj, "customId"),
batchImageNestedString(obj, "request", "key"),
)
if customID == "" {
return nil, ErrBatchImageIndexParseFailed.WithCause(fmt.Errorf("line %d: missing custom id", lineNumber))
}
parsed := &ParsedBatchImageResult{
CustomID: customID,
SourceLineNumber: lineNumber,
}
imageCount, mimeType := batchImageFindImageParts(obj)
if imageCount > 0 {
parsed.Status = BatchImageParsedStatusSucceeded
parsed.ImageCount = imageCount
parsed.MimeType = mimeType
parsed.FileExtension = batchImageFileExtension(mimeType)
return parsed, nil
}
if code, message, ok := batchImageFailureFromProviderFields(obj); ok {
parsed.Status = BatchImageParsedStatusFailed
parsed.ErrorCode = code
parsed.ErrorMessage = truncateBatchImageMessage(message, batchImageMaxErrorMessageLength)
return parsed, nil
}
if _, hasResponse := obj["response"]; hasResponse || batchImageHasCandidates(obj) {
parsed.Status = BatchImageParsedStatusFailed
parsed.ErrorCode = "EMPTY_IMAGE_OUTPUT"
parsed.ErrorMessage = "provider response contained no image output"
return parsed, nil
}
parsed.Status = BatchImageParsedStatusFailed
parsed.ErrorCode = "PROVIDER_ITEM_FAILED"
parsed.ErrorMessage = "provider result line contained no image output"
return parsed, nil
}
func batchImageFindImageParts(obj map[string]any) (int, string) {
count, mimeType := batchImageFindImagePartsInCandidates(batchImageNestedAny(obj, "response", "candidates"))
if count > 0 {
return count, mimeType
}
return batchImageFindImagePartsInCandidates(obj["candidates"])
}
func batchImageFindImagePartsInCandidates(raw any) (int, string) {
candidates, ok := raw.([]any)
if !ok {
return 0, ""
}
count := 0
firstMime := ""
for _, candidateRaw := range candidates {
candidate, ok := candidateRaw.(map[string]any)
if !ok {
continue
}
partsRaw := batchImageNestedAny(candidate, "content", "parts")
parts, ok := partsRaw.([]any)
if !ok {
continue
}
for _, partRaw := range parts {
part, ok := partRaw.(map[string]any)
if !ok {
continue
}
inline, ok := firstMap(part["inlineData"], part["inline_data"])
if !ok {
continue
}
data := strings.TrimSpace(batchImageMapString(inline, "data"))
mime := batchImageFirstNonEmptyString(batchImageMapString(inline, "mimeType"), batchImageMapString(inline, "mime_type"))
if data == "" || !strings.HasPrefix(strings.ToLower(strings.TrimSpace(mime)), "image/") {
continue
}
count++
if firstMime == "" {
firstMime = strings.TrimSpace(mime)
}
}
}
return count, firstMime
}
func batchImageFailureFromProviderFields(obj map[string]any) (string, string, bool) {
if status, ok := obj["status"].(map[string]any); ok {
message := batchImageFirstNonEmptyString(batchImageMapString(status, "message"), batchImageMapString(status, "details"))
code := batchImageFirstNonEmptyString(batchImageMapString(status, "code"), batchImageMapString(status, "status"))
return batchImageMapFailureCode(code, message), message, true
}
if errObj, ok := obj["error"].(map[string]any); ok {
message := batchImageFirstNonEmptyString(batchImageMapString(errObj, "message"), batchImageMapString(errObj, "details"))
code := batchImageFirstNonEmptyString(batchImageMapString(errObj, "code"), batchImageMapString(errObj, "status"))
return batchImageMapFailureCode(code, message), message, true
}
return "", "", false
}
func batchImageMapFailureCode(code, message string) string {
text := strings.ToLower(strings.TrimSpace(code + " " + message))
switch {
case strings.Contains(text, "safety"), strings.Contains(text, "policy"), strings.Contains(text, "blocked"), strings.Contains(text, "prohibited"):
return "SAFETY_BLOCKED"
case strings.Contains(text, "invalid_argument"), strings.Contains(text, "invalid argument"), strings.Contains(text, "bad request"):
return "INVALID_ARGUMENT"
case strings.Contains(text, "quota"), strings.Contains(text, "rate"), strings.Contains(text, "resource_exhausted"), strings.Contains(text, "too many requests"):
return "PROVIDER_RATE_LIMITED"
default:
return "PROVIDER_ITEM_FAILED"
}
}
func batchImageFileExtension(mimeType string) string {
switch strings.ToLower(strings.TrimSpace(mimeType)) {
case "image/png":
return "png"
case "image/jpeg", "image/jpg":
return "jpg"
case "image/webp":
return "webp"
default:
return ""
}
}
func batchImageHasCandidates(obj map[string]any) bool {
if _, ok := obj["candidates"]; ok {
return true
}
_, ok := batchImageNestedAny(obj, "response", "candidates").([]any)
return ok
}
func batchImageMapString(m map[string]any, key string) string {
if m == nil {
return ""
}
switch v := m[key].(type) {
case string:
return strings.TrimSpace(v)
case json.Number:
return v.String()
case float64:
return strconv.FormatInt(int64(v), 10)
default:
return ""
}
}
func batchImageNestedString(m map[string]any, keys ...string) string {
if nested, ok := batchImageNestedAny(m, keys...).(string); ok {
return strings.TrimSpace(nested)
}
return ""
}
func batchImageNestedAny(m map[string]any, keys ...string) any {
var current any = m
for _, key := range keys {
cm, ok := current.(map[string]any)
if !ok {
return nil
}
current = cm[key]
}
return current
}
func firstMap(values ...any) (map[string]any, bool) {
for _, value := range values {
if m, ok := value.(map[string]any); ok {
return m, true
}
}
return nil, false
}
func batchImageFirstNonEmptyString(values ...string) string {
for _, value := range values {
if strings.TrimSpace(value) != "" {
return strings.TrimSpace(value)
}
}
return ""
}
func batchImageDerefString(v *string) string {
if v == nil {
return ""
}
return strings.TrimSpace(*v)
}
func batchImageStringPtr(v string) *string {
return &v
}
func batchImageOptionalStringPtr(v string) *string {
v = strings.TrimSpace(v)
if v == "" {
return nil
}
return &v
}
func truncateBatchImageMessage(message string, limit int) string {
message = strings.TrimSpace(message)
if limit <= 0 || len(message) <= limit {
return message
}
return message[:limit]
}
@@ -0,0 +1,847 @@
//go:build unit
package service
import (
"context"
"errors"
"fmt"
"io"
"strings"
"testing"
"time"
infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors"
"github.com/stretchr/testify/require"
)
const batchImageTestData = "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJ"
func TestParseBatchImageResultLine_SuccessShapes(t *testing.T) {
tests := []struct {
name string
line string
wantID string
wantMime string
wantExt string
wantCount int
}{
{
name: "gemini_inlineData",
line: `{"key":"cover_001","response":{"candidates":[{"content":{"parts":[{"inlineData":{"mimeType":"image/png","data":"` + batchImageTestData + `"}}]}}]}}`,
wantID: "cover_001", wantMime: "image/png", wantExt: "png", wantCount: 1,
},
{
name: "snake_case_inline_data",
line: `{"custom_id":"cover_002","response":{"candidates":[{"content":{"parts":[{"inline_data":{"mime_type":"image/jpeg","data":"` + batchImageTestData + `"}}]}}]}}`,
wantID: "cover_002", wantMime: "image/jpeg", wantExt: "jpg", wantCount: 1,
},
{
name: "vertex_top_level_response",
line: `{"customId":"cover_003","response":{"candidates":[{"content":{"parts":[{"inlineData":{"mimeType":"image/webp","data":"` + batchImageTestData + `"}}]}}]}}`,
wantID: "cover_003", wantMime: "image/webp", wantExt: "webp", wantCount: 1,
},
{
name: "top_level_candidates",
line: `{"request":{"key":"cover_004"},"candidates":[{"content":{"parts":[{"inline_data":{"mime_type":"image/png","data":"` + batchImageTestData + `"}},{"inlineData":{"mimeType":"image/png","data":"` + batchImageTestData + `"}}]}}]}`,
wantID: "cover_004", wantMime: "image/png", wantExt: "png", wantCount: 2,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got, err := ParseBatchImageResultLine([]byte(tt.line), 7)
require.NoError(t, err)
require.Equal(t, tt.wantID, got.CustomID)
require.Equal(t, BatchImageParsedStatusSucceeded, got.Status)
require.Equal(t, tt.wantMime, got.MimeType)
require.Equal(t, tt.wantExt, got.FileExtension)
require.Equal(t, tt.wantCount, got.ImageCount)
require.Equal(t, 7, got.SourceLineNumber)
require.NotContains(t, fmt.Sprintf("%+v", got), batchImageTestData)
})
}
}
func TestParseBatchImageResultLine_FailureShapes(t *testing.T) {
tests := []struct {
name string
line string
wantCode string
}{
{name: "status_row", line: `{"key":"cover_001","status":{"code":3,"message":"invalid argument: bad prompt"}}`, wantCode: "INVALID_ARGUMENT"},
{name: "error_row", line: `{"key":"cover_002","error":{"code":"SAFETY","message":"blocked by safety policy"}}`, wantCode: "SAFETY_BLOCKED"},
{name: "quota_row", line: `{"key":"cover_003","error":{"code":"RESOURCE_EXHAUSTED","message":"quota exceeded"}}`, wantCode: "PROVIDER_RATE_LIMITED"},
{name: "empty_image_output", line: `{"key":"cover_004","response":{"candidates":[{"content":{"parts":[{"text":"no image"}]}}]}}`, wantCode: "EMPTY_IMAGE_OUTPUT"},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got, err := ParseBatchImageResultLine([]byte(tt.line), 1)
require.NoError(t, err)
require.Equal(t, BatchImageParsedStatusFailed, got.Status)
require.Equal(t, tt.wantCode, got.ErrorCode)
})
}
}
func TestParseBatchImageResultLine_RejectsMissingCustomIDAndDoesNotLeakData(t *testing.T) {
_, err := ParseBatchImageResultLine([]byte(`{"response":{"candidates":[{"content":{"parts":[{"inlineData":{"mimeType":"image/png","data":"`+batchImageTestData+`"}}]}}]}}`), 3)
require.ErrorIs(t, err, ErrBatchImageIndexParseFailed)
require.NotContains(t, err.Error(), batchImageTestData)
}
func TestBatchImageResultIndexer_WritesCountsAndReplacesItems(t *testing.T) {
output := strings.Join([]string{
`{"key":"ok","response":{"candidates":[{"content":{"parts":[{"inlineData":{"mimeType":"image/png","data":"` + batchImageTestData + `"}}]}}]}}`,
`{"key":"bad","error":{"code":"SAFETY","message":"blocked by safety policy"}}`,
}, "\n") + "\n"
repo := newFakeBatchImageRepository()
outputRef := "files/output"
job := &BatchImageJob{BatchID: "imgbatch_index", ProviderOutputRef: &outputRef}
provider := &fakeProcessorProvider{result: output}
result, err := (&BatchImageResultIndexer{Repo: repo}).Index(context.Background(), job, provider, &Account{})
require.NoError(t, err)
require.True(t, provider.openResultCalled)
require.Equal(t, 1, result.SuccessCount)
require.Equal(t, 1, result.FailCount)
require.Equal(t, 2, result.TotalCount)
require.Equal(t, 1, repo.replaceCalls)
require.Len(t, repo.items[job.BatchID], 2)
require.Equal(t, BatchImageItemStatusSuccess, repo.items[job.BatchID][0].Status)
require.Equal(t, BatchImageItemStatusFailed, repo.items[job.BatchID][1].Status)
require.Equal(t, BatchImageCounts{SuccessCount: 1, FailCount: 1}, repo.counts[job.BatchID])
require.NotContains(t, fmt.Sprintf("%+v", repo.items[job.BatchID]), batchImageTestData)
provider.result = `{"key":"ok2","response":{"candidates":[{"content":{"parts":[{"inlineData":{"mimeType":"image/webp","data":"` + batchImageTestData + `"}}]}}]}}` + "\n"
result, err = (&BatchImageResultIndexer{Repo: repo}).Index(context.Background(), job, provider, &Account{})
require.NoError(t, err)
require.Equal(t, 1, result.TotalCount)
require.Len(t, repo.items[job.BatchID], 1)
require.Equal(t, "ok2", repo.items[job.BatchID][0].CustomID)
}
func TestBatchImageResultIndexer_EmptyInvalidAndDuplicateOutput(t *testing.T) {
tests := []struct {
name string
body string
want error
}{
{name: "empty", body: "\n", want: ErrBatchImageIndexNoResultLines},
{name: "invalid_json", body: "{bad-json}\n", want: ErrBatchImageIndexParseFailed},
{name: "duplicate_custom_id", body: `{"key":"dup","error":{"message":"one"}}` + "\n" + `{"key":"dup","error":{"message":"two"}}` + "\n", want: ErrBatchImageDuplicateCustomID},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
repo := newFakeBatchImageRepository()
_, err := (&BatchImageResultIndexer{Repo: repo}).Index(context.Background(), &BatchImageJob{BatchID: "imgbatch_bad"}, &fakeProcessorProvider{result: tt.body}, &Account{})
require.ErrorIs(t, err, tt.want)
require.Empty(t, repo.items["imgbatch_bad"])
})
}
}
func TestBatchImageProviderProcessor_ValidationAndTerminalCases(t *testing.T) {
ctx := context.Background()
accountID := int64(10)
providerJob := "providers/job"
t.Run("terminal job returns without provider call", func(t *testing.T) {
repo := newFakeBatchImageRepository()
repo.jobs["imgbatch_done"] = &BatchImageJob{BatchID: "imgbatch_done", Status: BatchImageJobStatusFailed}
provider := &fakeProcessorProvider{}
got, err := (&BatchImageProviderProcessor{
Repo: repo, ProviderRegistry: NewBatchImageProviderRegistry(provider), AccountResolver: &fakeBatchImageAccountResolver{account: &Account{}},
}).Process(ctx, "imgbatch_done")
require.NoError(t, err)
require.True(t, got.Terminal)
require.False(t, provider.getCalled)
})
t.Run("missing provider", func(t *testing.T) {
repo := newFakeBatchImageRepository()
repo.jobs["imgbatch_missing_provider"] = &BatchImageJob{BatchID: "imgbatch_missing_provider", Status: BatchImageJobStatusSubmitted, Provider: "missing", AccountID: &accountID, ProviderJobName: &providerJob}
_, err := (&BatchImageProviderProcessor{Repo: repo, ProviderRegistry: NewBatchImageProviderRegistry(), AccountResolver: &fakeBatchImageAccountResolver{account: &Account{}}}).Process(ctx, "imgbatch_missing_provider")
require.ErrorIs(t, err, ErrBatchImageUnsupportedProvider)
})
t.Run("missing account id", func(t *testing.T) {
repo := newFakeBatchImageRepository()
repo.jobs["imgbatch_missing_account"] = &BatchImageJob{BatchID: "imgbatch_missing_account", Status: BatchImageJobStatusSubmitted, Provider: "fake", ProviderJobName: &providerJob}
_, err := (&BatchImageProviderProcessor{Repo: repo, ProviderRegistry: NewBatchImageProviderRegistry(&fakeProcessorProvider{}), AccountResolver: &fakeBatchImageAccountResolver{account: &Account{}}}).Process(ctx, "imgbatch_missing_account")
require.ErrorIs(t, err, ErrBatchImageMissingAccountID)
})
t.Run("missing provider job name", func(t *testing.T) {
repo := newFakeBatchImageRepository()
repo.jobs["imgbatch_missing_name"] = &BatchImageJob{BatchID: "imgbatch_missing_name", Status: BatchImageJobStatusSubmitted, Provider: "fake", AccountID: &accountID}
_, err := (&BatchImageProviderProcessor{Repo: repo, ProviderRegistry: NewBatchImageProviderRegistry(&fakeProcessorProvider{}), AccountResolver: &fakeBatchImageAccountResolver{account: &Account{}}}).Process(ctx, "imgbatch_missing_name")
require.ErrorIs(t, err, ErrBatchImageMissingProviderJobName)
})
}
func TestBatchImageProviderProcessor_StatusFlow(t *testing.T) {
ctx := context.Background()
accountID := int64(10)
providerJob := "providers/job"
newJob := func(status string) *BatchImageJob {
return &BatchImageJob{BatchID: "imgbatch_flow", Status: status, Provider: "fake", AccountID: &accountID, ProviderJobName: &providerJob}
}
t.Run("running status updates and requeues", func(t *testing.T) {
repo := newFakeBatchImageRepository()
repo.jobs["imgbatch_flow"] = newJob(BatchImageJobStatusSubmitted)
provider := &fakeProcessorProvider{status: &BatchProviderStatus{InternalState: BatchProviderStateRunning, RawState: "RUNNING", SuggestedRequeueAfter: 12 * time.Second}}
got, err := newTestBatchImageProcessor(repo, provider).Process(ctx, "imgbatch_flow")
require.NoError(t, err)
require.False(t, got.Terminal)
require.Equal(t, 12*time.Second, got.RequeueAfter)
require.Equal(t, BatchImageJobStatusRunning, repo.jobs["imgbatch_flow"].Status)
})
t.Run("queued status requeues", func(t *testing.T) {
repo := newFakeBatchImageRepository()
repo.jobs["imgbatch_flow"] = newJob(BatchImageJobStatusSubmitted)
provider := &fakeProcessorProvider{status: &BatchProviderStatus{InternalState: BatchProviderStateQueued}}
got, err := newTestBatchImageProcessor(repo, provider).Process(ctx, "imgbatch_flow")
require.NoError(t, err)
require.False(t, got.Terminal)
require.Equal(t, defaultBatchImageProcessorRequeue, got.RequeueAfter)
require.Equal(t, BatchImageJobStatusSubmitted, repo.jobs["imgbatch_flow"].Status)
})
t.Run("transient provider get error requeues", func(t *testing.T) {
repo := newFakeBatchImageRepository()
repo.jobs["imgbatch_flow"] = newJob(BatchImageJobStatusSubmitted)
provider := &fakeProcessorProvider{getErr: errors.New("temporary upstream failure")}
got, err := newTestBatchImageProcessor(repo, provider).Process(ctx, "imgbatch_flow")
require.NoError(t, err)
require.False(t, got.Terminal)
require.Equal(t, time.Minute, got.RequeueAfter)
})
t.Run("succeeded indexes and settles from submitted", func(t *testing.T) {
repo := newFakeBatchImageRepository()
repo.jobs["imgbatch_flow"] = newJob(BatchImageJobStatusSubmitted)
provider := &fakeProcessorProvider{
status: &BatchProviderStatus{InternalState: BatchProviderStateSucceeded, RawState: "SUCCEEDED", ProviderOutputRef: "files/output"},
result: `{"key":"ok","response":{"candidates":[{"content":{"parts":[{"inlineData":{"mimeType":"image/png","data":"` + batchImageTestData + `"}}]}}]}}` + "\n",
}
got, err := newTestBatchImageProcessor(repo, provider).Process(ctx, "imgbatch_flow")
require.NoError(t, err)
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"])
require.Equal(t, BatchImageCounts{SuccessCount: 1}, repo.counts["imgbatch_flow"])
})
t.Run("failed provider marks job failed", func(t *testing.T) {
repo := newFakeBatchImageRepository()
repo.jobs["imgbatch_flow"] = newJob(BatchImageJobStatusRunning)
provider := &fakeProcessorProvider{status: &BatchProviderStatus{InternalState: BatchProviderStateFailed, RawState: "FAILED", ErrorCode: "BAD_PROMPT", ErrorMessage: "bad prompt"}}
got, err := newTestBatchImageProcessor(repo, provider).Process(ctx, "imgbatch_flow")
require.NoError(t, err)
require.True(t, got.Terminal)
require.Equal(t, BatchImageJobStatusFailed, repo.jobs["imgbatch_flow"].Status)
require.Equal(t, "BAD_PROMPT", batchImageDerefString(repo.jobs["imgbatch_flow"].LastErrorCode))
})
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"}}
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)
})
}
func TestCanTransitionBatchImageJob_PR5DirectIndexing(t *testing.T) {
require.True(t, CanTransitionBatchImageJob(BatchImageJobStatusSubmitted, BatchImageJobStatusIndexing))
require.True(t, CanTransitionBatchImageJob(BatchImageJobStatusSubmitted, BatchImageJobStatusFailed))
require.True(t, CanTransitionBatchImageJob(BatchImageJobStatusIndexing, BatchImageJobStatusFailed))
}
func newTestBatchImageProcessor(repo *fakeBatchImageRepository, provider *fakeProcessorProvider) *BatchImageProviderProcessor {
return &BatchImageProviderProcessor{
Repo: repo,
ProviderRegistry: NewBatchImageProviderRegistry(provider),
AccountResolver: &fakeBatchImageAccountResolver{account: &Account{}},
Indexer: &BatchImageResultIndexer{Repo: repo},
}
}
type fakeBatchImageAccountResolver struct {
account *Account
err error
}
func (r *fakeBatchImageAccountResolver) ResolveBatchImageAccount(context.Context, int64) (*Account, error) {
if r.err != nil {
return nil, r.err
}
return r.account, nil
}
type fakeProcessorProvider struct {
status *BatchProviderStatus
getErr error
result string
getCalled bool
openResultCalled bool
}
func (p *fakeProcessorProvider) Name() string { return "fake" }
func (p *fakeProcessorProvider) SupportsAccount(*Account) bool {
return true
}
func (p *fakeProcessorProvider) Submit(context.Context, *BatchImageJob, *Account, BatchImageInput) (*BatchProviderJob, error) {
panic("Submit must not be called by PR5 processor")
}
func (p *fakeProcessorProvider) Get(context.Context, *BatchImageJob, *Account) (*BatchProviderStatus, error) {
p.getCalled = true
if p.getErr != nil {
return nil, p.getErr
}
if p.status == nil {
return &BatchProviderStatus{InternalState: BatchProviderStateQueued}, nil
}
return p.status, nil
}
func (p *fakeProcessorProvider) Cancel(context.Context, *BatchImageJob, *Account) error { return nil }
func (p *fakeProcessorProvider) OpenResult(context.Context, *BatchImageJob, *Account) (io.ReadCloser, string, error) {
p.openResultCalled = true
return io.NopCloser(strings.NewReader(p.result)), "application/jsonl", nil
}
func (p *fakeProcessorProvider) Cleanup(context.Context, *BatchImageJob, *Account, CleanupTarget) error {
return nil
}
type fakeBatchImageRepository struct {
jobs map[string]*BatchImageJob
items map[string][]CreateBatchImageItemParams
counts map[string]BatchImageCounts
transitions map[string][]string
events map[string][]string
transitionErr error
replaceCalls int
}
func newFakeBatchImageRepository() *fakeBatchImageRepository {
return &fakeBatchImageRepository{
jobs: make(map[string]*BatchImageJob),
items: make(map[string][]CreateBatchImageItemParams),
counts: make(map[string]BatchImageCounts),
transitions: make(map[string][]string),
events: make(map[string][]string),
}
}
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,
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
}
func (r *fakeBatchImageRepository) GetBatchImageJobByBatchID(_ context.Context, batchID string) (*BatchImageJob, error) {
job, ok := r.jobs[batchID]
if !ok {
return nil, ErrBatchImageJobNotFound
}
return job, nil
}
func (r *fakeBatchImageRepository) GetBatchImageJobByIdempotencyKey(_ context.Context, userID, apiKeyID int64, key string) (*BatchImageJob, error) {
for _, job := range r.jobs {
if job.UserID == userID && job.APIKeyID != nil && *job.APIKeyID == apiKeyID && batchImageDerefString(job.IdempotencyKey) == key {
return job, nil
}
}
return nil, ErrBatchImageJobNotFound
}
func (r *fakeBatchImageRepository) GetBatchImageJobByBatchIDForOwner(_ context.Context, userID, apiKeyID int64, batchID string) (*BatchImageJob, error) {
job, ok := r.jobs[batchID]
if !ok || job.UserID != userID || job.APIKeyID == nil || *job.APIKeyID != apiKeyID {
return nil, ErrBatchImageJobNotFound
}
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 {
return job, nil
}
}
return nil, ErrBatchImageJobNotFound
}
func (r *fakeBatchImageRepository) TransitionBatchImageJobStatus(_ context.Context, batchID, toStatus string, opts BatchImageTransitionOptions) error {
job, ok := r.jobs[batchID]
if !ok {
return ErrBatchImageJobNotFound
}
if !CanTransitionBatchImageJob(job.Status, toStatus) {
return ErrBatchImageInvalidTransition
}
if r.transitionErr != nil {
return r.transitionErr
}
job.Status = toStatus
job.LastErrorCode = opts.ErrorCode
job.LastErrorMessage = opts.ErrorMessage
r.transitions[batchID] = append(r.transitions[batchID], toStatus)
if opts.EventType != "" {
r.events[batchID] = append(r.events[batchID], opts.EventType)
}
return nil
}
func (r *fakeBatchImageRepository) UpdateBatchImageJobProviderOutputRef(_ context.Context, batchID, providerOutputRef string) error {
job, ok := r.jobs[batchID]
if !ok {
return ErrBatchImageJobNotFound
}
job.ProviderOutputRef = &providerOutputRef
return nil
}
func (r *fakeBatchImageRepository) UpdateBatchImageJobProviderSubmit(_ context.Context, params UpdateBatchImageJobProviderSubmitParams) error {
job, ok := r.jobs[params.BatchID]
if !ok {
return ErrBatchImageJobNotFound
}
if !CanTransitionBatchImageJob(job.Status, BatchImageJobStatusSubmitted) {
return ErrBatchImageInvalidTransition
}
job.Status = BatchImageJobStatusSubmitted
job.ProviderJobName = batchImageOptionalStringPtr(params.ProviderJobName)
job.ProviderInputRef = batchImageOptionalStringPtr(params.ProviderInputRef)
job.ProviderOutputRef = batchImageOptionalStringPtr(params.ProviderOutputRef)
job.GCSInputURI = batchImageOptionalStringPtr(params.GCSInputURI)
job.GCSOutputURI = batchImageOptionalStringPtr(params.GCSOutputURI)
now := time.Now()
job.SubmittedAt = &now
r.transitions[params.BatchID] = append(r.transitions[params.BatchID], BatchImageJobStatusSubmitted)
r.events[params.BatchID] = append(r.events[params.BatchID], "provider_submitted")
return nil
}
func (r *fakeBatchImageRepository) RecordBatchImageJobSubmitFailure(_ context.Context, batchID, code, message string, markFailed bool) error {
job, ok := r.jobs[batchID]
if !ok {
return ErrBatchImageJobNotFound
}
if markFailed {
job.Status = BatchImageJobStatusFailed
}
job.LastErrorCode = batchImageOptionalStringPtr(code)
job.LastErrorMessage = batchImageOptionalStringPtr(message)
eventType := "submit_failed"
if !markFailed {
eventType = "queue_failed"
}
r.events[batchID] = append(r.events[batchID], eventType)
return nil
}
func (r *fakeBatchImageRepository) MarkBatchImageJobSettled(_ context.Context, params MarkBatchImageJobSettledParams) error {
job, ok := r.jobs[params.BatchID]
if !ok {
return ErrBatchImageJobNotFound
}
if job.Status != BatchImageJobStatusSettling {
if job.Status == BatchImageJobStatusCompleted {
return ErrBatchImageAlreadySettled
}
return ErrBatchImageSettlementInvalidStatus
}
if batchImageDerefString(job.ManifestHash) != "" && batchImageDerefString(job.ManifestHash) != params.ManifestHash {
return ErrBatchImageSettlementManifestConflict
}
now := time.Now()
job.Status = BatchImageJobStatusCompleted
job.ActualCost = &params.ActualCost
job.ManifestHash = &params.ManifestHash
job.SettledAt = &now
if job.OutputExpiresAt == nil && params.OutputExpiresAt != nil {
job.OutputExpiresAt = params.OutputExpiresAt
}
r.transitions[params.BatchID] = append(r.transitions[params.BatchID], BatchImageJobStatusCompleted)
r.events[params.BatchID] = append(r.events[params.BatchID], "settlement_completed")
return nil
}
func (r *fakeBatchImageRepository) SetBatchImageJobSettlementFailed(_ context.Context, batchID, code, message string) (int, error) {
job, ok := r.jobs[batchID]
if !ok {
return 0, ErrBatchImageJobNotFound
}
job.LastErrorCode = batchImageStringPtr(code)
job.LastErrorMessage = batchImageOptionalStringPtr(message)
job.RetryCount++
r.events[batchID] = append(r.events[batchID], "settlement_failed")
return job.RetryCount, nil
}
func (r *fakeBatchImageRepository) CreateBatchImageItem(_ context.Context, params CreateBatchImageItemParams) (*BatchImageItem, error) {
r.items[params.JobID] = append(r.items[params.JobID], params)
return &BatchImageItem{JobID: params.JobID, CustomID: params.CustomID, Status: params.Status}, nil
}
func (r *fakeBatchImageRepository) BulkCreateBatchImageItems(ctx context.Context, params []CreateBatchImageItemParams) error {
for _, param := range params {
if _, err := r.CreateBatchImageItem(ctx, param); err != nil {
return err
}
}
return nil
}
func (r *fakeBatchImageRepository) ReplaceBatchImageItemsForJob(_ context.Context, batchID string, items []CreateBatchImageItemParams, counts BatchImageCounts) error {
r.replaceCalls++
copied := append([]CreateBatchImageItemParams(nil), items...)
for idx := range copied {
copied[idx].JobID = batchID
}
r.items[batchID] = copied
r.counts[batchID] = counts
if job, ok := r.jobs[batchID]; ok {
job.SuccessCount = counts.SuccessCount
job.FailCount = counts.FailCount
job.ItemCount = len(copied)
}
return nil
}
func (r *fakeBatchImageRepository) ListBatchImageItems(_ context.Context, batchID string, filter BatchImageItemFilter) ([]*BatchImageItem, error) {
limit := filter.Limit
if limit <= 0 || limit > 500 {
limit = 100
}
offset := filter.Offset
if offset < 0 {
offset = 0
}
var result []*BatchImageItem
for _, item := range r.items[batchID] {
if filter.Status != "" && item.Status != filter.Status {
continue
}
if offset > 0 {
offset--
continue
}
result = append(result, &BatchImageItem{
JobID: item.JobID,
CustomID: item.CustomID,
Status: item.Status,
RequestHash: item.RequestHash,
PromptPreview: item.PromptPreview,
ProviderSourceObject: item.ProviderSourceObject,
SourceLineNumber: item.SourceLineNumber,
SourceByteOffset: item.SourceByteOffset,
SourceByteLength: item.SourceByteLength,
MimeType: item.MimeType,
FileExtension: item.FileExtension,
ImageCount: item.ImageCount,
ErrorCode: item.ErrorCode,
ErrorMessage: item.ErrorMessage,
BilledAmount: item.BilledAmount,
IndexedAt: item.IndexedAt,
})
if len(result) >= limit {
break
}
}
return result, nil
}
func (r *fakeBatchImageRepository) ListBatchImageItemsForOwner(ctx context.Context, userID, apiKeyID int64, batchID string, filter BatchImageItemFilter) ([]*BatchImageItem, error) {
if _, err := r.GetBatchImageJobByBatchIDForOwner(ctx, userID, apiKeyID, batchID); err != nil {
return nil, err
}
return r.ListBatchImageItems(ctx, batchID, filter)
}
func (r *fakeBatchImageRepository) GetBatchImageJobForDownload(ctx context.Context, userID, apiKeyID int64, batchID string) (*BatchImageJob, error) {
return r.GetBatchImageJobByBatchIDForOwner(ctx, userID, apiKeyID, batchID)
}
func (r *fakeBatchImageRepository) GetBatchImageItemForDownload(_ context.Context, batchID, customID string) (*BatchImageItem, error) {
for _, item := range r.items[batchID] {
if item.CustomID != customID {
continue
}
return &BatchImageItem{
JobID: item.JobID,
CustomID: item.CustomID,
Status: item.Status,
RequestHash: item.RequestHash,
PromptPreview: item.PromptPreview,
ProviderSourceObject: item.ProviderSourceObject,
SourceLineNumber: item.SourceLineNumber,
SourceByteOffset: item.SourceByteOffset,
SourceByteLength: item.SourceByteLength,
MimeType: item.MimeType,
FileExtension: item.FileExtension,
ImageCount: item.ImageCount,
ErrorCode: item.ErrorCode,
ErrorMessage: item.ErrorMessage,
BilledAmount: item.BilledAmount,
IndexedAt: item.IndexedAt,
}, nil
}
return nil, ErrBatchImageItemNotFound
}
func (r *fakeBatchImageRepository) ListBatchImageItemsForDownload(ctx context.Context, batchID string, status string, limit int) ([]*BatchImageItem, error) {
return r.ListBatchImageItems(ctx, batchID, BatchImageItemFilter{Status: status, Limit: limit})
}
func (r *fakeBatchImageRepository) ListBatchImageJobsDueForInputCleanup(_ context.Context, cutoff time.Time, limit int) ([]*BatchImageJob, error) {
if limit <= 0 {
limit = 100
}
var jobs []*BatchImageJob
for _, job := range r.jobs {
if job.InputDeletedAt != nil || batchImageDerefString(job.ProviderInputRef) == "" || !IsTerminalBatchImageJobStatus(job.Status) {
continue
}
at := job.FinishedAt
if at == nil {
at = job.SettledAt
}
if at == nil {
at = &job.UpdatedAt
}
if at != nil && at.After(cutoff) {
continue
}
jobs = append(jobs, job)
if len(jobs) >= limit {
break
}
}
return jobs, nil
}
func (r *fakeBatchImageRepository) ListBatchImageJobsDueForOutputCleanup(_ context.Context, now time.Time, limit int) ([]*BatchImageJob, error) {
if limit <= 0 {
limit = 100
}
var jobs []*BatchImageJob
for _, job := range r.jobs {
if job.OutputDeletedAt != nil || batchImageDerefString(job.ProviderOutputRef) == "" || job.Status != BatchImageJobStatusCompleted || job.OutputExpiresAt == nil || job.OutputExpiresAt.After(now) {
continue
}
jobs = append(jobs, job)
if len(jobs) >= limit {
break
}
}
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 {
return ErrBatchImageJobNotFound
}
if job.InputDeletedAt == nil {
job.InputDeletedAt = &deletedAt
}
r.events[batchID] = append(r.events[batchID], "input_cleanup_completed")
return nil
}
func (r *fakeBatchImageRepository) MarkBatchImageOutputDeleted(_ context.Context, batchID string, deletedAt time.Time) error {
job, ok := r.jobs[batchID]
if !ok {
return ErrBatchImageJobNotFound
}
if job.OutputDeletedAt == nil {
job.OutputDeletedAt = &deletedAt
}
if job.Status == BatchImageJobStatusCompleted {
job.Status = BatchImageJobStatusOutputDeleted
}
r.events[batchID] = append(r.events[batchID], "output_cleanup_completed")
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 {
return ErrBatchImageJobNotFound
}
if job.OutputExpiresAt == nil {
job.OutputExpiresAt = &expiresAt
}
return nil
}
func (r *fakeBatchImageRepository) RecordBatchImageCleanupFailure(_ context.Context, batchID, code, message string) error {
job, ok := r.jobs[batchID]
if !ok {
return ErrBatchImageJobNotFound
}
job.LastErrorCode = batchImageStringPtr(code)
job.LastErrorMessage = batchImageOptionalStringPtr(message)
r.events[batchID] = append(r.events[batchID], "output_cleanup_failed")
return nil
}
func (r *fakeBatchImageRepository) AppendBatchImageEvent(_ context.Context, batchID, eventType string, _ any) error {
r.events[batchID] = append(r.events[batchID], eventType)
return nil
}
var _ BatchImageRepository = (*fakeBatchImageRepository)(nil)
var _ BatchImageProvider = (*fakeProcessorProvider)(nil)
var _ BatchImageAccountResolver = (*fakeBatchImageAccountResolver)(nil)
var _ = infraerrors.Reason
@@ -0,0 +1,180 @@
package service
import (
"context"
"fmt"
"io"
"net/http"
"strings"
"time"
"github.com/Wei-Shaw/sub2api/internal/config"
infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors"
)
type BatchImageProvider interface {
Name() string
SupportsAccount(account *Account) bool
Submit(ctx context.Context, job *BatchImageJob, account *Account, input BatchImageInput) (*BatchProviderJob, error)
Get(ctx context.Context, job *BatchImageJob, account *Account) (*BatchProviderStatus, error)
Cancel(ctx context.Context, job *BatchImageJob, account *Account) error
OpenResult(ctx context.Context, job *BatchImageJob, account *Account) (io.ReadCloser, string, error)
Cleanup(ctx context.Context, job *BatchImageJob, account *Account, target CleanupTarget) error
}
type BatchImageProviderRegistry struct {
providers map[string]BatchImageProvider
}
func NewBatchImageProviderRegistry(providers ...BatchImageProvider) *BatchImageProviderRegistry {
r := &BatchImageProviderRegistry{providers: make(map[string]BatchImageProvider, len(providers))}
for _, provider := range providers {
if provider == nil || strings.TrimSpace(provider.Name()) == "" {
continue
}
r.providers[provider.Name()] = provider
}
return r
}
func NewDefaultBatchImageProviderRegistry() *BatchImageProviderRegistry {
return NewBatchImageProviderRegistry(
NewGeminiAPIBatchImageProvider(nil),
NewVertexBatchImageProvider(VertexBatchImageProviderOptions{}, nil, nil, nil),
)
}
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
}
p, ok := r.providers[provider]
return p, ok
}
func (r *BatchImageProviderRegistry) MustGet(provider string) (BatchImageProvider, error) {
p, ok := r.Get(provider)
if !ok {
return nil, ErrBatchImageInvalidProvider
}
return p, nil
}
type BatchImageInput struct {
BatchID string
Model string
DisplayName string
Items []BatchImageInputItem
ResponseMimeType string
AspectRatio string
ImageSize string
Metadata map[string]string
}
type BatchImageInputItem struct {
CustomID string
Prompt string
ReferenceImages []BatchImageReference
}
type BatchImageReference struct {
ID string
Type string
MimeType string
Data []byte
FileURI string
}
type BatchProviderJob struct {
ProviderJobName string
ProviderInputRef string
ProviderOutputRef string
RawState string
}
type BatchProviderInternalState string
const (
BatchProviderStateQueued BatchProviderInternalState = "queued"
BatchProviderStateRunning BatchProviderInternalState = "running"
BatchProviderStateSucceeded BatchProviderInternalState = "succeeded"
BatchProviderStateFailed BatchProviderInternalState = "failed"
BatchProviderStateCancelled BatchProviderInternalState = "cancelled"
BatchProviderStateExpired BatchProviderInternalState = "expired"
)
type BatchProviderStatus struct {
RawState string
InternalState BatchProviderInternalState
Done bool
ProviderOutputRef string
ErrorCode string
ErrorMessage string
SuggestedRequeueAfter time.Duration
}
type CleanupTarget string
const (
CleanupTargetInput CleanupTarget = "input"
CleanupTargetOutput CleanupTarget = "output"
CleanupTargetAll CleanupTarget = "all"
)
var (
ErrBatchImageProviderUnsupportedAccount = infraerrors.New(http.StatusBadRequest, "BATCH_IMAGE_PROVIDER_UNSUPPORTED_ACCOUNT", "batch image provider does not support this account")
ErrBatchImageProviderMissingAPIKey = infraerrors.New(http.StatusBadRequest, "BATCH_IMAGE_PROVIDER_MISSING_API_KEY", "batch image provider account is missing api key")
ErrBatchImageProviderMissingServiceAccount = infraerrors.New(http.StatusBadRequest, "BATCH_IMAGE_PROVIDER_MISSING_SERVICE_ACCOUNT", "batch image provider account is missing service account credentials")
ErrBatchImageProviderMissingJobName = infraerrors.New(http.StatusBadRequest, "BATCH_IMAGE_PROVIDER_MISSING_JOB_NAME", "batch image provider job name is missing")
ErrBatchImageProviderMissingResultRef = infraerrors.New(http.StatusBadRequest, "BATCH_IMAGE_PROVIDER_MISSING_RESULT_REF", "batch image provider result reference is missing")
ErrBatchImageProviderInlineResultUnsupported = infraerrors.New(http.StatusBadRequest, "GEMINI_INLINE_BATCH_RESULT_UNSUPPORTED", "Gemini inline batch result is not supported")
ErrBatchImageProviderInvalidInput = infraerrors.New(http.StatusBadRequest, "BATCH_IMAGE_PROVIDER_INVALID_INPUT", "invalid batch image provider input")
ErrBatchImageProviderUnsafeCleanupPath = infraerrors.New(http.StatusBadRequest, "VERTEX_UNSAFE_CLEANUP_PATH", "unsafe batch image cleanup path")
ErrUnsupportedCleanupTarget = infraerrors.New(http.StatusBadRequest, "BATCH_IMAGE_PROVIDER_UNSUPPORTED_CLEANUP_TARGET", "unsupported batch image cleanup target")
)
func batchImageProviderJobName(job *BatchImageJob) string {
if job == nil || job.ProviderJobName == nil {
return ""
}
return strings.TrimSpace(*job.ProviderJobName)
}
func batchImageProviderInputRef(job *BatchImageJob) string {
if job == nil || job.ProviderInputRef == nil {
return ""
}
return strings.TrimSpace(*job.ProviderInputRef)
}
func batchImageProviderOutputRef(job *BatchImageJob) string {
if job == nil || job.ProviderOutputRef == nil {
return ""
}
return strings.TrimSpace(*job.ProviderOutputRef)
}
func batchImageProviderAPIKey(account *Account) string {
if account == nil {
return ""
}
return strings.TrimSpace(account.GetCredential("api_key"))
}
func batchImageProviderInputError(format string, args ...any) error {
return ErrBatchImageProviderInvalidInput.WithCause(fmt.Errorf(format, args...))
}
@@ -0,0 +1,680 @@
package service
import (
"bytes"
"context"
"encoding/base64"
"encoding/json"
"errors"
"fmt"
"io"
"mime/multipart"
"net/http"
"net/textproto"
"net/url"
"strings"
"time"
infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors"
"github.com/Wei-Shaw/sub2api/internal/pkg/geminicli"
)
const defaultGeminiBatchRequeueAfter = 30 * time.Second
type GeminiBatchClient interface {
UploadJSONL(ctx context.Context, apiKey string, displayName string, r io.Reader) (*GeminiUploadedFile, error)
CreateBatch(ctx context.Context, apiKey string, model string, fileName string, displayName string) (*GeminiBatchJob, error)
GetBatch(ctx context.Context, apiKey string, batchName string) (*GeminiBatchJob, error)
CancelBatch(ctx context.Context, apiKey string, batchName string) error
DownloadFile(ctx context.Context, apiKey string, fileName string) (io.ReadCloser, string, error)
DeleteFile(ctx context.Context, apiKey string, fileName string) error
}
type GeminiUploadedFile struct {
Name string `json:"name"`
DisplayName string `json:"displayName"`
URI string `json:"uri"`
MimeType string `json:"mimeType"`
}
type GeminiBatchJob struct {
Name string `json:"name"`
State string `json:"state"`
Dest *GeminiBatchDest `json:"dest"`
Response *GeminiBatchResponse `json:"response"`
Error *GeminiBatchError `json:"error"`
Raw map[string]any `json:"-"`
}
type GeminiBatchDest struct {
FileName string `json:"fileName"`
FileNameSnake string `json:"file_name"`
}
type GeminiBatchResponse struct {
ResponsesFile string `json:"responsesFile"`
ResponsesFileSnake string `json:"responses_file"`
InlinedResponses []any `json:"inlinedResponses"`
InlinedResponsesAlt []any `json:"inlined_responses"`
}
type GeminiBatchError struct {
Code string `json:"code"`
Message string `json:"message"`
Status string `json:"status"`
}
type GeminiAPIBatchImageProvider struct {
client GeminiBatchClient
}
func NewGeminiAPIBatchImageProvider(client GeminiBatchClient) *GeminiAPIBatchImageProvider {
if client == nil {
client = NewGeminiBatchHTTPClient("", nil)
}
return &GeminiAPIBatchImageProvider{client: client}
}
func (p *GeminiAPIBatchImageProvider) Name() string {
return BatchImageProviderGeminiAPI
}
func (p *GeminiAPIBatchImageProvider) SupportsAccount(account *Account) bool {
return account != nil &&
account.Platform == PlatformGemini &&
account.Type == AccountTypeAPIKey &&
batchImageProviderAPIKey(account) != ""
}
func (p *GeminiAPIBatchImageProvider) Submit(ctx context.Context, job *BatchImageJob, account *Account, input BatchImageInput) (*BatchProviderJob, error) {
if account == nil || account.Platform != PlatformGemini || account.Type != AccountTypeAPIKey {
return nil, ErrBatchImageProviderUnsupportedAccount
}
apiKey := batchImageProviderAPIKey(account)
if apiKey == "" {
return nil, ErrBatchImageProviderMissingAPIKey
}
if input.BatchID == "" && job != nil {
input.BatchID = job.BatchID
}
if input.Model == "" && job != nil {
input.Model = job.Model
}
jsonl, err := BuildGeminiBatchJSONL(input)
if err != nil {
return nil, err
}
displayName := strings.TrimSpace(input.DisplayName)
if displayName == "" {
displayName = strings.TrimSpace(input.BatchID)
}
uploaded, err := p.client.UploadJSONL(ctx, apiKey, displayName, bytes.NewReader(jsonl))
if err != nil {
return nil, mapGeminiClientError(err)
}
if uploaded == nil || strings.TrimSpace(uploaded.Name) == "" {
return nil, geminiProviderError("GEMINI_INVALID_RESPONSE", "Gemini upload response is missing file name", nil)
}
batch, err := p.client.CreateBatch(ctx, apiKey, input.Model, uploaded.Name, displayName)
if err != nil {
return nil, mapGeminiClientError(err)
}
if batch == nil || strings.TrimSpace(batch.Name) == "" {
return nil, geminiProviderError("GEMINI_INVALID_RESPONSE", "Gemini batch response is missing job name", nil)
}
return &BatchProviderJob{
ProviderJobName: batch.Name,
ProviderInputRef: uploaded.Name,
RawState: batch.State,
}, nil
}
func (p *GeminiAPIBatchImageProvider) Get(ctx context.Context, job *BatchImageJob, account *Account) (*BatchProviderStatus, error) {
if account == nil || account.Platform != PlatformGemini || account.Type != AccountTypeAPIKey {
return nil, ErrBatchImageProviderUnsupportedAccount
}
apiKey := batchImageProviderAPIKey(account)
if apiKey == "" {
return nil, ErrBatchImageProviderMissingAPIKey
}
jobName := batchImageProviderJobName(job)
if jobName == "" {
return nil, ErrBatchImageProviderMissingJobName
}
batch, err := p.client.GetBatch(ctx, apiKey, jobName)
if err != nil {
return nil, mapGeminiClientError(err)
}
if batch == nil {
return nil, geminiProviderError("GEMINI_INVALID_RESPONSE", "Gemini batch response is empty", nil)
}
status := mapGeminiBatchState(batch)
if status.InternalState == BatchProviderStateSucceeded {
if geminiBatchHasInlineResults(batch) {
return nil, ErrBatchImageProviderInlineResultUnsupported
}
outputRef := geminiBatchOutputRef(batch)
if outputRef == "" {
status.InternalState = BatchProviderStateFailed
status.Done = true
status.ErrorCode = "GEMINI_RESULT_FILE_MISSING"
status.ErrorMessage = "Gemini batch succeeded without a result file reference"
}
status.ProviderOutputRef = outputRef
}
return status, nil
}
func (p *GeminiAPIBatchImageProvider) Cancel(ctx context.Context, job *BatchImageJob, account *Account) error {
if account == nil || account.Platform != PlatformGemini || account.Type != AccountTypeAPIKey {
return ErrBatchImageProviderUnsupportedAccount
}
apiKey := batchImageProviderAPIKey(account)
if apiKey == "" {
return ErrBatchImageProviderMissingAPIKey
}
jobName := batchImageProviderJobName(job)
if jobName == "" {
return ErrBatchImageProviderMissingJobName
}
return mapGeminiClientError(p.client.CancelBatch(ctx, apiKey, jobName))
}
func (p *GeminiAPIBatchImageProvider) OpenResult(ctx context.Context, job *BatchImageJob, account *Account) (io.ReadCloser, string, error) {
if account == nil || account.Platform != PlatformGemini || account.Type != AccountTypeAPIKey {
return nil, "", ErrBatchImageProviderUnsupportedAccount
}
apiKey := batchImageProviderAPIKey(account)
if apiKey == "" {
return nil, "", ErrBatchImageProviderMissingAPIKey
}
outputRef := batchImageProviderOutputRef(job)
if outputRef == "" {
return nil, "", ErrBatchImageProviderMissingResultRef
}
r, contentType, err := p.client.DownloadFile(ctx, apiKey, outputRef)
return r, contentType, mapGeminiClientError(err)
}
func (p *GeminiAPIBatchImageProvider) Cleanup(ctx context.Context, job *BatchImageJob, account *Account, target CleanupTarget) error {
if account == nil || account.Platform != PlatformGemini || account.Type != AccountTypeAPIKey {
return ErrBatchImageProviderUnsupportedAccount
}
apiKey := batchImageProviderAPIKey(account)
if apiKey == "" {
return ErrBatchImageProviderMissingAPIKey
}
switch target {
case CleanupTargetInput:
return p.deleteGeminiFileIfPresent(ctx, apiKey, batchImageProviderInputRef(job))
case CleanupTargetOutput:
return p.deleteGeminiFileIfPresent(ctx, apiKey, batchImageProviderOutputRef(job))
case CleanupTargetAll:
if err := p.deleteGeminiFileIfPresent(ctx, apiKey, batchImageProviderInputRef(job)); err != nil {
return err
}
return p.deleteGeminiFileIfPresent(ctx, apiKey, batchImageProviderOutputRef(job))
default:
return ErrUnsupportedCleanupTarget
}
}
func (p *GeminiAPIBatchImageProvider) deleteGeminiFileIfPresent(ctx context.Context, apiKey, fileName string) error {
if strings.TrimSpace(fileName) == "" {
return nil
}
return mapGeminiClientError(p.client.DeleteFile(ctx, apiKey, fileName))
}
type geminiJSONLLine struct {
Key string `json:"key"`
Request geminiGenerateRequest `json:"request"`
}
type geminiGenerateRequest struct {
Contents []geminiContent `json:"contents"`
GenerationConfig geminiGenerationConfig `json:"generationConfig"`
}
type geminiContent struct {
Parts []geminiPart `json:"parts"`
}
type geminiPart struct {
Text string `json:"text,omitempty"`
InlineData *geminiInlineData `json:"inlineData,omitempty"`
FileData *geminiFileData `json:"fileData,omitempty"`
}
type geminiInlineData struct {
MimeType string `json:"mimeType"`
Data string `json:"data"`
}
type geminiFileData struct {
MimeType string `json:"mimeType"`
FileURI string `json:"fileUri"`
}
type geminiGenerationConfig struct {
ResponseModalities []string `json:"responseModalities"`
}
func BuildGeminiBatchJSONL(input BatchImageInput) ([]byte, error) {
if strings.TrimSpace(input.Model) == "" {
return nil, batchImageProviderInputError("model is required")
}
if len(input.Items) == 0 {
return nil, batchImageProviderInputError("at least one item is required")
}
seen := make(map[string]struct{}, len(input.Items))
var buf bytes.Buffer
enc := json.NewEncoder(&buf)
for _, item := range input.Items {
customID := strings.TrimSpace(item.CustomID)
if customID == "" {
return nil, batchImageProviderInputError("custom_id is required")
}
if _, ok := seen[customID]; ok {
return nil, batchImageProviderInputError("duplicate custom_id %q", customID)
}
seen[customID] = struct{}{}
prompt := strings.TrimSpace(item.Prompt)
if prompt == "" {
return nil, batchImageProviderInputError("prompt is required for custom_id %q", customID)
}
parts, err := batchImageGeminiParts(prompt, item.ReferenceImages)
if err != nil {
return nil, err
}
// TODO(batch-image): add response_mime_type/aspect_ratio/image_size once the
// Gemini batch image REST shape is stabilized for those options.
line := geminiJSONLLine{
Key: customID,
Request: geminiGenerateRequest{
Contents: []geminiContent{{
Parts: parts,
}},
GenerationConfig: geminiGenerationConfig{
ResponseModalities: []string{"TEXT", "IMAGE"},
},
},
}
if err := enc.Encode(line); err != nil {
return nil, err
}
}
return buf.Bytes(), nil
}
func batchImageGeminiParts(prompt string, refs []BatchImageReference) ([]geminiPart, error) {
parts := []geminiPart{{Text: prompt}}
for _, ref := range refs {
mimeType := normalizeBatchImageReferenceMimeType(ref.MimeType)
if mimeType == "" {
return nil, batchImageProviderInputError("reference image mime_type is required")
}
fileURI := strings.TrimSpace(ref.FileURI)
switch {
case len(ref.Data) > 0 && fileURI == "":
parts = append(parts, geminiPart{InlineData: &geminiInlineData{
MimeType: mimeType,
Data: base64.StdEncoding.EncodeToString(ref.Data),
}})
case len(ref.Data) == 0 && fileURI != "":
parts = append(parts, geminiPart{FileData: &geminiFileData{
MimeType: mimeType,
FileURI: fileURI,
}})
default:
return nil, batchImageProviderInputError("reference image must contain exactly one of data or file_uri")
}
}
return parts, nil
}
func mapGeminiBatchState(batch *GeminiBatchJob) *BatchProviderStatus {
state := strings.TrimSpace(batch.State)
normalized := strings.ToUpper(state)
status := &BatchProviderStatus{
RawState: state,
InternalState: BatchProviderStateRunning,
SuggestedRequeueAfter: defaultGeminiBatchRequeueAfter,
}
switch normalized {
case "JOB_STATE_PENDING", "JOB_STATE_QUEUED":
status.InternalState = BatchProviderStateQueued
case "JOB_STATE_RUNNING":
status.InternalState = BatchProviderStateRunning
case "JOB_STATE_SUCCEEDED":
status.InternalState = BatchProviderStateSucceeded
status.Done = true
case "JOB_STATE_FAILED":
status.InternalState = BatchProviderStateFailed
status.Done = true
status.ErrorCode = "GEMINI_BATCH_FAILED"
case "JOB_STATE_CANCELLED":
status.InternalState = BatchProviderStateCancelled
status.Done = true
status.ErrorCode = "GEMINI_BATCH_CANCELLED"
case "JOB_STATE_EXPIRED":
status.InternalState = BatchProviderStateExpired
status.Done = true
status.ErrorCode = "GEMINI_BATCH_EXPIRED"
default:
if batch.Error != nil && (strings.TrimSpace(batch.Error.Message) != "" || strings.TrimSpace(batch.Error.Code) != "") {
status.InternalState = BatchProviderStateFailed
status.Done = true
status.ErrorCode = "GEMINI_BATCH_FAILED"
}
}
if batch.Error != nil {
if code := strings.TrimSpace(batch.Error.Code); code != "" {
status.ErrorCode = code
} else if status.ErrorCode == "" && strings.TrimSpace(batch.Error.Status) != "" {
status.ErrorCode = strings.TrimSpace(batch.Error.Status)
}
status.ErrorMessage = strings.TrimSpace(batch.Error.Message)
}
return status
}
func geminiBatchOutputRef(batch *GeminiBatchJob) string {
if batch == nil {
return ""
}
if batch.Dest != nil {
if v := strings.TrimSpace(batch.Dest.FileName); v != "" {
return v
}
if v := strings.TrimSpace(batch.Dest.FileNameSnake); v != "" {
return v
}
}
if batch.Response != nil {
if v := strings.TrimSpace(batch.Response.ResponsesFile); v != "" {
return v
}
if v := strings.TrimSpace(batch.Response.ResponsesFileSnake); v != "" {
return v
}
}
return ""
}
func geminiBatchHasInlineResults(batch *GeminiBatchJob) bool {
return batch != nil &&
batch.Response != nil &&
(len(batch.Response.InlinedResponses) > 0 || len(batch.Response.InlinedResponsesAlt) > 0)
}
func geminiProviderError(reason, message string, cause error) error {
err := infraerrors.New(http.StatusBadGateway, reason, message)
if cause != nil {
return err.WithCause(cause)
}
return err
}
func mapGeminiClientError(err error) error {
if err == nil {
return nil
}
var apiErr *GeminiAPIError
if errors.As(err, &apiErr) {
switch apiErr.StatusCode {
case http.StatusUnauthorized, http.StatusForbidden:
return geminiProviderError("GEMINI_AUTH_FAILED", "Gemini authentication failed", nil)
case http.StatusTooManyRequests:
return geminiProviderError("GEMINI_RATE_LIMITED", "Gemini rate limit exceeded", nil)
case http.StatusNotFound:
return geminiProviderError("GEMINI_BATCH_NOT_FOUND", "Gemini batch resource was not found", nil)
default:
return geminiProviderError("GEMINI_INVALID_RESPONSE", "Gemini API request failed", nil)
}
}
return geminiProviderError("GEMINI_INVALID_RESPONSE", "Gemini API request failed", nil)
}
type GeminiBatchHTTPClient struct {
baseURL string
client *http.Client
}
func NewGeminiBatchHTTPClient(baseURL string, client *http.Client) *GeminiBatchHTTPClient {
baseURL = strings.TrimRight(strings.TrimSpace(baseURL), "/")
if baseURL == "" {
baseURL = geminicli.AIStudioBaseURL
}
if client == nil {
client = http.DefaultClient
}
return &GeminiBatchHTTPClient{baseURL: baseURL, client: client}
}
func (c *GeminiBatchHTTPClient) UploadJSONL(ctx context.Context, apiKey string, displayName string, r io.Reader) (*GeminiUploadedFile, error) {
var body bytes.Buffer
writer := multipart.NewWriter(&body)
metadataHeader := textproto.MIMEHeader{}
metadataHeader.Set("Content-Disposition", `form-data; name="metadata"`)
metadataHeader.Set("Content-Type", "application/json; charset=utf-8")
metadataPart, err := writer.CreatePart(metadataHeader)
if err != nil {
return nil, err
}
metadata := map[string]any{"file": map[string]any{"displayName": displayName, "mimeType": "application/jsonl"}}
if err := json.NewEncoder(metadataPart).Encode(metadata); err != nil {
return nil, err
}
fileHeader := textproto.MIMEHeader{}
fileHeader.Set("Content-Disposition", `form-data; name="file"; filename="batch.jsonl"`)
fileHeader.Set("Content-Type", "application/jsonl")
filePart, err := writer.CreatePart(fileHeader)
if err != nil {
return nil, err
}
if _, err := io.Copy(filePart, r); err != nil {
return nil, err
}
if err := writer.Close(); err != nil {
return nil, err
}
req, err := c.newRequest(ctx, http.MethodPost, "/upload/v1beta/files?uploadType=multipart", apiKey, &body)
if err != nil {
return nil, err
}
req.Header.Set("Content-Type", writer.FormDataContentType())
var resp struct {
File *GeminiUploadedFile `json:"file"`
*GeminiUploadedFile
}
if err := c.doJSON(req, &resp); err != nil {
return nil, err
}
if resp.File != nil {
return resp.File, nil
}
return resp.GeminiUploadedFile, nil
}
func (c *GeminiBatchHTTPClient) CreateBatch(ctx context.Context, apiKey string, model string, fileName string, displayName string) (*GeminiBatchJob, error) {
body := map[string]any{
"batch": map[string]any{
"displayName": displayName,
"inputConfig": map[string]any{
"fileName": fileName,
},
},
}
payload, _ := json.Marshal(body)
path := fmt.Sprintf("/v1beta/models/%s:batchGenerateContent", url.PathEscape(strings.TrimSpace(model)))
req, err := c.newRequest(ctx, http.MethodPost, path, apiKey, bytes.NewReader(payload))
if err != nil {
return nil, err
}
req.Header.Set("Content-Type", "application/json")
return c.doBatchJob(req)
}
func (c *GeminiBatchHTTPClient) GetBatch(ctx context.Context, apiKey string, batchName string) (*GeminiBatchJob, error) {
req, err := c.newRequest(ctx, http.MethodGet, "/v1beta/"+strings.TrimLeft(batchName, "/"), apiKey, nil)
if err != nil {
return nil, err
}
return c.doBatchJob(req)
}
func (c *GeminiBatchHTTPClient) CancelBatch(ctx context.Context, apiKey string, batchName string) error {
req, err := c.newRequest(ctx, http.MethodPost, "/v1beta/"+strings.TrimLeft(batchName, "/")+":cancel", apiKey, nil)
if err != nil {
return err
}
return c.doNoBody(req)
}
func (c *GeminiBatchHTTPClient) DownloadFile(ctx context.Context, apiKey string, fileName string) (io.ReadCloser, string, error) {
metaReq, err := c.newRequest(ctx, http.MethodGet, "/v1beta/"+strings.TrimLeft(fileName, "/"), apiKey, nil)
if err != nil {
return nil, "", err
}
var metadata struct {
DownloadURI string `json:"downloadUri"`
DownloadURL string `json:"download_url"`
MimeType string `json:"mimeType"`
}
if err := c.doJSON(metaReq, &metadata); err != nil {
return nil, "", err
}
downloadURL := strings.TrimSpace(metadata.DownloadURI)
if downloadURL == "" {
downloadURL = strings.TrimSpace(metadata.DownloadURL)
}
if downloadURL == "" {
downloadURL = c.baseURL + "/v1beta/" + strings.TrimLeft(fileName, "/") + ":download"
}
req, err := http.NewRequestWithContext(ctx, http.MethodGet, downloadURL, nil)
if err != nil {
return nil, "", err
}
req.Header.Set("x-goog-api-key", apiKey)
resp, err := c.client.Do(req)
if err != nil {
return nil, "", err
}
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
defer func() { _ = resp.Body.Close() }()
return nil, "", readGeminiAPIError(resp)
}
contentType := resp.Header.Get("Content-Type")
if contentType == "" {
contentType = metadata.MimeType
}
if contentType == "" {
contentType = "application/octet-stream"
}
return resp.Body, contentType, nil
}
func (c *GeminiBatchHTTPClient) DeleteFile(ctx context.Context, apiKey string, fileName string) error {
req, err := c.newRequest(ctx, http.MethodDelete, "/v1beta/"+strings.TrimLeft(fileName, "/"), apiKey, nil)
if err != nil {
return err
}
return c.doNoBody(req)
}
func (c *GeminiBatchHTTPClient) doBatchJob(req *http.Request) (*GeminiBatchJob, error) {
var job GeminiBatchJob
if err := c.doJSON(req, &job); err != nil {
return nil, err
}
job.Raw = map[string]any{}
return &job, nil
}
func (c *GeminiBatchHTTPClient) doNoBody(req *http.Request) error {
resp, err := c.client.Do(req)
if err != nil {
return err
}
defer func() { _ = resp.Body.Close() }()
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
return readGeminiAPIError(resp)
}
return nil
}
func (c *GeminiBatchHTTPClient) doJSON(req *http.Request, out any) error {
resp, err := c.client.Do(req)
if err != nil {
return err
}
defer func() { _ = resp.Body.Close() }()
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
return readGeminiAPIError(resp)
}
return json.NewDecoder(resp.Body).Decode(out)
}
func (c *GeminiBatchHTTPClient) newRequest(ctx context.Context, method, path, apiKey string, body io.Reader) (*http.Request, error) {
if strings.TrimSpace(apiKey) == "" {
return nil, ErrBatchImageProviderMissingAPIKey
}
req, err := http.NewRequestWithContext(ctx, method, c.baseURL+path, body)
if err != nil {
return nil, err
}
req.Header.Set("x-goog-api-key", apiKey)
return req, nil
}
type GeminiAPIError struct {
StatusCode int
Code string
Message string
}
func (e *GeminiAPIError) Error() string {
if e == nil {
return "<nil>"
}
if e.Code != "" {
return fmt.Sprintf("gemini api error: status=%d code=%s message=%s", e.StatusCode, e.Code, e.Message)
}
return fmt.Sprintf("gemini api error: status=%d message=%s", e.StatusCode, e.Message)
}
func readGeminiAPIError(resp *http.Response) error {
body, _ := io.ReadAll(io.LimitReader(resp.Body, 8192))
message := string(body)
var parsed struct {
Error struct {
Code any `json:"code"`
Message string `json:"message"`
Status string `json:"status"`
} `json:"error"`
}
if err := json.Unmarshal(body, &parsed); err == nil && parsed.Error.Message != "" {
message = parsed.Error.Message
return &GeminiAPIError{StatusCode: resp.StatusCode, Code: parsed.Error.Status, Message: message}
}
return &GeminiAPIError{StatusCode: resp.StatusCode, Message: message}
}
var _ BatchImageProvider = (*GeminiAPIBatchImageProvider)(nil)
var _ GeminiBatchClient = (*GeminiBatchHTTPClient)(nil)
@@ -0,0 +1,360 @@
//go:build unit
package service
import (
"bytes"
"context"
"encoding/json"
"errors"
"io"
"strings"
"testing"
infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors"
"github.com/stretchr/testify/require"
)
func TestBatchImageProviderRegistry_ReturnsGeminiAPI(t *testing.T) {
registry := NewDefaultBatchImageProviderRegistry()
provider, ok := registry.Get(BatchImageProviderGeminiAPI)
require.True(t, ok)
require.Equal(t, BatchImageProviderGeminiAPI, provider.Name())
must, err := registry.MustGet(BatchImageProviderGeminiAPI)
require.NoError(t, err)
require.Same(t, provider, must)
_, err = registry.MustGet("unknown_provider")
require.ErrorIs(t, err, ErrBatchImageInvalidProvider)
}
func TestGeminiProvider_SupportsOnlyGeminiAPIKeyWithSecret(t *testing.T) {
provider := NewGeminiAPIBatchImageProvider(&fakeGeminiBatchClient{})
require.True(t, provider.SupportsAccount(geminiAPIKeyAccount("sk-gemini")))
require.False(t, provider.SupportsAccount(&Account{Platform: PlatformGemini, Type: AccountTypeAPIKey, Credentials: map[string]any{}}))
require.False(t, provider.SupportsAccount(&Account{Platform: PlatformGemini, Type: AccountTypeOAuth, Credentials: map[string]any{"api_key": "sk"}}))
require.False(t, provider.SupportsAccount(&Account{Platform: PlatformOpenAI, Type: AccountTypeAPIKey, Credentials: map[string]any{"api_key": "sk"}}))
require.False(t, provider.SupportsAccount(nil))
}
func TestGeminiProvider_MissingAPIKeyRejected(t *testing.T) {
provider := NewGeminiAPIBatchImageProvider(&fakeGeminiBatchClient{})
_, err := provider.Submit(context.Background(), nil, &Account{Platform: PlatformGemini, Type: AccountTypeAPIKey}, validGeminiBatchInput())
require.ErrorIs(t, err, ErrBatchImageProviderMissingAPIKey)
}
func TestBuildGeminiBatchJSONL_WritesValidLinesAndPreservesCustomID(t *testing.T) {
input := validGeminiBatchInput()
input.Items = append(input.Items, BatchImageInputItem{CustomID: "cover_002", Prompt: "Second prompt"})
jsonl, err := BuildGeminiBatchJSONL(input)
require.NoError(t, err)
lines := strings.Split(strings.TrimSpace(string(jsonl)), "\n")
require.Len(t, lines, 2)
requireJSONLLine(t, lines[0], "cover_001", "A clean product hero image")
requireJSONLLine(t, lines[1], "cover_002", "Second prompt")
}
func TestBuildGeminiBatchJSONL_RejectsDuplicateCustomIDs(t *testing.T) {
input := validGeminiBatchInput()
input.Items = append(input.Items, BatchImageInputItem{CustomID: "cover_001", Prompt: "Duplicate"})
_, err := BuildGeminiBatchJSONL(input)
require.ErrorIs(t, err, ErrBatchImageProviderInvalidInput)
}
func TestBuildGeminiBatchJSONL_RejectsEmptyPrompt(t *testing.T) {
input := validGeminiBatchInput()
input.Items[0].Prompt = " "
_, err := BuildGeminiBatchJSONL(input)
require.ErrorIs(t, err, ErrBatchImageProviderInvalidInput)
}
func TestBuildGeminiBatchJSONL_WritesReferenceImages(t *testing.T) {
input := validGeminiBatchInput()
input.Items[0].ReferenceImages = []BatchImageReference{
{MimeType: "image/webp", Data: []byte("webp-bytes")},
{MimeType: "image/jpeg", FileURI: "gs://bucket/refs/style.jpg"},
}
jsonl, err := BuildGeminiBatchJSONL(input)
require.NoError(t, err)
lines := strings.Split(strings.TrimSpace(string(jsonl)), "\n")
require.Len(t, lines, 1)
var got map[string]any
require.NoError(t, json.Unmarshal([]byte(lines[0]), &got))
request := got["request"].(map[string]any)
contents := request["contents"].([]any)
parts := contents[0].(map[string]any)["parts"].([]any)
require.Len(t, parts, 3)
require.Equal(t, "A clean product hero image", parts[0].(map[string]any)["text"])
inlineData := parts[1].(map[string]any)["inlineData"].(map[string]any)
require.Equal(t, "image/webp", inlineData["mimeType"])
require.Equal(t, "d2VicC1ieXRlcw==", inlineData["data"])
fileData := parts[2].(map[string]any)["fileData"].(map[string]any)
require.Equal(t, "image/jpeg", fileData["mimeType"])
require.Equal(t, "gs://bucket/refs/style.jpg", fileData["fileUri"])
}
func TestGeminiProvider_SubmitUploadsJSONLThenCreatesBatch(t *testing.T) {
client := &fakeGeminiBatchClient{
uploaded: &GeminiUploadedFile{Name: "files/input-jsonl"},
created: &GeminiBatchJob{Name: "batches/job-123", State: "JOB_STATE_PENDING"},
}
provider := NewGeminiAPIBatchImageProvider(client)
got, err := provider.Submit(context.Background(), &BatchImageJob{BatchID: "imgbatch_123", Model: "gemini-3.1-flash-image"}, geminiAPIKeyAccount("sk-secret"), validGeminiBatchInput())
require.NoError(t, err)
require.Equal(t, []string{"upload", "create"}, client.calls)
require.Equal(t, "files/input-jsonl", got.ProviderInputRef)
require.Equal(t, "batches/job-123", got.ProviderJobName)
require.Empty(t, got.ProviderOutputRef)
require.NotContains(t, got.ProviderInputRef, "A clean product hero image")
require.NotContains(t, string(client.uploadedJSONL), "sk-secret")
}
func TestGeminiProvider_GetMapsStates(t *testing.T) {
tests := []struct {
name string
job *GeminiBatchJob
wantState BatchProviderInternalState
wantDone bool
wantRef string
wantCode string
}{
{name: "running", job: &GeminiBatchJob{Name: "batches/1", State: "JOB_STATE_RUNNING"}, wantState: BatchProviderStateRunning},
{name: "succeeded_dest_fileName", job: &GeminiBatchJob{Name: "batches/1", State: "JOB_STATE_SUCCEEDED", Dest: &GeminiBatchDest{FileName: "files/out"}}, wantState: BatchProviderStateSucceeded, wantDone: true, wantRef: "files/out"},
{name: "failed", job: &GeminiBatchJob{Name: "batches/1", State: "JOB_STATE_FAILED", Error: &GeminiBatchError{Code: "BAD_PROMPT", Message: "bad prompt"}}, wantState: BatchProviderStateFailed, wantDone: true, wantCode: "BAD_PROMPT"},
{name: "cancelled", job: &GeminiBatchJob{Name: "batches/1", State: "JOB_STATE_CANCELLED"}, wantState: BatchProviderStateCancelled, wantDone: true, wantCode: "GEMINI_BATCH_CANCELLED"},
{name: "expired", job: &GeminiBatchJob{Name: "batches/1", State: "JOB_STATE_EXPIRED"}, wantState: BatchProviderStateExpired, wantDone: true, wantCode: "GEMINI_BATCH_EXPIRED"},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
provider := NewGeminiAPIBatchImageProvider(&fakeGeminiBatchClient{got: tt.job})
got, err := provider.Get(context.Background(), jobWithProviderName("batches/1"), geminiAPIKeyAccount("sk-secret"))
require.NoError(t, err)
require.Equal(t, tt.wantState, got.InternalState)
require.Equal(t, tt.wantDone, got.Done)
require.Equal(t, tt.wantRef, got.ProviderOutputRef)
require.Equal(t, tt.wantCode, got.ErrorCode)
require.NotContains(t, got.ErrorMessage, "sk-secret")
})
}
}
func TestGeminiProvider_GetExtractsResponsesFileReference(t *testing.T) {
provider := NewGeminiAPIBatchImageProvider(&fakeGeminiBatchClient{
got: &GeminiBatchJob{
Name: "batches/1",
State: "JOB_STATE_SUCCEEDED",
Response: &GeminiBatchResponse{ResponsesFile: "files/responses-jsonl"},
},
})
got, err := provider.Get(context.Background(), jobWithProviderName("batches/1"), geminiAPIKeyAccount("sk-secret"))
require.NoError(t, err)
require.Equal(t, BatchProviderStateSucceeded, got.InternalState)
require.Equal(t, "files/responses-jsonl", got.ProviderOutputRef)
}
func TestGeminiProvider_GetRejectsInlineResultShape(t *testing.T) {
provider := NewGeminiAPIBatchImageProvider(&fakeGeminiBatchClient{
got: &GeminiBatchJob{
Name: "batches/1",
State: "JOB_STATE_SUCCEEDED",
Response: &GeminiBatchResponse{InlinedResponses: []any{map[string]any{"response": "large"}}},
},
})
_, err := provider.Get(context.Background(), jobWithProviderName("batches/1"), geminiAPIKeyAccount("sk-secret"))
require.ErrorIs(t, err, ErrBatchImageProviderInlineResultUnsupported)
}
func TestGeminiProvider_OpenResultStreamsResultFile(t *testing.T) {
client := &fakeGeminiBatchClient{downloadBody: "line1\n", downloadContentType: "application/jsonl"}
provider := NewGeminiAPIBatchImageProvider(client)
outputRef := "files/output-jsonl"
r, contentType, err := provider.OpenResult(context.Background(), &BatchImageJob{ProviderOutputRef: &outputRef}, geminiAPIKeyAccount("sk-secret"))
require.NoError(t, err)
defer r.Close()
body, err := io.ReadAll(r)
require.NoError(t, err)
require.Equal(t, "line1\n", string(body))
require.Equal(t, "application/jsonl", contentType)
require.Equal(t, "files/output-jsonl", client.downloadedFile)
}
func TestGeminiProvider_CancelCallsClient(t *testing.T) {
client := &fakeGeminiBatchClient{}
provider := NewGeminiAPIBatchImageProvider(client)
require.NoError(t, provider.Cancel(context.Background(), jobWithProviderName("batches/1"), geminiAPIKeyAccount("sk-secret")))
require.Equal(t, "batches/1", client.cancelledBatch)
}
func TestGeminiProvider_CleanupDeletesRefsOnlyWhenPresent(t *testing.T) {
inputRef := "files/input"
outputRef := "files/output"
client := &fakeGeminiBatchClient{}
provider := NewGeminiAPIBatchImageProvider(client)
err := provider.Cleanup(context.Background(), &BatchImageJob{ProviderInputRef: &inputRef, ProviderOutputRef: &outputRef}, geminiAPIKeyAccount("sk-secret"), CleanupTargetAll)
require.NoError(t, err)
require.Equal(t, []string{"files/input", "files/output"}, client.deletedFiles)
err = provider.Cleanup(context.Background(), &BatchImageJob{}, geminiAPIKeyAccount("sk-secret"), CleanupTargetAll)
require.NoError(t, err)
require.Equal(t, []string{"files/input", "files/output"}, client.deletedFiles)
}
func TestGeminiProvider_ErrorsDoNotExposeAPIKey(t *testing.T) {
apiKey := "sk-top-secret"
client := &fakeGeminiBatchClient{uploadErr: &GeminiAPIError{StatusCode: 401, Message: "upstream body should be hidden " + apiKey}}
provider := NewGeminiAPIBatchImageProvider(client)
_, err := provider.Submit(context.Background(), nil, geminiAPIKeyAccount(apiKey), validGeminiBatchInput())
require.Error(t, err)
require.Equal(t, "GEMINI_AUTH_FAILED", infraerrors.Reason(err))
require.NotContains(t, err.Error(), apiKey)
}
func TestGeminiProvider_MetadataDoesNotStoreImageBytesOrBase64(t *testing.T) {
client := &fakeGeminiBatchClient{
uploaded: &GeminiUploadedFile{Name: "files/input-jsonl"},
created: &GeminiBatchJob{Name: "batches/job-123", State: "JOB_STATE_PENDING"},
}
provider := NewGeminiAPIBatchImageProvider(client)
got, err := provider.Submit(context.Background(), nil, geminiAPIKeyAccount("sk-secret"), validGeminiBatchInput())
require.NoError(t, err)
require.NotContains(t, got.ProviderJobName, "base64")
require.NotContains(t, got.ProviderInputRef, "base64")
require.NotContains(t, got.ProviderOutputRef, "base64")
require.NotContains(t, got.ProviderJobName+got.ProviderInputRef+got.ProviderOutputRef, "iVBOR")
require.NotContains(t, got.ProviderJobName+got.ProviderInputRef+got.ProviderOutputRef, "A clean product hero image")
}
func requireJSONLLine(t *testing.T, line, wantKey, wantPrompt string) {
t.Helper()
var got map[string]any
require.NoError(t, json.Unmarshal([]byte(line), &got))
require.Equal(t, wantKey, got["key"])
request := got["request"].(map[string]any)
config := request["generationConfig"].(map[string]any)
require.Equal(t, []any{"TEXT", "IMAGE"}, config["responseModalities"])
contents := request["contents"].([]any)
parts := contents[0].(map[string]any)["parts"].([]any)
require.Equal(t, wantPrompt, parts[0].(map[string]any)["text"])
}
func validGeminiBatchInput() BatchImageInput {
return BatchImageInput{
BatchID: "imgbatch_123",
Model: "gemini-3.1-flash-image",
DisplayName: "test batch",
Items: []BatchImageInputItem{{
CustomID: "cover_001",
Prompt: "A clean product hero image",
}},
}
}
func geminiAPIKeyAccount(apiKey string) *Account {
return &Account{
Platform: PlatformGemini,
Type: AccountTypeAPIKey,
Credentials: map[string]any{"api_key": apiKey},
}
}
func jobWithProviderName(name string) *BatchImageJob {
return &BatchImageJob{ProviderJobName: &name}
}
type fakeGeminiBatchClient struct {
calls []string
uploaded *GeminiUploadedFile
created *GeminiBatchJob
got *GeminiBatchJob
uploadErr error
createErr error
getErr error
cancelErr error
downloadErr error
deleteErr error
uploadedJSONL []byte
createdFile string
cancelledBatch string
downloadedFile string
downloadBody string
downloadContentType string
deletedFiles []string
}
func (f *fakeGeminiBatchClient) UploadJSONL(_ context.Context, apiKey string, _ string, r io.Reader) (*GeminiUploadedFile, error) {
if strings.TrimSpace(apiKey) == "" {
return nil, errors.New("missing api key")
}
f.calls = append(f.calls, "upload")
f.uploadedJSONL, _ = io.ReadAll(r)
if f.uploadErr != nil {
return nil, f.uploadErr
}
if f.uploaded != nil {
return f.uploaded, nil
}
return &GeminiUploadedFile{Name: "files/input-jsonl"}, nil
}
func (f *fakeGeminiBatchClient) CreateBatch(_ context.Context, _ string, _ string, fileName string, _ string) (*GeminiBatchJob, error) {
f.calls = append(f.calls, "create")
f.createdFile = fileName
if f.createErr != nil {
return nil, f.createErr
}
if f.created != nil {
return f.created, nil
}
return &GeminiBatchJob{Name: "batches/job-123", State: "JOB_STATE_PENDING"}, nil
}
func (f *fakeGeminiBatchClient) GetBatch(_ context.Context, _ string, _ string) (*GeminiBatchJob, error) {
f.calls = append(f.calls, "get")
if f.getErr != nil {
return nil, f.getErr
}
return f.got, nil
}
func (f *fakeGeminiBatchClient) CancelBatch(_ context.Context, _ string, batchName string) error {
f.calls = append(f.calls, "cancel")
f.cancelledBatch = batchName
return f.cancelErr
}
func (f *fakeGeminiBatchClient) DownloadFile(_ context.Context, _ string, fileName string) (io.ReadCloser, string, error) {
f.calls = append(f.calls, "download")
f.downloadedFile = fileName
if f.downloadErr != nil {
return nil, "", f.downloadErr
}
contentType := f.downloadContentType
if contentType == "" {
contentType = "application/octet-stream"
}
return io.NopCloser(bytes.NewBufferString(f.downloadBody)), contentType, nil
}
func (f *fakeGeminiBatchClient) DeleteFile(_ context.Context, _ string, fileName string) error {
f.calls = append(f.calls, "delete")
f.deletedFiles = append(f.deletedFiles, fileName)
return f.deleteErr
}
@@ -0,0 +1,997 @@
package service
import (
"bytes"
"context"
"encoding/base64"
"encoding/json"
"errors"
"fmt"
"io"
"net/http"
"net/url"
"sort"
"strings"
"time"
"github.com/Wei-Shaw/sub2api/internal/config"
infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors"
)
const (
defaultVertexBatchRequeueAfter = 30 * time.Second
defaultVertexBatchLocation = "global"
defaultVertexManagedGCSPrefix = "batch-image/{env}/{batch_id}"
)
type VertexBatchImageProviderOptions struct {
Enabled bool
ProjectID string
Location string
ManagedGCSBucket string
ManagedGCSPrefix string
Environment string
InputRetentionHours int
OutputRetentionHours int
BatchPredictionBaseURL string
GCSBaseURL string
}
func NewVertexBatchImageProviderOptionsFromConfig(cfg *config.Config) VertexBatchImageProviderOptions {
if cfg == nil {
return VertexBatchImageProviderOptions{}
}
return VertexBatchImageProviderOptions{
Enabled: cfg.BatchImage.VertexEnabled,
ProjectID: cfg.BatchImage.VertexProjectID,
Location: cfg.BatchImage.VertexLocation,
ManagedGCSBucket: cfg.BatchImage.VertexManagedGCSBucket,
ManagedGCSPrefix: cfg.BatchImage.VertexManagedGCSPrefix,
Environment: cfg.Log.Environment,
InputRetentionHours: cfg.BatchImage.VertexInputRetentionHours,
OutputRetentionHours: cfg.BatchImage.VertexOutputRetentionHours,
BatchPredictionBaseURL: cfg.BatchImage.VertexBatchPredictionBaseURL,
GCSBaseURL: cfg.BatchImage.VertexGCSBaseURL,
}
}
type VertexBatchClient interface {
CreateBatchPredictionJob(ctx context.Context, accessToken string, req VertexCreateBatchPredictionJobRequest) (*VertexBatchPredictionJob, error)
GetBatchPredictionJob(ctx context.Context, accessToken string, name string) (*VertexBatchPredictionJob, error)
CancelBatchPredictionJob(ctx context.Context, accessToken string, name string) error
}
type VertexBatchObjectStore interface {
UploadJSONL(ctx context.Context, accessToken string, uri string, r io.Reader) error
ListJSONLObjects(ctx context.Context, accessToken string, prefixURI string) ([]string, error)
OpenObject(ctx context.Context, accessToken string, uri string) (io.ReadCloser, string, error)
DeleteObject(ctx context.Context, accessToken string, uri string) error
DeletePrefix(ctx context.Context, accessToken string, prefixURI string) error
}
type VertexCreateBatchPredictionJobRequest struct {
ProjectID string `json:"-"`
Location string `json:"-"`
DisplayName string `json:"displayName"`
Model string `json:"model"`
InputConfig VertexBatchInputConfig `json:"inputConfig"`
OutputConfig VertexBatchOutputConfig `json:"outputConfig"`
InstanceConfig *VertexBatchInstanceConfig `json:"instanceConfig,omitempty"`
}
type VertexBatchInputConfig struct {
InstancesFormat string `json:"instancesFormat"`
GCSSource VertexBatchGCSSource `json:"gcsSource"`
}
type VertexBatchGCSSource struct {
URIs []string `json:"uris"`
}
type VertexBatchOutputConfig struct {
PredictionsFormat string `json:"predictionsFormat"`
GCSDestination VertexBatchGCSDestination `json:"gcsDestination"`
}
type VertexBatchGCSDestination struct {
OutputURIPrefix string `json:"outputUriPrefix"`
}
type VertexBatchInstanceConfig struct {
KeyField string `json:"keyField"`
}
type VertexBatchPredictionJob struct {
Name string `json:"name"`
DisplayName string `json:"displayName"`
State string `json:"state"`
OutputConfig VertexBatchOutputConfig `json:"outputConfig"`
Error *VertexBatchJobError `json:"error"`
}
type VertexBatchJobError struct {
Code any `json:"code"`
Message string `json:"message"`
Status string `json:"status"`
}
type VertexBatchImageProvider struct {
opts VertexBatchImageProviderOptions
client VertexBatchClient
objectStore VertexBatchObjectStore
tokenCache GeminiTokenCache
}
func NewVertexBatchImageProvider(opts VertexBatchImageProviderOptions, client VertexBatchClient, objectStore VertexBatchObjectStore, tokenCache GeminiTokenCache) *VertexBatchImageProvider {
opts = normalizeVertexBatchImageProviderOptions(opts)
if client == nil {
client = NewVertexBatchHTTPClient(opts.BatchPredictionBaseURL, nil)
}
if objectStore == nil {
objectStore = NewVertexGCSObjectStore(opts.GCSBaseURL, nil)
}
return &VertexBatchImageProvider{
opts: opts,
client: client,
objectStore: objectStore,
tokenCache: tokenCache,
}
}
func NewVertexBatchImageProviderFromConfig(cfg *config.Config, client VertexBatchClient, objectStore VertexBatchObjectStore, tokenCache GeminiTokenCache) *VertexBatchImageProvider {
return NewVertexBatchImageProvider(NewVertexBatchImageProviderOptionsFromConfig(cfg), client, objectStore, tokenCache)
}
func normalizeVertexBatchImageProviderOptions(opts VertexBatchImageProviderOptions) VertexBatchImageProviderOptions {
opts.ProjectID = strings.TrimSpace(opts.ProjectID)
opts.Location = strings.TrimSpace(opts.Location)
if opts.Location == "" {
opts.Location = defaultVertexBatchLocation
}
opts.ManagedGCSBucket = strings.Trim(strings.TrimSpace(opts.ManagedGCSBucket), "/")
opts.ManagedGCSPrefix = strings.Trim(strings.TrimSpace(opts.ManagedGCSPrefix), "/")
if opts.ManagedGCSPrefix == "" {
opts.ManagedGCSPrefix = defaultVertexManagedGCSPrefix
}
opts.Environment = strings.TrimSpace(opts.Environment)
if opts.Environment == "" {
opts.Environment = "default"
}
opts.BatchPredictionBaseURL = strings.TrimRight(strings.TrimSpace(opts.BatchPredictionBaseURL), "/")
opts.GCSBaseURL = strings.TrimRight(strings.TrimSpace(opts.GCSBaseURL), "/")
return opts
}
func (p *VertexBatchImageProvider) Name() string {
return BatchImageProviderVertex
}
func (p *VertexBatchImageProvider) SupportsAccount(account *Account) bool {
if account == nil || account.Platform != PlatformGemini || account.Type != AccountTypeServiceAccount {
return false
}
_, err := parseVertexServiceAccountKey(account)
return err == nil
}
func (p *VertexBatchImageProvider) Submit(ctx context.Context, job *BatchImageJob, account *Account, input BatchImageInput) (*BatchProviderJob, error) {
if err := p.validateAccount(account); err != nil {
return nil, err
}
if strings.TrimSpace(p.opts.ManagedGCSBucket) == "" {
return nil, vertexProviderError("VERTEX_MANAGED_GCS_BUCKET_MISSING", "Vertex managed GCS bucket is not configured", nil)
}
if input.BatchID == "" && job != nil {
input.BatchID = job.BatchID
}
if input.Model == "" && job != nil {
input.Model = job.Model
}
jsonl, err := BuildVertexBatchJSONL(input)
if err != nil {
return nil, err
}
refs, err := p.managedRefs(input.BatchID)
if err != nil {
return nil, err
}
accessToken, err := p.accessToken(ctx, account)
if err != nil {
return nil, mapVertexClientError(err)
}
if err := p.objectStore.UploadJSONL(ctx, accessToken, refs.InputURI, bytes.NewReader(jsonl)); err != nil {
return nil, vertexProviderError("VERTEX_GCS_UPLOAD_FAILED", "Vertex managed GCS upload failed", nil)
}
projectID := strings.TrimSpace(p.opts.ProjectID)
if projectID == "" {
projectID = account.VertexProjectID()
}
if projectID == "" {
return nil, vertexProviderError("VERTEX_PROJECT_ID_MISSING", "Vertex project id is not configured", nil)
}
location := strings.TrimSpace(p.opts.Location)
if location == "" {
location = account.VertexLocation(input.Model)
}
req := VertexCreateBatchPredictionJobRequest{
ProjectID: projectID,
Location: location,
DisplayName: vertexBatchDisplayName(input),
Model: NormalizeVertexBatchModelPath(input.Model),
InputConfig: VertexBatchInputConfig{InstancesFormat: "jsonl", GCSSource: VertexBatchGCSSource{URIs: []string{refs.InputURI}}},
OutputConfig: VertexBatchOutputConfig{PredictionsFormat: "jsonl", GCSDestination: VertexBatchGCSDestination{OutputURIPrefix: refs.OutputPrefixURI}},
InstanceConfig: &VertexBatchInstanceConfig{KeyField: "key"},
}
created, err := p.client.CreateBatchPredictionJob(ctx, accessToken, req)
if err != nil {
return nil, mapVertexClientError(err)
}
if created == nil || strings.TrimSpace(created.Name) == "" {
return nil, vertexProviderError("VERTEX_INVALID_RESPONSE", "Vertex batch response is missing job name", nil)
}
return &BatchProviderJob{
ProviderJobName: created.Name,
ProviderInputRef: refs.InputURI,
ProviderOutputRef: refs.OutputPrefixURI,
RawState: created.State,
}, nil
}
func (p *VertexBatchImageProvider) Get(ctx context.Context, job *BatchImageJob, account *Account) (*BatchProviderStatus, error) {
if err := p.validateAccount(account); err != nil {
return nil, err
}
jobName := batchImageProviderJobName(job)
if jobName == "" {
return nil, ErrBatchImageProviderMissingJobName
}
accessToken, err := p.accessToken(ctx, account)
if err != nil {
return nil, mapVertexClientError(err)
}
vertexJob, err := p.client.GetBatchPredictionJob(ctx, accessToken, jobName)
if err != nil {
return nil, mapVertexClientError(err)
}
if vertexJob == nil {
return nil, vertexProviderError("VERTEX_INVALID_RESPONSE", "Vertex batch response is empty", nil)
}
status := mapVertexBatchState(vertexJob)
outputRef := strings.TrimSpace(vertexJob.OutputConfig.GCSDestination.OutputURIPrefix)
if outputRef == "" {
outputRef = batchImageProviderOutputRef(job)
}
if outputRef == "" && job != nil && job.GCSOutputURI != nil {
outputRef = strings.TrimSpace(*job.GCSOutputURI)
}
status.ProviderOutputRef = outputRef
return status, nil
}
func (p *VertexBatchImageProvider) Cancel(ctx context.Context, job *BatchImageJob, account *Account) error {
if err := p.validateAccount(account); err != nil {
return err
}
jobName := batchImageProviderJobName(job)
if jobName == "" {
return ErrBatchImageProviderMissingJobName
}
accessToken, err := p.accessToken(ctx, account)
if err != nil {
return mapVertexClientError(err)
}
return mapVertexClientError(p.client.CancelBatchPredictionJob(ctx, accessToken, jobName))
}
func (p *VertexBatchImageProvider) OpenResult(ctx context.Context, job *BatchImageJob, account *Account) (io.ReadCloser, string, error) {
if err := p.validateAccount(account); err != nil {
return nil, "", err
}
outputRef := batchImageProviderOutputRef(job)
if outputRef == "" && job != nil && job.GCSOutputURI != nil {
outputRef = strings.TrimSpace(*job.GCSOutputURI)
}
if outputRef == "" {
return nil, "", ErrBatchImageProviderMissingResultRef
}
accessToken, err := p.accessToken(ctx, account)
if err != nil {
return nil, "", mapVertexClientError(err)
}
objects, err := p.objectStore.ListJSONLObjects(ctx, accessToken, outputRef)
if err != nil {
return nil, "", vertexProviderError("VERTEX_GCS_LIST_FAILED", "Vertex managed GCS list failed", nil)
}
sort.Strings(objects)
if len(objects) == 0 {
return nil, "", vertexProviderError("VERTEX_RESULT_OBJECTS_MISSING", "Vertex result objects are missing", nil)
}
return &vertexCombinedJSONLReadCloser{
ctx: ctx,
accessToken: accessToken,
objects: objects,
store: p.objectStore,
}, "application/jsonl", nil
}
func (p *VertexBatchImageProvider) Cleanup(ctx context.Context, job *BatchImageJob, account *Account, target CleanupTarget) error {
if err := p.validateAccount(account); err != nil {
return err
}
accessToken, err := p.accessToken(ctx, account)
if err != nil {
return mapVertexClientError(err)
}
inputRef := batchImageProviderInputRef(job)
outputRef := batchImageProviderOutputRef(job)
if job != nil {
if inputRef == "" && job.GCSInputURI != nil {
inputRef = strings.TrimSpace(*job.GCSInputURI)
}
if outputRef == "" && job.GCSOutputURI != nil {
outputRef = strings.TrimSpace(*job.GCSOutputURI)
}
}
switch target {
case CleanupTargetInput:
return p.deleteManagedInput(ctx, accessToken, job, inputRef)
case CleanupTargetOutput:
return p.deleteManagedOutput(ctx, accessToken, job, outputRef)
case CleanupTargetAll:
if err := p.deleteManagedInput(ctx, accessToken, job, inputRef); err != nil {
return err
}
return p.deleteManagedOutput(ctx, accessToken, job, outputRef)
default:
return ErrUnsupportedCleanupTarget
}
}
func (p *VertexBatchImageProvider) validateAccount(account *Account) error {
if account == nil || account.Platform != PlatformGemini || account.Type != AccountTypeServiceAccount {
return ErrBatchImageProviderUnsupportedAccount
}
if _, err := parseVertexServiceAccountKey(account); err != nil {
return ErrBatchImageProviderMissingServiceAccount
}
return nil
}
func (p *VertexBatchImageProvider) accessToken(ctx context.Context, account *Account) (string, error) {
return getVertexServiceAccountAccessToken(ctx, p.tokenCache, account)
}
func (p *VertexBatchImageProvider) deleteManagedInput(ctx context.Context, accessToken string, job *BatchImageJob, uri string) error {
if strings.TrimSpace(uri) == "" {
return nil
}
if !p.isSafeManagedInput(job, uri) {
return ErrBatchImageProviderUnsafeCleanupPath
}
return mapVertexClientError(p.objectStore.DeleteObject(ctx, accessToken, uri))
}
func (p *VertexBatchImageProvider) deleteManagedOutput(ctx context.Context, accessToken string, job *BatchImageJob, uri string) error {
if strings.TrimSpace(uri) == "" {
return nil
}
if !p.isSafeManagedOutput(job, uri) {
return ErrBatchImageProviderUnsafeCleanupPath
}
return mapVertexClientError(p.objectStore.DeletePrefix(ctx, accessToken, uri))
}
func (p *VertexBatchImageProvider) isSafeManagedInput(job *BatchImageJob, uri string) bool {
if job == nil || strings.TrimSpace(job.BatchID) == "" {
return false
}
refs, err := p.managedRefs(job.BatchID)
return err == nil && strings.TrimSpace(uri) == refs.InputURI
}
func (p *VertexBatchImageProvider) isSafeManagedOutput(job *BatchImageJob, uri string) bool {
if job == nil || strings.TrimSpace(job.BatchID) == "" {
return false
}
refs, err := p.managedRefs(job.BatchID)
return err == nil && strings.HasPrefix(strings.TrimSpace(uri), refs.OutputPrefixURI)
}
type vertexManagedRefs struct {
Prefix string
InputURI string
OutputPrefixURI string
}
func (p *VertexBatchImageProvider) managedRefs(batchID string) (vertexManagedRefs, error) {
batchID = strings.TrimSpace(batchID)
if !IsValidBatchImageID(batchID) {
return vertexManagedRefs{}, batchImageProviderInputError("valid batch_id is required")
}
bucket := strings.Trim(strings.TrimSpace(p.opts.ManagedGCSBucket), "/")
if bucket == "" || strings.Contains(bucket, "://") {
return vertexManagedRefs{}, vertexProviderError("VERTEX_MANAGED_GCS_BUCKET_MISSING", "Vertex managed GCS bucket is not configured", nil)
}
prefix := buildVertexManagedGCSPrefix(p.opts.ManagedGCSPrefix, p.opts.Environment, batchID)
if !strings.Contains(prefix, batchID) {
return vertexManagedRefs{}, batchImageProviderInputError("managed GCS prefix must contain batch_id")
}
base := "gs://" + bucket + "/" + strings.Trim(prefix, "/")
return vertexManagedRefs{
Prefix: strings.Trim(prefix, "/"),
InputURI: base + "/input/requests.jsonl",
OutputPrefixURI: base + "/output/",
}, nil
}
func buildVertexManagedGCSPrefix(template, env, batchID string) string {
template = strings.Trim(strings.TrimSpace(template), "/")
if template == "" {
template = defaultVertexManagedGCSPrefix
}
env = sanitizeVertexGCSPathSegment(env)
batchID = sanitizeVertexGCSPathSegment(batchID)
prefix := strings.ReplaceAll(template, "{env}", env)
prefix = strings.ReplaceAll(prefix, "{batch_id}", batchID)
return strings.Trim(prefix, "/")
}
func sanitizeVertexGCSPathSegment(v string) string {
v = strings.TrimSpace(v)
if v == "" {
return "default"
}
var b strings.Builder
for _, r := range v {
switch {
case r >= 'a' && r <= 'z', r >= 'A' && r <= 'Z', r >= '0' && r <= '9':
_, _ = b.WriteRune(r)
case r == '-', r == '_', r == '.':
_, _ = b.WriteRune(r)
default:
_ = b.WriteByte('-')
}
}
return b.String()
}
func vertexBatchDisplayName(input BatchImageInput) string {
if v := strings.TrimSpace(input.DisplayName); v != "" {
return v
}
if v := strings.TrimSpace(input.BatchID); v != "" {
return "sub2api-" + v
}
return "sub2api-image-batch"
}
func BuildVertexBatchJSONL(input BatchImageInput) ([]byte, error) {
if strings.TrimSpace(input.Model) == "" {
return nil, batchImageProviderInputError("model is required")
}
if len(input.Items) == 0 {
return nil, batchImageProviderInputError("at least one item is required")
}
seen := make(map[string]struct{}, len(input.Items))
var buf bytes.Buffer
enc := json.NewEncoder(&buf)
for _, item := range input.Items {
customID := strings.TrimSpace(item.CustomID)
if customID == "" {
return nil, batchImageProviderInputError("custom_id is required")
}
if _, ok := seen[customID]; ok {
return nil, batchImageProviderInputError("duplicate custom_id %q", customID)
}
seen[customID] = struct{}{}
prompt := strings.TrimSpace(item.Prompt)
if prompt == "" {
return nil, batchImageProviderInputError("prompt is required for custom_id %q", customID)
}
parts, err := vertexBatchImageParts(prompt, item.ReferenceImages)
if err != nil {
return nil, err
}
line := map[string]any{
"key": customID,
"request": map[string]any{
"contents": []any{map[string]any{
"role": "user",
"parts": parts,
}},
"generationConfig": map[string]any{
"responseModalities": []string{"TEXT", "IMAGE"},
},
},
}
if err := enc.Encode(line); err != nil {
return nil, err
}
}
return buf.Bytes(), nil
}
func vertexBatchImageParts(prompt string, refs []BatchImageReference) ([]any, error) {
parts := []any{map[string]any{"text": prompt}}
for _, ref := range refs {
mimeType := normalizeBatchImageReferenceMimeType(ref.MimeType)
if mimeType == "" {
return nil, batchImageProviderInputError("reference image mime_type is required")
}
fileURI := strings.TrimSpace(ref.FileURI)
switch {
case len(ref.Data) > 0 && fileURI == "":
parts = append(parts, map[string]any{
"inlineData": map[string]any{
"mimeType": mimeType,
"data": base64.StdEncoding.EncodeToString(ref.Data),
},
})
case len(ref.Data) == 0 && fileURI != "":
parts = append(parts, map[string]any{
"fileData": map[string]any{
"mimeType": mimeType,
"fileUri": fileURI,
},
})
default:
return nil, batchImageProviderInputError("reference image must contain exactly one of data or file_uri")
}
}
return parts, nil
}
func NormalizeVertexBatchModelPath(model string) string {
model = strings.Trim(strings.TrimSpace(model), "/")
if strings.HasPrefix(model, "publishers/") || strings.HasPrefix(model, "projects/") {
return model
}
return "publishers/google/models/" + model
}
func BuildVertexBatchPredictionJobsEndpoint(baseURL, projectID, location string) (string, error) {
projectID = strings.TrimSpace(projectID)
location = strings.TrimSpace(location)
if projectID == "" {
return "", errors.New("vertex project_id is required")
}
if location == "" {
location = defaultVertexBatchLocation
}
if !vertexLocationPattern.MatchString(location) {
return "", fmt.Errorf("invalid vertex location: %s", location)
}
if strings.TrimSpace(baseURL) != "" {
return strings.TrimRight(strings.TrimSpace(baseURL), "/") + "/v1/projects/" + url.PathEscape(projectID) + "/locations/" + url.PathEscape(location) + "/batchPredictionJobs", nil
}
host := fmt.Sprintf("%s-aiplatform.googleapis.com", location)
if location == "global" {
host = "aiplatform.googleapis.com"
}
return fmt.Sprintf("https://%s/v1/projects/%s/locations/%s/batchPredictionJobs", host, url.PathEscape(projectID), url.PathEscape(location)), nil
}
func mapVertexBatchState(job *VertexBatchPredictionJob) *BatchProviderStatus {
state := strings.TrimSpace(job.State)
status := &BatchProviderStatus{
RawState: state,
InternalState: BatchProviderStateRunning,
SuggestedRequeueAfter: defaultVertexBatchRequeueAfter,
}
switch strings.ToUpper(state) {
case "JOB_STATE_PENDING", "JOB_STATE_QUEUED":
status.InternalState = BatchProviderStateQueued
case "JOB_STATE_RUNNING", "JOB_STATE_PAUSED":
status.InternalState = BatchProviderStateRunning
case "JOB_STATE_SUCCEEDED":
status.InternalState = BatchProviderStateSucceeded
status.Done = true
status.SuggestedRequeueAfter = 0
case "JOB_STATE_FAILED":
status.InternalState = BatchProviderStateFailed
status.Done = true
status.ErrorCode = "VERTEX_BATCH_FAILED"
status.SuggestedRequeueAfter = 0
case "JOB_STATE_CANCELLED":
status.InternalState = BatchProviderStateCancelled
status.Done = true
status.ErrorCode = "VERTEX_BATCH_CANCELLED"
status.SuggestedRequeueAfter = 0
case "JOB_STATE_EXPIRED":
status.InternalState = BatchProviderStateExpired
status.Done = true
status.ErrorCode = "VERTEX_BATCH_EXPIRED"
status.SuggestedRequeueAfter = 0
default:
if job.Error != nil && strings.TrimSpace(job.Error.Message) != "" {
status.InternalState = BatchProviderStateFailed
status.Done = true
status.ErrorCode = "VERTEX_BATCH_FAILED"
status.SuggestedRequeueAfter = 0
}
}
if job.Error != nil {
if code := strings.TrimSpace(job.Error.Status); code != "" {
status.ErrorCode = code
}
status.ErrorMessage = strings.TrimSpace(job.Error.Message)
}
return status
}
func vertexProviderError(reason, message string, cause error) error {
err := infraerrors.New(http.StatusBadGateway, reason, message)
if cause != nil {
return err.WithCause(cause)
}
return err
}
func mapVertexClientError(err error) error {
if err == nil {
return nil
}
if errors.Is(err, ErrBatchImageProviderMissingServiceAccount) ||
errors.Is(err, ErrBatchImageProviderMissingJobName) ||
errors.Is(err, ErrBatchImageProviderMissingResultRef) ||
errors.Is(err, ErrBatchImageProviderUnsafeCleanupPath) ||
errors.Is(err, ErrUnsupportedCleanupTarget) {
return err
}
var apiErr *VertexAPIError
if errors.As(err, &apiErr) {
switch apiErr.StatusCode {
case http.StatusUnauthorized:
return vertexProviderError("VERTEX_AUTH_FAILED", "Vertex authentication failed", nil)
case http.StatusForbidden:
return vertexProviderError("VERTEX_PERMISSION_DENIED", "Vertex permission denied", nil)
case http.StatusTooManyRequests:
return vertexProviderError("VERTEX_RATE_LIMITED", "Vertex rate limit exceeded", nil)
case http.StatusNotFound:
return vertexProviderError("VERTEX_BATCH_NOT_FOUND", "Vertex batch resource was not found", nil)
default:
return vertexProviderError("VERTEX_INVALID_RESPONSE", "Vertex API request failed", nil)
}
}
return vertexProviderError("VERTEX_INVALID_RESPONSE", "Vertex API request failed", err)
}
type vertexCombinedJSONLReadCloser struct {
ctx context.Context
accessToken string
objects []string
store VertexBatchObjectStore
index int
current io.ReadCloser
needBoundary bool
closed bool
}
func (r *vertexCombinedJSONLReadCloser) Read(p []byte) (int, error) {
if r.closed {
return 0, io.ErrClosedPipe
}
if r.needBoundary {
if len(p) == 0 {
return 0, nil
}
p[0] = '\n'
r.needBoundary = false
return 1, nil
}
for {
if r.current == nil {
if r.index >= len(r.objects) {
return 0, io.EOF
}
obj := r.objects[r.index]
r.index++
rc, _, err := r.store.OpenObject(r.ctx, r.accessToken, obj)
if err != nil {
return 0, err
}
r.current = rc
}
n, err := r.current.Read(p)
if err == io.EOF {
_ = r.current.Close()
r.current = nil
if r.index < len(r.objects) {
if n > 0 {
r.needBoundary = true
return n, nil
}
if len(p) == 0 {
return 0, nil
}
p[0] = '\n'
return 1, nil
}
if n > 0 {
return n, nil
}
continue
}
return n, err
}
}
func (r *vertexCombinedJSONLReadCloser) Close() error {
r.closed = true
if r.current != nil {
return r.current.Close()
}
return nil
}
type VertexBatchHTTPClient struct {
baseURL string
client *http.Client
}
func NewVertexBatchHTTPClient(baseURL string, client *http.Client) *VertexBatchHTTPClient {
if client == nil {
client = http.DefaultClient
}
return &VertexBatchHTTPClient{baseURL: strings.TrimRight(strings.TrimSpace(baseURL), "/"), client: client}
}
func (c *VertexBatchHTTPClient) CreateBatchPredictionJob(ctx context.Context, accessToken string, req VertexCreateBatchPredictionJobRequest) (*VertexBatchPredictionJob, error) {
endpoint, err := BuildVertexBatchPredictionJobsEndpoint(c.baseURL, req.ProjectID, req.Location)
if err != nil {
return nil, err
}
payload, err := json.Marshal(req)
if err != nil {
return nil, err
}
httpReq, err := http.NewRequestWithContext(ctx, http.MethodPost, endpoint, bytes.NewReader(payload))
if err != nil {
return nil, err
}
httpReq.Header.Set("Content-Type", "application/json")
httpReq.Header.Set("Authorization", "Bearer "+accessToken)
return doVertexJSON[VertexBatchPredictionJob](c.client, httpReq)
}
func (c *VertexBatchHTTPClient) GetBatchPredictionJob(ctx context.Context, accessToken string, name string) (*VertexBatchPredictionJob, error) {
endpoint := c.vertexResourceURL(name)
req, err := http.NewRequestWithContext(ctx, http.MethodGet, endpoint, nil)
if err != nil {
return nil, err
}
req.Header.Set("Authorization", "Bearer "+accessToken)
return doVertexJSON[VertexBatchPredictionJob](c.client, req)
}
func (c *VertexBatchHTTPClient) CancelBatchPredictionJob(ctx context.Context, accessToken string, name string) error {
endpoint := c.vertexResourceURL(name) + ":cancel"
req, err := http.NewRequestWithContext(ctx, http.MethodPost, endpoint, nil)
if err != nil {
return err
}
req.Header.Set("Authorization", "Bearer "+accessToken)
return doVertexNoBody(c.client, req)
}
func (c *VertexBatchHTTPClient) vertexResourceURL(name string) string {
name = strings.TrimLeft(strings.TrimSpace(name), "/")
if c.baseURL != "" {
return c.baseURL + "/v1/" + name
}
return "https://aiplatform.googleapis.com/v1/" + name
}
type VertexGCSObjectStore struct {
baseURL string
client *http.Client
}
func NewVertexGCSObjectStore(baseURL string, client *http.Client) *VertexGCSObjectStore {
if client == nil {
client = http.DefaultClient
}
baseURL = strings.TrimRight(strings.TrimSpace(baseURL), "/")
if baseURL == "" {
baseURL = "https://storage.googleapis.com"
}
return &VertexGCSObjectStore{baseURL: baseURL, client: client}
}
func (s *VertexGCSObjectStore) UploadJSONL(ctx context.Context, accessToken string, uri string, r io.Reader) error {
bucket, object, err := parseGCSURI(uri)
if err != nil {
return err
}
endpoint := fmt.Sprintf("%s/upload/storage/v1/b/%s/o?uploadType=media&name=%s", s.baseURL, url.PathEscape(bucket), url.QueryEscape(object))
req, err := http.NewRequestWithContext(ctx, http.MethodPost, endpoint, r)
if err != nil {
return err
}
req.Header.Set("Authorization", "Bearer "+accessToken)
req.Header.Set("Content-Type", "application/jsonl")
return doVertexNoBody(s.client, req)
}
func (s *VertexGCSObjectStore) ListJSONLObjects(ctx context.Context, accessToken string, prefixURI string) ([]string, error) {
return s.listObjects(ctx, accessToken, prefixURI, true)
}
func (s *VertexGCSObjectStore) listObjects(ctx context.Context, accessToken string, prefixURI string, jsonlOnly bool) ([]string, error) {
bucket, prefix, err := parseGCSURI(prefixURI)
if err != nil {
return nil, err
}
var objects []string
pageToken := ""
for {
endpoint := fmt.Sprintf("%s/storage/v1/b/%s/o?prefix=%s", s.baseURL, url.PathEscape(bucket), url.QueryEscape(prefix))
if pageToken != "" {
endpoint += "&pageToken=" + url.QueryEscape(pageToken)
}
req, err := http.NewRequestWithContext(ctx, http.MethodGet, endpoint, nil)
if err != nil {
return nil, err
}
req.Header.Set("Authorization", "Bearer "+accessToken)
var page struct {
Items []struct {
Name string `json:"name"`
} `json:"items"`
NextPageToken string `json:"nextPageToken"`
}
if err := doVertexDecodeJSON(s.client, req, &page); err != nil {
return nil, err
}
for _, item := range page.Items {
if !jsonlOnly || strings.HasSuffix(item.Name, ".jsonl") {
objects = append(objects, "gs://"+bucket+"/"+item.Name)
}
}
if page.NextPageToken == "" {
return objects, nil
}
pageToken = page.NextPageToken
}
}
func (s *VertexGCSObjectStore) OpenObject(ctx context.Context, accessToken string, uri string) (io.ReadCloser, string, error) {
bucket, object, err := parseGCSURI(uri)
if err != nil {
return nil, "", err
}
endpoint := fmt.Sprintf("%s/storage/v1/b/%s/o/%s?alt=media", s.baseURL, url.PathEscape(bucket), url.PathEscape(object))
req, err := http.NewRequestWithContext(ctx, http.MethodGet, endpoint, nil)
if err != nil {
return nil, "", err
}
req.Header.Set("Authorization", "Bearer "+accessToken)
resp, err := s.client.Do(req)
if err != nil {
return nil, "", err
}
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
defer func() { _ = resp.Body.Close() }()
return nil, "", readVertexAPIError(resp)
}
contentType := resp.Header.Get("Content-Type")
if contentType == "" {
contentType = "application/jsonl"
}
return resp.Body, contentType, nil
}
func (s *VertexGCSObjectStore) DeleteObject(ctx context.Context, accessToken string, uri string) error {
bucket, object, err := parseGCSURI(uri)
if err != nil {
return err
}
endpoint := fmt.Sprintf("%s/storage/v1/b/%s/o/%s", s.baseURL, url.PathEscape(bucket), url.PathEscape(object))
req, err := http.NewRequestWithContext(ctx, http.MethodDelete, endpoint, nil)
if err != nil {
return err
}
req.Header.Set("Authorization", "Bearer "+accessToken)
return doVertexNoBody(s.client, req)
}
func (s *VertexGCSObjectStore) DeletePrefix(ctx context.Context, accessToken string, prefixURI string) error {
objects, err := s.listObjects(ctx, accessToken, prefixURI, false)
if err != nil {
return err
}
for _, object := range objects {
if err := s.DeleteObject(ctx, accessToken, object); err != nil {
return err
}
}
return nil
}
func parseGCSURI(uri string) (bucket, object string, err error) {
uri = strings.TrimSpace(uri)
if !strings.HasPrefix(uri, "gs://") {
return "", "", fmt.Errorf("invalid gcs uri")
}
rest := strings.TrimPrefix(uri, "gs://")
parts := strings.SplitN(rest, "/", 2)
if len(parts) != 2 || strings.TrimSpace(parts[0]) == "" || strings.TrimSpace(parts[1]) == "" {
return "", "", fmt.Errorf("invalid gcs uri")
}
return parts[0], parts[1], nil
}
type VertexAPIError struct {
StatusCode int
Code string
Message string
}
func (e *VertexAPIError) Error() string {
if e == nil {
return "<nil>"
}
if e.Code != "" {
return fmt.Sprintf("vertex api error: status=%d code=%s message=%s", e.StatusCode, e.Code, e.Message)
}
return fmt.Sprintf("vertex api error: status=%d message=%s", e.StatusCode, e.Message)
}
func doVertexJSON[T any](client *http.Client, req *http.Request) (*T, error) {
var out T
if err := doVertexDecodeJSON(client, req, &out); err != nil {
return nil, err
}
return &out, nil
}
func doVertexDecodeJSON(client *http.Client, req *http.Request, out any) error {
resp, err := client.Do(req)
if err != nil {
return err
}
defer func() { _ = resp.Body.Close() }()
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
return readVertexAPIError(resp)
}
return json.NewDecoder(resp.Body).Decode(out)
}
func doVertexNoBody(client *http.Client, req *http.Request) error {
resp, err := client.Do(req)
if err != nil {
return err
}
defer func() { _ = resp.Body.Close() }()
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
return readVertexAPIError(resp)
}
return nil
}
func readVertexAPIError(resp *http.Response) error {
body, _ := io.ReadAll(io.LimitReader(resp.Body, 8192))
message := string(body)
code := ""
var parsed struct {
Error struct {
Code any `json:"code"`
Message string `json:"message"`
Status string `json:"status"`
} `json:"error"`
}
if err := json.Unmarshal(body, &parsed); err == nil && parsed.Error.Message != "" {
message = parsed.Error.Message
code = parsed.Error.Status
}
return &VertexAPIError{StatusCode: resp.StatusCode, Code: code, Message: message}
}
var _ BatchImageProvider = (*VertexBatchImageProvider)(nil)
var _ VertexBatchClient = (*VertexBatchHTTPClient)(nil)
var _ VertexBatchObjectStore = (*VertexGCSObjectStore)(nil)
@@ -0,0 +1,438 @@
//go:build unit
package service
import (
"bytes"
"context"
"encoding/json"
"errors"
"io"
"strings"
"testing"
"time"
infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors"
"github.com/stretchr/testify/require"
)
func TestBatchImageProviderRegistry_ReturnsVertex(t *testing.T) {
registry := NewDefaultBatchImageProviderRegistry()
provider, ok := registry.Get(BatchImageProviderVertex)
require.True(t, ok)
require.Equal(t, BatchImageProviderVertex, provider.Name())
}
func TestVertexProvider_SupportsOnlyGeminiServiceAccount(t *testing.T) {
provider := newTestVertexProvider(&fakeVertexBatchClient{}, &fakeVertexObjectStore{})
require.True(t, provider.SupportsAccount(vertexServiceAccount()))
require.False(t, provider.SupportsAccount(&Account{Platform: PlatformGemini, Type: AccountTypeAPIKey, Credentials: map[string]any{"api_key": "sk"}}))
require.False(t, provider.SupportsAccount(&Account{Platform: PlatformGemini, Type: AccountTypeOAuth, Credentials: map[string]any{"access_token": "tok"}}))
require.False(t, provider.SupportsAccount(&Account{Platform: PlatformAnthropic, Type: AccountTypeServiceAccount, Credentials: vertexServiceAccount().Credentials}))
require.False(t, provider.SupportsAccount(&Account{Platform: PlatformGemini, Type: AccountTypeServiceAccount, Credentials: map[string]any{}}))
}
func TestVertexProvider_MissingServiceAccountRejected(t *testing.T) {
provider := newTestVertexProvider(&fakeVertexBatchClient{}, &fakeVertexObjectStore{})
_, err := provider.Submit(context.Background(), nil, &Account{Platform: PlatformGemini, Type: AccountTypeServiceAccount, Credentials: map[string]any{}}, validVertexBatchInput())
require.ErrorIs(t, err, ErrBatchImageProviderMissingServiceAccount)
}
func TestVertexProvider_MissingManagedGCSBucketRejected(t *testing.T) {
provider := NewVertexBatchImageProvider(VertexBatchImageProviderOptions{ProjectID: "proj", Environment: "test"}, &fakeVertexBatchClient{}, &fakeVertexObjectStore{}, &fakeGeminiTokenCache{token: "token"})
_, err := provider.Submit(context.Background(), nil, vertexServiceAccount(), validVertexBatchInput())
require.Error(t, err)
require.Equal(t, "VERTEX_MANAGED_GCS_BUCKET_MISSING", infraerrors.Reason(err))
}
func TestBuildVertexBatchJSONL_WritesValidLinesAndPreservesCustomID(t *testing.T) {
input := validVertexBatchInput()
input.Items = append(input.Items, BatchImageInputItem{CustomID: "cover_002", Prompt: "Second prompt"})
jsonl, err := BuildVertexBatchJSONL(input)
require.NoError(t, err)
lines := strings.Split(strings.TrimSpace(string(jsonl)), "\n")
require.Len(t, lines, 2)
requireVertexJSONLLine(t, lines[0], "cover_001", "A clean product hero image")
requireVertexJSONLLine(t, lines[1], "cover_002", "Second prompt")
}
func TestBuildVertexBatchJSONL_RejectsDuplicateCustomIDs(t *testing.T) {
input := validVertexBatchInput()
input.Items = append(input.Items, BatchImageInputItem{CustomID: "cover_001", Prompt: "Duplicate"})
_, err := BuildVertexBatchJSONL(input)
require.ErrorIs(t, err, ErrBatchImageProviderInvalidInput)
}
func TestBuildVertexBatchJSONL_RejectsEmptyPrompt(t *testing.T) {
input := validVertexBatchInput()
input.Items[0].Prompt = " "
_, err := BuildVertexBatchJSONL(input)
require.ErrorIs(t, err, ErrBatchImageProviderInvalidInput)
}
func TestBuildVertexBatchJSONL_WritesReferenceImages(t *testing.T) {
input := validVertexBatchInput()
input.Items[0].ReferenceImages = []BatchImageReference{
{MimeType: "image/png", Data: []byte("png-bytes")},
{MimeType: "image/jpeg", FileURI: "gs://bucket/refs/style.jpg"},
}
jsonl, err := BuildVertexBatchJSONL(input)
require.NoError(t, err)
lines := strings.Split(strings.TrimSpace(string(jsonl)), "\n")
require.Len(t, lines, 1)
var got map[string]any
require.NoError(t, json.Unmarshal([]byte(lines[0]), &got))
request := got["request"].(map[string]any)
contents := request["contents"].([]any)
parts := contents[0].(map[string]any)["parts"].([]any)
require.Len(t, parts, 3)
require.Equal(t, "A clean product hero image", parts[0].(map[string]any)["text"])
inlineData := parts[1].(map[string]any)["inlineData"].(map[string]any)
require.Equal(t, "image/png", inlineData["mimeType"])
require.Equal(t, "cG5nLWJ5dGVz", inlineData["data"])
fileData := parts[2].(map[string]any)["fileData"].(map[string]any)
require.Equal(t, "image/jpeg", fileData["mimeType"])
require.Equal(t, "gs://bucket/refs/style.jpg", fileData["fileUri"])
}
func TestNormalizeVertexBatchModelPath(t *testing.T) {
require.Equal(t, "publishers/google/models/gemini-3.1-flash-image", NormalizeVertexBatchModelPath("gemini-3.1-flash-image"))
require.Equal(t, "publishers/google/models/gemini-2.5-flash-image", NormalizeVertexBatchModelPath("publishers/google/models/gemini-2.5-flash-image"))
require.Equal(t, "projects/p/locations/global/models/m", NormalizeVertexBatchModelPath("projects/p/locations/global/models/m"))
}
func TestBuildVertexBatchPredictionJobsEndpoint(t *testing.T) {
global, err := BuildVertexBatchPredictionJobsEndpoint("", "my-project", "global")
require.NoError(t, err)
require.Equal(t, "https://aiplatform.googleapis.com/v1/projects/my-project/locations/global/batchPredictionJobs", global)
regional, err := BuildVertexBatchPredictionJobsEndpoint("", "my-project", "asia-northeast1")
require.NoError(t, err)
require.Equal(t, "https://asia-northeast1-aiplatform.googleapis.com/v1/projects/my-project/locations/asia-northeast1/batchPredictionJobs", regional)
}
func TestVertexProvider_SubmitUploadsJSONLAndCreatesBatchPredictionJob(t *testing.T) {
vertexClient := &fakeVertexBatchClient{created: &VertexBatchPredictionJob{Name: "projects/proj/locations/global/batchPredictionJobs/job-1", State: "JOB_STATE_PENDING"}}
store := &fakeVertexObjectStore{}
provider := newTestVertexProvider(vertexClient, store)
got, err := provider.Submit(context.Background(), &BatchImageJob{BatchID: "imgbatch_abc123", Model: "gemini-3.1-flash-image"}, vertexServiceAccount(), validVertexBatchInput())
require.NoError(t, err)
require.Equal(t, "gs://managed-bucket/batch-image/test/imgbatch_abc123/input/requests.jsonl", store.uploadURI)
require.Equal(t, "projects/proj/locations/global/batchPredictionJobs/job-1", got.ProviderJobName)
require.Equal(t, store.uploadURI, got.ProviderInputRef)
require.Equal(t, "gs://managed-bucket/batch-image/test/imgbatch_abc123/output/", got.ProviderOutputRef)
require.Equal(t, "jsonl", vertexClient.createdReq.InputConfig.InstancesFormat)
require.Equal(t, "jsonl", vertexClient.createdReq.OutputConfig.PredictionsFormat)
require.Equal(t, got.ProviderOutputRef, vertexClient.createdReq.OutputConfig.GCSDestination.OutputURIPrefix)
require.Equal(t, "key", vertexClient.createdReq.InstanceConfig.KeyField)
require.NotContains(t, string(vertexClient.createdPayloadForAssert(t)), "serviceAccount")
require.NotContains(t, string(vertexClient.createdPayloadForAssert(t)), "encryptionSpec")
require.NotContains(t, got.ProviderInputRef+got.ProviderOutputRef+got.ProviderJobName, "A clean product hero image")
require.NotContains(t, string(store.uploadedJSONL), "private_key")
}
func TestVertexProvider_GetMapsStates(t *testing.T) {
tests := []struct {
name string
state string
err *VertexBatchJobError
wantState BatchProviderInternalState
wantDone bool
wantCode string
}{
{name: "pending", state: "JOB_STATE_PENDING", wantState: BatchProviderStateQueued},
{name: "queued", state: "JOB_STATE_QUEUED", wantState: BatchProviderStateQueued},
{name: "running", state: "JOB_STATE_RUNNING", wantState: BatchProviderStateRunning},
{name: "succeeded", state: "JOB_STATE_SUCCEEDED", wantState: BatchProviderStateSucceeded, wantDone: true},
{name: "failed", state: "JOB_STATE_FAILED", err: &VertexBatchJobError{Status: "INVALID_ARGUMENT", Message: "bad request"}, wantState: BatchProviderStateFailed, wantDone: true, wantCode: "INVALID_ARGUMENT"},
{name: "cancelled", state: "JOB_STATE_CANCELLED", wantState: BatchProviderStateCancelled, wantDone: true, wantCode: "VERTEX_BATCH_CANCELLED"},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
output := "gs://managed-bucket/batch-image/test/imgbatch_abc123/output/"
provider := newTestVertexProvider(&fakeVertexBatchClient{got: &VertexBatchPredictionJob{
Name: "projects/proj/locations/global/batchPredictionJobs/job-1",
State: tt.state,
Error: tt.err,
OutputConfig: VertexBatchOutputConfig{GCSDestination: VertexBatchGCSDestination{OutputURIPrefix: output}},
}}, &fakeVertexObjectStore{})
got, err := provider.Get(context.Background(), vertexJobWithName("projects/proj/locations/global/batchPredictionJobs/job-1"), vertexServiceAccount())
require.NoError(t, err)
require.Equal(t, tt.wantState, got.InternalState)
require.Equal(t, tt.wantDone, got.Done)
require.Equal(t, output, got.ProviderOutputRef)
require.Equal(t, tt.wantCode, got.ErrorCode)
})
}
}
func TestVertexProvider_OpenResultReturnsCombinedJSONLStream(t *testing.T) {
output := "gs://managed-bucket/batch-image/test/imgbatch_abc123/output/"
store := &fakeVertexObjectStore{
listed: []string{
output + "predictions_2.jsonl",
output + "predictions_1.jsonl",
},
objects: map[string]string{
output + "predictions_1.jsonl": `{"key":"1"}` + "\n",
output + "predictions_2.jsonl": `{"key":"2"}` + "\n",
},
}
provider := newTestVertexProvider(&fakeVertexBatchClient{}, store)
r, contentType, err := provider.OpenResult(context.Background(), &BatchImageJob{ProviderOutputRef: &output}, vertexServiceAccount())
require.NoError(t, err)
defer r.Close()
body, err := io.ReadAll(r)
require.NoError(t, err)
require.Equal(t, "application/jsonl", contentType)
require.Equal(t, "{\"key\":\"1\"}\n\n{\"key\":\"2\"}\n", string(body))
}
func TestVertexProvider_OpenResultMissingObjectsReturnsTypedError(t *testing.T) {
output := "gs://managed-bucket/batch-image/test/imgbatch_abc123/output/"
provider := newTestVertexProvider(&fakeVertexBatchClient{}, &fakeVertexObjectStore{})
_, _, err := provider.OpenResult(context.Background(), &BatchImageJob{ProviderOutputRef: &output}, vertexServiceAccount())
require.Error(t, err)
require.Equal(t, "VERTEX_RESULT_OBJECTS_MISSING", infraerrors.Reason(err))
}
func TestVertexProvider_CancelCallsClient(t *testing.T) {
vertexClient := &fakeVertexBatchClient{}
provider := newTestVertexProvider(vertexClient, &fakeVertexObjectStore{})
err := provider.Cancel(context.Background(), vertexJobWithName("projects/proj/locations/global/batchPredictionJobs/job-1"), vertexServiceAccount())
require.NoError(t, err)
require.Equal(t, "projects/proj/locations/global/batchPredictionJobs/job-1", vertexClient.cancelledName)
}
func TestVertexProvider_CleanupDeletesOnlyManagedPaths(t *testing.T) {
input := "gs://managed-bucket/batch-image/test/imgbatch_abc123/input/requests.jsonl"
output := "gs://managed-bucket/batch-image/test/imgbatch_abc123/output/"
store := &fakeVertexObjectStore{}
provider := newTestVertexProvider(&fakeVertexBatchClient{}, store)
err := provider.Cleanup(context.Background(), &BatchImageJob{BatchID: "imgbatch_abc123", ProviderInputRef: &input, ProviderOutputRef: &output}, vertexServiceAccount(), CleanupTargetAll)
require.NoError(t, err)
require.Equal(t, []string{input}, store.deletedObjects)
require.Equal(t, []string{output}, store.deletedPrefixes)
}
func TestVertexProvider_CleanupRejectsUnsafePath(t *testing.T) {
input := "gs://other-bucket/batch-image/test/imgbatch_abc123/input/requests.jsonl"
provider := newTestVertexProvider(&fakeVertexBatchClient{}, &fakeVertexObjectStore{})
err := provider.Cleanup(context.Background(), &BatchImageJob{BatchID: "imgbatch_abc123", ProviderInputRef: &input}, vertexServiceAccount(), CleanupTargetInput)
require.ErrorIs(t, err, ErrBatchImageProviderUnsafeCleanupPath)
}
func TestVertexProvider_ErrorsDoNotExposeServiceAccountSecrets(t *testing.T) {
privateKey := "-----BEGIN PRIVATE KEY-----secret-----END PRIVATE KEY-----"
account := vertexServiceAccount()
account.Credentials["service_account_json"] = map[string]any{
"type": "service_account",
"project_id": "proj",
"private_key": privateKey,
"client_email": "svc@proj.iam.gserviceaccount.com",
}
provider := newTestVertexProvider(&fakeVertexBatchClient{createErr: &VertexAPIError{StatusCode: 403, Message: "do not expose " + privateKey}}, &fakeVertexObjectStore{})
_, err := provider.Submit(context.Background(), nil, account, validVertexBatchInput())
require.Error(t, err)
require.Equal(t, "VERTEX_PERMISSION_DENIED", infraerrors.Reason(err))
require.NotContains(t, err.Error(), privateKey)
require.NotContains(t, err.Error(), "svc@proj")
}
func TestVertexProvider_MetadataDoesNotStoreImageBytesOrBase64(t *testing.T) {
vertexClient := &fakeVertexBatchClient{created: &VertexBatchPredictionJob{Name: "projects/proj/locations/global/batchPredictionJobs/job-1", State: "JOB_STATE_PENDING"}}
provider := newTestVertexProvider(vertexClient, &fakeVertexObjectStore{})
got, err := provider.Submit(context.Background(), nil, vertexServiceAccount(), validVertexBatchInput())
require.NoError(t, err)
metadata := got.ProviderJobName + got.ProviderInputRef + got.ProviderOutputRef
require.NotContains(t, metadata, "iVBOR")
require.NotContains(t, metadata, "base64")
require.NotContains(t, metadata, "A clean product hero image")
}
func validVertexBatchInput() BatchImageInput {
return BatchImageInput{
BatchID: "imgbatch_abc123",
Model: "gemini-3.1-flash-image",
DisplayName: "test vertex batch",
Items: []BatchImageInputItem{{
CustomID: "cover_001",
Prompt: "A clean product hero image",
}},
}
}
func requireVertexJSONLLine(t *testing.T, line, wantKey, wantPrompt string) {
t.Helper()
var got map[string]any
require.NoError(t, json.Unmarshal([]byte(line), &got))
require.Equal(t, wantKey, got["key"])
request := got["request"].(map[string]any)
contents := request["contents"].([]any)
require.Equal(t, "user", contents[0].(map[string]any)["role"])
parts := contents[0].(map[string]any)["parts"].([]any)
require.Equal(t, wantPrompt, parts[0].(map[string]any)["text"])
config := request["generationConfig"].(map[string]any)
require.Equal(t, []any{"TEXT", "IMAGE"}, config["responseModalities"])
}
func newTestVertexProvider(client *fakeVertexBatchClient, store *fakeVertexObjectStore) *VertexBatchImageProvider {
return NewVertexBatchImageProvider(VertexBatchImageProviderOptions{
ProjectID: "proj",
Location: "global",
ManagedGCSBucket: "managed-bucket",
ManagedGCSPrefix: "batch-image/{env}/{batch_id}",
Environment: "test",
}, client, store, &fakeGeminiTokenCache{token: "ya29.test-token"})
}
func vertexServiceAccount() *Account {
return &Account{
Platform: PlatformGemini,
Type: AccountTypeServiceAccount,
Credentials: map[string]any{
"service_account_json": map[string]any{
"type": "service_account",
"project_id": "proj",
"private_key": "-----BEGIN PRIVATE KEY-----\nabc\n-----END PRIVATE KEY-----\n",
"client_email": "svc@proj.iam.gserviceaccount.com",
},
},
}
}
func vertexJobWithName(name string) *BatchImageJob {
return &BatchImageJob{ProviderJobName: &name}
}
type fakeVertexBatchClient struct {
created *VertexBatchPredictionJob
got *VertexBatchPredictionJob
createErr error
getErr error
cancelErr error
createdReq VertexCreateBatchPredictionJobRequest
cancelledName string
}
func (f *fakeVertexBatchClient) CreateBatchPredictionJob(_ context.Context, accessToken string, req VertexCreateBatchPredictionJobRequest) (*VertexBatchPredictionJob, error) {
if strings.TrimSpace(accessToken) == "" {
return nil, errors.New("missing token")
}
f.createdReq = req
if f.createErr != nil {
return nil, f.createErr
}
if f.created != nil {
return f.created, nil
}
return &VertexBatchPredictionJob{Name: "projects/proj/locations/global/batchPredictionJobs/job-1", State: "JOB_STATE_PENDING"}, nil
}
func (f *fakeVertexBatchClient) GetBatchPredictionJob(_ context.Context, _ string, _ string) (*VertexBatchPredictionJob, error) {
if f.getErr != nil {
return nil, f.getErr
}
return f.got, nil
}
func (f *fakeVertexBatchClient) CancelBatchPredictionJob(_ context.Context, _ string, name string) error {
f.cancelledName = name
return f.cancelErr
}
func (f *fakeVertexBatchClient) createdPayloadForAssert(t *testing.T) []byte {
t.Helper()
b, err := json.Marshal(f.createdReq)
require.NoError(t, err)
return b
}
type fakeVertexObjectStore struct {
uploadURI string
uploadedJSONL []byte
uploadErr error
listed []string
objects map[string]string
listErr error
openErr error
deleteErr error
deletedObjects []string
deletedPrefixes []string
}
func (f *fakeVertexObjectStore) UploadJSONL(_ context.Context, _ string, uri string, r io.Reader) error {
f.uploadURI = uri
f.uploadedJSONL, _ = io.ReadAll(r)
return f.uploadErr
}
func (f *fakeVertexObjectStore) ListJSONLObjects(_ context.Context, _ string, _ string) ([]string, error) {
if f.listErr != nil {
return nil, f.listErr
}
out := make([]string, 0, len(f.listed))
for _, item := range f.listed {
if strings.HasSuffix(item, ".jsonl") {
out = append(out, item)
}
}
return out, nil
}
func (f *fakeVertexObjectStore) OpenObject(_ context.Context, _ string, uri string) (io.ReadCloser, string, error) {
if f.openErr != nil {
return nil, "", f.openErr
}
return io.NopCloser(bytes.NewBufferString(f.objects[uri])), "application/jsonl", nil
}
func (f *fakeVertexObjectStore) DeleteObject(_ context.Context, _ string, uri string) error {
f.deletedObjects = append(f.deletedObjects, uri)
return f.deleteErr
}
func (f *fakeVertexObjectStore) DeletePrefix(_ context.Context, _ string, uri string) error {
f.deletedPrefixes = append(f.deletedPrefixes, uri)
return f.deleteErr
}
type fakeGeminiTokenCache struct {
token string
}
func (f *fakeGeminiTokenCache) GetAccessToken(context.Context, string) (string, error) {
if strings.TrimSpace(f.token) == "" {
return "", errors.New("missing token")
}
return f.token, nil
}
func (f *fakeGeminiTokenCache) SetAccessToken(context.Context, string, string, time.Duration) error {
return nil
}
func (f *fakeGeminiTokenCache) DeleteAccessToken(context.Context, string) error {
return nil
}
func (f *fakeGeminiTokenCache) AcquireRefreshLock(context.Context, string, time.Duration) (bool, error) {
return false, nil
}
func (f *fakeGeminiTokenCache) ReleaseRefreshLock(context.Context, string) error {
return nil
}
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,961 @@
//go:build unit
package service
import (
"context"
"encoding/json"
"errors"
"io"
"strings"
"testing"
"time"
"github.com/Wei-Shaw/sub2api/internal/config"
infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors"
"github.com/stretchr/testify/require"
)
func TestBatchImagePublicService_Submit(t *testing.T) {
ctx := context.Background()
t.Run("rejects when disabled", func(t *testing.T) {
svc, _, _, _, _ := newTestBatchImagePublicService(false)
_, err := svc.Submit(ctx, testBatchImageOwner(), validBatchImageSubmitRequest(), "")
require.ErrorIs(t, err, ErrBatchImageDisabled)
})
t.Run("accepts valid request stores refs and enqueues once", func(t *testing.T) {
svc, repo, queue, gemini, _ := newTestBatchImagePublicService(true)
got, err := svc.Submit(ctx, testBatchImageOwner(), validBatchImageSubmitRequest(), "")
require.NoError(t, err)
require.Equal(t, "image.batch", got.Object)
require.Equal(t, "queued", got.Status)
require.Equal(t, BatchImageProviderGeminiAPI, got.Provider)
require.Equal(t, 2, got.ItemCount)
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)
require.Equal(t, "providers/gemini_api/job", batchImageDerefString(job.ProviderJobName))
require.Equal(t, "files/gemini_api/input", batchImageDerefString(job.ProviderInputRef))
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,
Platform: PlatformGemini,
RateMultiplier: 2.0,
AllowImageGeneration: true,
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,
Platform: PlatformGemini,
RateMultiplier: 1.0,
AllowImageGeneration: true,
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,
Platform: PlatformGemini,
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) {
svc, _, _, gemini, _ := newTestBatchImagePublicService(true)
req := validBatchImageSubmitRequest()
req.Items[0].CustomID = ""
req.Items[1].CustomID = ""
_, err := svc.Submit(ctx, testBatchImageOwner(), req, "")
require.NoError(t, err)
require.Len(t, gemini.submits, 1)
require.Equal(t, "item_000001", gemini.submits[0].Items[0].CustomID)
require.Equal(t, "item_000002", gemini.submits[0].Items[1].CustomID)
})
t.Run("expands output count into separate billable items", func(t *testing.T) {
svc, repo, _, gemini, _ := newTestBatchImagePublicService(true)
req := validBatchImageSubmitRequest()
req.Items = []BatchImageSubmitItem{
{CustomID: "cover", Prompt: "hero", OutputCount: 3, ReferenceImages: []BatchImageReferenceInput{{MimeType: "image/png", Data: []byte("ref")}}},
}
got, err := svc.Submit(ctx, testBatchImageOwner(), req, "")
require.NoError(t, err)
require.Equal(t, 3, got.ItemCount)
require.InDelta(t, 0.375, got.EstimatedCost, 1e-12)
require.Len(t, gemini.submits, 1)
require.Len(t, gemini.submits[0].Items, 3)
require.Equal(t, []string{"cover_01", "cover_02", "cover_03"}, []string{
gemini.submits[0].Items[0].CustomID,
gemini.submits[0].Items[1].CustomID,
gemini.submits[0].Items[2].CustomID,
})
require.Len(t, gemini.submits[0].Items[0].ReferenceImages, 1)
require.Len(t, repo.items[got.ID], 3)
})
t.Run("validates request fields", func(t *testing.T) {
tests := []struct {
name string
mutate func(*BatchImageSubmitRequest)
want error
}{
{name: "missing_model", mutate: func(r *BatchImageSubmitRequest) { r.Model = "" }, want: ErrBatchImageInvalidModel},
{name: "empty_items", mutate: func(r *BatchImageSubmitRequest) { r.Items = nil }, want: ErrBatchImageInvalidItems},
{name: "duplicate_custom_ids", mutate: func(r *BatchImageSubmitRequest) { r.Items[1].CustomID = r.Items[0].CustomID }, want: ErrBatchImageDuplicateCustomIDInRequest},
{name: "empty_prompt", mutate: func(r *BatchImageSubmitRequest) { r.Items[0].Prompt = " " }, want: ErrBatchImageInvalidItems},
{name: "prompt_too_long", mutate: func(r *BatchImageSubmitRequest) { r.Items[0].Prompt = strings.Repeat("x", 9) }, want: ErrBatchImagePromptTooLong},
{name: "unsupported_provider", mutate: func(r *BatchImageSubmitRequest) { r.Provider = "other" }, want: ErrBatchImageUnsupportedProvider},
{name: "vertex_rejects_2k", mutate: func(r *BatchImageSubmitRequest) { r.Provider = BatchImageProviderVertex; r.ImageSize = "2K" }, want: ErrBatchImageInvalidItems},
{name: "too_many_outputs_per_item", mutate: func(r *BatchImageSubmitRequest) {
r.Items[0].OutputCount = 5
}, want: ErrBatchImageInvalidItems},
{name: "too_many_reference_images_for_flash", mutate: func(r *BatchImageSubmitRequest) {
r.Model = "gemini-2.5-flash-image"
r.Items[0].ReferenceImages = []BatchImageReferenceInput{
{MimeType: "image/png", Data: []byte("1")},
{MimeType: "image/png", Data: []byte("2")},
{MimeType: "image/png", Data: []byte("3")},
{MimeType: "image/png", Data: []byte("4")},
}
}, want: ErrBatchImageTooManyReferenceImages},
{name: "bad_reference_mime", mutate: func(r *BatchImageSubmitRequest) {
r.Items[0].ReferenceImages = []BatchImageReferenceInput{{MimeType: "application/octet-stream", Data: []byte("x")}}
}, want: ErrBatchImageInvalidReferenceImage},
{name: "reference_requires_data_or_file_uri", mutate: func(r *BatchImageSubmitRequest) {
r.Items[0].ReferenceImages = []BatchImageReferenceInput{{MimeType: "image/png"}}
}, want: ErrBatchImageInvalidReferenceImage},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
svc, _, _, _, _ := newTestBatchImagePublicService(true)
req := validBatchImageSubmitRequest()
tt.mutate(&req)
_, err := svc.Submit(ctx, testBatchImageOwner(), req, "")
require.ErrorIs(t, err, tt.want)
})
}
})
t.Run("rejects too many items", func(t *testing.T) {
svc, _, _, _, _ := newTestBatchImagePublicService(true)
req := validBatchImageSubmitRequest()
req.Items = append(req.Items, BatchImageSubmitItem{CustomID: "too_many", Prompt: "x"})
_, err := svc.Submit(ctx, testBatchImageOwner(), req, "")
require.ErrorIs(t, err, ErrBatchImageInvalidItems)
})
t.Run("rejects too many output images", func(t *testing.T) {
svc, _, _, _, _ := newTestBatchImagePublicService(true)
svc.Config.BatchImage.MaxOutputImagesPerJob = 3
req := validBatchImageSubmitRequest()
req.Items[0].OutputCount = 2
req.Items[1].OutputCount = 2
_, err := svc.Submit(ctx, testBatchImageOwner(), req, "")
require.ErrorIs(t, err, ErrBatchImageTooManyOutputImages)
})
t.Run("rejects too many reference images across request", func(t *testing.T) {
svc, _, _, _, _ := newTestBatchImagePublicService(true)
svc.Config.BatchImage.MaxReferenceImagesPerJob = 3
req := validBatchImageSubmitRequest()
req.Model = "gemini-2.5-flash-image"
req.Items[0].ReferenceImages = []BatchImageReferenceInput{
{MimeType: "image/png", Data: []byte("1")},
{MimeType: "image/png", Data: []byte("2")},
}
req.Items[1].ReferenceImages = []BatchImageReferenceInput{
{MimeType: "image/png", Data: []byte("3")},
{MimeType: "image/png", Data: []byte("4")},
}
_, err := svc.Submit(ctx, testBatchImageOwner(), req, "")
require.ErrorIs(t, err, ErrBatchImageTooManyReferenceImages)
})
t.Run("rejects too much inline reference image data across request", func(t *testing.T) {
svc, _, _, _, _ := newTestBatchImagePublicService(true)
svc.Config.BatchImage.MaxReferenceImagesPerJob = 10
svc.Config.BatchImage.MaxReferenceInlineBytesPerJob = 4
req := validBatchImageSubmitRequest()
req.Model = "gemini-2.5-flash-image"
req.Items[0].ReferenceImages = []BatchImageReferenceInput{{MimeType: "image/png", Data: []byte("123")}}
req.Items[1].ReferenceImages = []BatchImageReferenceInput{{MimeType: "image/png", Data: []byte("456")}}
_, err := svc.Submit(ctx, testBatchImageOwner(), req, "")
require.ErrorIs(t, err, ErrBatchImageReferenceImagesTooLarge)
})
t.Run("selects requested provider", func(t *testing.T) {
svc, _, _, gemini, vertex := newTestBatchImagePublicService(true)
req := validBatchImageSubmitRequest()
req.Provider = BatchImageProviderVertex
got, err := svc.Submit(ctx, testBatchImageOwner(), req, "")
require.NoError(t, err)
require.Equal(t, BatchImageProviderVertex, got.Provider)
require.Empty(t, gemini.submits)
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)
require.Equal(t, "QUEUE_FAILED", batchImageDerefString(job.LastErrorCode))
require.Contains(t, repo.events[job.BatchID], "queue_failed")
}
})
t.Run("idempotency returns same batch without provider resubmit", func(t *testing.T) {
svc, _, queue, gemini, _ := newTestBatchImagePublicService(true)
req := validBatchImageSubmitRequest()
first, err := svc.Submit(ctx, testBatchImageOwner(), req, "client-key")
require.NoError(t, err)
second, err := svc.Submit(ctx, testBatchImageOwner(), req, "client-key")
require.NoError(t, err)
require.Equal(t, first.ID, second.ID)
require.Len(t, gemini.submits, 1)
require.Equal(t, []string{first.ID}, queue.enqueued)
})
t.Run("idempotency conflict rejects changed request", func(t *testing.T) {
svc, _, _, _, _ := newTestBatchImagePublicService(true)
req := validBatchImageSubmitRequest()
first, err := svc.Submit(ctx, testBatchImageOwner(), req, "client-key")
require.NoError(t, err)
req.Items[0].Prompt = "diff"
second, err := svc.Submit(ctx, testBatchImageOwner(), req, "client-key")
require.Nil(t, second)
require.ErrorIs(t, err, ErrBatchImageIdempotencyConflict)
require.NotEmpty(t, first.ID)
})
t.Run("public response does not expose internals", func(t *testing.T) {
svc, _, _, _, _ := newTestBatchImagePublicService(true)
got, err := svc.Submit(ctx, testBatchImageOwner(), validBatchImageSubmitRequest(), "")
require.NoError(t, err)
body, err := json.Marshal(got)
require.NoError(t, err)
requireBatchImagePublicJSONHasNoInternals(t, string(body))
})
}
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,
Platform: PlatformGemini,
RateMultiplier: 1,
AllowImageGeneration: true,
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()
t.Run("status is owner scoped and maps public status", func(t *testing.T) {
svc, repo, _, _, _ := newTestBatchImagePublicService(true)
apiKeyID := int64(22)
accountID := int64(101)
repo.jobs["imgbatch_status"] = &BatchImageJob{
BatchID: "imgbatch_status",
UserID: 11,
APIKeyID: &apiKeyID,
AccountID: &accountID,
Provider: BatchImageProviderGeminiAPI,
Model: "gemini-2.5-flash-image",
Status: BatchImageJobStatusIndexing,
ProviderJobName: batchImageStringPtr("providers/internal/job"),
CreatedAt: time.Now(),
}
got, err := svc.Get(ctx, testBatchImageOwner(), "imgbatch_status")
require.NoError(t, err)
require.Equal(t, "processing_results", got.Status)
body, err := json.Marshal(got)
require.NoError(t, err)
requireBatchImagePublicJSONHasNoInternals(t, string(body))
_, err = svc.Get(ctx, BatchImageOwner{UserID: 11, APIKeyID: 999}, "imgbatch_status")
require.ErrorIs(t, err, ErrBatchImageJobNotFound)
})
t.Run("items are filtered paginated and sanitized", func(t *testing.T) {
svc, repo, _, _, _ := newTestBatchImagePublicService(true)
apiKeyID := int64(22)
repo.jobs["imgbatch_items"] = &BatchImageJob{
BatchID: "imgbatch_items",
UserID: 11,
APIKeyID: &apiKeyID,
Provider: BatchImageProviderGeminiAPI,
Model: "gemini-2.5-flash-image",
Status: BatchImageJobStatusCompleted,
CreatedAt: time.Now(),
}
sourceObject := "gs://bucket/internal/output.jsonl"
mime := "image/png"
ext := "png"
code := "SAFETY_BLOCKED"
msg := "blocked in gs://bucket/internal/output.jsonl"
repo.items["imgbatch_items"] = []CreateBatchImageItemParams{
{JobID: "imgbatch_items", CustomID: "ok_1", Status: BatchImageItemStatusSuccess, ProviderSourceObject: &sourceObject, MimeType: &mime, FileExtension: &ext, ImageCount: 1},
{JobID: "imgbatch_items", CustomID: "bad_1", Status: BatchImageItemStatusFailed, ProviderSourceObject: &sourceObject, ErrorCode: &code, ErrorMessage: &msg},
{JobID: "imgbatch_items", CustomID: "ok_2", Status: BatchImageItemStatusSuccess, MimeType: &mime, FileExtension: &ext, ImageCount: 1},
}
page, err := svc.ListItems(ctx, testBatchImageOwner(), "imgbatch_items", BatchImageItemsQuery{Limit: 1})
require.NoError(t, err)
require.True(t, page.HasMore)
require.Len(t, page.Data, 1)
require.Equal(t, "ok_1", page.Data[0].CustomID)
filtered, err := svc.ListItems(ctx, testBatchImageOwner(), "imgbatch_items", BatchImageItemsQuery{Status: "failed", Limit: 100})
require.NoError(t, err)
require.False(t, filtered.HasMore)
require.Len(t, filtered.Data, 1)
require.Equal(t, "failed", filtered.Data[0].Status)
require.NotNil(t, filtered.Data[0].Error)
require.Equal(t, "upstream provider operation failed", filtered.Data[0].Error.Message)
body, err := json.Marshal(filtered)
require.NoError(t, err)
requireBatchImagePublicJSONHasNoInternals(t, string(body))
require.NotContains(t, string(body), "download_url")
_, err = svc.ListItems(ctx, BatchImageOwner{UserID: 12, APIKeyID: 22}, "imgbatch_items", BatchImageItemsQuery{})
require.ErrorIs(t, err, ErrBatchImageJobNotFound)
})
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,
APIKeyID: &apiKeyID,
AccountID: &accountID,
Provider: BatchImageProviderGeminiAPI,
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, "queued", got.Status)
require.Equal(t, 1, gemini.cancelCount)
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) {
svc, repo, _, gemini, _ := newTestBatchImagePublicService(true)
apiKeyID := int64(22)
repo.jobs["imgbatch_done"] = &BatchImageJob{
BatchID: "imgbatch_done",
UserID: 11,
APIKeyID: &apiKeyID,
Provider: BatchImageProviderGeminiAPI,
Model: "gemini-2.5-flash-image",
Status: BatchImageJobStatusCompleted,
CreatedAt: time.Now(),
}
got, err := svc.Cancel(ctx, testBatchImageOwner(), "imgbatch_done")
require.NoError(t, err)
require.Equal(t, "completed", got.Status)
require.Zero(t, gemini.cancelCount)
})
t.Run("cancel hides provider raw errors behind public error", func(t *testing.T) {
svc, repo, _, gemini, _ := newTestBatchImagePublicService(true)
gemini.cancelErr = errors.New("projects/secret-provider-job not found")
apiKeyID := int64(22)
accountID := int64(101)
repo.jobs["imgbatch_cancel_error"] = &BatchImageJob{
BatchID: "imgbatch_cancel_error",
UserID: 11,
APIKeyID: &apiKeyID,
AccountID: &accountID,
Provider: BatchImageProviderGeminiAPI,
Model: "gemini-2.5-flash-image",
Status: BatchImageJobStatusSubmitted,
ProviderJobName: batchImageStringPtr("providers/internal/job"),
CreatedAt: time.Now(),
}
_, err := svc.Cancel(ctx, testBatchImageOwner(), "imgbatch_cancel_error")
require.ErrorIs(t, err, ErrBatchImageCancelFailed)
require.Equal(t, "BATCH_IMAGE_CANCEL_FAILED", infraerrors.Reason(err))
require.NotContains(t, infraerrors.Message(err), "projects/")
})
}
func newTestBatchImagePublicService(enabled bool) (*BatchImagePublicService, *fakeBatchImageRepository, *publicBatchImageQueue, *publicBatchImageProvider, *publicBatchImageProvider) {
repo := newFakeBatchImageRepository()
queue := &publicBatchImageQueue{}
gemini := &publicBatchImageProvider{name: BatchImageProviderGeminiAPI}
vertex := &publicBatchImageProvider{name: BatchImageProviderVertex}
svc := &BatchImagePublicService{
Repo: repo,
AccountRepo: &publicBatchImageAccountRepo{accounts: []Account{testBatchImageAccount(101, AccountTypeAPIKey), testBatchImageAccount(202, AccountTypeServiceAccount)}},
Queue: queue,
ProviderRegistry: NewBatchImageProviderRegistry(
gemini,
vertex,
),
Pricing: &fakeBatchImagePricingResolver{unitPrice: 0.25},
BillingRepo: &fakeBatchImageBillingRepo{},
AuthCache: &fakeBatchImageAuthCacheInvalidator{},
Config: &config.Config{BatchImage: config.BatchImageConfig{
Enabled: enabled,
MaxItemsPerJobDefault: 2,
MaxPromptCharsPerItem: 8,
DefaultResponseMimeType: "image/png",
DefaultImageSize: "1K",
}},
}
return svc, repo, queue, gemini, vertex
}
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",
Provider: BatchImageProviderGeminiAPI,
ResponseMimeType: "image/png",
AspectRatio: "1:1",
ImageSize: "1K",
Metadata: map[string]string{"project": "campaign-a", "secret": strings.Repeat("x", 300)},
Items: []BatchImageSubmitItem{
{CustomID: "cover_001", Prompt: "hero"},
{CustomID: "cover_002", Prompt: "clean"},
},
}
}
func testBatchImageAccount(id int64, accountType string) Account {
return Account{
ID: id,
Platform: PlatformGemini,
Type: accountType,
Status: StatusActive,
Schedulable: true,
Priority: int(id),
Credentials: map[string]any{"api_key": "test-secret"},
Concurrency: 1,
RateLimitedAt: nil,
}
}
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{
"provider_job_name",
"provider_input_ref",
"provider_output_ref",
"gcs_input_uri",
"gcs_output_uri",
"account_id",
"service_account",
"api_key",
"download_url",
"providers/",
"files/",
"gs://",
} {
require.NotContains(t, body, forbidden)
}
}
type publicBatchImageAccountRepo struct {
accounts []Account
}
func (r *publicBatchImageAccountRepo) GetByID(_ context.Context, id int64) (*Account, error) {
for i := range r.accounts {
if r.accounts[i].ID == id {
return &r.accounts[i], nil
}
}
return nil, errors.New("account not found")
}
func (r *publicBatchImageAccountRepo) ListSchedulableByPlatform(_ context.Context, platform string) ([]Account, error) {
out := make([]Account, 0, len(r.accounts))
for _, account := range r.accounts {
if account.Platform == platform {
out = append(out, account)
}
}
return out, nil
}
func (r *publicBatchImageAccountRepo) ListSchedulableByGroupIDAndPlatform(ctx context.Context, _ int64, platform string) ([]Account, error) {
return r.ListSchedulableByPlatform(ctx, platform)
}
type publicBatchImageQueue struct {
enqueued []string
err error
}
func (q *publicBatchImageQueue) Enqueue(_ context.Context, batchID string) error {
if q.err != nil {
return q.err
}
for _, existing := range q.enqueued {
if existing == batchID {
return ErrBatchImageAlreadyQueued
}
}
q.enqueued = append(q.enqueued, batchID)
return nil
}
func (q *publicBatchImageQueue) Reserve(context.Context, time.Duration) (ReservedBatchImageJob, error) {
return ReservedBatchImageJob{}, ErrBatchImageQueueEmpty
}
func (q *publicBatchImageQueue) RequeueAfter(context.Context, string, time.Duration) error {
return nil
}
func (q *publicBatchImageQueue) Ack(context.Context, string) error {
return nil
}
func (q *publicBatchImageQueue) Heartbeat(context.Context, string) error {
return nil
}
func (q *publicBatchImageQueue) MoveDueDelayedToReady(context.Context, int) (int, error) {
return 0, nil
}
func (q *publicBatchImageQueue) RecoverStaleActive(context.Context, time.Duration, int) (int, error) {
return 0, nil
}
func (q *publicBatchImageQueue) TryAcquireJobLock(context.Context, string, time.Duration) (BatchImageJobLock, bool, error) {
return nil, false, nil
}
type publicBatchImageProvider struct {
name string
submits []BatchImageInput
submitErr error
cancelCount int
cancelErr error
result string
cleanupTargets []CleanupTarget
cleanupErr error
}
func (p *publicBatchImageProvider) Name() string { return p.name }
func (p *publicBatchImageProvider) SupportsAccount(*Account) bool { return true }
func (p *publicBatchImageProvider) Submit(_ context.Context, _ *BatchImageJob, _ *Account, input BatchImageInput) (*BatchProviderJob, error) {
p.submits = append(p.submits, input)
if p.submitErr != nil {
return nil, p.submitErr
}
return &BatchProviderJob{
ProviderJobName: "providers/" + p.name + "/job",
ProviderInputRef: "files/" + p.name + "/input",
ProviderOutputRef: "files/" + p.name + "/output",
}, nil
}
func (p *publicBatchImageProvider) Get(context.Context, *BatchImageJob, *Account) (*BatchProviderStatus, error) {
return &BatchProviderStatus{InternalState: BatchProviderStateQueued}, nil
}
func (p *publicBatchImageProvider) Cancel(context.Context, *BatchImageJob, *Account) error {
p.cancelCount++
return p.cancelErr
}
func (p *publicBatchImageProvider) OpenResult(context.Context, *BatchImageJob, *Account) (io.ReadCloser, string, error) {
return io.NopCloser(strings.NewReader(p.result)), "application/jsonl", nil
}
func (p *publicBatchImageProvider) Cleanup(_ context.Context, _ *BatchImageJob, _ *Account, target CleanupTarget) error {
p.cleanupTargets = append(p.cleanupTargets, target)
return p.cleanupErr
}
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)
@@ -0,0 +1,64 @@
package service
import (
"context"
"net/http"
"strings"
"time"
infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors"
)
var (
ErrBatchImageQueueEmpty = infraerrors.New(http.StatusNotFound, "BATCH_IMAGE_QUEUE_EMPTY", "batch image queue is empty")
ErrBatchImageAlreadyQueued = infraerrors.New(http.StatusConflict, "BATCH_IMAGE_ALREADY_QUEUED", "batch image job is already queued")
ErrBatchImageLockNotAcquired = infraerrors.New(http.StatusConflict, "BATCH_IMAGE_LOCK_NOT_ACQUIRED", "batch image job lock was not acquired")
ErrInvalidBatchImageQueuePayload = infraerrors.New(http.StatusBadRequest, "BATCH_IMAGE_QUEUE_INVALID_PAYLOAD", "invalid batch image queue payload")
)
type ReservedBatchImageJob struct {
BatchID string
}
type BatchImageJobLock interface {
Release(ctx context.Context) error
}
type BatchImageQueue interface {
Enqueue(ctx context.Context, batchID string) error
Reserve(ctx context.Context, blockTimeout time.Duration) (ReservedBatchImageJob, error)
RequeueAfter(ctx context.Context, batchID string, delay time.Duration) error
Ack(ctx context.Context, batchID string) error
Heartbeat(ctx context.Context, batchID string) error
MoveDueDelayedToReady(ctx context.Context, limit int) (int, error)
RecoverStaleActive(ctx context.Context, staleAfter time.Duration, limit int) (int, error)
TryAcquireJobLock(ctx context.Context, batchID string, ttl time.Duration) (BatchImageJobLock, bool, error)
}
type BatchImageService struct {
repo BatchImageRepository
queue BatchImageQueue
}
func NewBatchImageService(repo BatchImageRepository, queue BatchImageQueue) *BatchImageService {
return &BatchImageService{repo: repo, queue: queue}
}
func (s *BatchImageService) EnqueueBatchImageJob(ctx context.Context, batchID string) error {
if !IsValidBatchImageID(batchID) {
return ErrInvalidBatchImageQueuePayload
}
if s == nil || s.queue == nil {
return infraerrors.New(http.StatusInternalServerError, "BATCH_IMAGE_QUEUE_NOT_CONFIGURED", "batch image queue is not configured")
}
if s.repo != nil {
if _, err := s.repo.GetBatchImageJobByBatchID(ctx, batchID); err != nil {
return err
}
}
return s.queue.Enqueue(ctx, batchID)
}
func IsValidBatchImageID(batchID string) bool {
return strings.HasPrefix(batchID, "imgbatch_") && len(batchID) > len("imgbatch_")
}
@@ -0,0 +1,332 @@
package service
import (
"context"
"crypto/sha256"
"encoding/hex"
"errors"
"fmt"
"strconv"
"strings"
"time"
"github.com/Wei-Shaw/sub2api/internal/config"
)
const (
batchImageSettlementRequestPrefix = "batch_image_settlement:"
batchImageSettlementRetryDelay = time.Minute
batchImageSettlementMaxRetries = 5
batchImageCostEpsilon = 0.00000001
)
type BatchImagePricingResolver interface {
BatchImageUnitPrice(ctx context.Context, job *BatchImageJob) (float64, error)
}
type BatchImageModelPricingResolver struct {
Resolver *ModelPricingResolver
}
func (r *BatchImageModelPricingResolver) BatchImageUnitPrice(ctx context.Context, job *BatchImageJob) (float64, error) {
if r == nil || r.Resolver == nil || job == nil || strings.TrimSpace(job.Model) == "" {
return 0, ErrBatchImageSettlementPricingMissing
}
resolved := r.Resolver.Resolve(ctx, PricingInput{Model: job.Model})
if resolved == nil {
return 0, ErrBatchImageSettlementPricingMissing
}
switch resolved.Mode {
case BillingModeImage, BillingModePerRequest:
if resolved.DefaultPerRequestPrice > 0 {
return resolved.DefaultPerRequestPrice, nil
}
if len(resolved.RequestTiers) == 1 && resolved.RequestTiers[0].PerRequestPrice != nil && *resolved.RequestTiers[0].PerRequestPrice >= 0 {
return *resolved.RequestTiers[0].PerRequestPrice, nil
}
case BillingModeToken:
if resolved.BasePricing != nil && (resolved.BasePricing.ImageOutputPriceExplicit || resolved.BasePricing.ImageOutputPricePerToken > 0) {
return resolved.BasePricing.ImageOutputPricePerToken, nil
}
}
return 0, ErrBatchImageSettlementPricingMissing
}
type BatchImageSettlementService struct {
Repo BatchImageRepository
BillingRepo UsageBillingRepository
UsageLogRepo UsageLogRepository
Pricing BatchImagePricingResolver
AuthCache APIKeyAuthCacheInvalidator
Config *config.Config
}
type BatchImageSettlementResult struct {
BatchID string
SuccessCount int
FailCount int
ActualCost float64
ManifestHash string
RequestID string
AlreadySettled bool
}
func (s *BatchImageSettlementService) Settle(ctx context.Context, batchID string) (*BatchImageSettlementResult, error) {
if s == nil || s.Repo == nil || s.BillingRepo == nil || s.Pricing == nil {
return nil, ErrBatchImageSettlementBillingFailed.WithCause(errors.New("batch image settlement service is not configured"))
}
job, err := s.Repo.GetBatchImageJobByBatchID(ctx, batchID)
if err != nil {
return nil, err
}
manifestHash := BuildBatchImageSettlementManifestHash(job)
result := &BatchImageSettlementResult{
BatchID: job.BatchID,
SuccessCount: job.SuccessCount,
FailCount: job.FailCount,
ManifestHash: manifestHash,
RequestID: BatchImageCaptureRequestID(job.BatchID),
}
if job.ActualCost != nil {
result.ActualCost = *job.ActualCost
}
if job.Status == BatchImageJobStatusCompleted {
result.AlreadySettled = true
return result, nil
}
if job.Status != BatchImageJobStatusSettling {
return nil, ErrBatchImageSettlementInvalidStatus
}
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 {
return nil, ErrBatchImageSettlementManifestConflict
}
if job.APIKeyID == nil || *job.APIKeyID <= 0 {
return nil, ErrBatchImageSettlementMissingAPIKeyID
}
if job.AccountID == nil || *job.AccountID <= 0 {
return nil, ErrBatchImageSettlementMissingAccountID
}
if isBatchImageSettlementRetryExhausted(job) {
return nil, s.failExhaustedSettlement(ctx, job, manifestHash, "settlement billing retry limit reached")
}
unitPrice, err := s.settlementUnitPrice(ctx, job)
if err != nil {
return nil, err
}
if unitPrice < 0 {
return nil, ErrBatchImageSettlementPricingMissing
}
actualCost := float64(job.SuccessCount) * unitPrice
result.ActualCost = actualCost
holdAmount := job.EstimatedCost
if job.HoldAmount != nil {
holdAmount = *job.HoldAmount
}
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)
retryCount, recordErr := s.Repo.SetBatchImageJobSettlementFailed(ctx, job.BatchID, "SETTLEMENT_BILLING_FAILED", msg)
if recordErr == nil && retryCount >= batchImageSettlementMaxRetries {
job.RetryCount = retryCount
return nil, s.failExhaustedSettlement(ctx, job, manifestHash, msg)
}
return nil, err
}
s.invalidateAuthCache(ctx, job.UserID)
now := time.Now()
outputExpiresAt := now.Add(s.outputRetentionAfterTerminal())
if err := s.Repo.MarkBatchImageJobSettled(ctx, MarkBatchImageJobSettledParams{
BatchID: job.BatchID,
ActualCost: actualCost,
ManifestHash: manifestHash,
Now: &now,
OutputExpiresAt: &outputExpiresAt,
EventPayload: map[string]any{
"batch_id": job.BatchID,
"request_id": result.RequestID,
"success_count": job.SuccessCount,
"fail_count": job.FailCount,
"actual_cost": actualCost,
"manifest_hash": manifestHash,
},
}); err != nil {
return nil, err
}
s.recordUsageLog(ctx, job, actualCost, result.RequestID, now)
return result, nil
}
func isBatchImageSettlementRetryExhausted(job *BatchImageJob) bool {
return job != nil &&
job.Status == BatchImageJobStatusSettling &&
job.RetryCount >= batchImageSettlementMaxRetries &&
batchImageDerefString(job.LastErrorCode) == "SETTLEMENT_BILLING_FAILED"
}
func (s *BatchImageSettlementService) failExhaustedSettlement(ctx context.Context, job *BatchImageJob, manifestHash, message string) error {
if s == nil || s.Repo == nil {
return ErrBatchImageSettlementBillingFailed
}
if err := releaseBatchImageBalanceHold(ctx, s.BillingRepo, job, manifestHash); err != nil {
msg := truncateBatchImageMessage(err.Error(), batchImageMaxErrorMessageLength)
_, _ = s.Repo.SetBatchImageJobSettlementFailed(ctx, job.BatchID, "SETTLEMENT_RELEASE_FAILED", msg)
return ErrBatchImageSettlementBillingFailed.WithCause(err)
}
s.invalidateAuthCache(ctx, job.UserID)
msg := strings.TrimSpace(message)
if msg == "" {
msg = "settlement billing retry limit reached"
}
if err := s.Repo.TransitionBatchImageJobStatus(ctx, job.BatchID, BatchImageJobStatusFailed, BatchImageTransitionOptions{
ErrorCode: batchImageStringPtr("SETTLEMENT_BILLING_RETRY_EXHAUSTED"),
ErrorMessage: batchImageStringPtr(msg),
EventType: "settlement_retry_exhausted",
EventPayload: map[string]any{
"batch_id": job.BatchID,
"retry_count": job.RetryCount,
},
}); err != nil {
return err
}
return ErrBatchImageSettlementBillingFailed
}
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
}
return 72 * time.Hour
}
func BatchImageSettlementRequestID(batchID string) string {
return batchImageSettlementRequestPrefix + strings.TrimSpace(batchID)
}
func BuildBatchImageSettlementManifestHash(job *BatchImageJob) string {
if job == nil {
return ""
}
parts := []string{
strings.TrimSpace(job.BatchID),
strings.TrimSpace(job.Provider),
strings.TrimSpace(job.Model),
batchImageDerefString(job.ProviderJobName),
batchImageDerefString(job.ProviderOutputRef),
strconv.Itoa(job.SuccessCount),
strconv.Itoa(job.FailCount),
strconv.Itoa(job.ItemCount),
}
sum := sha256.Sum256([]byte(strings.Join(parts, "\x00")))
return hex.EncodeToString(sum[:])
}
type BatchImagePipelineProcessor struct {
ProviderProcessor *BatchImageProviderProcessor
SettlementService *BatchImageSettlementService
RetryDelay time.Duration
}
func (p *BatchImagePipelineProcessor) Process(ctx context.Context, batchID string) (BatchImageProcessResult, error) {
if p == nil || p.ProviderProcessor == nil {
return BatchImageProcessResult{}, errors.New("batch image pipeline processor is not configured")
}
job, err := p.ProviderProcessor.Repo.GetBatchImageJobByBatchID(ctx, batchID)
if err != nil {
return BatchImageProcessResult{}, err
}
if job.Status == BatchImageJobStatusSettling {
if p.SettlementService == nil {
return BatchImageProcessResult{Terminal: true}, nil
}
_, err := p.SettlementService.Settle(ctx, batchID)
if err != nil {
if errors.Is(err, ErrBatchImageSettlementBillingFailed) {
updated, getErr := p.ProviderProcessor.Repo.GetBatchImageJobByBatchID(ctx, batchID)
if getErr == nil && IsTerminalBatchImageJobStatus(updated.Status) {
return BatchImageProcessResult{Terminal: true}, nil
}
delay := p.RetryDelay
if delay <= 0 {
delay = batchImageSettlementRetryDelay
}
return BatchImageProcessResult{RequeueAfter: delay}, nil
}
return BatchImageProcessResult{}, err
}
return BatchImageProcessResult{Terminal: true}, nil
}
return p.ProviderProcessor.Process(ctx, batchID)
}
func (r *BatchImageSettlementResult) String() string {
if r == nil {
return ""
}
return fmt.Sprintf("batch_id=%s success=%d fail=%d actual_cost=%0.10f already_settled=%t",
r.BatchID, r.SuccessCount, r.FailCount, r.ActualCost, r.AlreadySettled)
}
@@ -0,0 +1,440 @@
//go:build unit
package service
import (
"context"
"errors"
"fmt"
"strings"
"testing"
"github.com/stretchr/testify/require"
)
func TestBatchImageSettlementService_SettlesAndChargesSuccessfulImagesOnly(t *testing.T) {
repo := newFakeBatchImageRepository()
job := testSettlingBatchImageJob("imgbatch_settle")
job.SuccessCount = 3
job.FailCount = 2
job.ItemCount = 5
repo.jobs[job.BatchID] = job
billing := &fakeBatchImageBillingRepo{}
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.75, result.ActualCost)
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.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) {
repo := newFakeBatchImageRepository()
job := testSettlingBatchImageJob("imgbatch_zero")
job.SuccessCount = 0
job.FailCount = 4
job.ItemCount = 4
repo.jobs[job.BatchID] = job
billing := &fakeBatchImageBillingRepo{}
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.0, result.ActualCost)
require.Equal(t, BatchImageJobStatusCompleted, repo.jobs[job.BatchID].Status)
require.Len(t, billing.captures, 1)
require.Equal(t, 0.0, billing.captures[0].ActualAmount)
}
func TestBatchImageSettlementService_CompletedJobReturnsAlreadySettledWithoutBilling(t *testing.T) {
repo := newFakeBatchImageRepository()
job := testSettlingBatchImageJob("imgbatch_done")
job.Status = BatchImageJobStatusCompleted
cost := 0.5
job.ActualCost = &cost
repo.jobs[job.BatchID] = job
billing := &fakeBatchImageBillingRepo{}
svc := &BatchImageSettlementService{Repo: repo, BillingRepo: billing, Pricing: &fakeBatchImagePricingResolver{unitPrice: 0.25}}
result, err := svc.Settle(context.Background(), job.BatchID)
require.NoError(t, err)
require.True(t, result.AlreadySettled)
require.Equal(t, 0.5, result.ActualCost)
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{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.captures, 1)
}
func TestBatchImageSettlementService_ValidationErrors(t *testing.T) {
tests := []struct {
name string
mutate func(*BatchImageJob)
pricing BatchImagePricingResolver
want error
}{
{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},
{name: "manifest_conflict", mutate: func(j *BatchImageJob) { v := "different"; j.ManifestHash = &v }, want: ErrBatchImageSettlementManifestConflict},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
repo := newFakeBatchImageRepository()
job := testSettlingBatchImageJob("imgbatch_" + tt.name)
if tt.mutate != nil {
tt.mutate(job)
}
repo.jobs[job.BatchID] = job
pricing := tt.pricing
if pricing == nil {
pricing = &fakeBatchImagePricingResolver{unitPrice: 0.25}
}
billing := &fakeBatchImageBillingRepo{}
svc := &BatchImageSettlementService{Repo: repo, BillingRepo: billing, Pricing: pricing}
_, err := svc.Settle(context.Background(), job.BatchID)
require.ErrorIs(t, err, tt.want)
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")
repo.jobs[job.BatchID] = job
billing := &fakeBatchImageBillingRepo{err: errors.New("temporary billing timeout with gs://hidden-output")}
svc := &BatchImageSettlementService{Repo: repo, BillingRepo: billing, Pricing: &fakeBatchImagePricingResolver{unitPrice: 0.25}}
_, err := svc.Settle(context.Background(), job.BatchID)
require.ErrorIs(t, err, ErrBatchImageSettlementBillingFailed)
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.captures[0])
}
func TestBatchImagePipelineProcessor_SettlesQueuedSettlingJob(t *testing.T) {
repo := newFakeBatchImageRepository()
job := testSettlingBatchImageJob("imgbatch_pipeline")
repo.jobs[job.BatchID] = job
billing := &fakeBatchImageBillingRepo{}
settlement := &BatchImageSettlementService{Repo: repo, BillingRepo: billing, Pricing: &fakeBatchImagePricingResolver{unitPrice: 0.25}}
processor := &BatchImagePipelineProcessor{
ProviderProcessor: &BatchImageProviderProcessor{Repo: repo, ProviderRegistry: NewBatchImageProviderRegistry(&fakeProcessorProvider{}), AccountResolver: &fakeBatchImageAccountResolver{account: &Account{}}},
SettlementService: settlement,
}
result, err := processor.Process(context.Background(), job.BatchID)
require.NoError(t, err)
require.True(t, result.Terminal)
require.Equal(t, BatchImageJobStatusCompleted, repo.jobs[job.BatchID].Status)
require.Len(t, billing.captures, 1)
}
func TestBatchImagePipelineProcessor_RequeuesTransientSettlementFailure(t *testing.T) {
repo := newFakeBatchImageRepository()
job := testSettlingBatchImageJob("imgbatch_pipeline_retry")
repo.jobs[job.BatchID] = job
settlement := &BatchImageSettlementService{Repo: repo, BillingRepo: &fakeBatchImageBillingRepo{err: errors.New("temporary")}, Pricing: &fakeBatchImagePricingResolver{unitPrice: 0.25}}
processor := &BatchImagePipelineProcessor{
ProviderProcessor: &BatchImageProviderProcessor{Repo: repo, ProviderRegistry: NewBatchImageProviderRegistry(&fakeProcessorProvider{}), AccountResolver: &fakeBatchImageAccountResolver{account: &Account{}}},
SettlementService: settlement,
}
result, err := processor.Process(context.Background(), job.BatchID)
require.NoError(t, err)
require.False(t, result.Terminal)
require.Equal(t, batchImageSettlementRetryDelay, result.RequeueAfter)
require.Equal(t, BatchImageJobStatusSettling, repo.jobs[job.BatchID].Status)
}
func TestBatchImagePipelineProcessor_FailsAndReleasesAfterSettlementRetryLimit(t *testing.T) {
repo := newFakeBatchImageRepository()
job := testSettlingBatchImageJob("imgbatch_pipeline_retry_exhausted")
job.RetryCount = batchImageSettlementMaxRetries - 1
repo.jobs[job.BatchID] = job
billing := &fakeBatchImageBillingRepo{captureErr: errors.New("temporary billing timeout")}
settlement := &BatchImageSettlementService{Repo: repo, BillingRepo: billing, Pricing: &fakeBatchImagePricingResolver{unitPrice: 0.25}}
processor := &BatchImagePipelineProcessor{
ProviderProcessor: &BatchImageProviderProcessor{Repo: repo, ProviderRegistry: NewBatchImageProviderRegistry(&fakeProcessorProvider{}), AccountResolver: &fakeBatchImageAccountResolver{account: &Account{}}},
SettlementService: settlement,
}
result, err := processor.Process(context.Background(), job.BatchID)
require.NoError(t, err)
require.True(t, result.Terminal)
require.Equal(t, BatchImageJobStatusFailed, repo.jobs[job.BatchID].Status)
require.Equal(t, "SETTLEMENT_BILLING_RETRY_EXHAUSTED", batchImageDerefString(repo.jobs[job.BatchID].LastErrorCode))
require.Len(t, billing.captures, 1)
require.Len(t, billing.releases, 1)
require.Equal(t, BatchImageReleaseRequestID(job.BatchID), billing.releases[0].RequestID)
}
func TestBatchImageSettlementRetryExhaustedReleaseIsIdempotentAfterTransitionFailure(t *testing.T) {
repo := newFakeBatchImageRepository()
job := testSettlingBatchImageJob("imgbatch_retry_exhausted_transition_fail")
job.RetryCount = batchImageSettlementMaxRetries
job.LastErrorCode = batchImageStringPtr("SETTLEMENT_BILLING_FAILED")
repo.jobs[job.BatchID] = job
repo.transitionErr = errors.New("temporary transition failure")
billing := &fakeBatchImageBillingRepo{}
svc := &BatchImageSettlementService{Repo: repo, BillingRepo: billing, Pricing: &fakeBatchImagePricingResolver{unitPrice: 0.25}}
_, err := svc.Settle(context.Background(), job.BatchID)
require.ErrorContains(t, err, "temporary transition failure")
require.Equal(t, BatchImageJobStatusSettling, repo.jobs[job.BatchID].Status)
require.Len(t, billing.releases, 1)
require.Len(t, billing.seen, 1)
repo.transitionErr = nil
_, err = svc.Settle(context.Background(), job.BatchID)
require.ErrorIs(t, err, ErrBatchImageSettlementBillingFailed)
require.Equal(t, BatchImageJobStatusFailed, repo.jobs[job.BatchID].Status)
require.Len(t, billing.releases, 2)
require.Equal(t, billing.releases[0].RequestID, billing.releases[1].RequestID)
require.Len(t, billing.seen, 1)
}
func TestBatchImageSettlementManifestHash(t *testing.T) {
job := testSettlingBatchImageJob("imgbatch_hash")
first := BuildBatchImageSettlementManifestHash(job)
job.CreatedAt = job.CreatedAt.AddDate(0, 0, 1)
job.UpdatedAt = job.UpdatedAt.AddDate(0, 0, 1)
require.Equal(t, first, BuildBatchImageSettlementManifestHash(job))
job.SuccessCount++
require.NotEqual(t, first, BuildBatchImageSettlementManifestHash(job))
job.SuccessCount--
promptOrBase64 := first + " prompt " + batchImageTestData
require.NotContains(t, BuildBatchImageSettlementManifestHash(job), promptOrBase64)
}
func TestBatchImageSettlementBillingRequestIDs(t *testing.T) {
repo := newFakeBatchImageRepository()
first := testSettlingBatchImageJob("imgbatch_unique_1")
second := testSettlingBatchImageJob("imgbatch_unique_2")
repo.jobs[first.BatchID] = first
repo.jobs[second.BatchID] = second
billing := &fakeBatchImageBillingRepo{}
svc := &BatchImageSettlementService{Repo: repo, BillingRepo: billing, Pricing: &fakeBatchImagePricingResolver{unitPrice: 0.25}}
_, err := svc.Settle(context.Background(), first.BatchID)
require.NoError(t, err)
_, err = svc.Settle(context.Background(), first.BatchID)
require.NoError(t, err)
_, err = svc.Settle(context.Background(), second.BatchID)
require.NoError(t, err)
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)
}
func testSettlingBatchImageJob(batchID string) *BatchImageJob {
apiKeyID := int64(321)
accountID := int64(654)
providerJobName := "providers/job"
outputRef := "files/output"
holdAmount := 1.25
holdID := BatchImageHoldRequestID(batchID)
return &BatchImageJob{
BatchID: batchID,
UserID: 123,
APIKeyID: &apiKeyID,
AccountID: &accountID,
Provider: BatchImageProviderGeminiAPI,
Model: "gemini-image",
Status: BatchImageJobStatusSettling,
ProviderJobName: &providerJobName,
ProviderOutputRef: &outputRef,
ItemCount: 3,
SuccessCount: 2,
FailCount: 1,
EstimatedCost: holdAmount,
HoldAmount: &holdAmount,
HoldID: &holdID,
}
}
type fakeBatchImagePricingResolver struct {
unitPrice float64
missingModels map[string]bool
err 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) {
if r.seen == nil {
r.seen = make(map[string]struct{})
}
if r.err != nil {
r.commands = append(r.commands, cmd)
return nil, r.err
}
if cmd != nil {
cmd.Normalize()
if _, ok := r.seen[cmd.RequestID]; ok || r.alreadyApplied[cmd.RequestID] {
r.commands = append(r.commands, cmd)
return &UsageBillingApplyResult{Applied: false}, nil
}
r.seen[cmd.RequestID] = struct{}{}
}
r.commands = append(r.commands, cmd)
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

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