Files
NianAIGC/backend/internal/orchestration/orchestration.go

378 lines
12 KiB
Go

// Package orchestration adapts jobs.Worker ports to the billing, usage,
// webhook, and asset modules. It contains cross-module policy but owns no
// persistence or external transport.
package orchestration
import (
"context"
"encoding/json"
"errors"
"fmt"
"path"
"strings"
"time"
"git.nianxx.cn/wangxuming/NianAIGC/backend/internal/assets"
"git.nianxx.cn/wangxuming/NianAIGC/backend/internal/billing"
"git.nianxx.cn/wangxuming/NianAIGC/backend/internal/jobs"
"git.nianxx.cn/wangxuming/NianAIGC/backend/internal/usage"
"git.nianxx.cn/wangxuming/NianAIGC/backend/internal/webhook"
)
type UsageSink interface {
Record(usage.Event) (*usage.Event, error)
}
type UsageRecorder struct {
sink UsageSink
newID func() string
now func() time.Time
}
func NewUsageRecorder(sink UsageSink, newID func() string, now func() time.Time) *UsageRecorder {
if now == nil {
now = time.Now
}
return &UsageRecorder{sink: sink, newID: newID, now: now}
}
func (r *UsageRecorder) Record(_ context.Context, job jobs.Job) error {
if job.Provider == "mock" || job.ExternalClientID != "" {
return nil
}
var contextValue usageContext
if len(job.UsageContext) != 0 && json.Unmarshal(job.UsageContext, &contextValue) != nil {
return errors.New("record generation usage")
}
if contextValue.Source == "api" {
return nil
}
if r == nil || r.sink == nil || r.newID == nil {
return errors.New("record generation usage")
}
var charge billingState
if len(job.Billing) != 0 && json.Unmarshal(job.Billing, &charge) != nil {
return errors.New("record generation usage")
}
quantity := charge.Quantity
if quantity <= 0 {
quantity = 1
}
unit := "job"
if charge.Unit == "image" || charge.Unit == "video_second" {
unit = charge.Unit
}
event := usage.Event{
ID: r.newID(), OwnerID: job.OwnerID, JobID: job.ID, Source: "platform",
Capability: job.Capability, Provider: job.Provider, ReqKey: job.ReqKey,
AccountUsername: contextValue.Username, AccountDisplayName: contextValue.DisplayName,
TenantID: contextValue.TenantID, OrganizationID: contextValue.OrganizationID,
OrganizationName: contextValue.OrganizationName, Quantity: quantity,
EstimatedUnit: unit, ChargedAmountFen: charge.amountPointer(), Currency: charge.Currency,
CreatedAt: r.now().UTC().Format(time.RFC3339Nano),
}
if _, err := r.sink.Record(event); err != nil {
return errors.New("record generation usage")
}
return nil
}
type RefundLedger interface {
Refund(context.Context, billing.RefundRequest) (billing.WalletPosting, error)
}
// JobStateWriter is the deliberately narrow persistence seam needed because
// jobs.Patch does not expose billing or output_asset_ids fields.
type JobStateWriter interface {
WriteBilling(context.Context, string, json.RawMessage) error
WriteOutputAssetIDs(context.Context, string, []string) error
}
type FencedJobStateWriter interface {
WriteBillingFenced(context.Context, string, json.RawMessage, jobs.Status, string) error
WriteOutputAssetIDsFenced(context.Context, string, []string, jobs.Status, string) error
}
type TerminalRefund struct {
ledger RefundLedger
state JobStateWriter
}
func NewTerminalRefund(ledger RefundLedger, state JobStateWriter) *TerminalRefund {
return &TerminalRefund{ledger: ledger, state: state}
}
func (r *TerminalRefund) Refund(ctx context.Context, job jobs.Job, reason string) (jobs.Job, error) {
if !job.Status.Terminal() || job.Status == jobs.StatusSucceeded {
return job, nil
}
var charge billingState
if len(job.Billing) == 0 {
return job, nil
}
if err := json.Unmarshal(job.Billing, &charge); err != nil {
return jobs.Job{}, errors.New("refund generation charge")
}
if charge.Status != "charged" || charge.QuotaExempt {
return job, nil
}
var use usageContext
if json.Unmarshal(job.UsageContext, &use) != nil || use.OrganizationID == "" || charge.AmountFen <= 0 || r == nil || r.ledger == nil || r.state == nil {
return jobs.Job{}, errors.New("refund generation charge")
}
posting, err := r.ledger.Refund(ctx, billing.RefundRequest{
OrganizationID: use.OrganizationID, AccountID: use.AccountID, JobID: job.ID,
AmountFen: charge.AmountFen, Description: capabilityLabel(job.Capability) + "失败退款 · " + reason,
Metadata: map[string]any{"chargeLedgerEntryId": charge.LedgerEntryID, "reason": reason, "quote": cloneMap(charge.raw)},
})
if err != nil {
return jobs.Job{}, errors.New("refund generation charge")
}
charge.raw["status"] = "refunded"
charge.raw["refundLedgerEntryId"] = posting.LedgerID
charge.raw["refundedAt"] = posting.CreatedAt.UTC().Format(time.RFC3339Nano)
charge.raw["refundReason"] = reason
encoded, err := json.Marshal(charge.raw)
if err != nil || writeBillingState(ctx, r.state, job, encoded) != nil {
return jobs.Job{}, errors.New("persist generation refund")
}
job.Billing = encoded
return job, nil
}
type WebhookDeliverer interface {
Deliver(context.Context, jobs.Job) (webhook.Result, error)
}
type WebhookBridge struct{ deliverer WebhookDeliverer }
func NewWebhookBridge(deliverer WebhookDeliverer) *WebhookBridge {
return &WebhookBridge{deliverer: deliverer}
}
func (b *WebhookBridge) Deliver(ctx context.Context, job jobs.Job) (jobs.WebhookResult, error) {
if b == nil || b.deliverer == nil {
return jobs.WebhookResult{}, errors.New("deliver generation webhook")
}
result, err := b.deliverer.Deliver(ctx, job)
if err != nil {
return jobs.WebhookResult{}, errors.New("deliver generation webhook")
}
return jobs.WebhookResult{Attempts: result.Attempts, LastStatus: result.LastStatus}, nil
}
type OutputRegistrar interface {
Register(context.Context, jobs.Job) ([]string, error)
}
type OutputRegisteringProcessor struct {
inner jobs.Processor
registrar OutputRegistrar
state JobStateWriter
}
func NewOutputRegisteringProcessor(inner jobs.Processor, registrar OutputRegistrar, state JobStateWriter) *OutputRegisteringProcessor {
return &OutputRegisteringProcessor{inner: inner, registrar: registrar, state: state}
}
func (p *OutputRegisteringProcessor) Advance(ctx context.Context, job jobs.Job) (jobs.Job, error) {
if job.Status == jobs.StatusSucceeded && len(job.OutputAssetIDs) != 0 {
return job, nil
}
if p == nil || p.inner == nil {
return jobs.Job{}, errors.New("advance generation job")
}
advanced, err := p.inner.Advance(ctx, job)
if err != nil || advanced.Status != jobs.StatusSucceeded || len(advanced.OutputAssetIDs) != 0 {
return advanced, err
}
if p.registrar == nil || p.state == nil {
return jobs.Job{}, errors.New("register generation outputs")
}
ids, err := p.registrar.Register(ctx, advanced)
if err != nil || len(ids) == 0 {
return jobs.Job{}, errors.New("register generation outputs")
}
if err := writeOutputState(ctx, p.state, advanced, ids); err != nil {
return jobs.Job{}, errors.New("persist generation outputs")
}
advanced.OutputAssetIDs = append([]string(nil), ids...)
return advanced, nil
}
func writeBillingState(ctx context.Context, state JobStateWriter, job jobs.Job, value json.RawMessage) error {
if fenced, ok := state.(FencedJobStateWriter); ok && job.LockedBy != "" {
return fenced.WriteBillingFenced(ctx, job.ID, value, job.Status, job.LockedBy)
}
return state.WriteBilling(ctx, job.ID, value)
}
func writeOutputState(ctx context.Context, state JobStateWriter, job jobs.Job, ids []string) error {
if fenced, ok := state.(FencedJobStateWriter); ok && job.LockedBy != "" {
return fenced.WriteOutputAssetIDsFenced(ctx, job.ID, ids, job.Status, job.LockedBy)
}
return state.WriteOutputAssetIDs(ctx, job.ID, ids)
}
type GeneratedAssetImporter interface {
List(context.Context, assets.Scope) ([]assets.Asset, error)
ImportGenerated(context.Context, assets.Scope, assets.ImportGeneratedCommand) (assets.Asset, error)
ImportMock(context.Context, assets.Scope, assets.ImportMockCommand) (assets.Asset, error)
}
type OutputURLResolver func(jobs.Job) ([]string, error)
type AssetOutputRegistrar struct {
assets GeneratedAssetImporter
resolve OutputURLResolver
}
func NewAssetOutputRegistrar(service GeneratedAssetImporter, resolve OutputURLResolver) *AssetOutputRegistrar {
return &AssetOutputRegistrar{assets: service, resolve: resolve}
}
func (r *AssetOutputRegistrar) Register(ctx context.Context, job jobs.Job) ([]string, error) {
if r == nil || r.assets == nil || r.resolve == nil {
return nil, errors.New("register generation outputs")
}
scope := assets.PlatformScope(job.OwnerID)
existing, err := r.assets.List(ctx, scope)
if err != nil {
return nil, errors.New("register generation outputs")
}
if job.Provider == "mock" {
if id := existingOutputID(existing, job.ID, "output:0"); id != "" {
return []string{id}, nil
}
kind := assets.KindImage
if job.Capability == "video.generate" {
kind = assets.KindVideo
}
created, createErr := r.assets.ImportMock(ctx, scope, assets.ImportMockCommand{
Capability: job.Capability, JobID: job.ID, Kind: kind,
Tags: []string{"generated", job.Capability, "job:" + job.ID, "output:0"},
Metadata: map[string]any{"capability": job.Capability, "jobId": job.ID, "index": 0},
})
if createErr != nil {
return nil, errors.New("register generation outputs")
}
return []string{created.ID}, nil
}
urls, err := r.resolve(job)
if err != nil || len(urls) == 0 {
return nil, errors.New("register generation outputs")
}
ids := make([]string, 0, len(urls))
for index, rawURL := range urls {
indexTag := fmt.Sprintf("output:%d", index)
if id := existingOutputID(existing, job.ID, indexTag); id != "" {
ids = append(ids, id)
continue
}
kind := assets.KindImage
if job.Capability == "video.generate" {
kind = assets.KindVideo
}
name := path.Base(strings.SplitN(rawURL, "?", 2)[0])
if name == "." || name == "/" || name == "" {
name = fmt.Sprintf("%s-%d", strings.ReplaceAll(job.Capability, ".", "-"), index+1)
}
created, createErr := r.assets.ImportGenerated(ctx, scope, assets.ImportGeneratedCommand{
URL: rawURL, Name: name, Kind: kind, Source: assets.SourceGenerated,
Tags: []string{"generated", job.Capability, "job:" + job.ID, indexTag},
Metadata: map[string]any{"capability": job.Capability, "jobId": job.ID, "index": index},
})
if createErr != nil {
return nil, errors.New("register generation outputs")
}
ids = append(ids, created.ID)
}
return ids, nil
}
func existingOutputID(existing []assets.Asset, jobID, indexTag string) string {
for _, asset := range existing {
if contains(asset.Tags, "job:"+jobID) && contains(asset.Tags, indexTag) {
return asset.ID
}
}
return ""
}
func contains(values []string, want string) bool {
for _, value := range values {
if value == want {
return true
}
}
return false
}
type usageContext struct {
Source string `json:"source"`
AccountID string `json:"accountId"`
Username string `json:"username"`
DisplayName string `json:"displayName"`
Role string `json:"role,omitempty"`
TenantID string `json:"tenantId"`
OrganizationID string `json:"organizationId"`
OrganizationName string `json:"organizationName"`
}
type billingState struct {
Status, Unit, Currency, LedgerEntryID string
Quantity int
AmountFen int64
QuotaExempt bool
raw map[string]any
}
func (s *billingState) UnmarshalJSON(data []byte) error {
type alias billingState
var decoded struct {
Status, Unit, Currency, LedgerEntryID string
Quantity float64
AmountFen int64
QuotaExempt bool
}
if err := json.Unmarshal(data, &decoded); err != nil {
return err
}
if err := json.Unmarshal(data, &s.raw); err != nil {
return err
}
s.Status, s.Unit, s.Currency, s.LedgerEntryID = decoded.Status, decoded.Unit, decoded.Currency, decoded.LedgerEntryID
s.Quantity, s.AmountFen, s.QuotaExempt = int(decoded.Quantity), decoded.AmountFen, decoded.QuotaExempt
return nil
}
func (s billingState) amountPointer() *int64 {
if _, exists := s.raw["amountFen"]; !exists {
return nil
}
value := s.AmountFen
return &value
}
func capabilityLabel(capability string) string {
if capability == "video.generate" {
return "视频生成"
}
return "图片生成"
}
func cloneMap(source map[string]any) map[string]any {
clone := make(map[string]any, len(source))
for key, value := range source {
clone[key] = value
}
return clone
}
var (
_ jobs.UsageRecorder = (*UsageRecorder)(nil)
_ jobs.TerminalRefund = (*TerminalRefund)(nil)
_ jobs.WebhookDelivery = (*WebhookBridge)(nil)
_ jobs.Processor = (*OutputRegisteringProcessor)(nil)
)