Files
NianAIGC/backend/internal/jobs/provider.go

559 lines
18 KiB
Go

package jobs
import (
"context"
"encoding/json"
"errors"
"fmt"
"math"
"net/url"
"strings"
"time"
"git.nianxx.cn/wangxuming/NianAIGC/backend/internal/prompt"
"git.nianxx.cn/wangxuming/NianAIGC/backend/internal/providers"
)
type ProviderRegistry map[string]providers.Adapter
type ProviderProcessor struct {
Providers ProviderRegistry
Store Store
}
func (p ProviderProcessor) Advance(ctx context.Context, job Job) (Job, error) {
// A persisted terminal provider result is the recovery checkpoint. Replaying
// it must run downstream idempotent finalizers, never submit/query again.
if job.Status.Terminal() {
return job, nil
}
adapter := p.Providers[job.Provider]
if adapter == nil {
return Job{}, fmt.Errorf("generation provider %q is unavailable", job.Provider)
}
var request providers.Request
if err := json.Unmarshal(job.RequestPayload, &request); err != nil {
return Job{}, errors.New("invalid provider request payload")
}
var result providers.Result
var err error
expectedStatus := job.Status
if job.ProviderTaskID == "" {
if job.ProviderDispatchStartedAt != nil {
failed := StatusFailed
failure := &JobError{Message: "provider submission outcome is unknown; refusing duplicate submission", Retryable: false}
if p.Store == nil {
job.Status, job.Error = failed, failure
return job, nil
}
return p.Store.UpdateJob(ctx, job.ID, workerPatch(job, Patch{Status: &failed, Error: failure}))
}
if p.Store != nil {
started := time.Now().UTC()
prepared, persistErr := p.Store.UpdateJob(ctx, job.ID, workerPatch(job, Patch{ProviderDispatchStartedAt: &started}))
if persistErr != nil {
return Job{}, persistErr
}
job = prepared
}
result, err = adapter.Submit(ctx, request)
} else {
if modeled, ok := adapter.(providers.ModelQueryAdapter); ok {
result, err = modeled.QueryModel(ctx, job.ProviderTaskID, job.ReqKey)
} else {
result, err = adapter.Query(ctx, job.ProviderTaskID)
}
}
if err != nil && job.ProviderTaskID == "" && p.Store != nil {
failed := StatusFailed
failure := &JobError{Message: "provider submission outcome is unknown; refusing duplicate submission", Retryable: false}
return p.Store.UpdateJob(ctx, job.ID, workerPatch(job, Patch{Status: &failed, Error: failure}))
}
if err != nil {
return Job{}, err
}
job.ProviderTaskID = result.TaskID
encoded, err := providers.EncodeResult(result)
if err != nil {
return Job{}, errors.New("encode provider result")
}
job.ResponsePayload = encoded
job.Status = Status(result.Status)
if result.ErrorMessage != "" {
job.Error = &JobError{Message: result.ErrorMessage}
}
if p.Store != nil {
patch := Patch{Status: &job.Status, ProviderTaskID: &job.ProviderTaskID, ResponsePayload: job.ResponsePayload, SetResponsePayload: true, ClearError: result.ErrorMessage == ""}
patch.ExpectedStatuses = []Status{expectedStatus}
if job.LockedBy != "" {
worker := job.LockedBy
patch.ExpectedLockedBy = &worker
}
if job.Error != nil {
patch.Error = job.Error
}
return p.Store.UpdateJob(ctx, job.ID, patch)
}
return job, nil
}
type ProviderJobBuilder struct {
ImageProvider, VideoProvider string
ImageModel, VideoModel string
ImageEngine, VideoEngine string
ImageEngines map[string]ProviderTarget
VideoEngines map[string]ProviderTarget
NewID func() string
}
// ProviderTarget is server-owned routing configuration for one UI engine.
// The request may select a known engine, but it cannot supply a provider or
// model directly.
type ProviderTarget struct {
Provider string
Model string
Settings map[string]any
}
func (b ProviderJobBuilder) Build(_ context.Context, owner, client, capability, idempotency string, body map[string]any) (CreateCommand, error) {
if capability != "image.generate" && capability != "video.generate" {
return CreateCommand{}, &Error{Kind: ErrorInvalid, Status: 400, Message: "Unsupported capability: " + capability}
}
target, engine, err := b.target(capability, body["engine"])
if err != nil {
return CreateCommand{}, err
}
if target.Provider == "" || target.Model == "" || b.NewID == nil {
return CreateCommand{}, errors.New("provider job builder is not configured")
}
prepared, err := prepareProviderRequest(capability, engine, target.Settings, body)
if err != nil {
return CreateCommand{}, err
}
prepared.request.Model = target.Model
raw, err := json.Marshal(prepared.request)
if err != nil {
return CreateCommand{}, errors.New("encode provider request")
}
priority := NormalizePriority(intFromAny(body["priority"]))
webhook, _ := body["webhookUrl"].(string)
webhook = strings.TrimSpace(webhook)
if client != "" && webhook != "" && !validWebhookURL(webhook) {
return CreateCommand{}, invalidPreparation("webhookUrl must be an HTTP or HTTPS URL")
}
return CreateCommand{Job: Job{ID: b.NewID(), OwnerID: owner, ExternalClientID: client, Capability: capability, Provider: target.Provider, ReqKey: target.Model, Status: StatusQueued, Prompt: prepared.request.Prompt, InputURLs: prepared.request.InputURLs, InputAssetIDs: prepared.assetIDs, OutputAssetIDs: []string{}, RequestPayload: raw, IdempotencyKey: idempotency, Priority: priority, WebhookURL: webhook}, IdempotencyBody: body}, nil
}
func validWebhookURL(value string) bool {
parsed, err := url.Parse(value)
return err == nil && parsed.IsAbs() && parsed.Host != "" && parsed.User == nil && (parsed.Scheme == "http" || parsed.Scheme == "https")
}
func (b ProviderJobBuilder) target(capability string, rawEngine any) (ProviderTarget, string, error) {
engine := engineName(rawEngine)
if capability == "image.generate" {
if rawEngine != nil && engine == "" {
return ProviderTarget{}, "", invalidPreparation("unsupported image engine")
}
if engine != "" {
if engine != "jimeng" && engine != "evolink" && engine != "bailian" {
return ProviderTarget{}, "", invalidPreparation("unsupported image engine")
}
target, ok := b.ImageEngines[engine]
if !ok || target.Provider == "" || target.Model == "" {
return ProviderTarget{}, "", errors.New("provider job builder is not configured")
}
return target, engine, nil
}
return ProviderTarget{Provider: b.ImageProvider, Model: b.ImageModel}, firstConfiguredEngine(b.ImageEngine, b.ImageProvider, "image"), nil
}
if rawEngine != nil && engine == "" {
return ProviderTarget{}, "", invalidPreparation("unsupported video engine")
}
if engine != "" {
if engine != "seedance" && engine != "bailian" {
return ProviderTarget{}, "", invalidPreparation("unsupported video engine")
}
target, ok := b.VideoEngines[engine]
if !ok || target.Provider == "" || target.Model == "" {
return ProviderTarget{}, "", errors.New("provider job builder is not configured")
}
return target, engine, nil
}
return ProviderTarget{Provider: b.VideoProvider, Model: b.VideoModel}, firstConfiguredEngine(b.VideoEngine, b.VideoProvider, "video"), nil
}
func firstConfiguredEngine(configured, provider, capabilityKind string) string {
if engine := engineName(configured); engine != "" {
return engine
}
if capabilityKind == "image" {
switch provider {
case "volcengine-visual":
return "jimeng"
case "evolink", "bailian":
return provider
}
} else if provider == "seedance" || provider == "bailian" {
return provider
}
return ""
}
type preparedProviderRequest struct {
request providers.Request
assetIDs []string
}
func prepareProviderRequest(capability, engine string, defaults map[string]any, body map[string]any) (preparedProviderRequest, error) {
materials, assembly, err := preparationMaterials(body, capability)
if err != nil {
return preparedProviderRequest{}, err
}
text := stringValue(body["prompt"])
if text == "" && assembly != nil {
text = prompt.Assemble(*assembly).Prompt
}
text = strings.TrimSpace(text)
if text == "" {
return preparedProviderRequest{}, invalidPreparation("prompt is required")
}
assetIDs := stringsFromAny(body["inputAssetIds"])
if len(assetIDs) == 0 {
for _, material := range materials {
if material.ID != "" {
assetIDs = append(assetIDs, material.ID)
}
}
}
var urls []string
var settings map[string]any
if capability == "image.generate" {
urls = stringsFromAny(firstNonNil(body["imageUrls"], body["inputUrls"]))
if len(urls) == 0 {
for _, material := range materials {
if material.Type == "image" {
urls = append(urls, material.URL)
}
}
}
if err := validateImageCoverage(text, materials, len(urls)); err != nil {
return preparedProviderRequest{}, err
}
settings, err = imageSettings(body)
if err == nil && engine == "bailian" {
err = validateBailianImage(urls, settings)
}
} else {
if err := validateMaterialCoverage(text, materials); err != nil {
return preparedProviderRequest{}, err
}
for _, material := range materials {
urls = append(urls, material.URL)
}
settings, err = videoSettings(engine, defaults, body["settings"], materials)
}
if err != nil {
return preparedProviderRequest{}, err
}
providerMaterials := make([]providers.Material, 0, len(materials))
for _, material := range materials {
materialType := providers.MaterialImage
switch material.Type {
case "video":
materialType = providers.MaterialVideo
case "audio":
materialType = providers.MaterialAudio
}
providerMaterials = append(providerMaterials, providers.Material{
URL: material.URL, Type: materialType, Role: material.Role, Label: material.Label,
})
}
return preparedProviderRequest{request: providers.Request{Capability: capability, Prompt: text, InputURLs: urls, Materials: providerMaterials, Settings: settings}, assetIDs: assetIDs}, nil
}
func validateImageCoverage(text string, materials []prompt.Material, imageURLCount int) error {
required := prompt.ExtractRequirements(text)
if required.Video > 0 || required.Audio > 0 {
return validateMaterialCoverage(text, materials)
}
if required.Image > imageURLCount {
return invalidPreparation(fmt.Sprintf("prompt requires @图片%d but only %d image materials were supplied", required.Image, imageURLCount))
}
return nil
}
func preparationMaterials(body map[string]any, capability string) ([]prompt.Material, *prompt.Input, error) {
assemblyRecord, hasAssembly := body["promptAssembly"].(map[string]any)
var assembly *prompt.Input
if hasAssembly {
var decoded prompt.Input
if err := decodeViaJSON(assemblyRecord, &decoded); err != nil {
return nil, nil, invalidPreparation("invalid promptAssembly")
}
if capability == "image.generate" {
decoded.Mode = "image"
} else {
decoded.Mode = "video"
}
assembly = &decoded
}
var materials []prompt.Material
if raw, ok := body["materials"]; ok {
if err := decodeViaJSON(raw, &materials); err != nil {
return nil, nil, invalidPreparation("invalid materials")
}
} else if assembly != nil {
materials = assembly.Materials
}
materials = prompt.NormalizeMaterials(materials)
if assembly != nil {
assembly.Materials = materials
}
return materials, assembly, nil
}
func validateMaterialCoverage(text string, materials []prompt.Material) error {
required := prompt.ExtractRequirements(text)
available := prompt.Requirements{}
for _, material := range materials {
switch material.Type {
case "video":
available.Video++
case "audio":
available.Audio++
default:
available.Image++
}
}
if required.Image > available.Image {
return invalidPreparation(fmt.Sprintf("prompt requires @图片%d but only %d image materials were supplied", required.Image, available.Image))
}
if required.Video > available.Video {
return invalidPreparation(fmt.Sprintf("prompt requires @视频%d but only %d video materials were supplied", required.Video, available.Video))
}
if required.Audio > available.Audio {
return invalidPreparation(fmt.Sprintf("prompt requires @音频%d but only %d audio materials were supplied", required.Audio, available.Audio))
}
return nil
}
func imageSettings(body map[string]any) (map[string]any, error) {
out := map[string]any{}
if nested, ok := body["settings"].(map[string]any); ok {
for _, key := range []string{"scale", "width", "height", "min_ratio", "max_ratio", "imageCount", "force_single", "quality"} {
if nested[key] != nil {
out[key] = nested[key]
}
}
}
for _, key := range []string{"scale", "width", "height", "min_ratio", "max_ratio", "imageCount"} {
raw := body[key]
if raw == nil {
raw = out[key]
}
if value, ok := finiteNumber(raw); ok {
out[key] = value
} else if raw != nil && raw != "" {
return nil, invalidPreparation("invalid image parameter: " + key)
} else {
delete(out, key)
}
}
forceSingle := body["force_single"]
if forceSingle == nil {
forceSingle = out["force_single"]
}
if value, ok := forceSingle.(bool); ok {
out["force_single"] = value
} else if forceSingle != nil {
return nil, invalidPreparation("invalid image parameter: force_single")
} else {
delete(out, "force_single")
}
qualityRaw := body["quality"]
if qualityRaw == nil {
qualityRaw = out["quality"]
}
delete(out, "quality")
if raw, ok := qualityRaw.(string); ok {
quality := strings.ToLower(strings.TrimSpace(raw))
if quality == "low" || quality == "medium" || quality == "high" {
out["quality"] = quality
}
}
if count, ok := out["imageCount"].(float64); ok && (count <= 0 || count > 9 || math.Trunc(count) != count) {
return nil, invalidPreparation("imageCount must be an integer between 1 and 9")
}
return out, nil
}
func validateBailianImage(urls []string, settings map[string]any) error {
if len(urls) > 9 {
return invalidPreparation("bailian supports at most 9 reference images")
}
width, hasWidth := settings["width"].(float64)
height, hasHeight := settings["height"].(float64)
if !hasWidth && !hasHeight {
return nil
}
if !hasWidth || !hasHeight || math.Trunc(width) != width || math.Trunc(height) != height || width <= 0 || height <= 0 {
return invalidPreparation("bailian image dimensions must be positive integers")
}
pixels := width * height
maximum := float64(4096 * 4096)
if len(urls) > 0 {
maximum = float64(2048 * 2048)
}
if pixels < float64(768*768) || pixels > maximum || width/height < 1.0/8.0 || width/height > 8 {
return invalidPreparation("bailian image dimensions are outside supported size constraints")
}
return nil
}
func videoSettings(engine string, defaults map[string]any, raw any, materials []prompt.Material) (map[string]any, error) {
input, _ := raw.(map[string]any)
merged := make(map[string]any, len(defaults)+len(input))
for key, value := range defaults {
merged[key] = value
}
for key, value := range input {
merged[key] = value
}
input = merged
settings := map[string]any{}
if engine == "bailian" {
if len(materials) < 1 || len(materials) > 2 {
return nil, invalidPreparation("bailian video requires 1 or 2 image materials")
}
for _, material := range materials {
if material.Type != "image" {
return nil, invalidPreparation("bailian video requires 1 or 2 image materials")
}
}
duration := float64(10)
if input["duration"] != nil {
var ok bool
duration, ok = finiteNumber(input["duration"])
if !ok || math.Trunc(duration) != duration || duration < 2 || duration > 15 {
return nil, invalidPreparation("bailian video duration must be an integer between 2 and 15 seconds")
}
}
resolution := strings.ToUpper(stringValue(input["resolution"]))
if resolution == "" {
resolution = "720P"
}
if resolution != "720P" && resolution != "1080P" {
return nil, invalidPreparation("bailian video resolution must be 720P or 1080P")
}
return map[string]any{"duration": duration, "resolution": resolution}, nil
}
if engine != "" && engine != "seedance" {
return nil, invalidPreparation("unsupported video engine")
}
ratio := stringValue(input["ratio"])
if ratio != "" && !oneOf(ratio, "9:16", "16:9", "1:1", "4:3", "3:4", "21:9", "adaptive") {
return nil, invalidPreparation("unsupported video ratio")
}
if ratio != "" {
settings["ratio"] = ratio
}
if input["duration"] != nil {
duration, ok := finiteNumber(input["duration"])
if !ok || math.Trunc(duration) != duration || duration < 4 || duration > 15 {
return nil, invalidPreparation("video duration must be an integer between 4 and 15 seconds")
}
settings["duration"] = duration
}
resolution := stringValue(input["resolution"])
if resolution != "" && !oneOf(resolution, "480p", "720p", "1080p", "4k") {
return nil, invalidPreparation("unsupported video resolution")
}
if resolution != "" {
settings["resolution"] = resolution
}
return settings, nil
}
func invalidPreparation(message string) error {
return &Error{Kind: ErrorInvalid, Status: 400, Message: message}
}
func decodeViaJSON(input, output any) error {
raw, err := json.Marshal(input)
if err != nil {
return err
}
return json.Unmarshal(raw, output)
}
func finiteNumber(value any) (float64, bool) {
var number float64
switch typed := value.(type) {
case float64:
number = typed
case float32:
number = float64(typed)
case int:
number = float64(typed)
case int64:
number = float64(typed)
default:
return 0, false
}
return number, !math.IsNaN(number) && !math.IsInf(number, 0)
}
func stringValue(value any) string {
text, _ := value.(string)
return strings.TrimSpace(text)
}
func engineName(value any) string { return strings.ToLower(stringValue(value)) }
func oneOf(value string, options ...string) bool {
for _, option := range options {
if value == option {
return true
}
}
return false
}
func stringsFromAny(v any) []string {
if values, ok := v.([]string); ok {
out := make([]string, 0, len(values))
for _, value := range values {
if value = strings.TrimSpace(value); value != "" {
out = append(out, value)
}
}
return out
}
values, _ := v.([]any)
out := []string{}
for _, x := range values {
if s, ok := x.(string); ok && strings.TrimSpace(s) != "" {
out = append(out, strings.TrimSpace(s))
}
}
return out
}
func intFromAny(v any) int {
switch x := v.(type) {
case float64:
return int(x)
case int:
return x
}
return 0
}
func firstNonNil(v ...any) any {
for _, x := range v {
if x != nil {
return x
}
}
return nil
}