Files
sub2api/backend/internal/service/batch_image_public_test.go
T
李建琦 6d655c9903
Release / update-version (push) Has been cancelled
Release / build-frontend (push) Has been cancelled
Release / release (push) Has been cancelled
Release / sync-version-file (push) Has been cancelled
CI / shell (push) Canceled after 0s
CI / test (push) Canceled after 0s
CI / frontend (push) Canceled after 0s
CI / golangci-lint (push) Canceled after 0s
Security Scan / backend-security (push) Canceled after 0s
Security Scan / frontend-security (push) Canceled after 0s
Sub2API v1.0 - AI API 网关(二开初始版本,基于上游 Wei-Shaw/sub2api)
2026-08-21 18:30:13 +08:00

970 lines
37 KiB
Go

//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)
req := validBatchImageSubmitRequest()
req.SessionID = batchImageStringPtr("batch-session-123")
got, err := svc.Submit(ctx, testBatchImageOwner(), req, "")
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)
require.Equal(t, "batch-session-123", batchImageDerefString(job.SessionID))
})
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)
// 配置的 hold(0.6) < discount(0.8) 属于会导致结算死锁的脏数据,
// 快照时被钳制为 discount,保证 holdAmount >= 实际成本上限。
require.InDelta(t, 0.8, job.HoldMultiplier, 1e-12)
require.InDelta(t, 0.125, job.BillableUnitPrice, 1e-12)
require.InDelta(t, 0.125, job.HoldUnitPrice, 1e-12)
require.InDelta(t, 0.25, *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, repo, queue, gemini, _ := newTestBatchImagePublicService(true)
req := validBatchImageSubmitRequest()
req.SessionID = batchImageStringPtr("original-session")
first, err := svc.Submit(ctx, testBatchImageOwner(), req, "client-key")
require.NoError(t, err)
req.SessionID = batchImageStringPtr("retry-session")
second, err := svc.Submit(ctx, testBatchImageOwner(), req, "client-key")
require.NoError(t, err)
require.Equal(t, first.ID, second.ID)
require.Equal(t, "original-session", batchImageDerefString(repo.jobs[first.ID].SessionID))
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)