feat: complete remaining Go backend modules
This commit is contained in:
1 parent
cea2751dc5
commit
aef5a97165
145 files changed
+18376
-199
No files matched your search
@@ -0,0 +1,377 @@
|
||||
// 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)
|
||||
)
|
||||
Reference in new issue
Block a user