// 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) )