mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-10-07 15:38:22 +08:00
Merge pull request #3768 from Turtle-Li/feature/batch-image-foundation
feat: add batch image generation MVP
This commit is contained in:
+6
-1
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -65,6 +65,8 @@ func TestProvideCleanup_WithMinimalDependencies_NoPanic(t *testing.T) {
|
||||
subscriptionExpirySvc,
|
||||
&service.UsageCleanupService{},
|
||||
idempotencyCleanupSvc,
|
||||
&service.BatchImageCleanupService{},
|
||||
nil, // batchImageWorker
|
||||
pricingSvc,
|
||||
emailQueueSvc,
|
||||
billingCacheSvc,
|
||||
|
||||
@@ -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()
|
||||
}
|
||||
@@ -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))
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
@@ -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
|
||||
@@ -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
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
@@ -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
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
@@ -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
@@ -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(", ")
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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))
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
File diff suppressed because it is too large
Load Diff
@@ -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
@@ -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()
|
||||
|
||||
@@ -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 <> ''")),
|
||||
}
|
||||
}
|
||||
@@ -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"),
|
||||
}
|
||||
}
|
||||
@@ -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"),
|
||||
}
|
||||
}
|
||||
@@ -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").
|
||||
|
||||
@@ -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").
|
||||
|
||||
@@ -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
@@ -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(", ")
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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))
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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,
|
||||
},
|
||||
})
|
||||
}
|
||||
@@ -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,
|
||||
|
||||
@@ -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"`
|
||||
|
||||
@@ -58,6 +58,7 @@ type Handlers struct {
|
||||
Payment *PaymentHandler
|
||||
PaymentWebhook *PaymentWebhookHandler
|
||||
AvailableChannel *AvailableChannelHandler
|
||||
BatchImage *BatchImageHandler
|
||||
}
|
||||
|
||||
// BuildInfo contains build-time information
|
||||
|
||||
@@ -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(¤t); 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(¤t); 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())
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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)))
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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 = ¶ms.ActualCost
|
||||
job.ManifestHash = ¶ms.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
Reference in New Issue
Block a user