378 lines
12 KiB
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)
|
|
)
|